use crate::backends::tlv::read_tlv;
use crate::traits::cert::CertView;
#[derive(Debug, PartialEq, Eq, Clone, Copy, thiserror::Error)]
pub enum IdentityError {
#[error("cert had no SubjectAltName extension")]
NoSan,
#[error("no SAN dNSName matched the requested hostname")]
HostnameMismatch,
#[error("cert SAN extension was malformed")]
MalformedSan,
#[error("pinned public key did not match the cert pubkey")]
PinMismatch,
#[error("pinned-key algorithm family did not match the cert")]
PinAlgorithmMismatch,
}
#[derive(Debug, Clone, Copy)]
pub enum PinnedPubkey<'a> {
Ed25519([u8; 32]),
#[cfg(feature = "rsa")]
Rsa { modulus: &'a [u8], exponent: u32 },
#[cfg(feature = "mldsa")]
MlDsa(&'a [u8]),
#[doc(hidden)]
_Phantom(core::marker::PhantomData<&'a ()>, core::convert::Infallible),
}
impl<'a> PinnedPubkey<'a> {
pub fn to_owned_pin(
&self,
) -> Result<crate::backends::PinnedPubkeyOwned, crate::backends::PinnedPubkeyOwnedError> {
match self {
PinnedPubkey::Ed25519(pk) => Ok(crate::backends::PinnedPubkeyOwned::ed25519(*pk)),
#[cfg(feature = "rsa")]
PinnedPubkey::Rsa { modulus, exponent } => {
crate::backends::PinnedPubkeyOwned::rsa(modulus, *exponent)
}
#[cfg(feature = "mldsa")]
PinnedPubkey::MlDsa(pk) => crate::backends::PinnedPubkeyOwned::mldsa(pk),
PinnedPubkey::_Phantom(_, never) => match *never {},
}
}
}
pub fn verify_hostname(cert_view: &CertView<'_>, hostname: &str) -> Result<(), IdentityError> {
let san = cert_view.san().ok_or(IdentityError::NoSan)?;
if looks_like_ip_literal(hostname) {
#[cfg(feature = "ip-host")]
{
let unbracketed = hostname
.strip_prefix('[')
.and_then(|s| s.strip_suffix(']'))
.unwrap_or(hostname);
if let Ok(v4) = unbracketed.parse::<core::net::Ipv4Addr>() {
return ip_host::match_ip_address(san, &v4.octets());
}
if let Ok(v6) = unbracketed.parse::<core::net::Ipv6Addr>() {
return ip_host::match_ip_address(san, &v6.octets());
}
}
return Err(IdentityError::HostnameMismatch);
}
let hostname_bytes = hostname.as_bytes();
for entry in san_dns_names(san) {
let entry = entry?;
if dns_name_matches(entry, hostname_bytes) {
return Ok(());
}
}
Err(IdentityError::HostnameMismatch)
}
fn looks_like_ip_literal(hostname: &str) -> bool {
let bytes = hostname.as_bytes();
if bytes.first() == Some(&b'[') {
return true;
}
let mut saw_dot = false;
let mut all_digit_or_dot = true;
for &b in bytes {
if b == b':' {
return true;
}
if b == b'.' {
saw_dot = true;
} else if !b.is_ascii_digit() {
all_digit_or_dot = false;
}
}
!bytes.is_empty() && all_digit_or_dot && saw_dot
}
pub fn san_dns_names(san_bytes: &[u8]) -> SanDnsIter<'_> {
SanDnsIter(SanEntryWalker {
rest: san_bytes,
target_tag: DNS_NAME_TAG,
})
}
const DNS_NAME_TAG: u8 = crate::backends::tlv::tag_ctx_primitive(2);
pub struct SanDnsIter<'a>(SanEntryWalker<'a>);
impl<'a> Iterator for SanDnsIter<'a> {
type Item = Result<&'a [u8], IdentityError>;
fn next(&mut self) -> Option<Self::Item> {
self.0.next()
}
}
struct SanEntryWalker<'a> {
rest: &'a [u8],
target_tag: u8,
}
impl<'a> SanEntryWalker<'a> {
fn next(&mut self) -> Option<Result<&'a [u8], IdentityError>> {
while !self.rest.is_empty() {
let t = match read_tlv(self.rest) {
Ok(t) => t,
Err(_) => {
self.rest = &[];
return Some(Err(IdentityError::MalformedSan));
}
};
self.rest = t.rest;
if t.tag == self.target_tag {
#[cfg(feature = "ip-host")]
if self.target_tag == ip_host::IP_ADDRESS_TAG
&& t.body.len() != 4
&& t.body.len() != 16
{
return Some(Err(IdentityError::MalformedSan));
}
return Some(Ok(t.body));
}
}
None
}
}
#[cfg(feature = "ip-host")]
mod ip_host {
use super::{IdentityError, SanEntryWalker};
pub(super) fn match_ip_address(san: &[u8], expected: &[u8]) -> Result<(), IdentityError> {
for entry in san_ip_addresses(san) {
if entry? == expected {
return Ok(());
}
}
Err(IdentityError::HostnameMismatch)
}
pub(super) fn san_ip_addresses(san_bytes: &[u8]) -> SanIpAddressIter<'_> {
SanIpAddressIter(SanEntryWalker {
rest: san_bytes,
target_tag: IP_ADDRESS_TAG,
})
}
pub(super) const IP_ADDRESS_TAG: u8 = crate::backends::tlv::tag_ctx_primitive(7);
pub(super) struct SanIpAddressIter<'a>(SanEntryWalker<'a>);
impl<'a> Iterator for SanIpAddressIter<'a> {
type Item = Result<&'a [u8], IdentityError>;
fn next(&mut self) -> Option<Self::Item> {
self.0.next()
}
}
}
fn dns_name_matches(pattern: &[u8], hostname: &[u8]) -> bool {
if let Some(suffix) = pattern.strip_prefix(b"*.") {
if suffix.is_empty() || suffix.contains(&b'*') {
return false;
}
if hostname.len() <= suffix.len() + 1 {
return false;
}
let split_at = hostname.len() - suffix.len();
let (label_with_dot, tail) = hostname.split_at(split_at);
if !ascii_eq_ignore_case(tail, suffix) {
return false;
}
if label_with_dot.last() != Some(&b'.') {
return false;
}
let label = &label_with_dot[..label_with_dot.len() - 1];
if label.is_empty() || label.contains(&b'.') {
return false;
}
true
} else {
if pattern.contains(&b'*') {
return false;
}
ascii_eq_ignore_case(pattern, hostname)
}
}
fn ascii_eq_ignore_case(a: &[u8], b: &[u8]) -> bool {
a.eq_ignore_ascii_case(b)
}
#[cfg(feature = "cert-der")]
mod validity {
use der::asn1::{GeneralizedTime, UtcTime};
use der::{Decode, Reader, SliceReader, Tag};
#[derive(Debug, PartialEq, Eq, Clone, Copy, thiserror::Error)]
pub enum ValidityError {
#[error("cert not yet valid: notBefore={not_before}, now={now}")]
NotYetValid { not_before: u64, now: u64 },
#[error("cert expired: notAfter={not_after}, now={now}")]
Expired { not_after: u64, now: u64 },
#[error("cert validity field did not decode")]
Malformed,
}
pub fn verify_validity<T: crate::traits::time::TimeSource + ?Sized>(
cert_view: &crate::traits::cert::CertView<'_>,
time: &T,
) -> Result<(), ValidityError> {
let (not_before, not_after) = parse_validity_der(cert_view.validity_der())?;
let now = time.now_unix_secs();
if now < not_before {
return Err(ValidityError::NotYetValid { not_before, now });
}
if now > not_after {
return Err(ValidityError::Expired { not_after, now });
}
Ok(())
}
fn parse_validity_der(der_bytes: &[u8]) -> Result<(u64, u64), ValidityError> {
let mut outer = SliceReader::new(der_bytes).map_err(|_| ValidityError::Malformed)?;
let result: der::Result<(u64, u64)> = outer.sequence(|inner| {
let nb = decode_time_to_unix(inner)?;
let na = decode_time_to_unix(inner)?;
if !inner.is_finished() {
return Err(inner.error(der::ErrorKind::TrailingData {
decoded: 0u8.into(),
remaining: 0u8.into(),
}));
}
Ok((nb, na))
});
let (nb, na) = result.map_err(|_| ValidityError::Malformed)?;
if !outer.is_finished() {
return Err(ValidityError::Malformed);
}
Ok((nb, na))
}
fn decode_time_to_unix<'a, R: der::Reader<'a>>(r: &mut R) -> der::Result<u64> {
let tag = Tag::peek(r)?;
let secs = match tag {
Tag::UtcTime => UtcTime::decode(r)?.to_unix_duration().as_secs(),
Tag::GeneralizedTime => GeneralizedTime::decode(r)?.to_unix_duration().as_secs(),
_ => {
return Err(r.error(der::ErrorKind::TagUnexpected {
expected: None,
actual: tag,
}));
}
};
Ok(secs)
}
}
#[cfg(feature = "cert-der")]
pub use validity::{ValidityError, verify_validity};
#[cfg(test)]
mod tests;