use crate::error::{Error, Result};
use ml_kem::array::Array;
use ml_kem::{B32, EncapsulateDeterministic, EncodedSizeUser, KemCore, MlKem768};
use zeroize::{Zeroize, ZeroizeOnDrop};
#[derive(Clone, PartialEq, Eq)]
pub struct PublicKey(pub(crate) Vec<u8>);
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct SecretKey(pub(crate) Vec<u8>);
#[derive(Clone, PartialEq, Eq)]
pub struct Ciphertext(pub(crate) Vec<u8>);
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct SharedSecret(pub(crate) [u8; 32]);
impl PublicKey {
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
if bytes.len() != PK_LEN {
return Err(Error::InvalidLength {
expected: PK_LEN,
got: bytes.len(),
});
}
Ok(Self(bytes))
}
pub(crate) fn from_bytes_unchecked(bytes: Vec<u8>) -> Self {
assert_eq!(
bytes.len(),
PK_LEN,
"mlkem::PublicKey::from_bytes_unchecked: wrong size"
);
Self(bytes)
}
}
impl SecretKey {
pub(crate) fn as_bytes(&self) -> &[u8] {
&self.0
}
pub(crate) fn from_bytes_unchecked(bytes: Vec<u8>) -> Self {
assert_eq!(
bytes.len(),
SK_LEN,
"mlkem::SecretKey::from_bytes_unchecked: wrong size"
);
Self(bytes)
}
}
impl Ciphertext {
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
if bytes.len() != CT_LEN {
return Err(Error::InvalidLength {
expected: CT_LEN,
got: bytes.len(),
});
}
Ok(Self(bytes))
}
pub(crate) fn from_bytes_unchecked(bytes: Vec<u8>) -> Self {
assert_eq!(
bytes.len(),
CT_LEN,
"mlkem::Ciphertext::from_bytes_unchecked: wrong size"
);
Self(bytes)
}
}
impl SharedSecret {
pub fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
const PK_LEN: usize = 1184;
const SK_LEN: usize = 2400;
const CT_LEN: usize = 1088;
pub const fn pk_len() -> usize {
PK_LEN
}
pub const fn sk_len() -> usize {
SK_LEN
}
pub const fn ct_len() -> usize {
CT_LEN
}
fn random_b32() -> B32 {
let mut buf = [0u8; 32];
super::random::random_bytes(&mut buf);
let arr = Array::from(buf);
buf.zeroize();
arr
}
pub fn keygen() -> Result<(PublicKey, SecretKey)> {
let mut d = random_b32();
let mut z = random_b32();
let (dk, ek) = MlKem768::generate_deterministic(&d, &z);
d.zeroize();
z.zeroize();
let pk_bytes = ek.as_bytes().to_vec();
let mut sk_bytes = zeroize::Zeroizing::new(dk.as_bytes().to_vec());
assert_eq!(
pk_bytes.len(),
PK_LEN,
"ML-KEM-768 EK size mismatch: got {}, expected {PK_LEN}",
pk_bytes.len()
);
assert_eq!(
sk_bytes.len(),
SK_LEN,
"ML-KEM-768 DK size mismatch: got {}, expected {SK_LEN}",
sk_bytes.len()
);
Ok((
PublicKey(pk_bytes),
SecretKey(std::mem::take(&mut *sk_bytes)),
))
}
pub fn encapsulate(pk: &PublicKey) -> Result<(Ciphertext, SharedSecret)> {
if pk.0.len() != PK_LEN {
return Err(Error::InvalidLength {
expected: PK_LEN,
got: pk.0.len(),
});
}
let ek_enc: &ml_kem::Encoded<<MlKem768 as KemCore>::EncapsulationKey> =
pk.0.as_slice().try_into().map_err(|_| Error::Internal)?;
let ek = <MlKem768 as KemCore>::EncapsulationKey::from_bytes(ek_enc);
let mut m = random_b32();
let (ct, mut ss) = ek
.encapsulate_deterministic(&m)
.map_err(|_| Error::Internal)?;
m.zeroize();
let ct_bytes: &[u8] = ct.as_ref();
let mut ss_bytes = [0u8; 32];
ss_bytes.copy_from_slice(ss.as_ref());
ss.zeroize();
let result = SharedSecret(ss_bytes);
ss_bytes.zeroize();
Ok((Ciphertext(ct_bytes.to_vec()), result))
}
pub fn decapsulate(sk: &SecretKey, ct: &Ciphertext) -> Result<SharedSecret> {
use kem::Decapsulate;
if sk.0.len() != SK_LEN {
return Err(Error::InvalidLength {
expected: SK_LEN,
got: sk.0.len(),
});
}
if ct.0.len() != CT_LEN {
return Err(Error::InvalidLength {
expected: CT_LEN,
got: ct.0.len(),
});
}
let dk_enc: &ml_kem::Encoded<<MlKem768 as KemCore>::DecapsulationKey> =
sk.0.as_slice().try_into().map_err(|_| Error::Internal)?;
let dk = <MlKem768 as KemCore>::DecapsulationKey::from_bytes(dk_enc);
let ct_inner: ml_kem::Ciphertext<MlKem768> =
ct.0.as_slice()
.try_into()
.map_err(|_| Error::InvalidLength {
expected: CT_LEN,
got: ct.0.len(),
})?;
let mut ss = dk
.decapsulate(&ct_inner)
.map_err(|_| Error::DecapsulationFailed)?;
let mut ss_bytes = [0u8; 32];
ss_bytes.copy_from_slice(ss.as_ref());
ss.zeroize();
let result = SharedSecret(ss_bytes);
ss_bytes.zeroize();
Ok(result)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
#[test]
fn keygen_sizes() {
let (pk, sk) = keygen().unwrap();
assert_eq!(pk.as_bytes().len(), 1184);
assert_eq!(sk.as_bytes().len(), 2400);
}
#[test]
fn round_trip() {
let (pk, sk) = keygen().unwrap();
let (ct, ss_enc) = encapsulate(&pk).unwrap();
let ss_dec = decapsulate(&sk, &ct).unwrap();
assert_eq!(ss_enc.as_bytes(), ss_dec.as_bytes());
}
#[test]
fn wrong_pk_size() {
assert!(matches!(
PublicKey::from_bytes(vec![0u8; 100]),
Err(Error::InvalidLength {
expected: 1184,
got: 100
})
));
}
#[test]
fn wrong_ct_size() {
assert!(matches!(
Ciphertext::from_bytes(vec![0u8; 100]),
Err(Error::InvalidLength {
expected: 1088,
got: 100
})
));
}
#[test]
fn tampered_ct() {
let (pk, sk) = keygen().unwrap();
let (ct, ss_enc) = encapsulate(&pk).unwrap();
let mut bad_ct_bytes = ct.as_bytes().to_vec();
bad_ct_bytes[0] ^= 0xFF;
let bad_ct = Ciphertext::from_bytes_unchecked(bad_ct_bytes);
let ss_dec = decapsulate(&sk, &bad_ct).unwrap();
assert_ne!(ss_enc.as_bytes(), ss_dec.as_bytes());
}
#[test]
fn round_trip_repeated() {
for _ in 0..256 {
let (pk, sk) = keygen().unwrap();
let (ct, ss_enc) = encapsulate(&pk).unwrap();
let ss_dec = decapsulate(&sk, &ct).unwrap();
assert_eq!(ss_enc.as_bytes(), ss_dec.as_bytes());
}
}
}