orca_proxy/acme/
resolver.rs1use 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
12const WARN_LIMIT: usize = 256;
15
16#[derive(Clone, Debug, Default)]
20pub struct DynCertResolver {
21 certs: Arc<RwLock<HashMap<String, Arc<CertifiedKey>>>>,
22 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 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 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 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 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 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 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 #[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}