1use std::path::PathBuf;
10
11use crate::{PemSource, TlsError};
12
13#[derive(Debug, Clone)]
29pub struct ClientTlsConfig {
30 pub ca: PemSource,
32 pub client_cert: Option<PemSource>,
34 pub client_key: Option<PemSource>,
36 pub alpn: Vec<Vec<u8>>,
38}
39
40impl ClientTlsConfig {
41 pub fn builder() -> ClientTlsConfigBuilder {
43 ClientTlsConfigBuilder::default()
44 }
45
46 pub fn into_rustls_config(self) -> Result<rustls::ClientConfig, TlsError> {
67 crate::ensure_default_provider();
68
69 let ca_bytes = self.ca.read()?;
70 let ca_certs = crate::load_certs_from_pem(ca_bytes.as_slice())?;
71 let mut roots = rustls::RootCertStore::empty();
72 for ca in ca_certs {
73 roots.add(ca)?;
74 }
75
76 let builder = rustls::ClientConfig::builder().with_root_certificates(roots);
77
78 let mut config = match (self.client_cert, self.client_key) {
79 (Some(cert_src), Some(key_src)) => {
80 let cert_bytes = cert_src.read()?;
81 let key_bytes = key_src.read()?;
82 let certs = crate::load_certs_from_pem(cert_bytes.as_slice())?;
83 let key = crate::load_key_from_pem(key_bytes.as_slice())?;
84 builder.with_client_auth_cert(certs, key)?
85 }
86 _ => builder.with_no_client_auth(),
87 };
88
89 config.alpn_protocols = self.alpn;
90 Ok(config)
91 }
92}
93
94#[derive(Debug, Default, Clone)]
96pub struct ClientTlsConfigBuilder {
97 client_cert: Option<PemSource>,
98 client_key: Option<PemSource>,
99 ca: Option<PemSource>,
100 alpn: Vec<Vec<u8>>,
101}
102
103impl ClientTlsConfigBuilder {
104 pub fn ca(mut self, src: PemSource) -> Self {
106 self.ca = Some(src);
107 self
108 }
109
110 pub fn client_cert(mut self, src: PemSource) -> Self {
112 self.client_cert = Some(src);
113 self
114 }
115
116 pub fn client_key(mut self, src: PemSource) -> Self {
118 self.client_key = Some(src);
119 self
120 }
121
122 pub fn with_alpn<I, S>(mut self, protocols: I) -> Self
126 where
127 I: IntoIterator<Item = S>,
128 S: Into<Vec<u8>>,
129 {
130 self.alpn = protocols.into_iter().map(Into::into).collect();
131 self
132 }
133
134 pub fn ca_pem_file(self, path: impl Into<PathBuf>) -> Self {
136 self.ca(PemSource::Path(path.into()))
137 }
138
139 pub fn ca_pem_bytes(self, bytes: impl Into<Vec<u8>>) -> Self {
141 self.ca(PemSource::Bytes(bytes.into()))
142 }
143
144 pub fn client_cert_pem_file(self, path: impl Into<PathBuf>) -> Self {
146 self.client_cert(PemSource::Path(path.into()))
147 }
148
149 pub fn client_cert_pem_bytes(self, bytes: impl Into<Vec<u8>>) -> Self {
151 self.client_cert(PemSource::Bytes(bytes.into()))
152 }
153
154 pub fn client_key_pem_file(self, path: impl Into<PathBuf>) -> Self {
156 self.client_key(PemSource::Path(path.into()))
157 }
158
159 pub fn client_key_pem_bytes(self, bytes: impl Into<Vec<u8>>) -> Self {
161 self.client_key(PemSource::Bytes(bytes.into()))
162 }
163
164 pub fn build(self) -> Result<ClientTlsConfig, TlsError> {
184 let ca = self.ca.ok_or(TlsError::MissingField("ca"))?;
185 match (&self.client_cert, &self.client_key) {
186 (Some(_), None) => return Err(TlsError::MissingField("client_key")),
187 (None, Some(_)) => return Err(TlsError::MissingField("client_cert")),
188 _ => {}
189 }
190 Ok(ClientTlsConfig {
191 ca,
192 client_cert: self.client_cert,
193 client_key: self.client_key,
194 alpn: self.alpn,
195 })
196 }
197}
198
199#[cfg(test)]
200mod tests {
201 use super::*;
202 use crate::PemSource;
203
204 #[test]
205 fn builder_returns_config_with_ca() {
206 let cfg = ClientTlsConfig::builder()
207 .ca_pem_bytes(b"--FAKE CA--".to_vec())
208 .build()
209 .unwrap();
210 assert!(matches!(cfg.ca, PemSource::Bytes(_)));
211 assert!(cfg.client_cert.is_none());
212 assert!(cfg.client_key.is_none());
213 assert!(cfg.alpn.is_empty());
214 }
215
216 #[test]
217 fn builder_errors_when_ca_is_missing() {
218 let err = ClientTlsConfig::builder().build().unwrap_err();
219 assert!(matches!(err, TlsError::MissingField("ca")));
220 }
221
222 #[test]
223 fn with_client_cert_pair_enables_mtls() {
224 let cfg = ClientTlsConfig::builder()
225 .ca_pem_bytes(vec![1])
226 .client_cert_pem_bytes(b"cert".to_vec())
227 .client_key_pem_bytes(b"key".to_vec())
228 .build()
229 .unwrap();
230 assert!(matches!(cfg.client_cert, Some(PemSource::Bytes(_))));
231 assert!(matches!(cfg.client_key, Some(PemSource::Bytes(_))));
232 }
233
234 #[test]
235 fn builder_errors_when_client_cert_without_key() {
236 let err = ClientTlsConfig::builder()
237 .ca_pem_bytes(vec![1])
238 .client_cert_pem_bytes(b"cert".to_vec())
239 .build()
240 .unwrap_err();
241 assert!(matches!(err, TlsError::MissingField("client_key")));
242 }
243
244 #[test]
245 fn builder_errors_when_client_key_without_cert() {
246 let err = ClientTlsConfig::builder()
247 .ca_pem_bytes(vec![1])
248 .client_key_pem_bytes(b"key".to_vec())
249 .build()
250 .unwrap_err();
251 assert!(matches!(err, TlsError::MissingField("client_cert")));
252 }
253
254 #[test]
255 fn with_alpn_sets_protocols() {
256 let cfg = ClientTlsConfig::builder()
257 .ca_pem_bytes(vec![1])
258 .with_alpn(["h2", "http/1.1"])
259 .build()
260 .unwrap();
261 assert_eq!(cfg.alpn, vec![b"h2".to_vec(), b"http/1.1".to_vec()]);
262 }
263
264 fn rcgen_self_signed() -> (Vec<u8>, Vec<u8>) {
265 let b = rcgen::generate_simple_self_signed(vec!["example.com".into()]).unwrap();
266 (
267 b.cert.pem().into_bytes(),
268 b.signing_key.serialize_pem().into_bytes(),
269 )
270 }
271
272 #[test]
273 fn into_rustls_config_succeeds_with_ca_only() {
274 let (ca, _) = rcgen_self_signed();
275 let cfg = ClientTlsConfig::builder().ca_pem_bytes(ca).build().unwrap();
276 let _rustls = cfg.into_rustls_config().unwrap();
277 }
278
279 #[test]
280 fn into_rustls_config_succeeds_with_mtls_client_cert() {
281 let (ca, _) = rcgen_self_signed();
282 let (cert, key) = rcgen_self_signed();
283 let cfg = ClientTlsConfig::builder()
284 .ca_pem_bytes(ca)
285 .client_cert_pem_bytes(cert)
286 .client_key_pem_bytes(key)
287 .build()
288 .unwrap();
289 let _rustls = cfg.into_rustls_config().unwrap();
290 }
291
292 #[test]
293 fn into_rustls_config_propagates_alpn_to_rustls() {
294 let (ca, _) = rcgen_self_signed();
295 let cfg = ClientTlsConfig::builder()
296 .ca_pem_bytes(ca)
297 .with_alpn(["h2"])
298 .build()
299 .unwrap();
300 let rustls = cfg.into_rustls_config().unwrap();
301 assert_eq!(rustls.alpn_protocols, vec![b"h2".to_vec()]);
302 }
303}