1use std::fmt;
18use std::path::PathBuf;
19use std::sync::{Arc, RwLock};
20
21use rustls::ServerConfig;
22use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject};
23use rustls::server::{ClientHello, ResolvesServerCert};
24use rustls::sign::CertifiedKey;
25
26use crate::error::{ServerError, ServerResult, TlsKind};
27
28pub fn load_certs(file_path: &str) -> ServerResult<Vec<CertificateDer<'static>>> {
32 let path = PathBuf::from(file_path);
33 let iter = CertificateDer::pem_file_iter(file_path).map_err(|e| ServerError::TlsLoad {
34 kind: TlsKind::Certificate,
35 path: path.clone(),
36 reason: e.to_string(),
37 })?;
38
39 let mut certs = Vec::new();
40 for (idx, item) in iter.enumerate() {
41 let cert = item.map_err(|e| ServerError::TlsLoad {
42 kind: TlsKind::Certificate,
43 path: path.clone(),
44 reason: format!("failed to parse certificate #{}: {}", idx + 1, e),
45 })?;
46 certs.push(cert);
47 }
48
49 if certs.is_empty() {
50 return Err(ServerError::TlsLoad {
51 kind: TlsKind::Certificate,
52 path,
53 reason: "no certificates found in PEM file".to_owned(),
54 });
55 }
56
57 Ok(certs)
58}
59
60pub fn load_private_key(file_path: &str) -> ServerResult<PrivateKeyDer<'static>> {
62 PrivateKeyDer::from_pem_file(file_path).map_err(|e| ServerError::TlsLoad {
63 kind: TlsKind::PrivateKey,
64 path: PathBuf::from(file_path),
65 reason: e.to_string(),
66 })
67}
68
69#[derive(Debug, Clone)]
73pub struct TlsReloadError {
74 pub reason: String,
75}
76
77impl fmt::Display for TlsReloadError {
78 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
79 write!(f, "TLS cert reload failed: {}", self.reason)
80 }
81}
82
83impl std::error::Error for TlsReloadError {}
84
85fn make_certified_key(
89 certs: Vec<CertificateDer<'static>>,
90 key: PrivateKeyDer<'static>,
91) -> Result<CertifiedKey, TlsReloadError> {
92 let signing_key =
93 rustls::crypto::ring::sign::any_supported_type(&key).map_err(|e| TlsReloadError {
94 reason: format!("unsupported private key type: {}", e),
95 })?;
96 Ok(CertifiedKey::new(certs, signing_key))
97}
98
99#[derive(Debug)]
115pub struct ReloadableCertResolver {
116 inner: RwLock<Arc<CertifiedKey>>,
117}
118
119impl ReloadableCertResolver {
120 pub fn new(
122 certs: Vec<CertificateDer<'static>>,
123 key: PrivateKeyDer<'static>,
124 ) -> Result<Self, TlsReloadError> {
125 let ck = make_certified_key(certs, key)?;
126 Ok(Self {
127 inner: RwLock::new(Arc::new(ck)),
128 })
129 }
130
131 pub fn reload_from_paths(&self, cert_path: &str, key_path: &str) -> Result<(), TlsReloadError> {
136 let certs = load_certs(cert_path).map_err(|e| TlsReloadError {
137 reason: e.to_string(),
138 })?;
139 let key = load_private_key(key_path).map_err(|e| TlsReloadError {
140 reason: e.to_string(),
141 })?;
142 let new_ck = make_certified_key(certs, key)?;
143
144 let mut guard = self.inner.write().map_err(|_| TlsReloadError {
145 reason: "cert RwLock poisoned".to_owned(),
146 })?;
147 *guard = Arc::new(new_ck);
148 log::info!("TLS certificate reloaded successfully from {}", cert_path);
149 Ok(())
150 }
151}
152
153impl ResolvesServerCert for ReloadableCertResolver {
154 fn resolve(&self, _client_hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
155 self.inner.read().ok().map(|g| Arc::clone(&*g))
156 }
157}
158
159pub fn build_server_config_static(
166 certs: Vec<CertificateDer<'static>>,
167 key: PrivateKeyDer<'static>,
168) -> Result<ServerConfig, String> {
169 let mut config = ServerConfig::builder()
170 .with_no_client_auth()
171 .with_single_cert(certs, key)
172 .map_err(|e| format!("failed to build rustls ServerConfig: {}", e))?;
173 config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
174 Ok(config)
175}
176
177pub fn build_server_config_reloadable(
182 certs: Vec<CertificateDer<'static>>,
183 key: PrivateKeyDer<'static>,
184) -> Result<(ServerConfig, Arc<ReloadableCertResolver>), TlsReloadError> {
185 let resolver = Arc::new(ReloadableCertResolver::new(certs, key)?);
186 let mut config = ServerConfig::builder()
187 .with_no_client_auth()
188 .with_cert_resolver(Arc::clone(&resolver) as Arc<dyn ResolvesServerCert>);
189 config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
190 Ok((config, resolver))
191}
192
193#[cfg(test)]
196mod tests {
197 use super::*;
198
199 fn write_pem_file(path: &str, content: &str) {
200 std::fs::write(path, content).unwrap();
201 }
202
203 const TEST_CERT_PEM: &str = "-----BEGIN CERTIFICATE-----\n\
207MIIBczCCARmgAwIBAgIUNNKjB+m5H6ZCjEPHNFEL5GYW3/UwCgYIKoZIzj0EAwIw\n\
208DzENMAsGA1UEAwwEdGVzdDAeFw0yNjA1MjIwMjQ0NTZaFw0zNjA1MTkwMjQ0NTZa\n\
209MA8xDTALBgNVBAMMBHRlc3QwWTATBgcqhkjOPQIBBggqhkjOPQMBBwNCAAQTo544\n\
210m3Yk+4kNlcFXR8RL5rtGVqrZohzvanN7oUiIYXzpofwYNBLqLg9AOZPeiX32aizX\n\
211wqEBuYMV4B6gBj1Ho1MwUTAdBgNVHQ4EFgQUxoL28LxPMcYmwNAvUCIaaZp02xAw\n\
212HwYDVR0jBBgwFoAUxoL28LxPMcYmwNAvUCIaaZp02xAwDwYDVR0TAQH/BAUwAwEB\n\
213/zAKBggqhkjOPQQDAgNIADBFAiEA2IO7sD+CIM4OWZkF0SMCmrnus/xQbNFBICXg\n\
214YNQ/K+oCIGlsqHA+PmxwUknuDDS5dQF26iNztRz2PY4diIfWxLNi\n\
215-----END CERTIFICATE-----\n";
216
217 const TEST_KEY_PEM: &str = "-----BEGIN EC PRIVATE KEY-----\n\
218MHcCAQEEIBK3C/2yAvhbvjxP7f5aCgVZN9udnXStns0xKk7LQ3RnoAoGCCqGSM49\n\
219AwEHoUQDQgAEE6OeOJt2JPuJDZXBV0fES+a7Rlaq2aIc72pze6FIiGF86aH8GDQS\n\
2206i4PQDmT3ol99mos18KhAbmDFeAeoAY9Rw==\n\
221-----END EC PRIVATE KEY-----\n";
222
223 #[test]
224 fn load_certs_returns_error_for_missing_file() {
225 let result = load_certs("/nonexistent/cert.pem");
226 assert!(result.is_err());
227 }
228
229 #[test]
230 fn load_private_key_returns_error_for_missing_file() {
231 let result = load_private_key("/nonexistent/key.pem");
232 assert!(result.is_err());
233 }
234
235 #[test]
236 fn reloadable_resolver_init_and_reload_bad_path_keeps_old_cert() {
237 let dir = tempfile::tempdir().unwrap();
238 let cert_path = dir.path().join("apimock_test_cert.pem");
239 let key_path = dir.path().join("apimock_test_key.pem");
240 let cert_path = cert_path.to_str().unwrap();
241 let key_path = key_path.to_str().unwrap();
242 write_pem_file(cert_path, TEST_CERT_PEM);
243 write_pem_file(key_path, TEST_KEY_PEM);
244
245 let certs = load_certs(cert_path).expect("load test cert");
246 let key = load_private_key(key_path).expect("load test key");
247 let resolver = ReloadableCertResolver::new(certs, key).expect("build resolver");
248
249 let result = resolver.reload_from_paths("/no/cert.pem", "/no/key.pem");
251 assert!(result.is_err(), "expected error for missing paths");
252
253 let guard = resolver.inner.read().unwrap();
255 drop(guard);
256 }
257
258 #[test]
259 fn reloadable_resolver_reload_from_same_files_succeeds() {
260 let dir = tempfile::tempdir().unwrap();
261 let cert_path = dir.path().join("apimock_test_cert2.pem");
262 let key_path = dir.path().join("apimock_test_key2.pem");
263 let cert_path = cert_path.to_str().unwrap();
264 let key_path = key_path.to_str().unwrap();
265 write_pem_file(cert_path, TEST_CERT_PEM);
266 write_pem_file(key_path, TEST_KEY_PEM);
267
268 let certs = load_certs(cert_path).expect("load test cert");
269 let key = load_private_key(key_path).expect("load test key");
270 let resolver = ReloadableCertResolver::new(certs, key).expect("build resolver");
271
272 let result = resolver.reload_from_paths(cert_path, key_path);
274 assert!(
275 result.is_ok(),
276 "reload from same valid files must succeed: {:?}",
277 result
278 );
279 }
280
281 #[test]
282 fn build_server_config_reloadable_returns_resolver() {
283 let dir = tempfile::tempdir().unwrap();
284 let cert_path = dir.path().join("apimock_test_cert3.pem");
285 let key_path = dir.path().join("apimock_test_key3.pem");
286 let cert_path = cert_path.to_str().unwrap();
287 let key_path = key_path.to_str().unwrap();
288 write_pem_file(cert_path, TEST_CERT_PEM);
289 write_pem_file(key_path, TEST_KEY_PEM);
290
291 let certs = load_certs(cert_path).unwrap();
292 let key = load_private_key(key_path).unwrap();
293 let result = build_server_config_reloadable(certs, key);
294 assert!(
295 result.is_ok(),
296 "build_server_config_reloadable failed: {:?}",
297 result
298 );
299 let (_config, resolver) = result.unwrap();
300 let reload = resolver.reload_from_paths(cert_path, key_path);
302 assert!(reload.is_ok());
303 }
304}