1use 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#[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(
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 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 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
162pub 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
186pub 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 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}