Skip to main content

moirai_http/
conn.rs

1//! Transport connection: a plain or TLS-wrapped Moirai TCP stream, unified as a
2//! single Moirai async stream via a static enum (no `dyn` on the hot path).
3
4use std::io;
5use std::pin::Pin;
6use std::task::{Context, Poll};
7
8use moirai_async::io::{AsyncRead, AsyncWrite};
9use moirai_async::net::TcpStream;
10use moirai_tls::{ServerName, TlsConnector, TlsStream};
11
12/// Connection target: scheme/host/port identifying a pool bucket.
13#[derive(Debug, Clone, PartialEq, Eq, Hash)]
14pub struct Origin {
15    /// True for `https`, false for `http`.
16    pub secure: bool,
17    /// Host name (used for TCP connect and TLS SNI / cert validation).
18    pub host: String,
19    /// TCP port.
20    pub port: u16,
21}
22
23impl Origin {
24    /// Address string `host:port` for [`TcpStream::connect`].
25    #[must_use]
26    pub fn authority(&self) -> String {
27        format!("{}:{}", self.host, self.port)
28    }
29
30    /// Value for the `Host` header: `host`, or `host:port` when the port is
31    /// non-default for the scheme (443 for https, 80 for http).
32    #[must_use]
33    pub fn host_header(&self) -> String {
34        let default = if self.secure { 443 } else { 80 };
35        if self.port == default {
36            self.host.clone()
37        } else {
38            format!("{}:{}", self.host, self.port)
39        }
40    }
41}
42
43/// An established HTTP transport connection.
44pub enum Conn {
45    /// Plaintext HTTP over TCP.
46    Plain(TcpStream),
47    /// HTTPS: TLS over TCP (boxed — the rustls stream is large).
48    Tls(Box<TlsStream<TcpStream>>),
49}
50
51impl Conn {
52    /// Open a new connection to `origin`, performing the TLS handshake when secure.
53    ///
54    /// # Errors
55    /// Propagates DNS/connect and TLS handshake failures.
56    pub async fn connect(origin: &Origin, tls: &TlsConnector) -> io::Result<Self> {
57        let tcp = TcpStream::connect(&origin.authority()).await?;
58        if origin.secure {
59            let domain = ServerName::try_from(origin.host.clone())
60                .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
61            let stream = tls.connect(domain, tcp).await?;
62            Ok(Conn::Tls(Box::new(stream)))
63        } else {
64            Ok(Conn::Plain(tcp))
65        }
66    }
67}
68
69impl AsyncRead for Conn {
70    #[inline]
71    fn poll_read(
72        self: Pin<&mut Self>,
73        cx: &mut Context<'_>,
74        buf: &mut [u8],
75    ) -> Poll<io::Result<usize>> {
76        match self.get_mut() {
77            Conn::Plain(s) => AsyncRead::poll_read(Pin::new(s), cx, buf),
78            Conn::Tls(s) => AsyncRead::poll_read(Pin::new(s.as_mut()), cx, buf),
79        }
80    }
81}
82
83impl AsyncWrite for Conn {
84    #[inline]
85    fn poll_write(
86        self: Pin<&mut Self>,
87        cx: &mut Context<'_>,
88        buf: &[u8],
89    ) -> Poll<io::Result<usize>> {
90        match self.get_mut() {
91            Conn::Plain(s) => AsyncWrite::poll_write(Pin::new(s), cx, buf),
92            Conn::Tls(s) => AsyncWrite::poll_write(Pin::new(s.as_mut()), cx, buf),
93        }
94    }
95
96    #[inline]
97    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
98        match self.get_mut() {
99            Conn::Plain(s) => AsyncWrite::poll_flush(Pin::new(s), cx),
100            Conn::Tls(s) => AsyncWrite::poll_flush(Pin::new(s.as_mut()), cx),
101        }
102    }
103
104    #[inline]
105    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
106        match self.get_mut() {
107            Conn::Plain(s) => AsyncWrite::poll_shutdown(Pin::new(s), cx),
108            Conn::Tls(s) => AsyncWrite::poll_shutdown(Pin::new(s.as_mut()), cx),
109        }
110    }
111}