use core::fmt;
use subtle::ConstantTimeEq;
use crate::HpkeError;
pub const MIN_PSK_LEN: usize = 32;
#[derive(Clone, Copy)]
pub struct Psk<'a> {
secret: &'a [u8],
id: &'a [u8],
}
impl PartialEq for Psk<'_> {
fn eq(&self, other: &Self) -> bool {
self.id == other.id && bool::from(self.secret.ct_eq(other.secret))
}
}
impl Eq for Psk<'_> {}
impl<'a> Psk<'a> {
pub fn new(secret: &'a [u8], id: &'a [u8]) -> Result<Self, HpkeError> {
if secret.is_empty() != id.is_empty() {
return Err(HpkeError::InconsistentPsk);
}
if secret.is_empty() {
return Err(HpkeError::MissingPsk);
}
if secret.len() < MIN_PSK_LEN {
return Err(HpkeError::InsecurePsk);
}
Ok(Self { secret, id })
}
#[must_use]
pub const fn secret(&self) -> &'a [u8] {
self.secret
}
#[must_use]
pub const fn id(&self) -> &'a [u8] {
self.id
}
}
impl fmt::Debug for Psk<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Psk")
.field("secret", &format_args!("<{} bytes>", self.secret.len()))
.field("id", &self.id)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::format;
#[test]
fn validation_matrix() {
let good = [0u8; MIN_PSK_LEN];
assert!(Psk::new(&good, b"id").is_ok());
assert_eq!(
Psk::new(&[0u8; MIN_PSK_LEN - 1], b"id"),
Err(HpkeError::InsecurePsk)
);
assert_eq!(Psk::new(&good, b""), Err(HpkeError::InconsistentPsk));
assert_eq!(Psk::new(b"", b"id"), Err(HpkeError::InconsistentPsk));
assert_eq!(Psk::new(b"", b""), Err(HpkeError::MissingPsk));
}
#[test]
fn accessors_round_trip() {
let secret = [0x5Au8; MIN_PSK_LEN];
let psk = Psk::new(&secret, b"the-id").unwrap();
assert_eq!(psk.secret(), &secret);
assert_eq!(psk.id(), b"the-id");
}
#[test]
fn equality_compares_both_fields() {
let a = [0x11u8; MIN_PSK_LEN];
let mut b = a;
b[MIN_PSK_LEN - 1] ^= 1;
assert_eq!(Psk::new(&a, b"id").unwrap(), Psk::new(&a, b"id").unwrap());
assert_ne!(Psk::new(&a, b"id").unwrap(), Psk::new(&b, b"id").unwrap());
assert_ne!(Psk::new(&a, b"id").unwrap(), Psk::new(&a, b"id-x").unwrap());
}
#[test]
fn debug_redacts_the_secret() {
let psk = Psk::new(&[0xAAu8; MIN_PSK_LEN], b"the-id").unwrap();
let rendered = format!("{psk:?}");
assert!(rendered.contains("<32 bytes>"), "{rendered}");
assert!(!rendered.contains("170"), "PSK bytes leaked: {rendered}");
assert!(rendered.contains("id: [116,"), "{rendered}");
}
}