Skip to main content

tollgate_server/
transport.rs

1//! TLS is owned by the listener. HTTP headers cannot manufacture a peer identity.
2
3use std::io;
4use std::net::{IpAddr, SocketAddr};
5use std::num::NonZeroUsize;
6use std::sync::Arc;
7use std::time::Duration;
8
9use axum::extract::connect_info::Connected;
10use axum::serve::{IncomingStream, Listener};
11use jiff::Timestamp;
12use rustls::pki_types::{CertificateDer, PrivateKeyDer, UnixTime, pem::PemObject};
13use rustls::server::danger::ClientCertVerifier;
14use sha2::{Digest, Sha256};
15use tokio::io::{AsyncRead, AsyncWrite};
16use tokio::net::{TcpListener, TcpStream};
17use tokio::task::JoinSet;
18use tokio_rustls::{TlsAcceptor, server::TlsStream};
19
20use crate::security::{SecurityError, ServerSecurity};
21
22/// Parsed and validated before binding or atomic replacement. Private keys are
23/// never rendered in diagnostics. Optional client authentication permits bearer
24/// users and unauthenticated health probes on the same encrypted listener.
25#[derive(Clone)]
26pub struct TlsConfig {
27    config: Arc<rustls::ServerConfig>,
28    verifier: Option<Arc<dyn ClientCertVerifier>>,
29    handshake_timeout: Duration,
30    max_handshakes: NonZeroUsize,
31}
32
33impl TlsConfig {
34    /// Builds a server configuration from the PEM certificate chain, its PEM
35    /// private key and, optionally, a PEM bundle of client CAs.
36    ///
37    /// The server uses rustls's safe default protocol versions and offers
38    /// only HTTP/1.1. It accepts no 0-RTT data and does not resume sessions.
39    /// With a client CA, the handshake requests a client certificate and
40    /// verifies any that is presented, but does not require one, so bearer
41    /// callers and probes share the listener. A verified certificate grants
42    /// nothing until its leaf fingerprint is mapped to an identity with
43    /// [`SecurityPolicy::with_certificate`](crate::security::SecurityPolicy::with_certificate).
44    ///
45    /// Handshake limits default to five seconds and 128 pending handshakes;
46    /// see [`with_handshake_limits`](Self::with_handshake_limits).
47    ///
48    /// # Errors
49    ///
50    /// Returns a [`SecurityError`] for empty or malformed PEM, an invalid
51    /// client CA, or a private key that does not match the certificate. The
52    /// error never contains key material.
53    pub fn from_pem(
54        certificates: &[u8],
55        private_key: &[u8],
56        client_ca: Option<&[u8]>,
57    ) -> Result<Self, SecurityError> {
58        let certificates = parse_certificates(certificates)?;
59        let private_key = PrivateKeyDer::from_pem_slice(private_key)
60            .map_err(|_| SecurityError("invalid TLS private key PEM"))?;
61        let provider = Arc::new(rustls::crypto::ring::default_provider());
62        let verifier = client_ca
63            .map(|pem| {
64                let mut roots = rustls::RootCertStore::empty();
65                for certificate in parse_certificates(pem)? {
66                    roots
67                        .add(certificate)
68                        .map_err(|_| SecurityError("invalid client CA certificate"))?;
69                }
70                rustls::server::WebPkiClientVerifier::builder_with_provider(
71                    Arc::new(roots),
72                    Arc::clone(&provider),
73                )
74                .allow_unauthenticated()
75                .build()
76                .map_err(|_| SecurityError("invalid client CA configuration"))
77            })
78            .transpose()?;
79        let builder = rustls::ServerConfig::builder_with_provider(provider)
80            .with_safe_default_protocol_versions()
81            .map_err(|_| SecurityError("no safe TLS protocol version"))?;
82        let builder = match &verifier {
83            Some(verifier) => builder.with_client_cert_verifier(Arc::clone(verifier)),
84            None => builder.with_no_client_auth(),
85        };
86        let mut config = builder
87            .with_single_cert(certificates, private_key)
88            .map_err(|_| SecurityError("TLS certificate and private key do not match"))?;
89        // HTTP/1.1 matches the current control-plane server. No 0-RTT replay of
90        // a non-idempotent deposit; session resumption is not enabled.
91        config.alpn_protocols = vec![b"http/1.1".to_vec()];
92        config.max_early_data_size = 0;
93        config.send_tls13_tickets = 0;
94        config.session_storage = Arc::new(rustls::server::NoServerSessionStorage {});
95        Ok(Self {
96            config: Arc::new(config),
97            verifier,
98            handshake_timeout: Duration::from_secs(5),
99            max_handshakes: NonZeroUsize::new(128).expect("128 is nonzero"),
100        })
101    }
102
103    /// A configurable connection-burst budget, independent of accounts or body
104    /// sizes. Excess connections wait in the OS backlog rather than spawning
105    /// unbounded handshake tasks. Each task has its own deadline.
106    pub fn with_handshake_limits(
107        mut self,
108        timeout: Duration,
109        max_pending: NonZeroUsize,
110    ) -> Result<Self, SecurityError> {
111        if timeout.is_zero() || std::time::Instant::now().checked_add(timeout).is_none() {
112            return Err(SecurityError(
113                "TLS handshake timeout must be positive and representable",
114            ));
115        }
116        self.handshake_timeout = timeout;
117        self.max_handshakes = max_pending;
118        Ok(self)
119    }
120
121    pub(crate) fn verifies_clients(&self) -> bool {
122        self.verifier.is_some()
123    }
124
125    pub(crate) fn verify(
126        &self,
127        chain: &[CertificateDer<'static>],
128        now: Timestamp,
129    ) -> Result<(), SecurityError> {
130        let verifier = self
131            .verifier
132            .as_ref()
133            .ok_or(SecurityError("client CA is not configured"))?;
134        let (leaf, intermediates) = chain
135            .split_first()
136            .ok_or(SecurityError("missing client certificate"))?;
137        let seconds = u64::try_from(now.as_second())
138            .map_err(|_| SecurityError("invalid certificate verification time"))?;
139        verifier
140            .verify_client_cert(
141                leaf,
142                intermediates,
143                UnixTime::since_unix_epoch(Duration::from_secs(seconds)),
144            )
145            .map_err(|_| SecurityError("client certificate no longer valid"))?;
146        Ok(())
147    }
148}
149
150pub(crate) fn parse_certificates(
151    pem: &[u8],
152) -> Result<Vec<CertificateDer<'static>>, SecurityError> {
153    let certificates: Vec<_> = CertificateDer::pem_slice_iter(pem)
154        .collect::<Result<_, _>>()
155        .map_err(|_| SecurityError("invalid certificate PEM"))?;
156    if certificates.is_empty() {
157        return Err(SecurityError("certificate PEM is empty"));
158    }
159    Ok(certificates)
160}
161
162/// The SHA-256 fingerprint of a client's DER leaf certificate, given as PEM.
163///
164/// This is the key [`SecurityPolicy::with_certificate`] maps to an identity;
165/// the security manifest derives it from each configured leaf file.
166///
167/// # Errors
168///
169/// Returns a [`SecurityError`] unless the PEM holds exactly one certificate.
170///
171/// [`SecurityPolicy::with_certificate`]: crate::security::SecurityPolicy::with_certificate
172pub fn certificate_fingerprint(pem: &[u8]) -> Result<[u8; 32], SecurityError> {
173    let certificates = parse_certificates(pem)?;
174    if certificates.len() != 1 {
175        return Err(SecurityError(
176            "identity requires exactly one leaf certificate",
177        ));
178    }
179    Ok(fingerprint(&certificates[0]))
180}
181
182pub(crate) fn fingerprint(certificate: &CertificateDer<'_>) -> [u8; 32] {
183    Sha256::digest(certificate.as_ref()).into()
184}
185
186/// Includes IPv4-mapped loopback addresses; all DNS names are resolved by the
187/// caller before the listener is checked. Prefixes such as `localhost.evil`
188/// are never trusted.
189pub fn is_loopback(address: IpAddr) -> bool {
190    match address {
191        IpAddr::V4(ip) => ip.is_loopback(),
192        IpAddr::V6(ip) => {
193            ip.is_loopback() || ip.to_ipv4_mapped().is_some_and(|ip| ip.is_loopback())
194        }
195    }
196}
197
198#[derive(Clone)]
199pub(crate) struct PeerIdentity {
200    address: SocketAddr,
201    pub certificates: Option<Arc<[CertificateDer<'static>]>>,
202}
203
204impl std::fmt::Debug for PeerIdentity {
205    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
206        f.debug_struct("PeerIdentity")
207            .field("address", &self.address)
208            .field("certificate_present", &self.certificates.is_some())
209            .finish()
210    }
211}
212
213impl Connected<IncomingStream<'_, SecureListener>> for PeerIdentity {
214    fn connect_info(stream: IncomingStream<'_, SecureListener>) -> Self {
215        stream.remote_addr().clone()
216    }
217}
218
219pub(crate) trait ServerIo: AsyncRead + AsyncWrite + Unpin + Send {}
220impl<T: AsyncRead + AsyncWrite + Unpin + Send> ServerIo for T {}
221
222pub(crate) struct SecureListener {
223    listener: TcpListener,
224    security: Arc<ServerSecurity>,
225    handshakes: JoinSet<(io::Result<TlsStream<TcpStream>>, SocketAddr)>,
226}
227
228impl SecureListener {
229    pub fn new(listener: TcpListener, security: Arc<ServerSecurity>) -> io::Result<Self> {
230        if !security.encrypted() && !is_loopback(listener.local_addr()?.ip()) {
231            return Err(io::Error::new(
232                io::ErrorKind::InvalidInput,
233                "non-loopback listeners require TLS",
234            ));
235        }
236        Ok(Self {
237            listener,
238            security,
239            handshakes: JoinSet::new(),
240        })
241    }
242}
243
244impl Listener for SecureListener {
245    type Io = Box<dyn ServerIo>;
246    type Addr = PeerIdentity;
247
248    async fn accept(&mut self) -> (Self::Io, Self::Addr) {
249        loop {
250            let bundle = self.security.current.load_full();
251            let capacity = bundle
252                .tls
253                .as_ref()
254                .map_or(1, |tls| tls.max_handshakes.get());
255            tokio::select! {
256                biased;
257                result = self.handshakes.join_next(), if !self.handshakes.is_empty() => {
258                    match result {
259                        Some(Ok((Ok(stream), address))) => {
260                            let certificates = stream.get_ref().1.peer_certificates().map(Arc::from);
261                            return (Box::new(stream), PeerIdentity { address, certificates });
262                        }
263                        Some(Ok((Err(error), address))) => tracing::debug!(%address, %error, "TLS handshake refused"),
264                        Some(Err(error)) => tracing::warn!(%error, "TLS handshake task failed"),
265                        None => {}
266                    }
267                }
268                accepted = self.listener.accept(), if self.handshakes.len() < capacity => {
269                    match accepted {
270                        Ok((stream, address)) => match &self.security.current.load_full().tls {
271                            None => return (Box::new(stream), PeerIdentity { address, certificates: None }),
272                            Some(tls) => {
273                                let acceptor = TlsAcceptor::from(Arc::clone(&tls.config));
274                                let timeout = tls.handshake_timeout;
275                                self.handshakes.spawn(async move {
276                                    let result = tokio::time::timeout(timeout, acceptor.accept(stream)).await
277                                        .unwrap_or_else(|_| Err(io::Error::new(io::ErrorKind::TimedOut, "TLS handshake deadline")));
278                                    (result, address)
279                                });
280                            }
281                        },
282                        Err(error) => {
283                            tracing::warn!(%error, "control-plane accept failed");
284                            tokio::time::sleep(Duration::from_millis(100)).await;
285                        }
286                    }
287                }
288            }
289        }
290    }
291
292    fn local_addr(&self) -> io::Result<Self::Addr> {
293        Ok(PeerIdentity {
294            address: self.listener.local_addr()?,
295            certificates: None,
296        })
297    }
298}
299
300#[cfg(test)]
301mod tests {
302    use super::*;
303    use crate::security::SecurityPolicy;
304
305    fn security(timeout: Duration, capacity: usize) -> Arc<ServerSecurity> {
306        let certificate = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
307        let tls = TlsConfig::from_pem(
308            certificate.cert.pem().as_bytes(),
309            certificate.signing_key.serialize_pem().as_bytes(),
310            None,
311        )
312        .unwrap()
313        .with_handshake_limits(timeout, NonZeroUsize::new(capacity).unwrap())
314        .unwrap();
315        ServerSecurity::new(SecurityPolicy::new(), Some(tls)).unwrap()
316    }
317
318    #[tokio::test]
319    async fn pending_tls_handshakes_are_bounded_expire_and_drop_with_the_listener() {
320        use tokio::io::AsyncReadExt;
321        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
322        let address = listener.local_addr().unwrap();
323        let mut listener =
324            SecureListener::new(listener, security(Duration::from_millis(50), 2)).unwrap();
325        let mut first = TcpStream::connect(address).await.unwrap();
326        let mut second = TcpStream::connect(address).await.unwrap();
327        let mut queued = TcpStream::connect(address).await.unwrap();
328        // No peer sends a ClientHello, so accept remains pending. Cancellation
329        // of accept retains its bounded tasks for the next call.
330        assert!(
331            tokio::time::timeout(Duration::from_millis(10), listener.accept())
332                .await
333                .is_err()
334        );
335        assert_eq!(listener.handshakes.len(), 2);
336        let mut byte = [0];
337        assert_eq!(
338            tokio::time::timeout(Duration::from_secs(1), first.read(&mut byte))
339                .await
340                .unwrap()
341                .unwrap(),
342            0
343        );
344        assert_eq!(
345            tokio::time::timeout(Duration::from_secs(1), second.read(&mut byte))
346                .await
347                .unwrap()
348                .unwrap(),
349            0
350        );
351        assert!(
352            tokio::time::timeout(Duration::from_millis(10), listener.accept())
353                .await
354                .is_err()
355        );
356        assert_eq!(listener.handshakes.len(), 1);
357        drop(listener);
358        assert_eq!(
359            tokio::time::timeout(Duration::from_secs(1), queued.read(&mut byte))
360                .await
361                .unwrap()
362                .unwrap(),
363            0
364        );
365    }
366}