use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::{fmt, io};
pub use crate::info::TlsConnectionInfo;
#[cfg(feature = "server")]
use crate::info::tls::TlsConnectionInfoReceiver;
use crate::info::{ConnectionInfo, HasConnectionInfo, HasTlsConnectionInfo};
use pin_project::pin_project;
use tokio::io::{AsyncRead, AsyncWrite};
#[pin_project]
#[derive(Debug)]
pub struct Handshaking<'a, T: ?Sized> {
inner: &'a mut T,
}
impl<'a, T: ?Sized> Handshaking<'a, T> {
fn new(inner: &'a mut T) -> Self {
Handshaking { inner }
}
}
impl<T> Future for Handshaking<'_, T>
where
T: TlsHandshakeStream + ?Sized,
{
type Output = Result<(), io::Error>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.project().inner.poll_handshake(cx)
}
}
pub trait TlsHandshakeStream {
fn poll_handshake(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>>;
fn finish_handshake(&mut self) -> Handshaking<'_, Self> {
Handshaking::new(self)
}
}
#[cfg(feature = "server")]
pub trait TlsHandshakeInfo: TlsHandshakeStream {
fn recv(&self) -> TlsConnectionInfoReceiver;
}
#[derive(Debug)]
#[pin_project(project=OptTlsProjection)]
pub enum OptTlsStream<Tls, NoTls> {
NoTls(#[pin] NoTls),
Tls(#[pin] Tls),
}
impl<Tls, NoTls> TlsHandshakeStream for OptTlsStream<Tls, NoTls>
where
Tls: TlsHandshakeStream + Unpin,
NoTls: AsyncRead + AsyncWrite + Unpin,
{
fn poll_handshake(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
match self {
OptTlsStream::NoTls(_) => Poll::Ready(Ok(())),
OptTlsStream::Tls(stream) => stream.poll_handshake(cx),
}
}
}
#[cfg(feature = "server")]
impl<Tls, NoTls> TlsHandshakeInfo for OptTlsStream<Tls, NoTls>
where
Tls: TlsHandshakeInfo + Unpin,
NoTls: AsyncRead + AsyncWrite + Send + Unpin,
{
fn recv(&self) -> TlsConnectionInfoReceiver {
match self {
OptTlsStream::NoTls(_) => TlsConnectionInfoReceiver::empty(),
OptTlsStream::Tls(stream) => stream.recv(),
}
}
}
impl<Tls, NoTls, A> HasConnectionInfo for OptTlsStream<Tls, NoTls>
where
A: fmt::Debug + fmt::Display + Send + 'static,
Tls: HasConnectionInfo<Addr = A>,
NoTls: HasConnectionInfo<Addr = A>,
{
type Addr = A;
fn info(&self) -> ConnectionInfo<A> {
match self {
OptTlsStream::NoTls(stream) => stream.info(),
OptTlsStream::Tls(stream) => stream.info(),
}
}
}
impl<Tls, NoTls, A> HasTlsConnectionInfo for OptTlsStream<Tls, NoTls>
where
A: fmt::Debug + fmt::Display + Send + 'static,
Tls: HasConnectionInfo<Addr = A> + HasTlsConnectionInfo,
NoTls: HasConnectionInfo<Addr = A>,
{
fn tls_info(&self) -> Option<&TlsConnectionInfo> {
match self {
OptTlsStream::NoTls(_) => None,
OptTlsStream::Tls(stream) => stream.tls_info(),
}
}
}
macro_rules! dispatch {
($driver:ident.$method:ident($($args:expr),+)) => {
match $driver.project() {
OptTlsProjection::NoTls(stream) => stream.$method($($args),+),
OptTlsProjection::Tls(stream) => stream.$method($($args),+),
}
};
}
impl<Tls, NoTls> AsyncRead for OptTlsStream<Tls, NoTls>
where
Tls: AsyncRead,
NoTls: AsyncRead,
{
fn poll_read(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> Poll<io::Result<()>> {
dispatch!(self.poll_read(cx, buf))
}
}
impl<Tls, NoTls> AsyncWrite for OptTlsStream<Tls, NoTls>
where
Tls: AsyncWrite,
NoTls: AsyncWrite,
{
fn poll_write(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, io::Error>> {
dispatch!(self.poll_write(cx, buf))
}
fn poll_flush(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), io::Error>> {
dispatch!(self.poll_flush(cx))
}
fn poll_shutdown(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), io::Error>> {
dispatch!(self.poll_shutdown(cx))
}
}
impl<Tls, NoTls> From<NoTls> for OptTlsStream<Tls, NoTls> {
fn from(stream: NoTls) -> Self {
Self::NoTls(stream)
}
}