Skip to main content

eggress_protocol_reverse/
tls.rs

1//! TLS/mTLS configuration for native reverse control channels.
2//!
3//! Wraps the existing `eggress-transport-tls` builders so reverse control
4//! traffic can be protected by Rustls without reverse-specific cryptography.
5//! Plaintext remains available when no TLS config is present; TLS is applied
6//! before reverse framing/authentication so credentials never cross in
7//! plaintext when configured.
8
9use std::sync::Arc;
10
11/// TLS material for a native reverse server (control listener).
12#[derive(Clone)]
13pub struct ReverseServerTlsConfig {
14    /// Server certificate chain PEM bytes.
15    pub cert_pem: Vec<u8>,
16    /// Server private key PEM bytes (PKCS#8).
17    pub key_pem: Vec<u8>,
18    /// Optional client CA roots PEM for mutual TLS.
19    pub client_ca_pem: Option<Vec<u8>>,
20    /// Require and validate a client certificate.
21    pub require_client_cert: bool,
22}
23
24impl std::fmt::Debug for ReverseServerTlsConfig {
25    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26        // Never print key material; presence flags only.
27        f.debug_struct("ReverseServerTlsConfig")
28            .field("has_cert", &!self.cert_pem.is_empty())
29            .field("has_key", &!self.key_pem.is_empty())
30            .field("has_client_ca", &self.client_ca_pem.is_some())
31            .field("require_client_cert", &self.require_client_cert)
32            .finish()
33    }
34}
35
36impl Drop for ReverseServerTlsConfig {
37    fn drop(&mut self) {
38        use zeroize::Zeroize;
39        self.key_pem.zeroize();
40        if let Some(ref mut ca) = self.client_ca_pem {
41            ca.zeroize();
42        }
43    }
44}
45
46impl ReverseServerTlsConfig {
47    /// Validate impossible combinations before runtime startup.
48    pub fn validate(&self) -> Result<(), crate::ProtocolError> {
49        if self.cert_pem.is_empty() {
50            return Err(crate::ProtocolError::ConfigInvalid(
51                "reverse server TLS requires a certificate".to_string(),
52            ));
53        }
54        if self.key_pem.is_empty() {
55            return Err(crate::ProtocolError::ConfigInvalid(
56                "reverse server TLS requires a private key".to_string(),
57            ));
58        }
59        if self.require_client_cert && self.client_ca_pem.is_none() {
60            return Err(crate::ProtocolError::ConfigInvalid(
61                "reverse server mTLS requires client CA roots when require_client_cert is set"
62                    .to_string(),
63            ));
64        }
65        Ok(())
66    }
67
68    /// Build a shared rustls `ServerConfig` via the shared TLS transport.
69    pub fn build_server_config(&self) -> Result<Arc<rustls::ServerConfig>, crate::ProtocolError> {
70        self.validate()?;
71        let mut builder = eggress_transport_tls::TlsServerConfigBuilder::new()
72            .with_certificate_pem(&self.cert_pem)
73            .map_err(|e| crate::ProtocolError::Tls(format!("invalid server certificate: {e}")))?
74            .with_key_pem(&self.key_pem)
75            .map_err(|e| crate::ProtocolError::Tls(format!("invalid server key: {e}")))?;
76        if let Some(ref ca_pem) = self.client_ca_pem {
77            builder = builder
78                .with_client_ca_pem(ca_pem)
79                .map_err(|e| crate::ProtocolError::Tls(format!("invalid client CA: {e}")))?;
80        }
81        if self.require_client_cert {
82            builder = builder.with_require_client_cert(true);
83        } else if self.client_ca_pem.is_some() {
84            // Verify client certs when presented, but do not require one.
85            // The builder treats any configured client CA as verify-when-present
86            // unless require is set; no extra flag needed beyond the CA.
87        }
88        builder
89            .build()
90            .map_err(|e| crate::ProtocolError::Tls(format!("invalid server TLS config: {e}")))
91    }
92}
93
94/// TLS material for a native reverse client (control dialer).
95#[derive(Clone)]
96pub struct ReverseClientTlsConfig {
97    /// Optional custom CA roots PEM. When `None`, system roots are used,
98    /// consistent with existing upstream TLS policy.
99    pub ca_pem: Option<Vec<u8>>,
100    /// SNI / server name for verification (required; server_addr is an IP
101    /// literal so SNI cannot be derived safely).
102    pub server_name: String,
103    /// Optional client certificate PEM for mutual TLS.
104    pub client_cert_pem: Option<Vec<u8>>,
105    /// Optional client private key PEM for mutual TLS.
106    pub client_key_pem: Option<Vec<u8>>,
107}
108
109impl std::fmt::Debug for ReverseClientTlsConfig {
110    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
111        f.debug_struct("ReverseClientTlsConfig")
112            .field("has_ca", &self.ca_pem.is_some())
113            .field("server_name", &self.server_name)
114            .field("has_client_cert", &self.client_cert_pem.is_some())
115            .field("has_client_key", &self.client_key_pem.is_some())
116            .finish()
117    }
118}
119
120impl Drop for ReverseClientTlsConfig {
121    fn drop(&mut self) {
122        use zeroize::Zeroize;
123        if let Some(ref mut key) = self.client_key_pem {
124            key.zeroize();
125        }
126    }
127}
128
129impl ReverseClientTlsConfig {
130    /// Validate impossible combinations before runtime startup.
131    pub fn validate(&self) -> Result<(), crate::ProtocolError> {
132        if self.server_name.is_empty() {
133            return Err(crate::ProtocolError::ConfigInvalid(
134                "reverse client TLS requires a server_name for SNI/verification".to_string(),
135            ));
136        }
137        // rustls ServerName parsing is the authority for validity; fail here
138        // rather than during the first reconnect attempt.
139        let _ =
140            rustls::pki_types::ServerName::try_from(self.server_name.clone()).map_err(|_| {
141                crate::ProtocolError::ConfigInvalid(format!(
142                    "reverse client TLS has an invalid server_name '{}'",
143                    self.server_name
144                ))
145            })?;
146        if self.client_cert_pem.is_some() != self.client_key_pem.is_some() {
147            return Err(crate::ProtocolError::ConfigInvalid(
148                "reverse client mTLS requires both client certificate and key".to_string(),
149            ));
150        }
151        Ok(())
152    }
153
154    /// Build a shared rustls `ClientConfig` via the shared TLS transport.
155    ///
156    /// The returned config is immutable and cheap to clone (`Arc`); reconnect
157    /// loops should reuse it rather than rebuilding per attempt.
158    pub fn build_client_config(&self) -> Result<Arc<rustls::ClientConfig>, crate::ProtocolError> {
159        self.validate()?;
160        let mut builder = eggress_transport_tls::TlsClientConfigBuilder::new();
161        builder = match self.ca_pem.as_deref() {
162            Some(ca_pem) => builder
163                .with_custom_ca_pem(ca_pem)
164                .map_err(|e| crate::ProtocolError::Tls(format!("invalid client CA: {e}")))?,
165            None => builder.with_system_roots().map_err(|e| {
166                crate::ProtocolError::Tls(format!("TLS system roots unavailable: {e}"))
167            })?,
168        };
169        if let (Some(cert_pem), Some(key_pem)) =
170            (self.client_cert_pem.as_ref(), self.client_key_pem.as_ref())
171        {
172            builder = builder.with_client_cert_pem(cert_pem, key_pem);
173        }
174        builder
175            .build()
176            .map_err(|e| crate::ProtocolError::Tls(format!("invalid client TLS config: {e}")))
177    }
178}
179
180#[cfg(test)]
181mod tests {
182    use super::*;
183
184    fn init_crypto() {
185        eggress_transport_tls::install_default_crypto_provider();
186    }
187
188    fn cert_for(names: Vec<String>) -> (String, String) {
189        let params = rcgen::CertificateParams::new(names).unwrap();
190        let key = rcgen::KeyPair::generate().unwrap();
191        let cert = params.self_signed(&key).unwrap();
192        (cert.pem(), key.serialize_pem())
193    }
194
195    #[test]
196    fn server_tls_debug_redacts_key_material() {
197        let (cert, key) = cert_for(vec!["localhost".to_string()]);
198        let cfg = ReverseServerTlsConfig {
199            cert_pem: cert.into_bytes(),
200            key_pem: key.into_bytes(),
201            client_ca_pem: None,
202            require_client_cert: false,
203        };
204        let rendered = format!("{cfg:?}");
205        assert!(!rendered.contains("BEGIN PRIVATE KEY"));
206        assert!(!rendered.contains("BEGIN CERTIFICATE"));
207    }
208
209    #[test]
210    fn server_tls_require_without_ca_rejected() {
211        let (cert, key) = cert_for(vec!["localhost".to_string()]);
212        let cfg = ReverseServerTlsConfig {
213            cert_pem: cert.into_bytes(),
214            key_pem: key.into_bytes(),
215            client_ca_pem: None,
216            require_client_cert: true,
217        };
218        assert!(cfg.validate().is_err());
219        assert!(cfg.build_server_config().is_err());
220    }
221
222    #[test]
223    fn server_tls_malformed_pem_rejected() {
224        let cfg = ReverseServerTlsConfig {
225            cert_pem: b"not pem".to_vec(),
226            key_pem: b"not pem".to_vec(),
227            client_ca_pem: None,
228            require_client_cert: false,
229        };
230        assert!(cfg.build_server_config().is_err());
231    }
232
233    #[test]
234    fn client_tls_cert_without_key_rejected() {
235        let (cert, _) = cert_for(vec!["localhost".to_string()]);
236        let cfg = ReverseClientTlsConfig {
237            ca_pem: None,
238            server_name: "localhost".to_string(),
239            client_cert_pem: Some(cert.into_bytes()),
240            client_key_pem: None,
241        };
242        assert!(cfg.validate().is_err());
243    }
244
245    #[test]
246    fn client_tls_missing_server_name_rejected() {
247        let cfg = ReverseClientTlsConfig {
248            ca_pem: None,
249            server_name: String::new(),
250            client_cert_pem: None,
251            client_key_pem: None,
252        };
253        assert!(cfg.validate().is_err());
254    }
255
256    #[test]
257    fn client_tls_debug_redacts_key_material() {
258        init_crypto();
259        let (cert, key) = cert_for(vec!["localhost".to_string()]);
260        let cfg = ReverseClientTlsConfig {
261            ca_pem: Some(cert.clone().into_bytes()),
262            server_name: "localhost".to_string(),
263            client_cert_pem: Some(cert.into_bytes()),
264            client_key_pem: Some(key.into_bytes()),
265        };
266        let rendered = format!("{cfg:?}");
267        assert!(!rendered.contains("BEGIN PRIVATE KEY"));
268    }
269}