Skip to main content

orca_proxy/acme/
resolver.rs

1//! SNI-based dynamic certificate resolver for multi-domain TLS.
2//!
3//! Allows hot-adding certificates at runtime when new domains are deployed.
4
5use std::collections::{HashMap, HashSet};
6use std::sync::{Arc, Mutex, RwLock};
7
8use rustls::server::{ClientHello, ResolvesServerCert};
9use rustls::sign::CertifiedKey;
10use tracing::{debug, warn};
11
12/// Distinct unknown hostnames to warn about before going quiet (scanners
13/// send random names).
14const WARN_LIMIT: usize = 256;
15
16/// Thread-safe certificate store that resolves certs by SNI hostname.
17///
18/// New certs can be added at runtime without restarting the TLS listener.
19#[derive(Clone, Debug, Default)]
20pub struct DynCertResolver {
21    certs: Arc<RwLock<HashMap<String, Arc<CertifiedKey>>>>,
22    /// Served when the client sends no SNI, an IP address, or a name with no
23    /// certificate (#206). Without it the handshake failed below HTTP: the
24    /// client saw `ERR_SSL_PROTOCOL_ERROR`, never an error page.
25    fallback: Option<Arc<CertifiedKey>>,
26    warned: Arc<Mutex<HashSet<String>>>,
27}
28
29impl DynCertResolver {
30    pub fn new() -> Self {
31        Self::default()
32    }
33
34    /// A resolver that answers every handshake, using a self-signed
35    /// certificate when it has none for the requested name. The client then
36    /// gets a certificate warning and, past it, orca's error page, instead
37    /// of a dead connection.
38    pub fn with_fallback() -> anyhow::Result<Self> {
39        let cert = rcgen::generate_simple_self_signed(vec!["orca-fallback.invalid".into()])?;
40        let key_der = rustls::pki_types::PrivatePkcs8KeyDer::from(cert.key_pair.serialize_der());
41        let signing_key = rustls::crypto::aws_lc_rs::sign::any_supported_type(&key_der.into())?;
42        Ok(Self {
43            fallback: Some(Arc::new(CertifiedKey::new(
44                vec![cert.cert.der().clone()],
45                signing_key,
46            ))),
47            ..Self::default()
48        })
49    }
50
51    /// Add or replace a certificate for a domain.
52    pub fn add_cert(&self, domain: &str, key: Arc<CertifiedKey>) {
53        self.certs
54            .write()
55            .expect("cert store poisoned")
56            .insert(domain.to_string(), key);
57    }
58
59    /// Check if a cert exists for the given domain.
60    pub fn has_cert(&self, domain: &str) -> bool {
61        self.certs
62            .read()
63            .expect("cert store poisoned")
64            .contains_key(domain)
65    }
66}
67
68impl ResolvesServerCert for DynCertResolver {
69    fn resolve(&self, client_hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
70        let Some(sni) = client_hello.server_name() else {
71            // No SNI, or an IP address (which rustls discards): scanners,
72            // monitoring probes against the bare IP.
73            debug!("TLS handshake without a usable SNI; using the fallback certificate");
74            return self.fallback.clone();
75        };
76        if let Some(key) = self.certs.read().expect("cert store poisoned").get(sni) {
77            return Some(key.clone());
78        }
79        // A hostname with no certificate: often a domain whose certificate
80        // never provisioned, so say so (once per name) instead of at debug.
81        let mut warned = self.warned.lock().unwrap_or_else(|e| e.into_inner());
82        if warned.len() < WARN_LIMIT && warned.insert(sni.to_string()) {
83            warn!(
84                sni,
85                "no certificate for this hostname; serving the fallback certificate"
86            );
87        }
88        self.fallback.clone()
89    }
90}
91
92#[cfg(test)]
93mod tests {
94    use std::sync::Arc;
95
96    use super::DynCertResolver;
97
98    /// A TLS client talking to a server that uses `resolver`, with SNI `sni`
99    /// (`None`: no SNI at all). Returns whether the handshake completed.
100    async fn handshake(resolver: DynCertResolver, sni: Option<&str>) -> bool {
101        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
102        let config = rustls::ServerConfig::builder()
103            .with_no_client_auth()
104            .with_cert_resolver(Arc::new(resolver));
105        let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(config));
106        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
107        let addr = listener.local_addr().unwrap();
108        tokio::spawn(async move {
109            let (tcp, _) = listener.accept().await.unwrap();
110            let _ = acceptor.accept(tcp).await;
111        });
112
113        let mut client = rustls::ClientConfig::builder()
114            .dangerous()
115            .with_custom_certificate_verifier(Arc::new(AcceptAnything))
116            .with_no_client_auth();
117        client.enable_sni = sni.is_some();
118        let connector = tokio_rustls::TlsConnector::from(Arc::new(client));
119        let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
120        let name = rustls::pki_types::ServerName::try_from(sni.unwrap_or("unused.invalid"))
121            .unwrap()
122            .to_owned();
123        connector.connect(name, tcp).await.is_ok()
124    }
125
126    /// #206: no SNI, or a name without a certificate, used to kill the
127    /// handshake before any HTTP was spoken.
128    #[tokio::test]
129    async fn the_fallback_answers_unknown_names_and_missing_sni() {
130        assert!(!handshake(DynCertResolver::new(), Some("cloud.example.com")).await);
131        let fallback = DynCertResolver::with_fallback().unwrap();
132        assert!(handshake(fallback.clone(), Some("cloud.example.com")).await);
133        assert!(handshake(fallback, None).await);
134    }
135
136    #[derive(Debug)]
137    struct AcceptAnything;
138
139    impl rustls::client::danger::ServerCertVerifier for AcceptAnything {
140        fn verify_server_cert(
141            &self,
142            _: &rustls::pki_types::CertificateDer<'_>,
143            _: &[rustls::pki_types::CertificateDer<'_>],
144            _: &rustls::pki_types::ServerName<'_>,
145            _: &[u8],
146            _: rustls::pki_types::UnixTime,
147        ) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
148            Ok(rustls::client::danger::ServerCertVerified::assertion())
149        }
150        fn verify_tls12_signature(
151            &self,
152            _: &[u8],
153            _: &rustls::pki_types::CertificateDer<'_>,
154            _: &rustls::DigitallySignedStruct,
155        ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
156            Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
157        }
158        fn verify_tls13_signature(
159            &self,
160            _: &[u8],
161            _: &rustls::pki_types::CertificateDer<'_>,
162            _: &rustls::DigitallySignedStruct,
163        ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
164            Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
165        }
166        fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
167            rustls::crypto::aws_lc_rs::default_provider()
168                .signature_verification_algorithms
169                .supported_schemes()
170        }
171    }
172}