use core::marker::PhantomData;
use core::time::Duration;
use crate::asn1::GeneralizedTime;
use crate::crypto::hash::Digest;
use crate::crypto::policy::VerificationPolicy;
use crate::crypto::x509::error::CertificateValidationError;
use crate::crypto::x509::utils::validate_certificate_expiry;
use crate::crypto::x509::Certificate;
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 FullValidator {
trust_chain: Vec<Certificate>,
}
impl FullValidator {
pub fn with_trust_chain(mut self, trust_chain: Vec<Certificate>) -> Self {
self.trust_chain = trust_chain;
self
}
}
impl CertificateValidation for FullValidator {
fn evaluate(&self, cert: &Certificate) -> Result<(), CertificateValidationError> {
validate_certificate_expiry(cert)
}
}
impl SignatureVerification for FullValidator {
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);
}
if cert.signature_algorithm.oid != cert.tbs_certificate.signature.oid {
return Err(CertificateValidationError::AlgorithmMismatch);
}
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:?}"),
}
}
}