use std::{
path::PathBuf,
pin::Pin,
task::{Context, Poll},
};
use anyhow::Context as _;
use tokio::{
io::{AsyncRead, AsyncWrite, ReadBuf},
net::TcpStream,
};
use tokio_rustls::{TlsAcceptor, server};
use crate::tls::{load_certs, load_private_key, make_tls_acceptor};
#[derive(Clone)]
pub struct Acceptor {
inner: Option<TlsAcceptor>,
}
#[non_exhaustive]
#[derive(Debug)]
pub enum MaybeTlsStream<S> {
Plain(S),
Tls(Box<server::TlsStream<S>>),
}
impl Acceptor {
pub fn new(cert: Option<String>, key: Option<String>) -> anyhow::Result<Self> {
match (cert, key) {
(Some(cert), Some(key)) => {
let certs = load_certs(PathBuf::from(cert)).context("load_certs failed")?;
let key =
load_private_key(PathBuf::from(key)).context("load_private_key failed")?;
let tls_acceptor = make_tls_acceptor(certs, key)?;
Ok(Self {
inner: Some(tls_acceptor),
})
}
_ => Ok(Self { inner: None }),
}
}
pub async fn accept(&self, ts: TcpStream) -> anyhow::Result<MaybeTlsStream<TcpStream>> {
match &self.inner {
Some(acceptor) => {
let tls_ts = acceptor.accept(ts).await?;
Ok(MaybeTlsStream::Tls(Box::new(tls_ts)))
}
_ => Ok(MaybeTlsStream::Plain(ts)),
}
}
}
impl<S> AsyncRead for MaybeTlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
match self.get_mut() {
MaybeTlsStream::Plain(s) => Pin::new(s).poll_read(cx, buf),
MaybeTlsStream::Tls(s) => Pin::new(s).poll_read(cx, buf),
}
}
}
impl<S> AsyncWrite for MaybeTlsStream<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
match self.get_mut() {
MaybeTlsStream::Plain(s) => Pin::new(s).poll_write(cx, buf),
MaybeTlsStream::Tls(s) => Pin::new(s).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), std::io::Error>> {
match self.get_mut() {
MaybeTlsStream::Plain(s) => Pin::new(s).poll_flush(cx),
MaybeTlsStream::Tls(s) => Pin::new(s).poll_flush(cx),
}
}
fn poll_shutdown(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
match self.get_mut() {
MaybeTlsStream::Plain(s) => Pin::new(s).poll_shutdown(cx),
MaybeTlsStream::Tls(s) => Pin::new(s).poll_shutdown(cx),
}
}
}
#[cfg(test)]
mod tests {
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::{Acceptor, MaybeTlsStream};
async fn tcp_pair() -> (TcpStream, TcpStream) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).await.unwrap();
let (server, _) = listener.accept().await.unwrap();
(server, client)
}
#[tokio::test]
async fn accept_without_tls_returns_plain_stream() {
let acceptor = Acceptor::new(None, None).unwrap();
let (server, _client) = tcp_pair().await;
let stream = acceptor.accept(server).await.unwrap();
assert!(matches!(stream, MaybeTlsStream::Plain(_)));
}
#[tokio::test]
async fn maybe_tls_plain_stream_reads_and_writes() {
let (server, mut client) = tcp_pair().await;
let mut stream = MaybeTlsStream::Plain(server);
stream.write_all(b"ping").await.unwrap();
let mut received = [0u8; 4];
client.read_exact(&mut received).await.unwrap();
assert_eq!(&received, b"ping");
client.write_all(b"pong").await.unwrap();
let mut buf = [0u8; 4];
stream.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"pong");
}
}