Skip to main content

astraea_server/
tls.rs

1//! TLS/mTLS support for AstraeaDB server.
2//!
3//! Provides mutual TLS authentication where both server and client present certificates.
4//! Supports optional client certificate verification for standard TLS or enforced mTLS.
5
6use 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/// Errors that can occur during TLS configuration and operations.
18#[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/// TLS configuration for the AstraeaDB server.
44#[derive(Debug, Clone)]
45pub struct TlsConfig {
46    /// Path to the server certificate chain (PEM format).
47    pub cert_path: PathBuf,
48    /// Path to the server private key (PEM format).
49    pub key_path: PathBuf,
50    /// Optional path to CA certificate for client verification.
51    /// When set, enables client certificate verification.
52    pub ca_cert_path: Option<PathBuf>,
53    /// Whether to require a valid client certificate (mTLS).
54    /// If false and ca_cert_path is set, client certs are optional but validated if present.
55    pub require_client_cert: bool,
56}
57
58impl TlsConfig {
59    /// Create a new TLS configuration for server-only TLS (no client verification).
60    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    /// Create a new TLS configuration with mutual TLS (client certificate required).
70    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    /// Load and build a rustls ServerConfig from this configuration.
84    pub fn load_server_config(&self) -> Result<ServerConfig, TlsError> {
85        // Ensure the crypto provider is installed
86        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
87
88        // Load server certificate chain
89        let certs = load_certs(&self.cert_path)?;
90
91        // Load server private key
92        let key = load_private_key(&self.key_path)?;
93
94        // Build the config
95        let builder = ServerConfig::builder();
96
97        let config = if let Some(ca_path) = &self.ca_cert_path {
98            // Load CA certificates for client verification
99            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            // Create client verifier
108            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            // No client verification
124            builder.with_no_client_auth().with_single_cert(certs, key)?
125        };
126
127        Ok(config)
128    }
129
130    /// Create a TlsAcceptor from this configuration.
131    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
137/// Load PEM-encoded certificates from a file.
138pub 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
159/// Load a PEM-encoded private key from a file.
160/// Supports RSA, EC, and PKCS#8 private keys.
161pub 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, // Skip other PEM items (certs, etc.)
183            None => break,
184        }
185    }
186
187    Err(TlsError::NoPrivateKey(path.to_path_buf()))
188}
189
190/// Extract the Common Name (CN) from the first certificate in a chain.
191/// Returns None if the certificate chain is empty or CN cannot be extracted.
192pub fn extract_client_cn(certs: &[CertificateDer<'_>]) -> Option<String> {
193    let cert = certs.first()?;
194
195    // Parse the X.509 certificate
196    let (_, parsed) = X509Certificate::from_der(cert.as_ref()).ok()?;
197
198    // Extract CN from subject
199    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
212/// Extract the Subject Alternative Names (SANs) from a certificate.
213/// Returns DNS names and IP addresses found in the SAN extension.
214pub 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                        // IPv6
228                        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
243/// Map a client certificate CN to a role name.
244/// This is a simple mapping that can be customized.
245///
246/// Default mapping:
247/// - CN ending with "-admin" -> "admin"
248/// - CN ending with "-writer" -> "writer"
249/// - All other CNs -> "reader"
250pub 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    /// Generate a self-signed CA certificate.
271    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    /// Generate a certificate signed by a CA.
289    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        // Sign with CA - clone params since self_signed takes ownership
314        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    /// Write content to a file and return the path.
321    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        // Write cert instead of key
374        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}