use std::fmt;
use std::sync::{Arc, PoisonError, RwLock};
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
use rustls::crypto::{
CryptoProvider, WebPkiSupportedAlgorithms, verify_tls12_signature, verify_tls13_signature,
};
use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, ServerName, UnixTime};
use rustls::server::danger::{ClientCertVerified, ClientCertVerifier};
use rustls::server::{ClientHello, ResolvesServerCert};
use rustls::sign::CertifiedKey;
use rustls::{CertificateError, DigitallySignedStruct, DistinguishedName, Error, SignatureScheme};
use sha2::{Digest, Sha256};
#[derive(Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct Spki([u8; 32]);
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SpkiError {
NotAPin(String),
NotACertificate(String),
}
impl fmt::Display for SpkiError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::NotAPin(s) => write!(f, "{s:?} is not a pin (sha256:<64 lowercase hex>)"),
Self::NotACertificate(why) => write!(f, "not a DER certificate: {why}"),
}
}
}
impl std::error::Error for SpkiError {}
impl Spki {
#[must_use]
pub fn of_spki_der(der: &[u8]) -> Self {
Self(Sha256::digest(der).into())
}
pub fn of_certificate(der: &[u8]) -> Result<Self, SpkiError> {
let (_, cert) = x509_parser::parse_x509_certificate(der)
.map_err(|e| SpkiError::NotACertificate(e.to_string()))?;
Ok(Self::of_spki_der(cert.tbs_certificate.subject_pki.raw))
}
pub fn parse(s: &str) -> Result<Self, SpkiError> {
let bad = || SpkiError::NotAPin(s.to_owned());
let hex = s.strip_prefix("sha256:").ok_or_else(bad)?;
if hex.len() != 64 {
return Err(bad());
}
let nibble = |c: u8| match c {
b'0'..=b'9' => Some(c - b'0'),
b'a'..=b'f' => Some(c - b'a' + 10),
_ => None,
};
let (pairs, _) = hex.as_bytes().as_chunks::<2>();
let mut out = [0u8; 32];
for (byte, [hi, lo]) in out.iter_mut().zip(pairs) {
*byte = (nibble(*hi).ok_or_else(bad)? << 4) | nibble(*lo).ok_or_else(bad)?;
}
Ok(Self(out))
}
}
impl fmt::Display for Spki {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("sha256:")?;
self.0.iter().try_for_each(|b| write!(f, "{b:02x}"))
}
}
impl fmt::Debug for Spki {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(self, f)
}
}
impl std::str::FromStr for Spki {
type Err = SpkiError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Self::parse(s)
}
}
impl serde::Serialize for Spki {
fn serialize<S: serde::Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
s.collect_str(self)
}
}
impl<'de> serde::Deserialize<'de> for Spki {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
let s = String::deserialize(d)?;
Self::parse(&s).map_err(serde::de::Error::custom)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct KeyError(pub String);
impl fmt::Display for KeyError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl std::error::Error for KeyError {}
pub struct KeyMaterial {
key: rcgen::KeyPair,
}
impl fmt::Debug for KeyMaterial {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("KeyMaterial")
.field("spki", &self.spki())
.finish_non_exhaustive()
}
}
impl KeyMaterial {
pub fn generate() -> Result<Self, KeyError> {
rcgen::KeyPair::generate_for(&rcgen::PKCS_ED25519)
.map(|key| Self { key })
.map_err(|e| KeyError(e.to_string()))
}
pub fn from_pem(pem: &str) -> Result<Self, KeyError> {
let key = rcgen::KeyPair::from_pem(pem).map_err(|e| KeyError(e.to_string()))?;
if !key.is_compatible(&rcgen::PKCS_ED25519) {
return Err(KeyError("the key is not an ed25519 key".into()));
}
Ok(Self { key })
}
#[must_use]
pub fn to_pem(&self) -> String {
self.key.serialize_pem()
}
#[must_use]
pub fn spki(&self) -> Spki {
Spki::of_spki_der(&self.key.public_key_der())
}
pub fn certificate(
&self,
name: &str,
) -> Result<(CertificateDer<'static>, PrivateKeyDer<'static>), KeyError> {
let params = rcgen::CertificateParams::new(vec![name.to_owned()])
.map_err(|e| KeyError(e.to_string()))?;
let cert = params
.self_signed(&self.key)
.map_err(|e| KeyError(e.to_string()))?;
Ok((
cert.der().clone(),
PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(self.key.serialize_der())),
))
}
}
fn provider() -> Arc<CryptoProvider> {
Arc::new(rustls::crypto::ring::default_provider())
}
const VERSIONS: &[&rustls::SupportedProtocolVersion] = &[&rustls::version::TLS13];
pub trait PinSet: Send + Sync + fmt::Debug {
fn admits(&self, spki: &Spki) -> bool;
}
fn pinned(cert: &CertificateDer<'_>, admits: impl Fn(&Spki) -> bool) -> Result<Spki, Error> {
let spki = Spki::of_certificate(cert)
.map_err(|_| Error::InvalidCertificate(CertificateError::BadEncoding))?;
if admits(&spki) {
Ok(spki)
} else {
Err(Error::InvalidCertificate(
CertificateError::ApplicationVerificationFailure,
))
}
}
#[derive(Debug, Clone, Copy)]
struct Signatures(WebPkiSupportedAlgorithms);
impl Signatures {
fn tls12(
self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
verify_tls12_signature(message, cert, dss, &self.0)
}
fn tls13(
self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
verify_tls13_signature(message, cert, dss, &self.0)
}
fn schemes(self) -> Vec<SignatureScheme> {
self.0.supported_schemes()
}
}
#[derive(Debug)]
struct ClientPins {
pins: Arc<dyn PinSet>,
signatures: Signatures,
}
impl ClientCertVerifier for ClientPins {
fn offer_client_auth(&self) -> bool {
true
}
fn client_auth_mandatory(&self) -> bool {
true
}
fn root_hint_subjects(&self) -> &[DistinguishedName] {
&[]
}
fn verify_client_cert(
&self,
end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_now: UnixTime,
) -> Result<ClientCertVerified, Error> {
pinned(end_entity, |spki| self.pins.admits(spki)).map(|_| ClientCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
self.signatures.tls12(message, cert, dss)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
self.signatures.tls13(message, cert, dss)
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
self.signatures.schemes()
}
}
#[derive(Debug)]
struct ServerPins {
pins: Vec<Spki>,
signatures: Signatures,
}
impl ServerCertVerifier for ServerPins {
fn verify_server_cert(
&self,
end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp: &[u8],
_now: UnixTime,
) -> Result<ServerCertVerified, Error> {
pinned(end_entity, |spki| self.pins.contains(spki)).map(|_| ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
self.signatures.tls12(message, cert, dss)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
self.signatures.tls13(message, cert, dss)
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
self.signatures.schemes()
}
}
pub const CERT_NAME: &str = "engenho-control";
#[derive(Debug)]
pub struct Presented {
current: RwLock<Arc<CertifiedKey>>,
}
impl Presented {
pub fn new(identity: &KeyMaterial) -> Result<Self, KeyError> {
Ok(Self {
current: RwLock::new(certified(identity)?),
})
}
pub fn present(&self, identity: &KeyMaterial) -> Result<(), KeyError> {
let next = certified(identity)?;
*self.current.write().unwrap_or_else(PoisonError::into_inner) = next;
Ok(())
}
}
impl ResolvesServerCert for Presented {
fn resolve(&self, _: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
Some(Arc::clone(
&self.current.read().unwrap_or_else(PoisonError::into_inner),
))
}
}
fn certified(identity: &KeyMaterial) -> Result<Arc<CertifiedKey>, KeyError> {
let (cert, key) = identity.certificate(CERT_NAME)?;
CertifiedKey::from_der(vec![cert], key, &provider())
.map(Arc::new)
.map_err(|e| KeyError(e.to_string()))
}
pub fn server_config(
identity: Arc<Presented>,
clients: Arc<dyn PinSet>,
) -> Result<rustls::ServerConfig, KeyError> {
let provider = provider();
let signatures = Signatures(provider.signature_verification_algorithms);
Ok(rustls::ServerConfig::builder_with_provider(provider)
.with_protocol_versions(VERSIONS)
.map_err(|e| KeyError(e.to_string()))?
.with_client_cert_verifier(Arc::new(ClientPins {
pins: clients,
signatures,
}))
.with_cert_resolver(identity))
}
pub fn client_config(
key: &KeyMaterial,
server: Vec<Spki>,
) -> Result<rustls::ClientConfig, KeyError> {
let provider = provider();
let signatures = Signatures(provider.signature_verification_algorithms);
let (cert, private) = key.certificate(CERT_NAME)?;
rustls::ClientConfig::builder_with_provider(provider)
.with_protocol_versions(VERSIONS)
.map_err(|e| KeyError(e.to_string()))?
.dangerous()
.with_custom_certificate_verifier(Arc::new(ServerPins {
pins: server,
signatures,
}))
.with_client_auth_cert(vec![cert], private)
.map_err(|e| KeyError(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_pin_is_spelled_sha256_hex_and_round_trips() {
let key = KeyMaterial::generate().unwrap();
let spki = key.spki();
let text = spki.to_string();
assert!(text.starts_with("sha256:") && text.len() == 71, "{text}");
assert_eq!(Spki::parse(&text), Ok(spki));
assert_eq!(
serde_json::from_value::<Spki>(serde_json::to_value(spki).unwrap()).unwrap(),
spki
);
for bad in ["", "sha256:", "sha1:00", &text.to_uppercase(), &text[..70]] {
assert!(Spki::parse(bad).is_err(), "{bad:?}");
}
}
#[test]
fn the_pin_of_a_key_is_the_pin_of_its_certificate() {
let key = KeyMaterial::generate().unwrap();
let (cert, _) = key.certificate(CERT_NAME).unwrap();
assert_eq!(Spki::of_certificate(&cert), Ok(key.spki()));
let reread = KeyMaterial::from_pem(&key.to_pem()).unwrap();
assert_eq!(reread.spki(), key.spki());
assert!(
!format!("{key:?}").contains("PRIVATE"),
"Debug never shows the key"
);
}
#[test]
fn a_non_ed25519_key_is_refused() {
let p256 = rcgen::KeyPair::generate_for(&rcgen::PKCS_ECDSA_P256_SHA256).unwrap();
assert!(KeyMaterial::from_pem(&p256.serialize_pem()).is_err());
}
#[derive(Debug)]
struct Admits(Vec<Spki>);
impl PinSet for Admits {
fn admits(&self, spki: &Spki) -> bool {
self.0.contains(spki)
}
}
fn handshake(
server: &rustls::ServerConfig,
client: &rustls::ClientConfig,
) -> Result<(), String> {
let say = |e: &dyn std::error::Error| e.to_string();
let name = rustls::pki_types::ServerName::try_from(CERT_NAME)
.expect("a name")
.to_owned();
let mut c =
rustls::ClientConnection::new(Arc::new(client.clone()), name).map_err(|e| say(&e))?;
let mut s = rustls::ServerConnection::new(Arc::new(server.clone())).map_err(|e| say(&e))?;
for _ in 0..16 {
if !c.is_handshaking() && !s.is_handshaking() {
return Ok(());
}
let mut to_server = Vec::new();
while c.wants_write() {
c.write_tls(&mut to_server).map_err(|e| say(&e))?;
}
if !to_server.is_empty() {
s.read_tls(&mut to_server.as_slice()).map_err(|e| say(&e))?;
s.process_new_packets().map_err(|e| say(&e))?;
}
let mut to_client = Vec::new();
while s.wants_write() {
s.write_tls(&mut to_client).map_err(|e| say(&e))?;
}
if !to_client.is_empty() {
c.read_tls(&mut to_client.as_slice()).map_err(|e| say(&e))?;
c.process_new_packets().map_err(|e| say(&e))?;
}
}
Err("the handshake never finished".to_owned())
}
#[test]
fn what_the_server_presents_is_what_it_was_last_given() {
let client_key = KeyMaterial::generate().unwrap();
let (first, second) = (
KeyMaterial::generate().unwrap(),
KeyMaterial::generate().unwrap(),
);
let presented = Arc::new(Presented::new(&first).unwrap());
let admits: Arc<dyn PinSet> = Arc::new(Admits(vec![client_key.spki()]));
let server = server_config(Arc::clone(&presented), admits).unwrap();
let pinned_to = |key: &KeyMaterial| client_config(&client_key, vec![key.spki()]).unwrap();
assert!(handshake(&server, &pinned_to(&first)).is_ok());
assert!(
handshake(&server, &pinned_to(&second)).is_err(),
"a client pinned to a key the server never held got through"
);
presented.present(&second).unwrap();
assert!(
handshake(&server, &pinned_to(&second)).is_ok(),
"the replaced key is not presented"
);
assert!(
handshake(&server, &pinned_to(&first)).is_err(),
"the old key is still presented"
);
}
#[test]
fn a_client_the_server_does_not_pin_is_refused() {
let identity = KeyMaterial::generate().unwrap();
let (known, stranger) = (
KeyMaterial::generate().unwrap(),
KeyMaterial::generate().unwrap(),
);
let admits: Arc<dyn PinSet> = Arc::new(Admits(vec![known.spki()]));
let server = server_config(
Arc::new(Presented::new(&identity).unwrap()),
Arc::clone(&admits),
)
.unwrap();
let pin = vec![identity.spki()];
assert!(handshake(&server, &client_config(&known, pin.clone()).unwrap()).is_ok());
assert!(handshake(&server, &client_config(&stranger, pin).unwrap()).is_err());
}
}