1use std::fs::File;
7use std::io::BufReader;
8use std::path::{Path, PathBuf};
9use std::sync::Arc;
10
11use rustls::pki_types::{CertificateDer, PrivateKeyDer};
12use rustls::server::WebPkiClientVerifier;
13use rustls::{RootCertStore, ServerConfig};
14use tokio_rustls::TlsAcceptor;
15use x509_parser::prelude::*;
16
17#[derive(Debug, thiserror::Error)]
19pub enum TlsError {
20 #[error("I/O error reading {path}: {source}")]
21 Io {
22 path: PathBuf,
23 #[source]
24 source: std::io::Error,
25 },
26
27 #[error("no certificates found in {0}")]
28 NoCertificates(PathBuf),
29
30 #[error("no private key found in {0}")]
31 NoPrivateKey(PathBuf),
32
33 #[error("failed to build TLS config: {0}")]
34 ConfigBuild(String),
35
36 #[error("invalid certificate: {0}")]
37 InvalidCertificate(String),
38
39 #[error("TLS error: {0}")]
40 Rustls(#[from] rustls::Error),
41}
42
43#[derive(Debug, Clone)]
45pub struct TlsConfig {
46 pub cert_path: PathBuf,
48 pub key_path: PathBuf,
50 pub ca_cert_path: Option<PathBuf>,
53 pub require_client_cert: bool,
56}
57
58impl TlsConfig {
59 pub fn new(cert_path: impl Into<PathBuf>, key_path: impl Into<PathBuf>) -> Self {
61 Self {
62 cert_path: cert_path.into(),
63 key_path: key_path.into(),
64 ca_cert_path: None,
65 require_client_cert: false,
66 }
67 }
68
69 pub fn with_mtls(
71 cert_path: impl Into<PathBuf>,
72 key_path: impl Into<PathBuf>,
73 ca_cert_path: impl Into<PathBuf>,
74 ) -> Self {
75 Self {
76 cert_path: cert_path.into(),
77 key_path: key_path.into(),
78 ca_cert_path: Some(ca_cert_path.into()),
79 require_client_cert: true,
80 }
81 }
82
83 pub fn load_server_config(&self) -> Result<ServerConfig, TlsError> {
85 let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
87
88 let certs = load_certs(&self.cert_path)?;
90
91 let key = load_private_key(&self.key_path)?;
93
94 let builder = ServerConfig::builder();
96
97 let config = if let Some(ca_path) = &self.ca_cert_path {
98 let ca_certs = load_certs(ca_path)?;
100 let mut root_store = RootCertStore::empty();
101 for cert in ca_certs {
102 root_store
103 .add(cert)
104 .map_err(|e| TlsError::ConfigBuild(format!("failed to add CA cert: {e}")))?;
105 }
106
107 let verifier = if self.require_client_cert {
109 WebPkiClientVerifier::builder(Arc::new(root_store))
110 .build()
111 .map_err(|e| TlsError::ConfigBuild(format!("failed to build verifier: {e}")))?
112 } else {
113 WebPkiClientVerifier::builder(Arc::new(root_store))
114 .allow_unauthenticated()
115 .build()
116 .map_err(|e| TlsError::ConfigBuild(format!("failed to build verifier: {e}")))?
117 };
118
119 builder
120 .with_client_cert_verifier(verifier)
121 .with_single_cert(certs, key)?
122 } else {
123 builder.with_no_client_auth().with_single_cert(certs, key)?
125 };
126
127 Ok(config)
128 }
129
130 pub fn build_acceptor(&self) -> Result<TlsAcceptor, TlsError> {
132 let config = self.load_server_config()?;
133 Ok(TlsAcceptor::from(Arc::new(config)))
134 }
135}
136
137pub fn load_certs(path: &Path) -> Result<Vec<CertificateDer<'static>>, TlsError> {
139 let file = File::open(path).map_err(|e| TlsError::Io {
140 path: path.to_path_buf(),
141 source: e,
142 })?;
143 let mut reader = BufReader::new(file);
144
145 let certs: Vec<CertificateDer<'static>> = rustls_pemfile::certs(&mut reader)
146 .collect::<Result<Vec<_>, _>>()
147 .map_err(|e| TlsError::Io {
148 path: path.to_path_buf(),
149 source: e,
150 })?;
151
152 if certs.is_empty() {
153 return Err(TlsError::NoCertificates(path.to_path_buf()));
154 }
155
156 Ok(certs)
157}
158
159pub fn load_private_key(path: &Path) -> Result<PrivateKeyDer<'static>, TlsError> {
162 let file = File::open(path).map_err(|e| TlsError::Io {
163 path: path.to_path_buf(),
164 source: e,
165 })?;
166 let mut reader = BufReader::new(file);
167
168 loop {
169 match rustls_pemfile::read_one(&mut reader).map_err(|e| TlsError::Io {
170 path: path.to_path_buf(),
171 source: e,
172 })? {
173 Some(rustls_pemfile::Item::Pkcs1Key(key)) => {
174 return Ok(PrivateKeyDer::Pkcs1(key));
175 }
176 Some(rustls_pemfile::Item::Pkcs8Key(key)) => {
177 return Ok(PrivateKeyDer::Pkcs8(key));
178 }
179 Some(rustls_pemfile::Item::Sec1Key(key)) => {
180 return Ok(PrivateKeyDer::Sec1(key));
181 }
182 Some(_) => continue, None => break,
184 }
185 }
186
187 Err(TlsError::NoPrivateKey(path.to_path_buf()))
188}
189
190pub fn extract_client_cn(certs: &[CertificateDer<'_>]) -> Option<String> {
193 let cert = certs.first()?;
194
195 let (_, parsed) = X509Certificate::from_der(cert.as_ref()).ok()?;
197
198 for rdn in parsed.subject().iter_rdn() {
200 for attr in rdn.iter() {
201 if attr.attr_type() == &oid_registry::OID_X509_COMMON_NAME
202 && let Ok(cn) = attr.attr_value().as_str()
203 {
204 return Some(cn.to_string());
205 }
206 }
207 }
208
209 None
210}
211
212pub fn extract_sans(cert: &CertificateDer<'_>) -> Vec<String> {
215 let mut sans = Vec::new();
216
217 if let Ok((_, parsed)) = X509Certificate::from_der(cert.as_ref())
218 && let Ok(Some(ext)) = parsed.subject_alternative_name()
219 {
220 for name in &ext.value.general_names {
221 match name {
222 GeneralName::DNSName(dns) => sans.push(dns.to_string()),
223 GeneralName::IPAddress(ip) => {
224 if ip.len() == 4 {
225 sans.push(format!("{}.{}.{}.{}", ip[0], ip[1], ip[2], ip[3]));
226 } else if ip.len() == 16 {
227 let parts: Vec<String> = ip
229 .chunks(2)
230 .map(|c| format!("{:02x}{:02x}", c[0], c[1]))
231 .collect();
232 sans.push(parts.join(":"));
233 }
234 }
235 _ => {}
236 }
237 }
238 }
239
240 sans
241}
242
243pub fn cn_to_role(cn: &str) -> &'static str {
251 if cn.ends_with("-admin") {
252 "admin"
253 } else if cn.ends_with("-writer") {
254 "writer"
255 } else {
256 "reader"
257 }
258}
259
260#[cfg(test)]
261mod tests {
262 use super::*;
263 use rcgen::{
264 BasicConstraints, CertificateParams, DnType, ExtendedKeyUsagePurpose, IsCa, KeyPair,
265 KeyUsagePurpose, SanType,
266 };
267 use std::io::Write;
268 use tempfile::TempDir;
269
270 fn generate_ca() -> (String, String, CertificateParams, KeyPair) {
272 let mut params = CertificateParams::default();
273 params
274 .distinguished_name
275 .push(DnType::CommonName, "Test CA");
276 params
277 .distinguished_name
278 .push(DnType::OrganizationName, "AstraeaDB Test");
279 params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained);
280 params.key_usages = vec![KeyUsagePurpose::KeyCertSign, KeyUsagePurpose::CrlSign];
281
282 let key_pair = KeyPair::generate().unwrap();
283 let cert = params.clone().self_signed(&key_pair).unwrap();
284
285 (cert.pem(), key_pair.serialize_pem(), params, key_pair)
286 }
287
288 fn generate_signed_cert(
290 cn: &str,
291 ca_params: &CertificateParams,
292 ca_key: &KeyPair,
293 is_server: bool,
294 ) -> (String, String) {
295 let mut params = CertificateParams::default();
296 params.distinguished_name.push(DnType::CommonName, cn);
297 params
298 .distinguished_name
299 .push(DnType::OrganizationName, "AstraeaDB Test");
300
301 if is_server {
302 params.subject_alt_names = vec![
303 SanType::DnsName("localhost".try_into().unwrap()),
304 SanType::IpAddress(std::net::IpAddr::V4(std::net::Ipv4Addr::new(127, 0, 0, 1))),
305 ];
306 params.extended_key_usages = vec![ExtendedKeyUsagePurpose::ServerAuth];
307 } else {
308 params.extended_key_usages = vec![ExtendedKeyUsagePurpose::ClientAuth];
309 }
310
311 let key_pair = KeyPair::generate().unwrap();
312
313 let ca_cert = ca_params.clone().self_signed(ca_key).unwrap();
315 let cert = params.signed_by(&key_pair, &ca_cert, ca_key).unwrap();
316
317 (cert.pem(), key_pair.serialize_pem())
318 }
319
320 fn write_file(dir: &TempDir, name: &str, content: &str) -> PathBuf {
322 let path = dir.path().join(name);
323 let mut file = File::create(&path).unwrap();
324 file.write_all(content.as_bytes()).unwrap();
325 path
326 }
327
328 #[test]
329 fn test_load_certs_valid() {
330 let (ca_pem, _, _, _) = generate_ca();
331 let dir = TempDir::new().unwrap();
332 let path = write_file(&dir, "ca.pem", &ca_pem);
333
334 let certs = load_certs(&path).unwrap();
335 assert_eq!(certs.len(), 1);
336 }
337
338 #[test]
339 fn test_load_certs_missing_file() {
340 let result = load_certs(Path::new("/nonexistent/path/cert.pem"));
341 assert!(matches!(result, Err(TlsError::Io { .. })));
342 }
343
344 #[test]
345 fn test_load_certs_empty_file() {
346 let dir = TempDir::new().unwrap();
347 let path = write_file(&dir, "empty.pem", "");
348
349 let result = load_certs(&path);
350 assert!(matches!(result, Err(TlsError::NoCertificates(_))));
351 }
352
353 #[test]
354 fn test_load_private_key_valid() {
355 let (_, key_pem, _, _) = generate_ca();
356 let dir = TempDir::new().unwrap();
357 let path = write_file(&dir, "key.pem", &key_pem);
358
359 let key = load_private_key(&path);
360 assert!(key.is_ok());
361 }
362
363 #[test]
364 fn test_load_private_key_missing_file() {
365 let result = load_private_key(Path::new("/nonexistent/path/key.pem"));
366 assert!(matches!(result, Err(TlsError::Io { .. })));
367 }
368
369 #[test]
370 fn test_load_private_key_no_key() {
371 let (ca_pem, _, _, _) = generate_ca();
372 let dir = TempDir::new().unwrap();
373 let path = write_file(&dir, "cert.pem", &ca_pem);
375
376 let result = load_private_key(&path);
377 assert!(matches!(result, Err(TlsError::NoPrivateKey(_))));
378 }
379
380 #[test]
381 fn test_extract_client_cn() {
382 let (_, _, ca_params, ca_key) = generate_ca();
383 let (client_pem, _) = generate_signed_cert("test-client-admin", &ca_params, &ca_key, false);
384
385 let dir = TempDir::new().unwrap();
386 let path = write_file(&dir, "client.pem", &client_pem);
387 let certs = load_certs(&path).unwrap();
388
389 let cn = extract_client_cn(&certs);
390 assert_eq!(cn, Some("test-client-admin".to_string()));
391 }
392
393 #[test]
394 fn test_extract_client_cn_empty() {
395 let certs: Vec<CertificateDer<'static>> = vec![];
396 assert_eq!(extract_client_cn(&certs), None);
397 }
398
399 #[test]
400 fn test_extract_sans() {
401 let (_, _, ca_params, ca_key) = generate_ca();
402 let (server_pem, _) = generate_signed_cert("test-server", &ca_params, &ca_key, true);
403
404 let dir = TempDir::new().unwrap();
405 let path = write_file(&dir, "server.pem", &server_pem);
406 let certs = load_certs(&path).unwrap();
407
408 let sans = extract_sans(&certs[0]);
409 assert!(sans.contains(&"localhost".to_string()));
410 assert!(sans.contains(&"127.0.0.1".to_string()));
411 }
412
413 #[test]
414 fn test_cn_to_role() {
415 assert_eq!(cn_to_role("service-admin"), "admin");
416 assert_eq!(cn_to_role("service-writer"), "writer");
417 assert_eq!(cn_to_role("service-reader"), "reader");
418 assert_eq!(cn_to_role("some-other-service"), "reader");
419 }
420
421 #[test]
422 fn test_tls_config_new() {
423 let config = TlsConfig::new("/path/to/cert.pem", "/path/to/key.pem");
424 assert_eq!(config.cert_path, PathBuf::from("/path/to/cert.pem"));
425 assert_eq!(config.key_path, PathBuf::from("/path/to/key.pem"));
426 assert!(config.ca_cert_path.is_none());
427 assert!(!config.require_client_cert);
428 }
429
430 #[test]
431 fn test_tls_config_with_mtls() {
432 let config =
433 TlsConfig::with_mtls("/path/to/cert.pem", "/path/to/key.pem", "/path/to/ca.pem");
434 assert!(config.ca_cert_path.is_some());
435 assert!(config.require_client_cert);
436 }
437
438 #[test]
439 fn test_load_server_config_no_client_auth() {
440 let (_ca_pem, _, ca_params, ca_key) = generate_ca();
441 let (server_pem, server_key_pem) =
442 generate_signed_cert("test-server", &ca_params, &ca_key, true);
443
444 let dir = TempDir::new().unwrap();
445 let cert_path = write_file(&dir, "server.pem", &server_pem);
446 let key_path = write_file(&dir, "server-key.pem", &server_key_pem);
447
448 let config = TlsConfig::new(&cert_path, &key_path);
449 let server_config = config.load_server_config();
450 assert!(server_config.is_ok());
451 }
452
453 #[test]
454 fn test_load_server_config_with_mtls() {
455 let (ca_pem, _, ca_params, ca_key) = generate_ca();
456 let (server_pem, server_key_pem) =
457 generate_signed_cert("test-server", &ca_params, &ca_key, true);
458
459 let dir = TempDir::new().unwrap();
460 let cert_path = write_file(&dir, "server.pem", &server_pem);
461 let key_path = write_file(&dir, "server-key.pem", &server_key_pem);
462 let ca_path = write_file(&dir, "ca.pem", &ca_pem);
463
464 let config = TlsConfig::with_mtls(&cert_path, &key_path, &ca_path);
465 let server_config = config.load_server_config();
466 assert!(server_config.is_ok());
467 }
468
469 #[test]
470 fn test_load_server_config_invalid_cert() {
471 let dir = TempDir::new().unwrap();
472 let cert_path = write_file(&dir, "invalid.pem", "not a certificate");
473 let key_path = write_file(&dir, "key.pem", "not a key");
474
475 let config = TlsConfig::new(&cert_path, &key_path);
476 let result = config.load_server_config();
477 assert!(result.is_err());
478 }
479
480 #[test]
481 fn test_build_acceptor() {
482 let (_, _, ca_params, ca_key) = generate_ca();
483 let (server_pem, server_key_pem) =
484 generate_signed_cert("test-server", &ca_params, &ca_key, true);
485
486 let dir = TempDir::new().unwrap();
487 let cert_path = write_file(&dir, "server.pem", &server_pem);
488 let key_path = write_file(&dir, "server-key.pem", &server_key_pem);
489
490 let config = TlsConfig::new(&cert_path, &key_path);
491 let acceptor = config.build_acceptor();
492 assert!(acceptor.is_ok());
493 }
494}