1use std::sync::Arc;
2
3use rustls_pki_types::{CertificateDer, PrivateKeyDer};
4use unb_transport::webtransport::wtransport;
5
6use crate::host::HostError;
7
8pub(crate) fn install_crypto_provider() {
9 let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
10}
11
12pub struct ServerIdentity {
13 certificate_chain: Vec<CertificateDer<'static>>,
14 private_key: PrivateKeyDer<'static>,
15}
16
17impl ServerIdentity {
18 pub fn from_pem(chain_pem: &[u8], key_pem: &[u8]) -> Result<ServerIdentity, HostError> {
19 let certificate_chain = rustls_pemfile::certs(&mut &*chain_pem)
20 .collect::<Result<Vec<_>, _>>()
21 .map_err(|error| HostError::Identity(error.to_string()))?;
22 if certificate_chain.is_empty() {
23 return Err(HostError::Identity(
24 "certificate chain PEM contains no certificates".into(),
25 ));
26 }
27 let private_key = rustls_pemfile::private_key(&mut &*key_pem)
28 .map_err(|error| HostError::Identity(error.to_string()))?
29 .ok_or_else(|| HostError::Identity("private key PEM contains no key".into()))?;
30 Ok(ServerIdentity {
31 certificate_chain,
32 private_key,
33 })
34 }
35
36 pub fn tcp_rustls(&self) -> Result<Arc<rustls::ServerConfig>, HostError> {
37 install_crypto_provider();
38 let mut config = rustls::ServerConfig::builder()
39 .with_no_client_auth()
40 .with_single_cert(self.certificate_chain.clone(), self.private_key.clone_key())
41 .map_err(|error| HostError::Identity(error.to_string()))?;
42 config.alpn_protocols = vec![b"http/1.1".to_vec()];
43 Ok(Arc::new(config))
44 }
45
46 pub fn webtransport(&self) -> Result<wtransport::Identity, HostError> {
47 let chain = self
48 .certificate_chain
49 .iter()
50 .map(|certificate| {
51 wtransport::tls::Certificate::from_der(certificate.as_ref().to_vec())
52 .map_err(|error| HostError::Identity(error.to_string()))
53 })
54 .collect::<Result<Vec<_>, _>>()?;
55 let key = match &self.private_key {
56 PrivateKeyDer::Pkcs8(key) => {
57 wtransport::tls::PrivateKey::from_der_pkcs8(key.secret_pkcs8_der().to_vec())
58 }
59 _ => {
60 return Err(HostError::Identity(
61 "webtransport identities require a PKCS#8 private key".into(),
62 ))
63 }
64 };
65 Ok(wtransport::Identity::new(
66 wtransport::tls::CertificateChain::new(chain),
67 key,
68 ))
69 }
70}
71
72#[cfg(test)]
73mod tests {
74 use super::*;
75
76 fn pem_pair() -> (Vec<u8>, Vec<u8>) {
77 let identity =
78 unb_transport::webtransport::SelfSignedIdentity::generate(["localhost"]).unwrap();
79 let chain = identity
80 .identity()
81 .certificate_chain()
82 .as_slice()
83 .iter()
84 .map(wtransport::tls::Certificate::to_pem)
85 .collect::<String>();
86 let key = identity.identity().private_key().to_secret_pem();
87 (chain.into_bytes(), key.into_bytes())
88 }
89
90 #[test]
91 fn a_pem_credential_projects_into_both_socket_configs() {
92 let (chain, key) = pem_pair();
93 let identity = ServerIdentity::from_pem(&chain, &key).unwrap();
94 let tcp = identity.tcp_rustls().unwrap();
95 assert_eq!(tcp.alpn_protocols, vec![b"http/1.1".to_vec()]);
96 let webtransport = identity.webtransport().unwrap();
97 assert_eq!(webtransport.certificate_chain().as_slice().len(), 1);
98 }
99
100 #[test]
101 fn malformed_pem_is_a_typed_identity_error() {
102 let error = match ServerIdentity::from_pem(b"not-pem", b"also-not-pem") {
103 Err(error) => error,
104 Ok(_) => panic!("malformed pem must fail"),
105 };
106 assert!(matches!(error, HostError::Identity(_)));
107 }
108
109 #[test]
110 fn an_empty_chain_is_rejected_before_any_bind() {
111 let (_, key) = pem_pair();
112 assert!(matches!(
113 ServerIdentity::from_pem(b"", &key),
114 Err(HostError::Identity(_))
115 ));
116 }
117}