borer-core 0.5.9

network borer
Documentation
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};

/// Accepts plain TCP or upgrades accepted sockets to TLS.
#[derive(Clone)]
pub struct Acceptor {
    inner: Option<TlsAcceptor>,
}

#[non_exhaustive]
#[derive(Debug)]
/// Stream returned by [`Acceptor`], either plain TCP or TLS-wrapped TCP.
pub enum MaybeTlsStream<S> {
    Plain(S),
    Tls(Box<server::TlsStream<S>>),
}

impl Acceptor {
    /// Create a plain or TLS acceptor from optional certificate and key paths.
    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");
    }
}