use crate::error::CryptError;
use crate::protocol::version::{MAGIC, VERSION_V2};
pub const HEADER_SIZE: usize = 14;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[repr(u8)]
pub enum KemAlgId {
MlKem512 = 1,
MlKem768 = 2,
MlKem1024 = 3,
}
impl KemAlgId {
pub fn from_byte(b: u8) -> Result<Self, CryptError> {
match b {
1 => Ok(Self::MlKem512),
2 => Ok(Self::MlKem768),
3 => Ok(Self::MlKem1024),
_ => Err(CryptError::UnsupportedAlgorithm),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[repr(u8)]
pub enum AeadAlgId {
AesCbc = 1,
AesGcmSiv = 2,
AesCtr = 3,
AesXts = 4,
XChaCha20 = 5,
XChaCha20Poly1305 = 6,
}
impl AeadAlgId {
pub fn is_aead(self) -> bool {
matches!(self, Self::AesGcmSiv | Self::XChaCha20Poly1305)
}
pub fn from_byte(b: u8) -> Result<Self, CryptError> {
match b {
1 => Ok(Self::AesCbc),
2 => Ok(Self::AesGcmSiv),
3 => Ok(Self::AesCtr),
4 => Ok(Self::AesXts),
5 => Ok(Self::XChaCha20),
6 => Ok(Self::XChaCha20Poly1305),
_ => Err(CryptError::UnsupportedAlgorithm),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[repr(u8)]
pub enum KdfAlgId {
HkdfSha256 = 1,
HkdfSha512 = 2,
}
impl KdfAlgId {
pub fn from_byte(b: u8) -> Result<Self, CryptError> {
match b {
1 => Ok(Self::HkdfSha256),
2 => Ok(Self::HkdfSha512),
_ => Err(CryptError::UnsupportedAlgorithm),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Header {
pub magic: [u8; 4],
pub version: u16,
pub kem_alg: KemAlgId,
pub aead_alg: AeadAlgId,
pub kdf_alg: KdfAlgId,
pub flags: u8,
}
impl Header {
pub fn new(kem_alg: KemAlgId, aead_alg: AeadAlgId, kdf_alg: KdfAlgId) -> Self {
Self {
magic: MAGIC,
version: VERSION_V2,
kem_alg,
aead_alg,
kdf_alg,
flags: 0,
}
}
pub fn to_bytes(self) -> [u8; HEADER_SIZE] {
let mut out = [0u8; HEADER_SIZE];
out[0..4].copy_from_slice(&self.magic);
out[4..6].copy_from_slice(&self.version.to_le_bytes());
out[6] = self.kem_alg as u8;
out[7] = self.aead_alg as u8;
out[8] = self.kdf_alg as u8;
out[9] = self.flags;
out
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, CryptError> {
if bytes.len() < HEADER_SIZE {
return Err(CryptError::InvalidEnvelope);
}
let mut magic = [0u8; 4];
magic.copy_from_slice(&bytes[0..4]);
if magic != MAGIC {
return Err(CryptError::InvalidEnvelope);
}
let version = u16::from_le_bytes([bytes[4], bytes[5]]);
if version != VERSION_V2 {
return Err(CryptError::UnsupportedEnvelopeVersion);
}
let kem_alg = KemAlgId::from_byte(bytes[6])?;
let aead_alg = AeadAlgId::from_byte(bytes[7])?;
let kdf_alg = KdfAlgId::from_byte(bytes[8])?;
let flags = bytes[9];
Ok(Self {
magic,
version,
kem_alg,
aead_alg,
kdf_alg,
flags,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_header_roundtrip() {
let hdr = Header::new(
KemAlgId::MlKem768,
AeadAlgId::XChaCha20Poly1305,
KdfAlgId::HkdfSha256,
);
let bytes = hdr.to_bytes();
assert_eq!(bytes.len(), HEADER_SIZE);
let parsed = Header::from_bytes(&bytes).unwrap();
assert_eq!(hdr, parsed);
}
#[test]
fn test_header_bad_magic() {
let mut bytes = Header::new(
KemAlgId::MlKem512,
AeadAlgId::AesGcmSiv,
KdfAlgId::HkdfSha512,
)
.to_bytes();
bytes[0] = 0xFF;
assert!(matches!(
Header::from_bytes(&bytes),
Err(CryptError::InvalidEnvelope)
));
}
#[test]
fn test_header_bad_version() {
let mut bytes = Header::new(
KemAlgId::MlKem512,
AeadAlgId::AesGcmSiv,
KdfAlgId::HkdfSha256,
)
.to_bytes();
bytes[4] = 9;
bytes[5] = 0;
assert!(matches!(
Header::from_bytes(&bytes),
Err(CryptError::UnsupportedEnvelopeVersion)
));
}
#[test]
fn test_header_short_bytes() {
assert!(matches!(
Header::from_bytes(&[0u8; 5]),
Err(CryptError::InvalidEnvelope)
));
}
#[test]
fn test_aead_alg_is_aead() {
assert!(AeadAlgId::AesGcmSiv.is_aead());
assert!(AeadAlgId::XChaCha20Poly1305.is_aead());
assert!(!AeadAlgId::AesCbc.is_aead());
assert!(!AeadAlgId::AesCtr.is_aead());
assert!(!AeadAlgId::XChaCha20.is_aead());
}
}