use ml_kem::array::Array;
use ml_kem::{Decapsulate, Encapsulate};
use ml_kem::{DecapsulationKey, EncapsulationKey, Kem, KeyExport, MlKem768};
pub const ENCAPS_KEY_LEN: usize = 1184;
pub const CIPHERTEXT_LEN: usize = 1088;
pub const SEED_LEN: usize = 64;
pub const SHARED_SECRET_LEN: usize = 32;
pub struct MlKem768Keypair {
dk: DecapsulationKey<MlKem768>,
ek: EncapsulationKey<MlKem768>,
}
impl MlKem768Keypair {
pub fn generate() -> Self {
let (dk, ek) = MlKem768::generate_keypair();
Self { dk, ek }
}
pub fn from_seed(seed: [u8; SEED_LEN]) -> Self {
let dk = DecapsulationKey::<MlKem768>::from_seed(Array::from(seed));
let ek = dk.encapsulation_key().clone();
Self { dk, ek }
}
pub fn seed(&self) -> Option<[u8; SEED_LEN]> {
self.dk.to_seed().map(|s| {
let mut out = [0u8; SEED_LEN];
out.copy_from_slice(s.as_slice());
out
})
}
pub fn encapsulation_key(&self) -> Vec<u8> {
self.ek.to_bytes().as_slice().to_vec()
}
pub fn decapsulate(&self, ciphertext: &[u8]) -> Result<[u8; SHARED_SECRET_LEN], &'static str> {
let ss = self
.dk
.decapsulate_slice(ciphertext)
.map_err(|_| "ciphertext wrong length")?;
let mut out = [0u8; SHARED_SECRET_LEN];
out.copy_from_slice(ss.as_slice());
Ok(out)
}
}
pub fn encapsulate(encaps_key: &[u8]) -> Result<(Vec<u8>, [u8; SHARED_SECRET_LEN]), &'static str> {
let ek_arr = Array::try_from(encaps_key).map_err(|_| "encapsulation key wrong length")?;
let ek = EncapsulationKey::<MlKem768>::new(&ek_arr).map_err(|_| "invalid encapsulation key")?;
let (ct, ss) = ek.encapsulate();
let mut secret = [0u8; SHARED_SECRET_LEN];
secret.copy_from_slice(ss.as_slice());
Ok((ct.as_slice().to_vec(), secret))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip() {
let kp = MlKem768Keypair::generate();
let ek = kp.encapsulation_key();
assert_eq!(ek.len(), ENCAPS_KEY_LEN);
let (ct, ss_send) = encapsulate(&ek).expect("encapsulate");
assert_eq!(ct.len(), CIPHERTEXT_LEN);
let ss_recv = kp.decapsulate(&ct).expect("decapsulate");
assert_eq!(
ss_send, ss_recv,
"both parties derive the same shared secret"
);
}
#[test]
fn seed_is_deterministic() {
let kp = MlKem768Keypair::generate();
let seed = kp.seed().expect("seed");
let kp2 = MlKem768Keypair::from_seed(seed);
assert_eq!(kp.encapsulation_key(), kp2.encapsulation_key());
}
#[test]
fn wrong_length_rejected() {
let kp = MlKem768Keypair::generate();
assert!(kp.decapsulate(&[0u8; 10]).is_err());
assert!(encapsulate(&[0u8; 10]).is_err());
}
}