use std::sync::Arc;
use rustls_pki_types::pem::PemObject as _;
use rustls_pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime};
use tokio_rustls::rustls::{ClientConfig, RootCertStore, ServerConfig};
use tokio_rustls::{TlsAcceptor, TlsConnector};
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum TlsError {
#[error("{0} is not a name a certificate can be checked against")]
UnusableName(String),
#[error("reading {what}: {detail}")]
Material {
what: String,
detail: String,
},
#[error("tls configuration: {0}")]
Config(String),
#[error("tls handshake with {peer}: {detail}")]
Handshake {
peer: String,
detail: String,
},
}
#[derive(Clone)]
pub struct ClientTls {
config: Arc<ClientConfig>,
}
impl std::fmt::Debug for ClientTls {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("ClientTls { .. }")
}
}
#[derive(Debug, Clone, Default)]
pub struct TrustAnchors {
extra: Vec<CertificateDer<'static>>,
system: bool,
}
impl TrustAnchors {
#[must_use]
pub fn system() -> Self {
Self {
extra: Vec::new(),
system: true,
}
}
#[must_use]
pub fn only() -> Self {
Self {
extra: Vec::new(),
system: false,
}
}
pub fn add_pem(&mut self, pem: &[u8]) -> Result<(), TlsError> {
let certs: Vec<CertificateDer<'static>> = CertificateDer::pem_slice_iter(pem)
.collect::<Result<_, _>>()
.map_err(|error| TlsError::Material {
what: "trust anchor".to_owned(),
detail: error.to_string(),
})?;
if certs.is_empty() {
return Err(TlsError::Material {
what: "trust anchor".to_owned(),
detail: "no certificate found in the PEM data".to_owned(),
});
}
self.extra.extend(certs);
Ok(())
}
fn store(&self) -> Result<RootCertStore, TlsError> {
let mut store = RootCertStore::empty();
if self.system {
let loaded = rustls_native_certs::load_native_certs();
for error in &loaded.errors {
tracing::warn!(%error, "could not read part of the platform trust store");
}
for cert in loaded.certs {
if let Err(error) = store.add(cert) {
tracing::debug!(%error, "skipping an unusable root from the platform store");
}
}
}
for cert in &self.extra {
store
.add(cert.clone())
.map_err(|error| TlsError::Config(error.to_string()))?;
}
if store.is_empty() {
return Err(TlsError::Config(
"no trust anchors: every certificate would be refused".to_owned(),
));
}
Ok(store)
}
}
impl ClientTls {
pub fn new(anchors: &TrustAnchors) -> Result<Self, TlsError> {
Self::with_identity(anchors, None)
}
pub fn with_identity(
anchors: &TrustAnchors,
identity: Option<Identity>,
) -> Result<Self, TlsError> {
let roots = anchors.store()?;
let builder = ClientConfig::builder().with_root_certificates(roots);
let config = match identity {
Some(identity) => builder
.with_client_auth_cert(identity.chain, identity.key)
.map_err(|error| TlsError::Config(error.to_string()))?,
None => builder.with_no_client_auth(),
};
Ok(Self {
config: Arc::new(config),
})
}
#[must_use]
pub fn connector(&self) -> TlsConnector {
TlsConnector::from(Arc::clone(&self.config))
}
#[cfg(feature = "quic")]
pub(crate) fn quic_config(&self) -> Result<quinn::ClientConfig, TlsError> {
let config = self.quic_rustls_config();
let crypto = quinn::crypto::rustls::QuicClientConfig::try_from(config)
.map_err(|error| TlsError::Config(error.to_string()))?;
Ok(quinn::ClientConfig::new(Arc::new(crypto)))
}
#[cfg(feature = "quic")]
fn quic_rustls_config(&self) -> ClientConfig {
let mut config = (*self.config).clone();
config.alpn_protocols = vec![b"sip/2".to_vec()];
config.enable_early_data = false;
config
}
}
pub struct Identity {
chain: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
}
impl std::fmt::Debug for Identity {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("Identity { .. }")
}
}
impl Identity {
pub fn from_pem(cert_pem: &[u8], key_pem: &[u8]) -> Result<Self, TlsError> {
let chain: Vec<CertificateDer<'static>> = CertificateDer::pem_slice_iter(cert_pem)
.collect::<Result<_, _>>()
.map_err(|error| TlsError::Material {
what: "certificate".to_owned(),
detail: error.to_string(),
})?;
if chain.is_empty() {
return Err(TlsError::Material {
what: "certificate".to_owned(),
detail: "no certificate found in the PEM data".to_owned(),
});
}
let key = PrivateKeyDer::from_pem_slice(key_pem).map_err(|error| TlsError::Material {
what: "private key".to_owned(),
detail: error.to_string(),
})?;
Ok(Self { chain, key })
}
fn validate_server_chain(&self) -> Result<(), TlsError> {
let Some((leaf, issuers)) = self.chain.split_first() else {
return Err(TlsError::Config(
"server certificate chain has no leaf".to_owned(),
));
};
let end_entity = webpki::EndEntityCert::try_from(leaf).map_err(|error| {
TlsError::Config(format!("invalid server certificate leaf: {error}"))
})?;
let Some((anchor_certificate, intermediates)) = issuers.split_last() else {
return Ok(());
};
let anchor = webpki::anchor_from_trusted_cert(anchor_certificate).map_err(|error| {
TlsError::Config(format!("invalid server certificate chain anchor: {error}"))
})?;
let provider = tokio_rustls::rustls::crypto::ring::default_provider();
let anchors = [anchor];
let verified = end_entity
.verify_for_usage(
provider.signature_verification_algorithms.all,
&anchors,
intermediates,
UnixTime::now(),
webpki::KeyUsage::server_auth(),
None,
None,
)
.map_err(|error| {
TlsError::Config(format!("invalid server certificate chain: {error}"))
})?;
let supplied_in_order = verified
.intermediate_certificates()
.map(webpki::Cert::der)
.eq(intermediates.iter().cloned());
if !supplied_in_order {
return Err(TlsError::Config(
"server certificate chain contains an unrelated or out-of-order certificate"
.to_owned(),
));
}
Ok(())
}
}
#[derive(Clone)]
pub struct ServerTls {
config: Arc<ServerConfig>,
}
impl std::fmt::Debug for ServerTls {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("ServerTls { .. }")
}
}
impl ServerTls {
pub fn new(identity: Identity) -> Result<Self, TlsError> {
identity.validate_server_chain()?;
let config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(identity.chain, identity.key)
.map_err(|error| TlsError::Config(error.to_string()))?;
Ok(Self {
config: Arc::new(config),
})
}
#[must_use]
pub fn acceptor(&self) -> TlsAcceptor {
TlsAcceptor::from(Arc::clone(&self.config))
}
#[cfg(feature = "quic")]
pub(crate) fn quic_config(&self) -> Result<quinn::ServerConfig, TlsError> {
let config = self.quic_rustls_config();
let crypto = quinn::crypto::rustls::QuicServerConfig::try_from(config)
.map_err(|error| TlsError::Config(error.to_string()))?;
Ok(quinn::ServerConfig::with_crypto(Arc::new(crypto)))
}
#[cfg(feature = "quic")]
fn quic_rustls_config(&self) -> ServerConfig {
let mut config = (*self.config).clone();
config.alpn_protocols = vec![b"sip/2".to_vec()];
config.max_early_data_size = 0;
config
}
}
pub fn verification_name(uri_host: &str) -> Result<ServerName<'static>, TlsError> {
ServerName::try_from(uri_host.to_owned())
.map_err(|_| TlsError::UnusableName(uri_host.to_owned()))
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::indexing_slicing
)]
mod tests {
use super::*;
#[test]
fn a_hostname_is_a_usable_verification_name() {
assert!(verification_name("sip.example.com").is_ok());
assert!(verification_name("example.com").is_ok());
}
#[test]
fn an_address_is_accepted_as_a_name() {
assert!(verification_name("192.0.2.1").is_ok());
}
#[test]
fn something_that_is_not_a_name_is_refused_by_name() {
let error = verification_name("not a hostname!").expect_err("refused");
assert!(error.to_string().contains("not a hostname!"), "{error}");
}
#[test]
fn a_client_with_no_anchors_is_refused_at_construction() {
let error = ClientTls::new(&TrustAnchors::only()).expect_err("refused");
assert!(error.to_string().contains("no trust anchors"), "{error}");
}
#[test]
fn the_system_anchors_are_enough_to_build_a_client() {
assert!(ClientTls::new(&TrustAnchors::system()).is_ok());
}
#[test]
fn pem_that_holds_no_certificate_is_refused_by_name() {
let mut anchors = TrustAnchors::only();
let error = anchors.add_pem(b"not a certificate").expect_err("refused");
assert!(error.to_string().contains("no certificate"), "{error}");
}
#[test]
fn an_identity_needs_both_halves() {
let error = Identity::from_pem(b"", b"").expect_err("refused");
assert!(error.to_string().contains("certificate"), "{error}");
}
#[test]
fn debug_output_does_not_leak_key_material() {
let client = ClientTls::new(&TrustAnchors::system()).expect("builds");
let printed = format!("{client:?}");
assert_eq!(printed, "ClientTls { .. }");
let ca = sipx_testkit::certs::Ca::new();
let (certificate, key) = ca.issue_for("localhost");
let identity =
Identity::from_pem(certificate.as_bytes(), key.as_bytes()).expect("identity");
assert_eq!(format!("{identity:?}"), "Identity { .. }");
}
#[cfg(feature = "quic")]
#[test]
fn quic_requires_sip2_and_refuses_early_data_in_both_directions() {
let client = ClientTls::new(&TrustAnchors::system()).expect("client");
let client = client.quic_rustls_config();
assert_eq!(client.alpn_protocols, [b"sip/2".to_vec()]);
assert!(!client.enable_early_data);
let ca = sipx_testkit::certs::Ca::new();
let (certificate, key) = ca.issue_for("localhost");
let identity =
Identity::from_pem(certificate.as_bytes(), key.as_bytes()).expect("identity");
let server = ServerTls::new(identity)
.expect("server")
.quic_rustls_config();
assert_eq!(server.alpn_protocols, [b"sip/2".to_vec()]);
assert_eq!(server.max_early_data_size, 0);
}
#[cfg(feature = "quic")]
#[tokio::test]
async fn a_resumed_client_cannot_deliver_early_data_to_a_sipx_server() {
let ca = sipx_testkit::certs::Ca::new();
let (certificate, key) = ca.issue_for("localhost");
let identity =
Identity::from_pem(certificate.as_bytes(), key.as_bytes()).expect("identity");
let server_policy = ServerTls::new(identity).expect("server policy");
let mut permissive = server_policy.quic_rustls_config();
permissive.max_early_data_size = u32::MAX;
let permissive = quinn::crypto::rustls::QuicServerConfig::try_from(permissive)
.map(|crypto| quinn::ServerConfig::with_crypto(Arc::new(crypto)))
.expect("permissive ticket server");
let rejecting = server_policy.quic_config().expect("sipx QUIC server");
let server =
quinn::Endpoint::server(permissive, "127.0.0.1:0".parse().expect("server address"))
.expect("server endpoint");
let server_addr = server.local_addr().expect("server address");
let mut anchors = TrustAnchors::only();
anchors
.add_pem(ca.pem().as_bytes())
.expect("test authority");
let mut client_tls = ClientTls::new(&anchors)
.expect("client policy")
.quic_rustls_config();
client_tls.enable_early_data = true;
let client_crypto = quinn::crypto::rustls::QuicClientConfig::try_from(client_tls)
.expect("early-data client");
let mut client_config = quinn::ClientConfig::new(Arc::new(client_crypto));
client_config.transport_config(crate::quic::transport_config());
let mut client = quinn::Endpoint::client("127.0.0.1:0".parse().expect("client address"))
.expect("client endpoint");
client.set_default_client_config(client_config);
let (ready, configured) = tokio::sync::oneshot::channel();
let server_task = tokio::spawn(async move {
let first = server
.accept()
.await
.expect("first connection")
.await
.expect("first handshake");
let (mut marker, _unused) = first.open_bi().await.expect("1-RTT marker stream");
marker.write_all(b"ready").await.expect("1-RTT marker");
marker.finish().expect("1-RTT marker finishes");
server.set_server_config(Some(rejecting));
ready.send(()).expect("client waits for configuration");
let second = server
.accept()
.await
.expect("resumed connection")
.await
.expect("resumption handshake");
let (_reply, mut request) = second.accept_bi().await.expect("request stream");
let was_early = request.is_0rtt();
let bytes = request.read_to_end(1024).await.expect("request bytes");
(was_early, bytes)
});
let first = client
.connect(server_addr, "localhost")
.expect("first connect starts")
.await
.expect("first connect");
let (_unused, mut marker) = first.accept_bi().await.expect("server marker");
assert_eq!(
marker.read_to_end(16).await.expect("marker bytes"),
b"ready"
);
drop(first);
configured.await.expect("rejecting server installed");
let (resumed, accepted) = client
.connect(server_addr, "localhost")
.expect("resumption starts")
.into_0rtt()
.expect("client has early-data keys");
let (mut request, _reply) = resumed.open_bi().await.expect("early request stream");
request
.write_all(b"SIP request attempted as early data")
.await
.expect("early write is queued");
request.finish().expect("request finishes");
assert!(!accepted.await, "sipx accepted replayable early data");
let (mut request, _reply) = resumed.open_bi().await.expect("1-RTT request stream");
request
.write_all(b"SIP request attempted as early data")
.await
.expect("1-RTT retry writes");
request.finish().expect("1-RTT retry finishes");
let (was_early, bytes) = server_task.await.expect("server task");
assert!(!was_early, "the server exposed a 0-RTT request stream");
assert_eq!(bytes, b"SIP request attempted as early data");
}
}