use std::fmt::{self, Debug};
use serde::{Deserialize, Serialize};
use zeroize::{Zeroize, ZeroizeOnDrop};
use crate::encoding::Base64UrlBytes;
use crate::error::{IncompatibleKeyError, InvalidKeyError, Result};
#[derive(Clone, PartialEq, Eq, Hash, Serialize, Deserialize, Zeroize, ZeroizeOnDrop)]
#[non_exhaustive]
pub struct SymmetricParams {
pub k: Base64UrlBytes,
}
impl SymmetricParams {
#[must_use]
pub fn new(k: Base64UrlBytes) -> Self {
Self { k }
}
pub fn is_public_key_only(&self) -> bool {
false
}
pub fn has_private_key(&self) -> bool {
true
}
pub fn key_size_bits(&self) -> usize {
self.k.len() * 8
}
pub fn validate(&self) -> Result<()> {
if self.k.is_empty() {
return Err(InvalidKeyError::MissingParameter("k").into());
}
Ok(())
}
pub fn validate_min_size(&self, min_bits: usize, context: &'static str) -> Result<()> {
self.validate()?;
let actual_bits = self.key_size_bits();
if actual_bits < min_bits {
return Err(IncompatibleKeyError::InsufficientKeyStrength {
minimum_bits: min_bits,
actual_bits,
context,
}
.into());
}
Ok(())
}
pub fn validate_exact_size(&self, exact_bits: usize, context: &'static str) -> Result<()> {
self.validate()?;
let actual_bits = self.key_size_bits();
if actual_bits != exact_bits {
return Err(IncompatibleKeyError::KeySizeMismatch {
required_bits: exact_bits,
actual_bits,
context,
}
.into());
}
Ok(())
}
#[inline]
pub fn ct_eq(&self, other: &Self) -> bool {
self.k.ct_eq(&other.k)
}
}
impl Debug for SymmetricParams {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SymmetricParams")
.field("key_size_bits", &self.key_size_bits())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_key_size() {
let params = SymmetricParams::new(Base64UrlBytes::new(vec![0; 32]));
assert_eq!(params.key_size_bits(), 256);
}
#[test]
fn test_validate_empty() {
let params = SymmetricParams::new(Base64UrlBytes::new(vec![]));
assert!(params.validate().is_err());
}
#[test]
fn test_validate_min_size() {
let small_key = SymmetricParams::new(Base64UrlBytes::new(vec![0; 16]));
assert!(small_key.validate_min_size(128, "test").is_ok());
assert!(small_key.validate_min_size(256, "test").is_err());
}
#[test]
fn test_validate_exact_size() {
let key_128 = SymmetricParams::new(Base64UrlBytes::new(vec![0; 16]));
assert!(key_128.validate_exact_size(128, "AES-128").is_ok());
assert!(key_128.validate_exact_size(256, "AES-256").is_err());
}
#[test]
fn test_json_roundtrip() {
let original = SymmetricParams::new(Base64UrlBytes::new(vec![1, 2, 3, 4, 5]));
let json = serde_json::to_string(&original).unwrap();
let decoded: SymmetricParams = serde_json::from_str(&json).unwrap();
assert_eq!(original, decoded);
}
#[test]
fn test_constant_time_equality() {
let a = SymmetricParams::new(Base64UrlBytes::new(vec![0; 32]));
let b = SymmetricParams::new(Base64UrlBytes::new(vec![0; 32]));
let mut different = vec![0; 32];
different[31] = 1;
let c = SymmetricParams::new(Base64UrlBytes::new(different));
assert!(a.ct_eq(&b));
assert!(!a.ct_eq(&c));
}
}