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