use std::fmt;
use std::path::PathBuf;
use std::sync::{Arc, RwLock};
use rustls::ServerConfig;
use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject};
use rustls::server::{ClientHello, ResolvesServerCert};
use rustls::sign::CertifiedKey;
use crate::error::{ServerError, ServerResult, TlsKind};
#[allow(clippy::result_large_err)]
pub fn load_certs(file_path: &str) -> ServerResult<Vec<CertificateDer<'static>>> {
let path = PathBuf::from(file_path);
let iter = CertificateDer::pem_file_iter(file_path).map_err(|e| ServerError::TlsLoad {
kind: TlsKind::Certificate,
path: path.clone(),
reason: e.to_string(),
})?;
let mut certs = Vec::new();
for (idx, item) in iter.enumerate() {
let cert = item.map_err(|e| ServerError::TlsLoad {
kind: TlsKind::Certificate,
path: path.clone(),
reason: format!("failed to parse certificate #{}: {}", idx + 1, e),
})?;
certs.push(cert);
}
if certs.is_empty() {
return Err(ServerError::TlsLoad {
kind: TlsKind::Certificate,
path,
reason: "no certificates found in PEM file".to_owned(),
});
}
Ok(certs)
}
#[allow(clippy::result_large_err)]
pub fn load_private_key(file_path: &str) -> ServerResult<PrivateKeyDer<'static>> {
PrivateKeyDer::from_pem_file(file_path).map_err(|e| ServerError::TlsLoad {
kind: TlsKind::PrivateKey,
path: PathBuf::from(file_path),
reason: e.to_string(),
})
}
#[derive(Debug, Clone)]
pub struct TlsReloadError {
pub reason: String,
}
impl fmt::Display for TlsReloadError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "TLS cert reload failed: {}", self.reason)
}
}
impl std::error::Error for TlsReloadError {}
fn make_certified_key(
certs: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
) -> Result<CertifiedKey, TlsReloadError> {
let signing_key =
rustls::crypto::ring::sign::any_supported_type(&key).map_err(|e| TlsReloadError {
reason: format!("unsupported private key type: {}", e),
})?;
Ok(CertifiedKey::new(certs, signing_key))
}
#[derive(Debug)]
pub struct ReloadableCertResolver {
inner: RwLock<Arc<CertifiedKey>>,
}
impl ReloadableCertResolver {
pub fn new(
certs: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
) -> Result<Self, TlsReloadError> {
let ck = make_certified_key(certs, key)?;
Ok(Self {
inner: RwLock::new(Arc::new(ck)),
})
}
pub fn reload_from_paths(&self, cert_path: &str, key_path: &str) -> Result<(), TlsReloadError> {
let certs = load_certs(cert_path).map_err(|e| TlsReloadError {
reason: e.to_string(),
})?;
let key = load_private_key(key_path).map_err(|e| TlsReloadError {
reason: e.to_string(),
})?;
let new_ck = make_certified_key(certs, key)?;
let mut guard = self.inner.write().map_err(|_| TlsReloadError {
reason: "cert RwLock poisoned".to_owned(),
})?;
*guard = Arc::new(new_ck);
log::info!("TLS certificate reloaded successfully from {}", cert_path);
Ok(())
}
}
impl ResolvesServerCert for ReloadableCertResolver {
fn resolve(&self, _client_hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
self.inner.read().ok().map(|g| Arc::clone(&*g))
}
}
pub fn build_server_config_static(
certs: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
) -> Result<ServerConfig, String> {
let mut config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.map_err(|e| format!("failed to build rustls ServerConfig: {}", e))?;
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
Ok(config)
}
pub fn build_server_config_reloadable(
certs: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
) -> Result<(ServerConfig, Arc<ReloadableCertResolver>), TlsReloadError> {
let resolver = Arc::new(ReloadableCertResolver::new(certs, key)?);
let mut config = ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(Arc::clone(&resolver) as Arc<dyn ResolvesServerCert>);
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
Ok((config, resolver))
}
#[cfg(test)]
mod tests {
use super::*;
fn write_pem_file(path: &str, content: &str) {
std::fs::write(path, content).unwrap();
}
const TEST_CERT_PEM: &str = "-----BEGIN CERTIFICATE-----\n\
MIIBczCCARmgAwIBAgIUNNKjB+m5H6ZCjEPHNFEL5GYW3/UwCgYIKoZIzj0EAwIw\n\
DzENMAsGA1UEAwwEdGVzdDAeFw0yNjA1MjIwMjQ0NTZaFw0zNjA1MTkwMjQ0NTZa\n\
MA8xDTALBgNVBAMMBHRlc3QwWTATBgcqhkjOPQIBBggqhkjOPQMBBwNCAAQTo544\n\
m3Yk+4kNlcFXR8RL5rtGVqrZohzvanN7oUiIYXzpofwYNBLqLg9AOZPeiX32aizX\n\
wqEBuYMV4B6gBj1Ho1MwUTAdBgNVHQ4EFgQUxoL28LxPMcYmwNAvUCIaaZp02xAw\n\
HwYDVR0jBBgwFoAUxoL28LxPMcYmwNAvUCIaaZp02xAwDwYDVR0TAQH/BAUwAwEB\n\
/zAKBggqhkjOPQQDAgNIADBFAiEA2IO7sD+CIM4OWZkF0SMCmrnus/xQbNFBICXg\n\
YNQ/K+oCIGlsqHA+PmxwUknuDDS5dQF26iNztRz2PY4diIfWxLNi\n\
-----END CERTIFICATE-----\n";
const TEST_KEY_PEM: &str = "-----BEGIN EC PRIVATE KEY-----\n\
MHcCAQEEIBK3C/2yAvhbvjxP7f5aCgVZN9udnXStns0xKk7LQ3RnoAoGCCqGSM49\n\
AwEHoUQDQgAEE6OeOJt2JPuJDZXBV0fES+a7Rlaq2aIc72pze6FIiGF86aH8GDQS\n\
6i4PQDmT3ol99mos18KhAbmDFeAeoAY9Rw==\n\
-----END EC PRIVATE KEY-----\n";
#[test]
fn load_certs_returns_error_for_missing_file() {
let result = load_certs("/nonexistent/cert.pem");
assert!(result.is_err());
}
#[test]
fn load_private_key_returns_error_for_missing_file() {
let result = load_private_key("/nonexistent/key.pem");
assert!(result.is_err());
}
#[test]
fn reloadable_resolver_init_and_reload_bad_path_keeps_old_cert() {
let cert_path = "/tmp/apimock_test_cert.pem";
let key_path = "/tmp/apimock_test_key.pem";
write_pem_file(cert_path, TEST_CERT_PEM);
write_pem_file(key_path, TEST_KEY_PEM);
let certs = load_certs(cert_path).expect("load test cert");
let key = load_private_key(key_path).expect("load test key");
let resolver = ReloadableCertResolver::new(certs, key).expect("build resolver");
let result = resolver.reload_from_paths("/no/cert.pem", "/no/key.pem");
assert!(result.is_err(), "expected error for missing paths");
let guard = resolver.inner.read().unwrap();
drop(guard);
}
#[test]
fn reloadable_resolver_reload_from_same_files_succeeds() {
let cert_path = "/tmp/apimock_test_cert2.pem";
let key_path = "/tmp/apimock_test_key2.pem";
write_pem_file(cert_path, TEST_CERT_PEM);
write_pem_file(key_path, TEST_KEY_PEM);
let certs = load_certs(cert_path).expect("load test cert");
let key = load_private_key(key_path).expect("load test key");
let resolver = ReloadableCertResolver::new(certs, key).expect("build resolver");
let result = resolver.reload_from_paths(cert_path, key_path);
assert!(
result.is_ok(),
"reload from same valid files must succeed: {:?}",
result
);
}
#[test]
fn build_server_config_reloadable_returns_resolver() {
let cert_path = "/tmp/apimock_test_cert3.pem";
let key_path = "/tmp/apimock_test_key3.pem";
write_pem_file(cert_path, TEST_CERT_PEM);
write_pem_file(key_path, TEST_KEY_PEM);
let certs = load_certs(cert_path).unwrap();
let key = load_private_key(key_path).unwrap();
let result = build_server_config_reloadable(certs, key);
assert!(
result.is_ok(),
"build_server_config_reloadable failed: {:?}",
result
);
let (_config, resolver) = result.unwrap();
let reload = resolver.reload_from_paths(cert_path, key_path);
assert!(reload.is_ok());
}
}