Skip to main content

unb_server/
identity.rs

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}