use core::task::{Context, Poll, ready};
use std::pin::Pin;
use std::sync::Arc;
use pin_project::pin_project;
use rustls::ServerConfig;
use tokio::io::{AsyncRead, AsyncWrite};
use crate::info::HasConnectionInfo;
use crate::server::conn::Accept;
#[derive(Debug)]
#[pin_project]
pub struct TlsAcceptor<A> {
config: Arc<ServerConfig>,
#[pin]
incoming: A,
}
pub(super) use super::TlsStream;
impl<A> TlsAcceptor<A> {
pub fn new(config: Arc<ServerConfig>, incoming: A) -> Self {
TlsAcceptor { config, incoming }
}
}
impl<A> Accept for TlsAcceptor<A>
where
A: Accept,
A::Connection: AsyncRead + AsyncWrite + HasConnectionInfo,
<A::Connection as HasConnectionInfo>::Addr: Clone + Unpin + Send + Sync + 'static,
{
type Connection = TlsStream<A::Connection>;
type Error = A::Error;
fn poll_accept(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<Self::Connection, Self::Error>> {
let this = self.project();
match ready!(this.incoming.poll_accept(cx)) {
Ok(stream) => {
let accept =
tokio_rustls::TlsAcceptor::from(Arc::clone(this.config)).accept(stream);
Poll::Ready(Ok(TlsStream::new(accept)))
}
Err(e) => Poll::Ready(Err(e)),
}
}
}
pub trait TlsAcceptExt: Accept {
fn with_tls(self, config: Arc<ServerConfig>) -> TlsAcceptor<Self>
where
Self: Sized,
{
TlsAcceptor::new(config, self)
}
}
impl<A> TlsAcceptExt for A where A: Accept {}