use std::io;
use std::sync::Arc;
use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName};
use rustls::{ClientConfig as RustlsConfig, RootCertStore};
use tokio_rustls::TlsConnector;
use crate::error::{Error, Result};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TrustAnchors {
System,
Pem(Vec<u8>),
SystemAndPem(Vec<u8>),
}
#[derive(Clone)]
pub struct ClientCertificate {
pub chain_pem: Vec<u8>,
pub key_pem: Vec<u8>,
}
impl std::fmt::Debug for ClientCertificate {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ClientCertificate")
.field(
"chain_pem",
&format_args!("<{} bytes>", self.chain_pem.len()),
)
.field("key_pem", &"<redacted>")
.finish()
}
}
#[derive(Debug, Clone)]
pub struct TlsConfig {
pub anchors: TrustAnchors,
pub client_certificate: Option<ClientCertificate>,
pub server_name_override: Option<String>,
}
impl Default for TlsConfig {
fn default() -> Self {
Self {
anchors: TrustAnchors::System,
client_certificate: None,
server_name_override: None,
}
}
}
impl TlsConfig {
pub fn system() -> Self {
Self::default()
}
pub fn with_ca_pem(pem: impl Into<Vec<u8>>) -> Self {
Self {
anchors: TrustAnchors::Pem(pem.into()),
..Self::default()
}
}
pub fn with_client_certificate(
mut self,
chain_pem: impl Into<Vec<u8>>,
key_pem: impl Into<Vec<u8>>,
) -> Self {
self.client_certificate = Some(ClientCertificate {
chain_pem: chain_pem.into(),
key_pem: key_pem.into(),
});
self
}
pub fn with_server_name(mut self, name: impl Into<String>) -> Self {
self.server_name_override = Some(name.into());
self
}
pub fn connector(&self) -> Result<TlsConnector> {
let provider = Arc::new(rustls::crypto::ring::default_provider());
let builder = RustlsConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.map_err(|e| Error::Unsupported(format!("rustls protocol versions: {e}")))?
.with_root_certificates(self.root_store()?);
let config = match &self.client_certificate {
None => builder.with_no_client_auth(),
Some(cert) => {
let chain = parse_certs(&cert.chain_pem)?;
let key = parse_key(&cert.key_pem)?;
builder
.with_client_auth_cert(chain, key)
.map_err(|e| Error::Unsupported(format!("client certificate rejected: {e}")))?
}
};
Ok(TlsConnector::from(Arc::new(config)))
}
pub fn server_name(&self, host: &str) -> Result<ServerName<'static>> {
let name = self.server_name_override.as_deref().unwrap_or(host);
ServerName::try_from(name.to_owned())
.map_err(|e| Error::InvalidRequest(format!("invalid TLS server name {name:?}: {e}")))
}
fn root_store(&self) -> Result<RootCertStore> {
let mut store = RootCertStore::empty();
let (use_system, pem) = match &self.anchors {
TrustAnchors::System => (true, None),
TrustAnchors::Pem(pem) => (false, Some(pem)),
TrustAnchors::SystemAndPem(pem) => (true, Some(pem)),
};
if use_system {
let loaded = rustls_native_certs::load_native_certs();
for error in &loaded.errors {
tracing::warn!(%error, "skipping unreadable system trust anchor");
}
for cert in loaded.certs {
if let Err(error) = store.add(cert) {
tracing::warn!(%error, "skipping unusable system trust anchor");
}
}
}
if let Some(pem) = pem {
let certs = parse_certs(pem)?;
if certs.is_empty() {
return Err(Error::InvalidRequest(
"TLS trust anchor PEM contained no certificates".to_owned(),
));
}
for cert in certs {
store
.add(cert)
.map_err(|e| Error::InvalidRequest(format!("invalid CA certificate: {e}")))?;
}
}
if store.is_empty() {
return Err(Error::InvalidRequest(
"no usable TLS trust anchors were configured".to_owned(),
));
}
Ok(store)
}
}
fn parse_certs(pem: &[u8]) -> Result<Vec<CertificateDer<'static>>> {
let mut reader = io::BufReader::new(pem);
rustls_pemfile::certs(&mut reader)
.collect::<std::result::Result<Vec<_>, _>>()
.map_err(|e| Error::InvalidRequest(format!("could not parse certificate PEM: {e}")))
}
fn parse_key(pem: &[u8]) -> Result<PrivateKeyDer<'static>> {
let mut reader = io::BufReader::new(pem);
rustls_pemfile::private_key(&mut reader)
.map_err(|e| Error::InvalidRequest(format!("could not parse private key PEM: {e}")))?
.ok_or_else(|| Error::InvalidRequest("private key PEM contained no key".to_owned()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn an_empty_pem_bundle_is_rejected_rather_than_trusting_nothing() {
let cfg = TlsConfig::with_ca_pem(b"not a certificate".to_vec());
assert!(cfg.connector().is_err());
}
#[test]
fn the_server_name_override_wins() {
let cfg = TlsConfig::system().with_server_name("broker.internal");
let name = cfg.server_name("10.0.0.1").expect("valid name");
assert!(format!("{name:?}").contains("broker.internal"));
}
#[test]
fn without_an_override_the_host_is_used() {
let cfg = TlsConfig::system();
assert!(cfg.server_name("broker.example.com").is_ok());
assert!(cfg.server_name("not a host name").is_err());
}
#[test]
fn client_certificate_debug_never_prints_the_key() {
let cert = ClientCertificate {
chain_pem: b"chain".to_vec(),
key_pem: b"SUPER SECRET".to_vec(),
};
let rendered = format!("{cert:?}");
assert!(!rendered.contains("SUPER SECRET"), "{rendered}");
}
}