use arc_swap::ArcSwap;
use rustls::pki_types::pem::PemObject;
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use rustls::server::{ServerConfig, WebPkiClientVerifier};
use std::sync::Arc;
pub fn build_mtls_config(
cert_path: &str,
key_path: &str,
client_ca_path: &str,
) -> anyhow::Result<Arc<ServerConfig>> {
let certs: Vec<CertificateDer<'static>> =
CertificateDer::pem_file_iter(cert_path)?.collect::<Result<_, _>>()?;
let key = PrivateKeyDer::from_pem_file(key_path)
.map_err(|e| anyhow::anyhow!("no usable private key in {key_path}: {e}"))?;
let mut roots = rustls::RootCertStore::empty();
for c in CertificateDer::pem_file_iter(client_ca_path)? {
roots.add(c?)?;
}
let verifier = WebPkiClientVerifier::builder(Arc::new(roots)).build()?;
let mut cfg = ServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_single_cert(certs, key)?;
cfg.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
Ok(Arc::new(cfg))
}
pub struct ReloadingTls {
inner: Arc<ArcSwap<ServerConfig>>,
paths: (String, String, String),
}
impl ReloadingTls {
pub fn new(cert: &str, key: &str, ca: &str) -> anyhow::Result<Arc<Self>> {
let cfg = build_mtls_config(cert, key, ca)?;
let s = Arc::new(Self {
inner: Arc::new(ArcSwap::new(cfg)),
paths: (cert.into(), key.into(), ca.into()),
});
Self::watch(s.clone());
Ok(s)
}
pub fn current(&self) -> Arc<ServerConfig> {
self.inner.load_full()
}
pub fn reload(&self) -> anyhow::Result<()> {
let cfg = build_mtls_config(&self.paths.0, &self.paths.1, &self.paths.2)?;
self.inner.store(cfg);
Ok(())
}
fn watch(s: Arc<Self>) {
tokio::spawn(async move {
let mut sig =
match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::user_defined1())
{
Ok(s) => s,
Err(e) => {
tracing::error!(?e, "failed to register SIGUSR1 handler");
return;
}
};
while sig.recv().await.is_some() {
match s.reload() {
Ok(()) => tracing::info!("mTLS config reloaded"),
Err(e) => tracing::error!(?e, "mTLS reload failed; keeping old"),
}
}
});
}
}