use std::fmt::Display;
use smallvec::SmallVec;
use thiserror::Error;
use valuable::Valuable;
use x509_parser::prelude::{ParsedExtension, X509Certificate};
use crate::{
cached_string_repr::CachedStringRepr, certificate::id::DangerousComparableId, hex::colon_string,
};
#[derive(Debug, Error)]
#[error("no Subject Key Identifier found")]
pub struct ErrorNoSKI;
#[derive(Debug, Error)]
#[error("certificate ID is an invalid length")]
pub struct InvalidLength {
actual_len: usize,
}
#[derive(Debug, Error)]
pub enum InvalidCertId {
#[error(transparent)]
NoSKI(#[from] ErrorNoSKI),
#[error(transparent)]
InvalidLength(#[from] InvalidLength),
}
#[derive(Debug, Hash, Clone)] pub struct CertId {
bytes: SmallVec<[u8; 20]>,
rendered: CachedStringRepr,
}
impl CertId {
const MIN_LENGTH: usize = 16;
const MAX_LENGTH: usize = 64;
pub fn as_hex_str(&self) -> &str {
self.rendered.get_or_init(|| colon_string(&self.bytes))
}
pub(super) fn as_bytes(&self) -> &[u8] {
&self.bytes
}
pub fn as_dangerous_comparable(&self) -> DangerousComparableId<'_, Self> {
DangerousComparableId::from(self)
}
}
impl Display for CertId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_hex_str())
}
}
impl From<CertId> for Vec<u8> {
fn from(id: CertId) -> Vec<u8> {
id.bytes.into_vec()
}
}
impl TryFrom<&[u8]> for CertId {
type Error = InvalidLength;
fn try_from(value: &[u8]) -> Result<Self, Self::Error> {
if value.len() < Self::MIN_LENGTH || value.len() > Self::MAX_LENGTH {
return Err(InvalidLength {
actual_len: value.len(),
});
}
Ok(Self {
bytes: SmallVec::from_slice(value),
rendered: Default::default(),
})
}
}
impl TryFrom<Vec<u8>> for CertId {
type Error = InvalidLength;
fn try_from(value: Vec<u8>) -> Result<Self, Self::Error> {
if value.len() < Self::MIN_LENGTH || value.len() > Self::MAX_LENGTH {
return Err(InvalidLength {
actual_len: value.len(),
});
}
Ok(Self {
bytes: SmallVec::from_vec(value),
rendered: Default::default(),
})
}
}
impl<'a> TryFrom<&X509Certificate<'a>> for CertId {
type Error = InvalidCertId;
fn try_from(cert: &X509Certificate<'a>) -> Result<Self, Self::Error> {
let bytes = cert
.iter_extensions()
.find_map(|v| match v.parsed_extension() {
ParsedExtension::SubjectKeyIdentifier(ski) => Some(ski.0),
_ => None,
})
.ok_or(ErrorNoSKI)?;
Ok(CertId::try_from(bytes)?)
}
}
impl Valuable for CertId {
fn as_value(&self) -> valuable::Value<'_> {
valuable::Value::String(self.as_hex_str())
}
fn visit(&self, visit: &mut dyn valuable::Visit) {
visit.visit_value(self.as_value());
}
}
#[cfg(test)]
mod tests {
use proptest::prelude::*;
use rc_x509_test_helpers::assert_valuable_repr;
use static_assertions::assert_not_impl_any;
use x509_parser::prelude::FromDer;
use super::*;
use crate::certificate::tests::cert_fixture;
const FIXTURE_SKI_STR: &str = "dc:8d:b6:27:52:78:58:4c:fd:a2:43:db:cb:2b:e0:57:68:6e:2b:8e";
assert_not_impl_any!(CertId: PartialEq, Eq);
fn fixture_ski() -> CertId {
let der = cert_fixture().as_der();
let cert = X509Certificate::from_der(&der).expect("valid DER").1;
CertId::try_from(&cert).expect("extract SKI")
}
#[test]
fn test_fixture() {
let aki = fixture_ski();
assert_eq!(aki.as_hex_str(), FIXTURE_SKI_STR,);
}
#[test]
fn test_valuable_repr() {
let aki = fixture_ski();
assert_valuable_repr(&aki, FIXTURE_SKI_STR);
}
#[test]
fn test_danger_eq() {
let ski = fixture_ski();
assert_eq!(ski.as_dangerous_comparable(), ski);
}
proptest! {
#[test]
fn prop_length_bounds_enforced(
ski in prop::collection::vec(any::<u8>(), 0..(CertId::MAX_LENGTH + 20)),
) {
let in_bounds = (CertId::MIN_LENGTH..=CertId::MAX_LENGTH).contains(&ski.len());
assert_eq!(CertId::try_from(ski.as_slice()).is_ok(), in_bounds);
assert_eq!(CertId::try_from(ski).is_ok(), in_bounds);
}
}
}