use std::fs::File;
use std::io::BufReader;
use std::path::Path;
use std::sync::Arc;
use rustls::crypto::ring::default_provider;
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use rustls::server::WebPkiClientVerifier;
use rustls::{RootCertStore, ServerConfig};
pub struct TlsConfig {
pub(crate) cert_chain: Vec<CertificateDer<'static>>,
pub(crate) private_key: PrivateKeyDer<'static>,
pub(crate) client_auth: ClientAuth,
pub(crate) alpn_protocols: Vec<Vec<u8>>,
}
#[derive(Clone)]
pub enum ClientAuth {
None,
Optional {
root_certs: Arc<RootCertStore>,
},
Required {
root_certs: Arc<RootCertStore>,
},
}
impl Clone for TlsConfig {
fn clone(&self) -> Self {
Self {
cert_chain: self.cert_chain.clone(),
private_key: self.private_key.clone_key(),
client_auth: self.client_auth.clone(),
alpn_protocols: self.alpn_protocols.clone(),
}
}
}
impl TlsConfig {
pub fn new(
cert_chain: Vec<CertificateDer<'static>>,
private_key: PrivateKeyDer<'static>,
) -> Self {
Self {
cert_chain,
private_key,
client_auth: ClientAuth::None,
alpn_protocols: Vec::new(),
}
}
pub fn builder() -> TlsConfigBuilder {
TlsConfigBuilder::default()
}
pub(crate) fn build_server_config(&self) -> Result<ServerConfig, TlsConfigError> {
let provider = Arc::new(default_provider());
let builder = ServerConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.map_err(|e| TlsConfigError::InvalidConfig(e.to_string()))?;
let server_config = match &self.client_auth {
ClientAuth::None => builder
.with_no_client_auth()
.with_single_cert(self.cert_chain.clone(), self.private_key.clone_key())
.map_err(|e| TlsConfigError::InvalidCertificate(e.to_string()))?,
ClientAuth::Optional { root_certs } => {
let verifier = WebPkiClientVerifier::builder(root_certs.clone())
.allow_unauthenticated()
.build()
.map_err(|e| TlsConfigError::InvalidConfig(e.to_string()))?;
builder
.with_client_cert_verifier(verifier)
.with_single_cert(self.cert_chain.clone(), self.private_key.clone_key())
.map_err(|e| TlsConfigError::InvalidCertificate(e.to_string()))?
}
ClientAuth::Required { root_certs } => {
let verifier = WebPkiClientVerifier::builder(root_certs.clone())
.build()
.map_err(|e| TlsConfigError::InvalidConfig(e.to_string()))?;
builder
.with_client_cert_verifier(verifier)
.with_single_cert(self.cert_chain.clone(), self.private_key.clone_key())
.map_err(|e| TlsConfigError::InvalidCertificate(e.to_string()))?
}
};
let mut config = server_config;
if !self.alpn_protocols.is_empty() {
config.alpn_protocols = self.alpn_protocols.clone();
}
Ok(config)
}
}
#[derive(Default)]
pub struct TlsConfigBuilder {
cert_chain: Option<Vec<CertificateDer<'static>>>,
private_key: Option<PrivateKeyDer<'static>>,
client_auth: Option<ClientAuth>,
alpn_protocols: Vec<Vec<u8>>,
}
impl TlsConfigBuilder {
pub fn cert_chain_file(mut self, path: impl AsRef<Path>) -> Result<Self, TlsConfigError> {
let file = File::open(path.as_ref()).map_err(|e| TlsConfigError::IoError(e))?;
let mut reader = BufReader::new(file);
let certs = rustls_pemfile::certs(&mut reader)
.collect::<Result<Vec<_>, _>>()
.map_err(|e| TlsConfigError::InvalidCertificate(e.to_string()))?;
if certs.is_empty() {
return Err(TlsConfigError::InvalidCertificate(
"No certificates found in file".into(),
));
}
self.cert_chain = Some(certs);
Ok(self)
}
pub fn cert_chain_pem(mut self, pem: &[u8]) -> Result<Self, TlsConfigError> {
let certs = rustls_pemfile::certs(&mut &pem[..])
.collect::<Result<Vec<_>, _>>()
.map_err(|e| TlsConfigError::InvalidCertificate(e.to_string()))?;
if certs.is_empty() {
return Err(TlsConfigError::InvalidCertificate(
"No certificates found".into(),
));
}
self.cert_chain = Some(certs);
Ok(self)
}
pub fn private_key_file(mut self, path: impl AsRef<Path>) -> Result<Self, TlsConfigError> {
let file = File::open(path.as_ref()).map_err(|e| TlsConfigError::IoError(e))?;
let mut reader = BufReader::new(file);
let key = rustls_pemfile::private_key(&mut reader)
.map_err(|e| TlsConfigError::InvalidKey(e.to_string()))?
.ok_or_else(|| TlsConfigError::InvalidKey("No private key found in file".into()))?;
self.private_key = Some(key);
Ok(self)
}
pub fn private_key_pem(mut self, pem: &[u8]) -> Result<Self, TlsConfigError> {
let key = rustls_pemfile::private_key(&mut &pem[..])
.map_err(|e| TlsConfigError::InvalidKey(e.to_string()))?
.ok_or_else(|| TlsConfigError::InvalidKey("No private key found".into()))?;
self.private_key = Some(key);
Ok(self)
}
pub fn require_client_auth(
mut self,
ca_certs_path: impl AsRef<Path>,
) -> Result<Self, TlsConfigError> {
let root_certs = Self::load_root_certs(ca_certs_path)?;
self.client_auth = Some(ClientAuth::Required { root_certs });
Ok(self)
}
pub fn optional_client_auth(
mut self,
ca_certs_path: impl AsRef<Path>,
) -> Result<Self, TlsConfigError> {
let root_certs = Self::load_root_certs(ca_certs_path)?;
self.client_auth = Some(ClientAuth::Optional { root_certs });
Ok(self)
}
pub fn alpn_protocol(mut self, protocol: impl Into<Vec<u8>>) -> Self {
self.alpn_protocols.push(protocol.into());
self
}
pub fn alpn_protocols(mut self, protocols: &[&str]) -> Self {
for proto in protocols {
self.alpn_protocols.push(proto.as_bytes().to_vec());
}
self
}
pub fn build(self) -> Result<TlsConfig, TlsConfigError> {
let cert_chain = self
.cert_chain
.ok_or_else(|| TlsConfigError::MissingField("certificate chain"))?;
let private_key = self
.private_key
.ok_or_else(|| TlsConfigError::MissingField("private key"))?;
Ok(TlsConfig {
cert_chain,
private_key,
client_auth: self.client_auth.unwrap_or(ClientAuth::None),
alpn_protocols: self.alpn_protocols,
})
}
fn load_root_certs(path: impl AsRef<Path>) -> Result<Arc<RootCertStore>, TlsConfigError> {
let file = File::open(path.as_ref()).map_err(|e| TlsConfigError::IoError(e))?;
let mut reader = BufReader::new(file);
let certs = rustls_pemfile::certs(&mut reader)
.collect::<Result<Vec<_>, _>>()
.map_err(|e| TlsConfigError::InvalidCertificate(e.to_string()))?;
let mut root_store = RootCertStore::empty();
let (added, _ignored) = root_store.add_parsable_certificates(certs);
if added == 0 {
return Err(TlsConfigError::InvalidCertificate(
"No valid CA certificates found".into(),
));
}
Ok(Arc::new(root_store))
}
}
#[derive(Debug)]
pub enum TlsConfigError {
IoError(std::io::Error),
InvalidCertificate(String),
InvalidKey(String),
InvalidConfig(String),
MissingField(&'static str),
}
impl std::fmt::Display for TlsConfigError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::IoError(e) => write!(f, "I/O error: {}", e),
Self::InvalidCertificate(e) => write!(f, "Invalid certificate: {}", e),
Self::InvalidKey(e) => write!(f, "Invalid private key: {}", e),
Self::InvalidConfig(e) => write!(f, "Invalid TLS configuration: {}", e),
Self::MissingField(field) => write!(f, "Missing required field: {}", field),
}
}
}
impl std::error::Error for TlsConfigError {}