1use std::sync::{Arc, Mutex};
27
28pub use rustls;
30
31use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
32use rustls::crypto::CryptoProvider;
33use rustls::{
34 ClientConfig, DigitallySignedStruct, Error as TlsError, RootCertStore, SignatureScheme,
35};
36use rustls_pki_types::pem::PemObject;
37use rustls_pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime};
38use sha2::{Digest, Sha256};
39
40pub fn fingerprint_of_der(der: &[u8]) -> String {
46 hex::encode(Sha256::digest(der))
47}
48
49pub fn normalize_fingerprint(input: &str) -> String {
54 input
55 .trim()
56 .trim_start_matches("sha256:")
57 .trim_start_matches("SHA256:")
58 .replace([':', ' '], "")
59 .to_ascii_lowercase()
60}
61
62#[derive(Debug, Clone)]
64pub enum TlsPolicy {
65 Ca,
67 Pin(String),
69 PrivateCa(Vec<CertificateDer<'static>>),
71}
72
73impl TlsPolicy {
74 pub fn verifies(&self) -> bool {
81 true
82 }
83}
84
85fn provider() -> Arc<CryptoProvider> {
86 Arc::new(rustls::crypto::aws_lc_rs::default_provider())
87}
88
89pub fn client_config(policy: &TlsPolicy) -> Result<ClientConfig, String> {
91 client_config_with(policy, false)
92}
93
94pub fn client_config_tls13(policy: &TlsPolicy) -> Result<ClientConfig, String> {
98 client_config_with(policy, true)
99}
100
101fn client_config_with(policy: &TlsPolicy, tls13_only: bool) -> Result<ClientConfig, String> {
102 let versions = |b: rustls::ConfigBuilder<ClientConfig, rustls::WantsVersions>| {
103 if tls13_only {
104 b.with_protocol_versions(&[&rustls::version::TLS13])
105 } else {
106 b.with_safe_default_protocol_versions()
107 }
108 .map_err(|e| e.to_string())
109 };
110 match policy {
111 TlsPolicy::Ca => {
112 let mut roots = RootCertStore::empty();
113 roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
114 Ok(versions(ClientConfig::builder_with_provider(provider()))?
115 .with_root_certificates(roots)
116 .with_no_client_auth())
117 }
118 TlsPolicy::Pin(expected) => {
119 let p = provider();
120 Ok(versions(ClientConfig::builder_with_provider(p.clone()))?
121 .dangerous()
122 .with_custom_certificate_verifier(Arc::new(FingerprintVerifier {
123 expected: normalize_fingerprint(expected),
124 provider: p,
125 }))
126 .with_no_client_auth())
127 }
128 TlsPolicy::PrivateCa(certs) => {
129 let mut roots = RootCertStore::empty();
130 for c in certs {
131 roots
132 .add(c.clone())
133 .map_err(|e| format!("Not a usable CA certificate: {e}"))?;
134 }
135 if roots.is_empty() {
136 return Err("No certificates found in the supplied CA file".into());
137 }
138 Ok(versions(ClientConfig::builder_with_provider(provider()))?
139 .with_root_certificates(roots)
140 .with_no_client_auth())
141 }
142 }
143}
144
145pub fn self_signed(names: Vec<String>) -> Result<(String, String, String), String> {
149 use rcgen::{CertificateParams, KeyPair};
150 let key = KeyPair::generate().map_err(|e| e.to_string())?;
151 let mut params = CertificateParams::new(names).map_err(|e| e.to_string())?;
152 params.not_after = time::OffsetDateTime::now_utc()
153 .checked_add(time::Duration::days(365 * 5))
154 .ok_or("date overflow")?;
155 let cert = params.self_signed(&key).map_err(|e| e.to_string())?;
156 Ok((
157 cert.pem(),
158 key.serialize_pem(),
159 fingerprint_of_der(cert.der().as_ref()),
160 ))
161}
162
163pub fn fingerprint_of_pem(cert_pem: &str) -> Result<String, String> {
165 let certs = certs_from_pem(cert_pem.as_bytes())?;
166 Ok(fingerprint_of_der(certs[0].as_ref()))
167}
168
169pub fn server_config_tls13(cert_pem: &str, key_pem: &str) -> Result<rustls::ServerConfig, String> {
172 let chain = certs_from_pem(cert_pem.as_bytes())?;
173 let key = match PrivateKeyDer::from_pem_slice(key_pem.as_bytes()) {
174 Ok(k) => k,
175 Err(rustls_pki_types::pem::Error::NoItemsFound) => {
176 return Err("The node's TLS key file holds no private key".into())
177 }
178 Err(e) => return Err(format!("Cannot read the node's TLS key: {e}")),
179 };
180 rustls::ServerConfig::builder_with_provider(provider())
181 .with_protocol_versions(&[&rustls::version::TLS13])
182 .map_err(|e| e.to_string())?
183 .with_no_client_auth()
184 .with_single_cert(chain, key)
185 .map_err(|e| e.to_string())
186}
187
188pub fn certs_from_pem(pem: &[u8]) -> Result<Vec<CertificateDer<'static>>, String> {
190 let certs: Result<Vec<_>, _> = CertificateDer::pem_slice_iter(pem).collect();
191 let certs = certs.map_err(|e| format!("Cannot read CA file: {e}"))?;
192 if certs.is_empty() {
193 return Err("CA file contains no CERTIFICATE block".into());
194 }
195 Ok(certs)
196}
197
198pub type ProbeConfig = (ClientConfig, Arc<Mutex<Option<String>>>);
201
202pub fn capturing_config() -> Result<ProbeConfig, String> {
210 let seen = Arc::new(Mutex::new(None::<String>));
211 let p = provider();
212 let cfg = ClientConfig::builder_with_provider(p.clone())
213 .with_safe_default_protocol_versions()
214 .map_err(|e| e.to_string())?
215 .dangerous()
216 .with_custom_certificate_verifier(Arc::new(CapturingVerifier {
217 seen: seen.clone(),
218 provider: p,
219 }))
220 .with_no_client_auth();
221 Ok((cfg, seen))
222}
223
224#[derive(Debug)]
230struct FingerprintVerifier {
231 expected: String,
232 provider: Arc<CryptoProvider>,
233}
234
235impl ServerCertVerifier for FingerprintVerifier {
236 fn verify_server_cert(
237 &self,
238 end_entity: &CertificateDer<'_>,
239 _intermediates: &[CertificateDer<'_>],
240 _server_name: &ServerName<'_>,
241 _ocsp_response: &[u8],
242 _now: UnixTime,
243 ) -> Result<ServerCertVerified, TlsError> {
244 let actual = fingerprint_of_der(end_entity.as_ref());
245 if actual == self.expected {
246 Ok(ServerCertVerified::assertion())
247 } else {
248 Err(TlsError::General(
249 "TLS certificate fingerprint does not match the pinned value".into(),
250 ))
251 }
252 }
253
254 fn verify_tls12_signature(
255 &self,
256 message: &[u8],
257 cert: &CertificateDer<'_>,
258 dss: &DigitallySignedStruct,
259 ) -> Result<HandshakeSignatureValid, TlsError> {
260 rustls::crypto::verify_tls12_signature(
261 message,
262 cert,
263 dss,
264 &self.provider.signature_verification_algorithms,
265 )
266 }
267
268 fn verify_tls13_signature(
269 &self,
270 message: &[u8],
271 cert: &CertificateDer<'_>,
272 dss: &DigitallySignedStruct,
273 ) -> Result<HandshakeSignatureValid, TlsError> {
274 rustls::crypto::verify_tls13_signature(
275 message,
276 cert,
277 dss,
278 &self.provider.signature_verification_algorithms,
279 )
280 }
281
282 fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
283 self.provider
284 .signature_verification_algorithms
285 .supported_schemes()
286 }
287}
288
289#[derive(Debug)]
292struct CapturingVerifier {
293 seen: Arc<Mutex<Option<String>>>,
294 provider: Arc<CryptoProvider>,
295}
296
297impl ServerCertVerifier for CapturingVerifier {
298 fn verify_server_cert(
299 &self,
300 end_entity: &CertificateDer<'_>,
301 _intermediates: &[CertificateDer<'_>],
302 _server_name: &ServerName<'_>,
303 _ocsp_response: &[u8],
304 _now: UnixTime,
305 ) -> Result<ServerCertVerified, TlsError> {
306 *self.seen.lock().unwrap() = Some(fingerprint_of_der(end_entity.as_ref()));
307 Ok(ServerCertVerified::assertion())
308 }
309
310 fn verify_tls12_signature(
311 &self,
312 message: &[u8],
313 cert: &CertificateDer<'_>,
314 dss: &DigitallySignedStruct,
315 ) -> Result<HandshakeSignatureValid, TlsError> {
316 rustls::crypto::verify_tls12_signature(
317 message,
318 cert,
319 dss,
320 &self.provider.signature_verification_algorithms,
321 )
322 }
323
324 fn verify_tls13_signature(
325 &self,
326 message: &[u8],
327 cert: &CertificateDer<'_>,
328 dss: &DigitallySignedStruct,
329 ) -> Result<HandshakeSignatureValid, TlsError> {
330 rustls::crypto::verify_tls13_signature(
331 message,
332 cert,
333 dss,
334 &self.provider.signature_verification_algorithms,
335 )
336 }
337
338 fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
339 self.provider
340 .signature_verification_algorithms
341 .supported_schemes()
342 }
343}
344
345#[cfg(test)]
346mod tests {
347 use super::*;
348
349 #[test]
350 fn fingerprint_is_stable_and_lowercase_hex() {
351 let fp = fingerprint_of_der(b"not really a certificate");
352 assert_eq!(fp.len(), 64);
353 assert_eq!(fp, fp.to_ascii_lowercase());
354 assert_eq!(fp, fingerprint_of_der(b"not really a certificate"));
355 assert_ne!(fp, fingerprint_of_der(b"a different certificate"));
356 }
357
358 #[test]
359 fn normalize_accepts_the_forms_people_paste() {
360 let canonical = "ab12cd34";
362 for input in [
363 "AB:12:CD:34",
364 "sha256:AB12CD34",
365 " ab12cd34 ",
366 "AB 12 CD 34",
367 ] {
368 assert_eq!(normalize_fingerprint(input), canonical, "input: {input}");
369 }
370 }
371
372 #[test]
373 fn pin_config_builds_and_normalizes_its_expectation() {
374 let cfg = client_config(&TlsPolicy::Pin("AB:CD".into()));
377 assert!(cfg.is_ok());
378 }
379
380 #[test]
381 fn private_ca_rejects_a_file_with_no_certificates() {
382 let err = certs_from_pem(b"-----BEGIN PRIVATE KEY-----\nzzz\n-----END PRIVATE KEY-----\n")
383 .unwrap_err();
384 assert!(err.contains("no CERTIFICATE"), "{err}");
385 }
386
387 #[test]
391 fn tls_identity_load_refuses_bad_keys_and_certificates() {
392 let (cert, key, _) = self_signed(vec!["node.test".into()]).unwrap();
393 assert!(server_config_tls13(&cert, &key).is_ok());
394
395 let err = server_config_tls13(&cert, &cert).unwrap_err();
397 assert!(err.contains("no private key"), "{err}");
398 let err = server_config_tls13(&cert, "").unwrap_err();
400 assert!(err.contains("no private key"), "{err}");
401 let broken_key = key.replacen("MIG", "M!G", 1).replacen("MC4", "M!4", 1);
403 let broken_key = if broken_key == key {
404 let mut lines: Vec<String> = key.lines().map(str::to_owned).collect();
406 lines[1] = format!("!!{}", lines[1]);
407 lines.join("\n")
408 } else {
409 broken_key
410 };
411 assert!(server_config_tls13(&cert, &broken_key).is_err());
412 let mut lines: Vec<String> = cert.lines().map(str::to_owned).collect();
414 lines[1] = format!("!!{}", lines[1]);
415 let err = certs_from_pem(lines.join("\n").as_bytes()).unwrap_err();
416 assert!(err.contains("Cannot read CA file"), "{err}");
417 let (_, other_key, _) = self_signed(vec!["other.test".into()]).unwrap();
419 assert!(server_config_tls13(&cert, &other_key).is_err());
420 }
421
422 #[test]
423 fn a_pinned_client_reaches_the_listener_only_with_the_right_certificate() {
424 use std::io::{Read, Write};
425 let (cert, key, fp) = self_signed(vec!["node.test".into()]).unwrap();
426 assert_eq!(fingerprint_of_pem(&cert).unwrap(), fp);
427 let cfg = Arc::new(server_config_tls13(&cert, &key).unwrap());
428 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
429 let port = listener.local_addr().unwrap().port();
430 std::thread::spawn(move || {
431 for _ in 0..3 {
432 let Ok((mut tcp, _)) = listener.accept() else {
433 return;
434 };
435 let mut conn = rustls::ServerConnection::new(cfg.clone()).unwrap();
436 let mut tls = rustls::Stream::new(&mut conn, &mut tcp);
437 let mut buf = [0u8; 4];
438 if tls.read_exact(&mut buf).is_ok() {
439 let _ = tls.write_all(b"pong");
440 }
441 }
442 });
443 let talk = |pin: &str, tls12: bool| -> Result<Vec<u8>, String> {
444 let cfg = if tls12 {
445 client_config(&TlsPolicy::Pin(pin.into()))
446 } else {
447 client_config_tls13(&TlsPolicy::Pin(pin.into()))
448 }?;
449 let name = ServerName::try_from("node.test").unwrap();
450 let mut conn =
451 rustls::ClientConnection::new(Arc::new(cfg), name).map_err(|e| e.to_string())?;
452 let mut tcp =
453 std::net::TcpStream::connect(("127.0.0.1", port)).map_err(|e| e.to_string())?;
454 let mut tls = rustls::Stream::new(&mut conn, &mut tcp);
455 tls.write_all(b"ping").map_err(|e| e.to_string())?;
456 let mut out = vec![0u8; 4];
457 tls.read_exact(&mut out).map_err(|e| e.to_string())?;
458 Ok(out)
459 };
460 assert_eq!(talk(&fp, false).unwrap(), b"pong");
461 assert!(
462 talk(&"00".repeat(32), false).is_err(),
463 "a wrong pin must not connect"
464 );
465 assert_eq!(
467 talk(&fp, true).unwrap(),
468 b"pong",
469 "a client offering 1.2 and 1.3 lands on 1.3"
470 );
471 }
472}