unb-server 2.0.3

unb inbound server: Node, request/subscribe handlers, catalog, relay orchestration, accept
Documentation
use std::sync::Arc;

use rustls_pki_types::{CertificateDer, PrivateKeyDer};
use unb_transport::webtransport::wtransport;

use crate::host::HostError;

pub(crate) fn install_crypto_provider() {
    let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
}

pub struct ServerIdentity {
    certificate_chain: Vec<CertificateDer<'static>>,
    private_key: PrivateKeyDer<'static>,
}

impl ServerIdentity {
    pub fn from_pem(chain_pem: &[u8], key_pem: &[u8]) -> Result<ServerIdentity, HostError> {
        let certificate_chain = rustls_pemfile::certs(&mut &*chain_pem)
            .collect::<Result<Vec<_>, _>>()
            .map_err(|error| HostError::Identity(error.to_string()))?;
        if certificate_chain.is_empty() {
            return Err(HostError::Identity(
                "certificate chain PEM contains no certificates".into(),
            ));
        }
        let private_key = rustls_pemfile::private_key(&mut &*key_pem)
            .map_err(|error| HostError::Identity(error.to_string()))?
            .ok_or_else(|| HostError::Identity("private key PEM contains no key".into()))?;
        Ok(ServerIdentity {
            certificate_chain,
            private_key,
        })
    }

    pub fn tcp_rustls(&self) -> Result<Arc<rustls::ServerConfig>, HostError> {
        install_crypto_provider();
        let mut config = rustls::ServerConfig::builder()
            .with_no_client_auth()
            .with_single_cert(self.certificate_chain.clone(), self.private_key.clone_key())
            .map_err(|error| HostError::Identity(error.to_string()))?;
        config.alpn_protocols = vec![b"http/1.1".to_vec()];
        Ok(Arc::new(config))
    }

    pub fn webtransport(&self) -> Result<wtransport::Identity, HostError> {
        let chain = self
            .certificate_chain
            .iter()
            .map(|certificate| {
                wtransport::tls::Certificate::from_der(certificate.as_ref().to_vec())
                    .map_err(|error| HostError::Identity(error.to_string()))
            })
            .collect::<Result<Vec<_>, _>>()?;
        let key = match &self.private_key {
            PrivateKeyDer::Pkcs8(key) => {
                wtransport::tls::PrivateKey::from_der_pkcs8(key.secret_pkcs8_der().to_vec())
            }
            _ => {
                return Err(HostError::Identity(
                    "webtransport identities require a PKCS#8 private key".into(),
                ))
            }
        };
        Ok(wtransport::Identity::new(
            wtransport::tls::CertificateChain::new(chain),
            key,
        ))
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    fn pem_pair() -> (Vec<u8>, Vec<u8>) {
        let identity =
            unb_transport::webtransport::SelfSignedIdentity::generate(["localhost"]).unwrap();
        let chain = identity
            .identity()
            .certificate_chain()
            .as_slice()
            .iter()
            .map(wtransport::tls::Certificate::to_pem)
            .collect::<String>();
        let key = identity.identity().private_key().to_secret_pem();
        (chain.into_bytes(), key.into_bytes())
    }

    #[test]
    fn a_pem_credential_projects_into_both_socket_configs() {
        let (chain, key) = pem_pair();
        let identity = ServerIdentity::from_pem(&chain, &key).unwrap();
        let tcp = identity.tcp_rustls().unwrap();
        assert_eq!(tcp.alpn_protocols, vec![b"http/1.1".to_vec()]);
        let webtransport = identity.webtransport().unwrap();
        assert_eq!(webtransport.certificate_chain().as_slice().len(), 1);
    }

    #[test]
    fn malformed_pem_is_a_typed_identity_error() {
        let error = match ServerIdentity::from_pem(b"not-pem", b"also-not-pem") {
            Err(error) => error,
            Ok(_) => panic!("malformed pem must fail"),
        };
        assert!(matches!(error, HostError::Identity(_)));
    }

    #[test]
    fn an_empty_chain_is_rejected_before_any_bind() {
        let (_, key) = pem_pair();
        assert!(matches!(
            ServerIdentity::from_pem(b"", &key),
            Err(HostError::Identity(_))
        ));
    }
}