use crate::RuntimeError;
use rustls::{
ClientConfig, RootCertStore, ServerConfig,
crypto::CryptoProvider,
server::WebPkiClientVerifier,
pki_types::{
CertificateDer, PrivateKeyDer,
pem::{self, PemObject},
},
};
use rustls_platform_verifier::Verifier;
use std::{
path::PathBuf,
sync::{Arc, OnceLock},
};
fn provider() -> Arc<CryptoProvider> {
static PROVIDER: OnceLock<Arc<CryptoProvider>> = OnceLock::new();
PROVIDER
.get_or_init(|| Arc::new(rustls::crypto::aws_lc_rs::default_provider()))
.clone()
}
pub(crate) fn client() -> Result<Arc<ClientConfig>, RuntimeError> {
static CLIENT: OnceLock<Result<Arc<ClientConfig>, RuntimeError>> = OnceLock::new();
CLIENT
.get_or_init(|| build_client(Verifier::new(provider())))
.clone()
}
#[derive(Debug, Clone)]
pub(crate) enum Keys {
Files(PathBuf, PathBuf),
Pem(Arc<[u8]>, Arc<[u8]>),
}
#[derive(Debug, Clone, Default)]
pub(crate) struct ClientSettings {
pub(crate) roots: Option<Arc<[u8]>>,
pub(crate) alpn: Arc<[Vec<u8>]>,
pub(crate) identity: Option<(Arc<[u8]>, Arc<[u8]>)>,
}
pub(crate) fn client_with(settings: &ClientSettings) -> Result<Arc<ClientConfig>, RuntimeError> {
if settings.roots.is_none() && settings.alpn.is_empty() && settings.identity.is_none() {
return client();
}
let verifier = match &settings.roots {
Some(pem) => Verifier::new_with_extra_roots(certificates(pem)?, provider()),
None => Verifier::new(provider()),
};
let builder = ClientConfig::builder_with_provider(provider())
.with_safe_default_protocol_versions()
.map_err(tls_error)?
.dangerous()
.with_custom_certificate_verifier(Arc::new(verifier.map_err(tls_error)?));
let mut config = match &settings.identity {
Some((cert, key)) => builder
.with_client_auth_cert(certificates(cert)?, private_key(key)?)
.map_err(|_| RuntimeError::BadCertificate)?,
None => builder.with_no_client_auth(),
};
config.alpn_protocols = settings.alpn.to_vec();
Ok(Arc::new(config))
}
fn certificates(pem: &[u8]) -> Result<Vec<CertificateDer<'static>>, RuntimeError> {
let found: Vec<CertificateDer<'static>> = CertificateDer::pem_slice_iter(pem)
.collect::<Result<_, _>>()
.map_err(|_| RuntimeError::BadCertificate)?;
match found.is_empty() {
true => Err(RuntimeError::BadCertificate),
false => Ok(found),
}
}
fn private_key(pem: &[u8]) -> Result<PrivateKeyDer<'static>, RuntimeError> {
PrivateKeyDer::from_pem_slice(pem).map_err(|_| RuntimeError::BadCertificate)
}
fn build_client(
verifier: Result<Verifier, rustls::Error>,
) -> Result<Arc<ClientConfig>, RuntimeError> {
let verifier = verifier.map_err(tls_error)?;
let config = ClientConfig::builder_with_provider(provider())
.with_safe_default_protocol_versions()
.map_err(tls_error)?
.dangerous()
.with_custom_certificate_verifier(Arc::new(verifier))
.with_no_client_auth();
Ok(Arc::new(config))
}
#[derive(Debug, Clone)]
pub(crate) struct ServerSettings {
pub(crate) keys: Keys,
pub(crate) alpn: Arc<[Vec<u8>]>,
pub(crate) client_roots: Option<Arc<[u8]>>,
}
pub(crate) fn server_with(settings: &ServerSettings) -> Result<Arc<ServerConfig>, RuntimeError> {
let (chain, key) = match &settings.keys {
Keys::Files(cert, key) => {
let chain: Vec<CertificateDer<'static>> = CertificateDer::pem_file_iter(cert)
.map_err(pem_error)?
.collect::<Result<_, _>>()
.map_err(pem_error)?;
(chain, PrivateKeyDer::from_pem_file(key).map_err(pem_error)?)
}
Keys::Pem(cert, key) => (certificates(cert)?, private_key(key)?),
};
if chain.is_empty() {
return Err(RuntimeError::BadCertificate);
}
let builder = ServerConfig::builder_with_provider(provider())
.with_safe_default_protocol_versions()
.map_err(tls_error)?;
let builder = match &settings.client_roots {
Some(pem) => {
let mut roots = RootCertStore::empty();
for root in certificates(pem)? {
roots.add(root).map_err(|_| RuntimeError::BadCertificate)?;
}
let verifier = WebPkiClientVerifier::builder_with_provider(Arc::new(roots), provider())
.build()
.map_err(|_| RuntimeError::BadCertificate)?;
builder.with_client_cert_verifier(verifier)
}
None => builder.with_no_client_auth(),
};
let mut config = builder
.with_single_cert(chain, key)
.map_err(|_| RuntimeError::BadCertificate)?;
config.alpn_protocols = settings.alpn.to_vec();
Ok(Arc::new(config))
}
pub(crate) fn tls_error(error: rustls::Error) -> RuntimeError {
match error {
rustls::Error::InvalidCertificate(_) => RuntimeError::BadCertificate,
_ => RuntimeError::TlsFailed,
}
}
fn pem_error(error: pem::Error) -> RuntimeError {
match error {
pem::Error::Io(error) => RuntimeError::CheckError(error.raw_os_error()),
_ => RuntimeError::BadCertificate,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_refused_certificate_is_told_apart() {
assert_eq!(
tls_error(rustls::Error::InvalidCertificate(
rustls::CertificateError::UnknownIssuer
)),
RuntimeError::BadCertificate,
);
assert_eq!(
tls_error(rustls::Error::HandshakeNotComplete),
RuntimeError::TlsFailed,
);
}
#[test]
fn roots_that_are_not_certificates_are_refused() {
for pem in [&b""[..], b"not a certificate"] {
let settings = ClientSettings {
roots: Some(Arc::from(pem)),
..ClientSettings::default()
};
assert_eq!(
client_with(&settings).map(|_| ()),
Err(RuntimeError::BadCertificate)
);
}
}
#[test]
fn a_missing_certificate_file_is_not_found() {
let missing = PathBuf::from("/nonexistent/atap/cert.pem");
let settings = ServerSettings {
keys: Keys::Files(missing.clone(), missing),
alpn: Arc::from([]),
client_roots: None,
};
assert_eq!(
server_with(&settings).map(|_| ()),
Err(RuntimeError::CheckError(Some(libc::ENOENT))),
);
}
}