use aws_lc_rs::digest;
use rand::random;
use rcgen::{CertificateParams, DistinguishedName, DnType, IsCa, KeyPair, PKCS_ECDSA_P256_SHA256};
use std::fmt;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CertificateError {
InvalidFormat,
FingerprintMismatch,
GenerationFailed,
}
impl fmt::Display for CertificateError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
CertificateError::InvalidFormat => write!(f, "Invalid certificate format"),
CertificateError::FingerprintMismatch => write!(f, "Fingerprint mismatch"),
CertificateError::GenerationFailed => write!(f, "Certificate generation failed"),
}
}
}
impl std::error::Error for CertificateError {}
pub use crate::DtlsCertificate;
pub fn generate_self_signed_certificate() -> Result<DtlsCertificate, CertificateError> {
let key_pair = KeyPair::generate_for(&PKCS_ECDSA_P256_SHA256)
.map_err(|_| CertificateError::GenerationFailed)?;
let mut params = CertificateParams::new(Vec::<String>::new())
.map_err(|_| CertificateError::GenerationFailed)?;
let mut distinguished_name = DistinguishedName::new();
distinguished_name.push(DnType::OrganizationName, "DTLS".to_string());
distinguished_name.push(DnType::CommonName, "DTLS Peer".to_string());
params.distinguished_name = distinguished_name;
params.is_ca = IsCa::NoCa;
let not_before = time::OffsetDateTime::now_utc();
let not_after = not_before + time::Duration::days(365);
params.not_before = not_before;
params.not_after = not_after;
let serial_buf: [u8; 16] = random();
params.serial_number = Some(serial_buf.to_vec().into());
let cert = params
.self_signed(&key_pair)
.map_err(|_| CertificateError::GenerationFailed)?;
let cert_der = cert.der().to_vec();
let key_der = key_pair.serialize_der();
Ok(DtlsCertificate {
certificate: cert_der,
private_key: key_der,
})
}
pub fn calculate_fingerprint(cert_der: &[u8]) -> Vec<u8> {
digest::digest(&digest::SHA256, cert_der).as_ref().to_vec()
}
pub fn format_fingerprint(fingerprint: &[u8]) -> String {
fingerprint
.iter()
.map(|byte| format!("{:02X}", byte))
.collect::<Vec<String>>()
.join(":")
}
impl DtlsCertificate {
pub fn fingerprint(&self) -> Vec<u8> {
calculate_fingerprint(&self.certificate)
}
pub fn fingerprint_str(&self) -> String {
format_fingerprint(&self.fingerprint())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_self_signed_certificate() {
let cert = generate_self_signed_certificate().unwrap();
assert!(!cert.certificate.is_empty());
assert!(!cert.private_key.is_empty());
assert_eq!(cert.fingerprint().len(), 32);
}
#[test]
fn test_unique_serial_numbers() {
let cert1 = generate_self_signed_certificate().unwrap();
let cert2 = generate_self_signed_certificate().unwrap();
assert_ne!(cert1.fingerprint(), cert2.fingerprint());
use x509_parser::prelude::*;
let (_, parsed1) = X509Certificate::from_der(&cert1.certificate).unwrap();
let (_, parsed2) = X509Certificate::from_der(&cert2.certificate).unwrap();
assert_ne!(
parsed1.serial, parsed2.serial,
"Serial numbers must be unique for Firefox compatibility"
);
}
#[test]
fn test_fingerprint_formatting() {
let test_fingerprint = vec![0xAF, 0x12, 0xF6, 0x38, 0x2A];
let formatted = format_fingerprint(&test_fingerprint);
assert_eq!(formatted, "AF:12:F6:38:2A");
let cert = generate_self_signed_certificate().unwrap();
let formatted = format_fingerprint(&cert.fingerprint());
assert_eq!(formatted.len(), 95); assert!(formatted.contains(':'));
for segment in formatted.split(':') {
assert_eq!(segment.len(), 2);
assert!(u8::from_str_radix(segment, 16).is_ok());
}
}
}