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(
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 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 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
157pub 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 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}