use core::marker::PhantomData;
use core::time::Duration;
#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
use crate::crypto::hash::Digest;
use crate::crypto::policy::VerificationPolicy;
use crate::crypto::x509::error::CertificateValidationError;
use crate::crypto::x509::utils::{ensure_signature_algorithm_consistency, validate_certificate_expiry};
use crate::crypto::x509::Certificate;
use crate::der::asn1::GeneralizedTime;
use crate::der::Encode;
pub trait CertificateValidation: Send + Sync {
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError>;
}
pub trait SignatureVerification: CertificateValidation {
fn verify_with_policy(
&self,
cert: &Certificate,
curr_time: u64,
issuer_pub_key: &[u8],
policy: &dyn VerificationPolicy,
) -> Result<(), CertificateValidationError>;
}
#[derive(Debug, Clone, Copy, Default)]
pub struct ExpiryValidator;
impl CertificateValidation for ExpiryValidator {
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
validate_certificate_expiry(cert)
}
}
#[derive(Debug, Clone, Copy)]
pub struct PublicKeyPinning<const N: usize> {
allowed_keys: [&'static [u8]; N],
}
impl<const N: usize> PublicKeyPinning<N> {
pub const fn new(keys: [&'static [u8]; N]) -> Self {
Self { allowed_keys: keys }
}
}
impl<const N: usize> CertificateValidation for PublicKeyPinning<N> {
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
let pub_key_bytes = cert.tbs_certificate.subject_public_key_info.subject_public_key.raw_bytes();
self.allowed_keys
.contains(&pub_key_bytes)
.then_some(())
.ok_or(CertificateValidationError::PublicKeyNotPinned)
}
}
#[derive(Debug, Clone, Copy)]
pub struct FingerprintPinning<D, const N: usize> {
allowed_fingerprints: [&'static [u8]; N],
_digest: PhantomData<D>,
}
impl<D, const N: usize> FingerprintPinning<D, N>
where
D: Digest,
{
pub const fn new(fingerprints: [&'static [u8]; N]) -> Self {
Self { allowed_fingerprints: fingerprints, _digest: PhantomData }
}
}
impl<D, const N: usize> CertificateValidation for FingerprintPinning<D, N>
where
D: Digest + Send + Sync,
{
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
let cert_der = cert.to_der()?;
let fingerprint = D::digest(&cert_der);
self.allowed_fingerprints
.iter()
.any(|fp| *fp == fingerprint.as_ref())
.then_some(())
.ok_or(CertificateValidationError::CertificateNotPinned)
}
}
#[derive(Debug, Clone, Copy)]
pub struct FingerprintDenylist<D, const N: usize> {
denied_fingerprints: [&'static [u8]; N],
_digest: PhantomData<D>,
}
impl<D, const N: usize> FingerprintDenylist<D, N>
where
D: Digest,
{
pub const fn new(fingerprints: [&'static [u8]; N]) -> Self {
Self { denied_fingerprints: fingerprints, _digest: PhantomData }
}
}
impl<D, const N: usize> CertificateValidation for FingerprintDenylist<D, N>
where
D: Digest + Send + Sync,
{
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
let cert_der = cert.to_der()?;
let fingerprint = D::digest(&cert_der);
let is_denied = self.denied_fingerprints.iter().any(|fp| *fp == fingerprint.as_ref());
if is_denied {
Err(CertificateValidationError::CertificateDenied)
} else {
Ok(())
}
}
}
#[cfg(feature = "std")]
#[derive(Clone)]
pub struct RuntimeCertificatePinning<D> {
fingerprints: Vec<Vec<u8>>,
_digest: PhantomData<D>,
}
#[cfg(feature = "std")]
impl<D> RuntimeCertificatePinning<D>
where
D: Digest,
{
pub fn from_certificates(certs: impl IntoIterator<Item = Certificate>) -> Result<Self, CertificateValidationError> {
let fingerprints = certs
.into_iter()
.map(|cert| cert.to_der().map(|der| D::digest(&der).to_vec()))
.collect::<Result<Vec<_>, _>>()?;
Ok(Self { fingerprints, _digest: PhantomData })
}
pub fn from_fingerprints(fingerprints: impl IntoIterator<Item = Vec<u8>>) -> Self {
Self { fingerprints: fingerprints.into_iter().collect(), _digest: PhantomData }
}
}
#[cfg(feature = "std")]
impl<D> CertificateValidation for RuntimeCertificatePinning<D>
where
D: Digest + Send + Sync,
{
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
let cert_der = cert.to_der()?;
let fingerprint = D::digest(&cert_der);
self.fingerprints
.iter()
.any(|fp| fp.as_slice() == fingerprint.as_ref())
.then_some(())
.ok_or(CertificateValidationError::CertificateNotPinned)
}
}
#[derive(Default, Clone)]
pub struct DirectTrustValidator {
trust_chain: Vec<Certificate>,
}
impl DirectTrustValidator {
pub fn with_trust_chain(mut self, trust_chain: Vec<Certificate>) -> Self {
self.trust_chain = trust_chain;
self
}
}
impl CertificateValidation for DirectTrustValidator {
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
validate_certificate_expiry(cert)?;
let cert_der = cert.to_der()?;
for anchor in &self.trust_chain {
if anchor.to_der()? == cert_der {
return Ok(());
}
}
Err(CertificateValidationError::CertificateNotTrusted)
}
}
impl SignatureVerification for DirectTrustValidator {
fn verify_with_policy(
&self,
cert: &Certificate,
curr_time: u64,
public_key_der: &[u8],
policy: &dyn VerificationPolicy,
) -> Result<(), CertificateValidationError> {
let not_before = cert.tbs_certificate.validity.not_before.to_unix_duration();
let not_after = cert.tbs_certificate.validity.not_after.to_unix_duration();
let now_duration = GeneralizedTime::from_unix_duration(Duration::from_secs(curr_time))
.map_err(|_| CertificateValidationError::InvalidTimestamp)?
.to_unix_duration();
if now_duration < not_before {
return Err(CertificateValidationError::NotYetValid);
}
if now_duration > not_after {
return Err(CertificateValidationError::Expired);
}
let subject_public_key = cert.tbs_certificate.subject_public_key_info.subject_public_key.raw_bytes();
if subject_public_key.is_empty() {
return Err(CertificateValidationError::EmptyPublicKey);
}
let signature_bytes = cert.signature.raw_bytes();
if signature_bytes.is_empty() {
return Err(CertificateValidationError::EmptySignature);
}
ensure_signature_algorithm_consistency(cert)?;
let message = cert.tbs_certificate.to_der()?;
policy.verify_signature(&cert.signature_algorithm.oid, public_key_der, &message, signature_bytes)
}
}
#[cfg(feature = "std")]
#[derive(Default)]
pub struct ChainValidator {
validators: Vec<Box<dyn CertificateValidation>>,
}
#[cfg(feature = "std")]
impl ChainValidator {
#[allow(clippy::should_implement_trait)]
pub fn add(mut self, validator: Box<dyn CertificateValidation>) -> Self {
self.validators.push(validator);
self
}
}
#[cfg(feature = "std")]
impl CertificateValidation for ChainValidator {
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
self.validators.iter().try_for_each(|v| v.evaluate(cert))
}
}
#[cfg(test)]
mod tests {
use crate::crypto::x509::error::CertificateValidationError;
use crate::crypto::x509::policy::{CertificateValidation, ExpiryValidator};
use crate::testing::create_expired_test_certificate;
#[test]
fn test_expiry_validator_rejects_expired_cert() {
let expired_cert = create_expired_test_certificate();
let validator = ExpiryValidator;
let result = validator.evaluate(&expired_cert);
assert!(result.is_err(), "Expired certificate should be rejected");
match result {
Err(CertificateValidationError::Expired) => {
}
other => panic!("Expected Expired error, got: {other:?}"),
}
}
#[cfg(all(feature = "secp256k1", feature = "signature", feature = "x509", feature = "std"))]
mod direct_trust {
use k256::ecdsa::SigningKey;
use crate::crypto::policy::Secp256k1Policy;
use crate::crypto::x509::error::CertificateValidationError;
use crate::crypto::x509::policy::{CertificateValidation, DirectTrustValidator, SignatureVerification};
use crate::crypto::x509::Certificate;
use crate::oids::SIGNER_ECDSA_WITH_SHA256;
use crate::spki::EncodePublicKey;
use crate::testing::utils::{create_test_certificate, create_test_signing_key};
fn test_cert() -> Certificate {
create_test_certificate(&create_test_signing_key())
}
fn verify(signing_key: &SigningKey, cert: &Certificate) -> Result<(), CertificateValidationError> {
let issuer_pub = signing_key.verifying_key().to_public_key_der()?;
DirectTrustValidator::default().verify_with_policy(cert, 1_000, issuer_pub.as_bytes(), &Secp256k1Policy)
}
#[test]
fn accepts_configured_anchor() {
let cert = test_cert();
let trust_chain = vec![cert.clone()];
let validator = DirectTrustValidator::default().with_trust_chain(trust_chain);
assert!(validator.evaluate(&cert).is_ok());
}
#[test]
fn rejects_non_anchor() -> Result<(), Box<dyn core::error::Error>> {
let validator = DirectTrustValidator::default().with_trust_chain(vec![test_cert()]);
let other = create_test_certificate(&SigningKey::from_bytes(&[7u8; 32].into())?);
assert!(matches!(
validator.evaluate(&other),
Err(CertificateValidationError::CertificateNotTrusted)
));
Ok(())
}
#[test]
fn without_anchor_fails_closed() {
assert!(matches!(
DirectTrustValidator::default().evaluate(&test_cert()),
Err(CertificateValidationError::CertificateNotTrusted)
));
}
#[test]
fn rejects_algorithm_mismatch() {
let key = create_test_signing_key();
let mut cert = create_test_certificate(&key);
cert.signature_algorithm.oid = SIGNER_ECDSA_WITH_SHA256;
assert!(matches!(
verify(&key, &cert),
Err(CertificateValidationError::AlgorithmMismatch)
));
}
#[test]
fn rejects_foreign_algorithm() {
let key = create_test_signing_key();
let mut cert = create_test_certificate(&key);
cert.signature_algorithm.oid = SIGNER_ECDSA_WITH_SHA256;
cert.tbs_certificate.signature.oid = SIGNER_ECDSA_WITH_SHA256;
assert!(matches!(
verify(&key, &cert),
Err(CertificateValidationError::UnsupportedAlgorithm(_))
));
}
}
}