This repository has no description
1use std::io::BufReader;
2use std::path::Path;
3use std::sync::Arc;
4
5use arc_swap::ArcSwap;
6use base64::Engine;
7use quinn::crypto::rustls::QuicServerConfig;
8use rustls::RootCertStore;
9use rustls::ServerConfig;
10use rustls::crypto::aws_lc_rs;
11use rustls::pki_types::{CertificateDer, PrivateKeyDer, UnixTime};
12use rustls::server::danger::{ClientCertVerified, ClientCertVerifier};
13use rustls::server::{ClientHello, ResolvesServerCert, WebPkiClientVerifier};
14use rustls::sign::CertifiedKey;
15use sha2::{Digest, Sha256};
16use tokio_util::sync::CancellationToken;
17
18use crate::limits::ListenLimits;
19use crate::zerortt::EarlyDataPolicy;
20
21pub const ACME_TLS_ALPN: &[u8] = rustls_acme::acme::ACME_TLS_ALPN_NAME;
22
23#[derive(Debug, thiserror::Error)]
24pub enum TlsError {
25 #[error("reading {path}: {source}")]
26 Read {
27 path: String,
28 source: std::io::Error,
29 },
30 #[error("parsing {path}: {message}")]
31 Parse { path: String, message: String },
32 #[error("no certificates found in {0}")]
33 NoCertificates(String),
34 #[error("no private key found in {0}")]
35 NoPrivateKey(String),
36 #[error("unusable private key: {0}")]
37 SigningKey(String),
38 #[error("building server config: {0}")]
39 Config(String),
40 #[error("certificate and private key don't match: {0}")]
41 KeyMismatch(String),
42 #[error("session ticketer: {0}")]
43 Ticketer(String),
44 #[error("client certificate verifier for {path}: {message}")]
45 ClientVerifier { path: String, message: String },
46 #[error("admin SPKI pin: {0}")]
47 SpkiPin(String),
48}
49
50#[derive(Clone)]
51pub struct SpkiPin([u8; 32]);
52
53impl PartialEq for SpkiPin {
54 fn eq(&self, other: &Self) -> bool {
55 self.0
56 .iter()
57 .zip(other.0.iter())
58 .fold(0u8, |acc, (left, right)| acc | (left ^ right))
59 == 0
60 }
61}
62
63impl Eq for SpkiPin {}
64
65impl SpkiPin {
66 pub fn from_base64(encoded: &str) -> Result<Self, TlsError> {
67 let bytes = base64::engine::general_purpose::STANDARD
68 .decode(encoded.trim())
69 .map_err(|error| TlsError::SpkiPin(error.to_string()))?;
70 let array: [u8; 32] = bytes.try_into().map_err(|bytes: Vec<u8>| {
71 TlsError::SpkiPin(format!("expected 32 bytes, got {}", bytes.len()))
72 })?;
73 Ok(Self(array))
74 }
75
76 fn of_certificate(cert: &CertificateDer<'_>) -> Result<Self, rustls::Error> {
77 let (_, parsed) = x509_parser::parse_x509_certificate(cert.as_ref()).map_err(|error| {
78 rustls::Error::General(format!("parse client certificate: {error}"))
79 })?;
80 Ok(Self(
81 Sha256::digest(parsed.tbs_certificate.subject_pki.raw).into(),
82 ))
83 }
84}
85
86impl std::fmt::Debug for SpkiPin {
87 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
88 f.debug_tuple("SpkiPin").finish_non_exhaustive()
89 }
90}
91
92pub(crate) fn fuzz_of_certificate(data: &[u8]) {
93 let cert = CertificateDer::from(data.to_vec());
94 let _ = SpkiPin::of_certificate(&cert);
95}
96
97pub struct ReloadableCertResolver {
98 current: ArcSwap<CertifiedKey>,
99}
100
101impl std::fmt::Debug for ReloadableCertResolver {
102 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
103 f.debug_struct("ReloadableCertResolver")
104 .finish_non_exhaustive()
105 }
106}
107
108impl ReloadableCertResolver {
109 pub fn new(initial: CertifiedKey) -> Self {
110 Self {
111 current: ArcSwap::from_pointee(initial),
112 }
113 }
114
115 pub fn store(&self, key: CertifiedKey) {
116 self.current.store(Arc::new(key));
117 }
118}
119
120impl ResolvesServerCert for ReloadableCertResolver {
121 fn resolve(&self, _client_hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
122 Some(self.current.load_full())
123 }
124}
125
126pub fn spawn_cert_reload(
127 resolver: Arc<ReloadableCertResolver>,
128 paths: crate::StaticCertPaths,
129 shutdown: CancellationToken,
130) {
131 #[cfg(unix)]
132 tokio::spawn(async move {
133 use tokio::signal::unix::{SignalKind, signal};
134 let mut hangup = match signal(SignalKind::hangup()) {
135 Ok(stream) => stream,
136 Err(error) => {
137 tracing::error!("install SIGHUP handler: {error}");
138 return;
139 }
140 };
141 loop {
142 tokio::select! {
143 () = shutdown.cancelled() => break,
144 received = hangup.recv() => {
145 if received.is_none() {
146 break;
147 }
148 match load_certified_key(&paths) {
149 Ok(key) => {
150 resolver.store(key);
151 tracing::info!("reloaded TLS certificate on SIGHUP");
152 }
153 Err(error) => {
154 tracing::warn!("SIGHUP reload kept the existing certificate: {error}");
155 }
156 }
157 }
158 }
159 }
160 });
161
162 #[cfg(not(unix))]
163 let _ = (resolver, paths, shutdown);
164}
165
166pub fn load_certified_key(paths: &crate::StaticCertPaths) -> Result<CertifiedKey, TlsError> {
167 let certs = load_certs(paths.cert_path.as_path())?;
168 let key = load_private_key(paths.key_path.as_path())?;
169 let signing_key = aws_lc_rs::sign::any_supported_type(&key)
170 .map_err(|error| TlsError::SigningKey(error.to_string()))?;
171 let certified = CertifiedKey::new(certs, signing_key);
172 certified
173 .keys_match()
174 .map_err(|error| TlsError::KeyMismatch(error.to_string()))?;
175 Ok(certified)
176}
177
178fn load_certs(path: &Path) -> Result<Vec<CertificateDer<'static>>, TlsError> {
179 let bytes = std::fs::read(path).map_err(|source| TlsError::Read {
180 path: path.display().to_string(),
181 source,
182 })?;
183 let mut reader = BufReader::new(bytes.as_slice());
184 let certs = rustls_pemfile::certs(&mut reader)
185 .collect::<Result<Vec<_>, _>>()
186 .map_err(|error| TlsError::Parse {
187 path: path.display().to_string(),
188 message: error.to_string(),
189 })?;
190 match certs.is_empty() {
191 true => Err(TlsError::NoCertificates(path.display().to_string())),
192 false => Ok(certs),
193 }
194}
195
196fn load_private_key(path: &Path) -> Result<PrivateKeyDer<'static>, TlsError> {
197 let bytes = std::fs::read(path).map_err(|source| TlsError::Read {
198 path: path.display().to_string(),
199 source,
200 })?;
201 let mut reader = BufReader::new(bytes.as_slice());
202 rustls_pemfile::private_key(&mut reader)
203 .map_err(|error| TlsError::Parse {
204 path: path.display().to_string(),
205 message: error.to_string(),
206 })?
207 .ok_or_else(|| TlsError::NoPrivateKey(path.display().to_string()))
208}
209
210fn ticketer() -> Result<Arc<dyn rustls::server::ProducesTickets>, TlsError> {
211 aws_lc_rs::Ticketer::new().map_err(|error| TlsError::Ticketer(error.to_string()))
212}
213
214fn tcp_alpn(extra: &[&[u8]]) -> Vec<Vec<u8>> {
215 [b"h2".as_slice(), b"http/1.1".as_slice()]
216 .into_iter()
217 .chain(extra.iter().copied())
218 .map(<[u8]>::to_vec)
219 .collect()
220}
221
222pub fn build_tls_server_config(
223 resolver: Arc<dyn ResolvesServerCert>,
224 extra_alpn: &[&[u8]],
225) -> Result<ServerConfig, TlsError> {
226 let provider = Arc::new(aws_lc_rs::default_provider());
227 let mut config = ServerConfig::builder_with_provider(provider)
228 .with_safe_default_protocol_versions()
229 .map_err(|error| TlsError::Config(error.to_string()))?
230 .with_no_client_auth()
231 .with_cert_resolver(resolver);
232 config.alpn_protocols = tcp_alpn(extra_alpn);
233 config.ticketer = ticketer()?;
234 Ok(config)
235}
236
237pub fn build_mtls_server_config(
238 resolver: Arc<dyn ResolvesServerCert>,
239 client_ca: &crate::ClientCaPath,
240 pin: SpkiPin,
241) -> Result<ServerConfig, TlsError> {
242 let provider = Arc::new(aws_lc_rs::default_provider());
243 let roots = load_client_ca(client_ca.as_path())?;
244 let webpki =
245 WebPkiClientVerifier::builder_with_provider(Arc::new(roots), Arc::clone(&provider))
246 .build()
247 .map_err(|error| TlsError::ClientVerifier {
248 path: client_ca.as_path().display().to_string(),
249 message: error.to_string(),
250 })?;
251 let verifier = Arc::new(PinnedClientVerifier { inner: webpki, pin });
252 let mut config = ServerConfig::builder_with_provider(provider)
253 .with_safe_default_protocol_versions()
254 .map_err(|error| TlsError::Config(error.to_string()))?
255 .with_client_cert_verifier(verifier)
256 .with_cert_resolver(resolver);
257 config.alpn_protocols = tcp_alpn(&[]);
258 config.ticketer = ticketer()?;
259 Ok(config)
260}
261
262fn load_client_ca(path: &Path) -> Result<RootCertStore, TlsError> {
263 let certs = load_certs(path)?;
264 let mut roots = RootCertStore::empty();
265 let (added, _) = roots.add_parsable_certificates(certs);
266 match added {
267 0 => Err(TlsError::NoCertificates(path.display().to_string())),
268 _ => Ok(roots),
269 }
270}
271
272#[derive(Debug)]
273struct PinnedClientVerifier {
274 inner: Arc<dyn ClientCertVerifier>,
275 pin: SpkiPin,
276}
277
278impl ClientCertVerifier for PinnedClientVerifier {
279 fn root_hint_subjects(&self) -> &[rustls::DistinguishedName] {
280 self.inner.root_hint_subjects()
281 }
282
283 fn offer_client_auth(&self) -> bool {
284 self.inner.offer_client_auth()
285 }
286
287 fn client_auth_mandatory(&self) -> bool {
288 self.inner.client_auth_mandatory()
289 }
290
291 fn verify_client_cert(
292 &self,
293 end_entity: &CertificateDer<'_>,
294 intermediates: &[CertificateDer<'_>],
295 now: UnixTime,
296 ) -> Result<ClientCertVerified, rustls::Error> {
297 let verified = self
298 .inner
299 .verify_client_cert(end_entity, intermediates, now)?;
300 match SpkiPin::of_certificate(end_entity)? == self.pin {
301 true => Ok(verified),
302 false => Err(rustls::Error::General(
303 "client certificate SPKI doesn't match the pinned admin identity".to_string(),
304 )),
305 }
306 }
307
308 fn verify_tls12_signature(
309 &self,
310 message: &[u8],
311 cert: &CertificateDer<'_>,
312 dss: &rustls::DigitallySignedStruct,
313 ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
314 self.inner.verify_tls12_signature(message, cert, dss)
315 }
316
317 fn verify_tls13_signature(
318 &self,
319 message: &[u8],
320 cert: &CertificateDer<'_>,
321 dss: &rustls::DigitallySignedStruct,
322 ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
323 self.inner.verify_tls13_signature(message, cert, dss)
324 }
325
326 fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
327 self.inner.supported_verify_schemes()
328 }
329}
330
331pub fn build_quic_server_config(
332 resolver: Arc<dyn ResolvesServerCert>,
333 limits: ListenLimits,
334 early_data: EarlyDataPolicy,
335) -> Result<quinn::ServerConfig, TlsError> {
336 let provider = Arc::new(aws_lc_rs::default_provider());
337 let mut crypto = ServerConfig::builder_with_provider(provider)
338 .with_protocol_versions(&[&rustls::version::TLS13])
339 .map_err(|error| TlsError::Config(error.to_string()))?
340 .with_no_client_auth()
341 .with_cert_resolver(resolver);
342 crypto.alpn_protocols = vec![b"h3".to_vec()];
343 crypto.max_early_data_size = early_data.max_early_data_size();
344
345 let quic_crypto =
346 QuicServerConfig::try_from(crypto).map_err(|error| TlsError::Config(error.to_string()))?;
347 let mut config = quinn::ServerConfig::with_crypto(Arc::new(quic_crypto));
348
349 let budget = limits.connection_budget();
350 let mut transport = quinn::TransportConfig::default();
351 transport.max_concurrent_bidi_streams(quinn::VarInt::from_u32(
352 budget.max_concurrent_streams().get(),
353 ));
354 transport.stream_receive_window(quinn::VarInt::from_u32(budget.stream_receive_window()));
355 transport.receive_window(quinn::VarInt::from_u32(budget.connection_receive_window()));
356 let idle = limits.idle_timeout().get();
357 transport.max_idle_timeout(Some(
358 quinn::IdleTimeout::try_from(idle).map_err(|error| TlsError::Config(error.to_string()))?,
359 ));
360 transport.keep_alive_interval(Some(idle / 2));
361 config.transport_config(Arc::new(transport));
362 Ok(config)
363}
364
365#[cfg(test)]
366pub(crate) mod test_support {
367 use std::num::{NonZeroU32, NonZeroU64};
368
369 use rustls::pki_types::{PrivateKeyDer, PrivatePkcs8KeyDer};
370 use rustls::sign::CertifiedKey;
371
372 use super::*;
373 use crate::limits::ListenLimits;
374
375 pub(crate) fn self_signed() -> CertifiedKey {
376 let generated = rcgen::generate_simple_self_signed(vec!["localhost".to_string()]).unwrap();
377 let cert_der = generated.cert.der().clone();
378 let key_der = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(
379 generated.signing_key.serialize_der(),
380 ));
381 let signing_key = aws_lc_rs::sign::any_supported_type(&key_der).unwrap();
382 CertifiedKey::new(vec![cert_der], signing_key)
383 }
384
385 pub(crate) fn resolver() -> Arc<ReloadableCertResolver> {
386 Arc::new(ReloadableCertResolver::new(self_signed()))
387 }
388
389 pub(crate) fn limits() -> ListenLimits {
390 ListenLimits::new(
391 crate::limits::HeaderTimeout::from_millis(NonZeroU64::new(5_000).unwrap()),
392 crate::limits::IdleTimeout::from_millis(NonZeroU64::new(30_000).unwrap()),
393 NonZeroU32::new(64).unwrap(),
394 )
395 }
396
397 pub(crate) fn classical_only_provider() -> rustls::crypto::CryptoProvider {
398 let mut provider = aws_lc_rs::default_provider();
399 provider.kx_groups = vec![aws_lc_rs::kx_group::X25519];
400 provider
401 }
402
403 #[derive(Debug)]
404 pub(crate) struct AcceptAnyServerCert;
405
406 impl rustls::client::danger::ServerCertVerifier for AcceptAnyServerCert {
407 fn verify_server_cert(
408 &self,
409 _end_entity: &rustls::pki_types::CertificateDer<'_>,
410 _intermediates: &[rustls::pki_types::CertificateDer<'_>],
411 _server_name: &rustls::pki_types::ServerName<'_>,
412 _ocsp_response: &[u8],
413 _now: rustls::pki_types::UnixTime,
414 ) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
415 Ok(rustls::client::danger::ServerCertVerified::assertion())
416 }
417
418 fn verify_tls12_signature(
419 &self,
420 message: &[u8],
421 cert: &rustls::pki_types::CertificateDer<'_>,
422 dss: &rustls::DigitallySignedStruct,
423 ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
424 rustls::crypto::verify_tls12_signature(
425 message,
426 cert,
427 dss,
428 &aws_lc_rs::default_provider().signature_verification_algorithms,
429 )
430 }
431
432 fn verify_tls13_signature(
433 &self,
434 message: &[u8],
435 cert: &rustls::pki_types::CertificateDer<'_>,
436 dss: &rustls::DigitallySignedStruct,
437 ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
438 rustls::crypto::verify_tls13_signature(
439 message,
440 cert,
441 dss,
442 &aws_lc_rs::default_provider().signature_verification_algorithms,
443 )
444 }
445
446 fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
447 aws_lc_rs::default_provider()
448 .signature_verification_algorithms
449 .supported_schemes()
450 }
451 }
452}
453
454#[cfg(test)]
455mod tests {
456 use super::*;
457
458 #[test]
459 fn the_tcp_alpn_set_offers_h2_and_http1_but_not_h3() {
460 let config = build_tls_server_config(test_support::resolver(), &[]).unwrap();
461 assert_eq!(
462 config.alpn_protocols,
463 vec![b"h2".to_vec(), b"http/1.1".to_vec()],
464 "h3 is QUIC-only and must never appear in the TCP ALPN set"
465 );
466 }
467
468 #[test]
469 fn the_acme_challenge_alpn_joins_only_when_requested() {
470 let config = build_tls_server_config(test_support::resolver(), &[ACME_TLS_ALPN]).unwrap();
471 assert_eq!(
472 config.alpn_protocols,
473 vec![b"h2".to_vec(), b"http/1.1".to_vec(), ACME_TLS_ALPN.to_vec()],
474 "acme-tls/1 must trail h2 and http/1.1 so normal clients never select it"
475 );
476 }
477
478 #[test]
479 fn the_tcp_config_installs_an_enabled_session_ticketer() {
480 let config = build_tls_server_config(test_support::resolver(), &[]).unwrap();
481 assert!(
482 config.ticketer.enabled(),
483 "session resumption requires an enabled ticketer"
484 );
485 }
486
487 #[test]
488 fn an_spki_pin_round_trips_through_base64() {
489 let encoded = base64::engine::general_purpose::STANDARD.encode([9u8; 32]);
490 assert_eq!(SpkiPin::from_base64(&encoded).unwrap(), SpkiPin([9u8; 32]));
491 }
492
493 #[test]
494 fn the_spki_pin_matches_the_standard_openssl_recipe() {
495 const CERT_PEM: &str = "-----BEGIN CERTIFICATE-----\n\
496MIIBhDCCASugAwIBAgIUSkWE4CZvV8B8z9phedKo1PbahDUwCgYIKoZIzj0EAwIw\n\
497GDEWMBQGA1UEAwwNYW5lbW9uZS5hZG1pbjAeFw0yNjA2MjExOTA5MDdaFw0zNjA2\n\
498MTgxOTA5MDdaMBgxFjAUBgNVBAMMDWFuZW1vbmUuYWRtaW4wWTATBgcqhkjOPQIB\n\
499BggqhkjOPQMBBwNCAASFLKd70MtSGSyI2UjdpQyjaJrvXLofac41nI346wK0lC9G\n\
500PjZH/NKqo1iwQn+UfZB7gotfezWrDmAUz5OgT6Rlo1MwUTAdBgNVHQ4EFgQUis7S\n\
501XEFGpe4gQwWnzX/uzjpt274wHwYDVR0jBBgwFoAUis7SXEFGpe4gQwWnzX/uzjpt\n\
502274wDwYDVR0TAQH/BAUwAwEB/zAKBggqhkjOPQQDAgNHADBEAiBhJgE4cMP5/FJw\n\
503imc3fYQxOhm5nO59cfG06+0vuDIV1QIgMsKZjFjsch8rbLRNiJL5+bDmlgO7MD14\n\
5040PAyPOyjb+w=\n\
505-----END CERTIFICATE-----\n";
506 const OPENSSL_PIN: &str = "EhmM1HyWzC54br06EDvoAaqt1q1h+je3vVFJTcZ9e1U=";
507
508 let der = rustls_pemfile::certs(&mut CERT_PEM.as_bytes())
509 .next()
510 .unwrap()
511 .unwrap();
512 assert_eq!(
513 SpkiPin::of_certificate(&der).unwrap(),
514 SpkiPin::from_base64(OPENSSL_PIN).unwrap(),
515 "of_certificate must hash the same SubjectPublicKeyInfo bytes as openssl pkey -pubin -outform DER | dgst -sha256"
516 );
517 }
518
519 #[test]
520 fn an_spki_pin_of_the_wrong_length_is_rejected() {
521 let encoded = base64::engine::general_purpose::STANDARD.encode([9u8; 16]);
522 assert!(matches!(
523 SpkiPin::from_base64(&encoded),
524 Err(TlsError::SpkiPin(_))
525 ));
526 }
527
528 #[test]
529 fn the_quic_alpn_set_offers_only_h3() {
530 let quic = build_quic_server_config(
531 test_support::resolver(),
532 test_support::limits(),
533 EarlyDataPolicy::Disabled,
534 );
535 assert!(quic.is_ok());
536 }
537
538 #[test]
539 fn the_default_provider_prefers_post_quantum_key_exchange() {
540 use rustls::NamedGroup;
541
542 let provider = aws_lc_rs::default_provider();
543 let first = provider.kx_groups.first().expect("a key exchange group");
544 assert_eq!(
545 first.name(),
546 NamedGroup::X25519MLKEM768,
547 "prefer-post-quantum must order X25519MLKEM768 first for both the TCP and QUIC configs"
548 );
549 }
550
551 #[test]
552 fn a_mismatched_certificate_and_key_is_rejected() {
553 let first = test_support::self_signed();
554 let second = test_support::self_signed();
555 let mismatched = CertifiedKey::new(first.cert.clone(), second.key.clone());
556 assert!(mismatched.keys_match().is_err());
557 }
558
559 struct ClientIdentity {
560 ca_path: crate::ClientCaPath,
561 chain: Vec<CertificateDer<'static>>,
562 key_der: Vec<u8>,
563 pin: SpkiPin,
564 }
565
566 fn issue_client_identity() -> ClientIdentity {
567 let ca_key = rcgen::KeyPair::generate().unwrap();
568 let mut ca_params = rcgen::CertificateParams::new(Vec::<String>::new()).unwrap();
569 ca_params.is_ca = rcgen::IsCa::Ca(rcgen::BasicConstraints::Unconstrained);
570 let ca_cert = ca_params.self_signed(&ca_key).unwrap();
571 let issuer = rcgen::Issuer::new(ca_params, ca_key);
572
573 let client_key = rcgen::KeyPair::generate().unwrap();
574 let client_params = rcgen::CertificateParams::new(vec!["admin.knot".to_string()]).unwrap();
575 let client_cert = client_params.signed_by(&client_key, &issuer).unwrap();
576 let client_der = client_cert.der().clone();
577
578 let ca_path = std::env::temp_dir().join(format!(
579 "knot_edge_mtls_ca_{}_{:p}.pem",
580 std::process::id(),
581 &client_der as *const _
582 ));
583 std::fs::write(&ca_path, ca_cert.pem()).unwrap();
584
585 ClientIdentity {
586 ca_path: crate::ClientCaPath::new(ca_path),
587 pin: SpkiPin::of_certificate(&client_der).unwrap(),
588 chain: vec![client_der],
589 key_der: client_key.serialize_der(),
590 }
591 }
592
593 async fn mtls_handshake(
594 server_config: ServerConfig,
595 identity: Option<&ClientIdentity>,
596 ) -> bool {
597 use rustls::pki_types::{PrivateKeyDer, PrivatePkcs8KeyDer, ServerName};
598 use tokio::net::{TcpListener, TcpStream};
599 use tokio_rustls::{TlsAcceptor, TlsConnector};
600
601 let acceptor = TlsAcceptor::from(Arc::new(server_config));
602 let listener = TcpListener::bind("[::1]:0").await.unwrap();
603 let addr = listener.local_addr().unwrap();
604 let server = tokio::spawn(async move {
605 let (tcp, _) = listener.accept().await.unwrap();
606 acceptor.accept(tcp).await.is_ok()
607 });
608
609 let verifier =
610 rustls::ClientConfig::builder_with_provider(Arc::new(aws_lc_rs::default_provider()))
611 .with_safe_default_protocol_versions()
612 .unwrap()
613 .dangerous()
614 .with_custom_certificate_verifier(Arc::new(test_support::AcceptAnyServerCert));
615 let mut client_config = match identity {
616 Some(identity) => verifier
617 .with_client_auth_cert(
618 identity.chain.clone(),
619 PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(identity.key_der.clone())),
620 )
621 .unwrap(),
622 None => verifier.with_no_client_auth(),
623 };
624 client_config.alpn_protocols = vec![b"h2".to_vec()];
625 let connector = TlsConnector::from(Arc::new(client_config));
626 let tcp = TcpStream::connect(addr).await.unwrap();
627 let client_ok = connector
628 .connect(ServerName::try_from("localhost").unwrap(), tcp)
629 .await
630 .is_ok();
631 let server_ok = server.await.unwrap();
632 client_ok && server_ok
633 }
634
635 #[tokio::test]
636 async fn mtls_rejects_a_client_presenting_no_certificate() {
637 let identity = issue_client_identity();
638 let config = build_mtls_server_config(
639 test_support::resolver(),
640 &identity.ca_path,
641 identity.pin.clone(),
642 )
643 .unwrap();
644 assert!(
645 !mtls_handshake(config, None).await,
646 "the mandatory mTLS verifier must reject a client that presents no certificate"
647 );
648 }
649
650 #[tokio::test]
651 async fn mtls_admits_the_pinned_admin_certificate() {
652 let identity = issue_client_identity();
653 let config = build_mtls_server_config(
654 test_support::resolver(),
655 &identity.ca_path,
656 identity.pin.clone(),
657 )
658 .unwrap();
659 assert!(
660 mtls_handshake(config, Some(&identity)).await,
661 "a client presenting the pinned admin certificate must complete the mTLS handshake"
662 );
663 }
664
665 #[tokio::test]
666 async fn mtls_rejects_a_ca_trusted_client_whose_spki_is_not_pinned() {
667 let identity = issue_client_identity();
668 let config = build_mtls_server_config(
669 test_support::resolver(),
670 &identity.ca_path,
671 SpkiPin([0u8; 32]),
672 )
673 .unwrap();
674 assert!(
675 !mtls_handshake(config, Some(&identity)).await,
676 "a client trusted by the CA but failing the SPKI pin must be rejected"
677 );
678 }
679}