use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex, RwLock};
use rustls::server::{ClientHello, ResolvesServerCert};
use rustls::sign::CertifiedKey;
use tracing::{debug, warn};
const WARN_LIMIT: usize = 256;
#[derive(Clone, Debug, Default)]
pub struct DynCertResolver {
certs: Arc<RwLock<HashMap<String, Arc<CertifiedKey>>>>,
fallback: Option<Arc<CertifiedKey>>,
warned: Arc<Mutex<HashSet<String>>>,
}
impl DynCertResolver {
pub fn new() -> Self {
Self::default()
}
pub fn with_fallback() -> anyhow::Result<Self> {
let cert = rcgen::generate_simple_self_signed(vec!["orca-fallback.invalid".into()])?;
let key_der = rustls::pki_types::PrivatePkcs8KeyDer::from(cert.key_pair.serialize_der());
let signing_key = rustls::crypto::aws_lc_rs::sign::any_supported_type(&key_der.into())?;
Ok(Self {
fallback: Some(Arc::new(CertifiedKey::new(
vec![cert.cert.der().clone()],
signing_key,
))),
..Self::default()
})
}
pub fn add_cert(&self, domain: &str, key: Arc<CertifiedKey>) {
self.certs
.write()
.expect("cert store poisoned")
.insert(domain.to_string(), key);
}
pub fn has_cert(&self, domain: &str) -> bool {
self.certs
.read()
.expect("cert store poisoned")
.contains_key(domain)
}
}
impl ResolvesServerCert for DynCertResolver {
fn resolve(&self, client_hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
let Some(sni) = client_hello.server_name() else {
debug!("TLS handshake without a usable SNI; using the fallback certificate");
return self.fallback.clone();
};
if let Some(key) = self.certs.read().expect("cert store poisoned").get(sni) {
return Some(key.clone());
}
let mut warned = self.warned.lock().unwrap_or_else(|e| e.into_inner());
if warned.len() < WARN_LIMIT && warned.insert(sni.to_string()) {
warn!(
sni,
"no certificate for this hostname; serving the fallback certificate"
);
}
self.fallback.clone()
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::DynCertResolver;
async fn handshake(resolver: DynCertResolver, sni: Option<&str>) -> bool {
let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
let config = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(Arc::new(resolver));
let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(config));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (tcp, _) = listener.accept().await.unwrap();
let _ = acceptor.accept(tcp).await;
});
let mut client = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(AcceptAnything))
.with_no_client_auth();
client.enable_sni = sni.is_some();
let connector = tokio_rustls::TlsConnector::from(Arc::new(client));
let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
let name = rustls::pki_types::ServerName::try_from(sni.unwrap_or("unused.invalid"))
.unwrap()
.to_owned();
connector.connect(name, tcp).await.is_ok()
}
#[tokio::test]
async fn the_fallback_answers_unknown_names_and_missing_sni() {
assert!(!handshake(DynCertResolver::new(), Some("cloud.example.com")).await);
let fallback = DynCertResolver::with_fallback().unwrap();
assert!(handshake(fallback.clone(), Some("cloud.example.com")).await);
assert!(handshake(fallback, None).await);
}
#[derive(Debug)]
struct AcceptAnything;
impl rustls::client::danger::ServerCertVerifier for AcceptAnything {
fn verify_server_cert(
&self,
_: &rustls::pki_types::CertificateDer<'_>,
_: &[rustls::pki_types::CertificateDer<'_>],
_: &rustls::pki_types::ServerName<'_>,
_: &[u8],
_: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_: &[u8],
_: &rustls::pki_types::CertificateDer<'_>,
_: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_: &[u8],
_: &rustls::pki_types::CertificateDer<'_>,
_: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
rustls::crypto::aws_lc_rs::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
}
}