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(_))
));
}
}