use crate::{QsshError, Result};
use kem::{Decapsulate, Encapsulate};
use ml_kem::{EncodedSizeUser, KemCore, MlKem768, MlKem1024};
use ml_kem::kem::{DecapsulationKey, EncapsulationKey};
use zeroize::{Zeroize, ZeroizeOnDrop};
pub mod mlkem768 {
pub const EK_SIZE: usize = 1184;
pub const DK_SIZE: usize = 2400;
pub const CT_SIZE: usize = 1088;
pub const SS_SIZE: usize = 32;
}
pub mod mlkem1024 {
pub const EK_SIZE: usize = 1568;
pub const DK_SIZE: usize = 3168;
pub const CT_SIZE: usize = 1568;
pub const SS_SIZE: usize = 32;
}
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct MlKem768KeyPair {
dk_bytes: Vec<u8>,
#[zeroize(skip)]
ek_bytes: Vec<u8>,
}
impl MlKem768KeyPair {
pub fn generate() -> Result<Self> {
let mut rng = rand::thread_rng();
let (dk, ek) = MlKem768::generate(&mut rng);
let ek_bytes = EncodedSizeUser::as_bytes(&ek).to_vec();
let dk_bytes = EncodedSizeUser::as_bytes(&dk).to_vec();
Ok(Self { dk_bytes, ek_bytes })
}
pub fn encapsulation_key(&self) -> &[u8] {
&self.ek_bytes
}
pub fn decapsulation_key(&self) -> &[u8] {
&self.dk_bytes
}
pub fn decapsulate(&self, ciphertext: &[u8]) -> Result<Vec<u8>> {
if ciphertext.len() != mlkem768::CT_SIZE {
return Err(QsshError::Crypto(format!(
"Invalid ML-KEM-768 ciphertext size: expected {}, got {}",
mlkem768::CT_SIZE,
ciphertext.len()
)));
}
let dk_array: ml_kem::Encoded<DecapsulationKey<ml_kem::MlKem768Params>> =
self.dk_bytes.as_slice().try_into().map_err(|_| {
QsshError::Crypto("Invalid ML-KEM-768 decapsulation key".into())
})?;
let dk = DecapsulationKey::<ml_kem::MlKem768Params>::from_bytes(&dk_array);
let ct_array: [u8; mlkem768::CT_SIZE] = ciphertext.try_into()
.map_err(|_| QsshError::Crypto("Invalid ML-KEM-768 ciphertext size".into()))?;
let ct = ml_kem::array::Array::from(ct_array);
let ss = dk.decapsulate(&ct).map_err(|_| {
QsshError::Crypto("ML-KEM-768 decapsulation failed".into())
})?;
Ok(ss.as_slice().to_vec())
}
}
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct MlKem1024KeyPair {
dk_bytes: Vec<u8>,
#[zeroize(skip)]
ek_bytes: Vec<u8>,
}
impl MlKem1024KeyPair {
pub fn generate() -> Result<Self> {
let mut rng = rand::thread_rng();
let (dk, ek) = MlKem1024::generate(&mut rng);
let ek_bytes = EncodedSizeUser::as_bytes(&ek).to_vec();
let dk_bytes = EncodedSizeUser::as_bytes(&dk).to_vec();
Ok(Self { dk_bytes, ek_bytes })
}
pub fn encapsulation_key(&self) -> &[u8] {
&self.ek_bytes
}
pub fn decapsulation_key(&self) -> &[u8] {
&self.dk_bytes
}
pub fn decapsulate(&self, ciphertext: &[u8]) -> Result<Vec<u8>> {
if ciphertext.len() != mlkem1024::CT_SIZE {
return Err(QsshError::Crypto(format!(
"Invalid ML-KEM-1024 ciphertext size: expected {}, got {}",
mlkem1024::CT_SIZE,
ciphertext.len()
)));
}
let dk_array: ml_kem::Encoded<DecapsulationKey<ml_kem::MlKem1024Params>> =
self.dk_bytes.as_slice().try_into().map_err(|_| {
QsshError::Crypto("Invalid ML-KEM-1024 decapsulation key".into())
})?;
let dk = DecapsulationKey::<ml_kem::MlKem1024Params>::from_bytes(&dk_array);
let ct_array: [u8; mlkem1024::CT_SIZE] = ciphertext.try_into()
.map_err(|_| QsshError::Crypto("Invalid ML-KEM-1024 ciphertext size".into()))?;
let ct = ml_kem::array::Array::from(ct_array);
let ss = dk.decapsulate(&ct).map_err(|_| {
QsshError::Crypto("ML-KEM-1024 decapsulation failed".into())
})?;
Ok(ss.as_slice().to_vec())
}
}
pub fn mlkem768_encapsulate(ek_bytes: &[u8]) -> Result<(Vec<u8>, Vec<u8>)> {
if ek_bytes.len() != mlkem768::EK_SIZE {
return Err(QsshError::Crypto(format!(
"Invalid ML-KEM-768 encapsulation key size: expected {}, got {}",
mlkem768::EK_SIZE,
ek_bytes.len()
)));
}
let mut rng = rand::thread_rng();
let ek_array: ml_kem::Encoded<EncapsulationKey<ml_kem::MlKem768Params>> =
ek_bytes.try_into().map_err(|_| {
QsshError::Crypto("Invalid ML-KEM-768 encapsulation key".into())
})?;
let ek = EncapsulationKey::<ml_kem::MlKem768Params>::from_bytes(&ek_array);
let (ct, ss) = ek.encapsulate(&mut rng).map_err(|_| {
QsshError::Crypto("ML-KEM-768 encapsulation failed".into())
})?;
Ok((ss.as_slice().to_vec(), ct.as_slice().to_vec()))
}
pub fn mlkem1024_encapsulate(ek_bytes: &[u8]) -> Result<(Vec<u8>, Vec<u8>)> {
if ek_bytes.len() != mlkem1024::EK_SIZE {
return Err(QsshError::Crypto(format!(
"Invalid ML-KEM-1024 encapsulation key size: expected {}, got {}",
mlkem1024::EK_SIZE,
ek_bytes.len()
)));
}
let mut rng = rand::thread_rng();
let ek_array: ml_kem::Encoded<EncapsulationKey<ml_kem::MlKem1024Params>> =
ek_bytes.try_into().map_err(|_| {
QsshError::Crypto("Invalid ML-KEM-1024 encapsulation key".into())
})?;
let ek = EncapsulationKey::<ml_kem::MlKem1024Params>::from_bytes(&ek_array);
let (ct, ss) = ek.encapsulate(&mut rng).map_err(|_| {
QsshError::Crypto("ML-KEM-1024 encapsulation failed".into())
})?;
Ok((ss.as_slice().to_vec(), ct.as_slice().to_vec()))
}
pub fn derive_session_material(
shared_secret: &[u8],
client_random: &[u8],
server_random: &[u8],
) -> Vec<u8> {
use sha3::{Sha3_256, Digest};
let mut hasher = Sha3_256::new();
hasher.update(b"QSSH-MLKEM-v1");
hasher.update(shared_secret);
hasher.update(client_random);
hasher.update(server_random);
hasher.finalize().to_vec()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mlkem768_roundtrip() {
let keypair = MlKem768KeyPair::generate().unwrap();
let (ss_sender, ct) = mlkem768_encapsulate(keypair.encapsulation_key()).unwrap();
let ss_receiver = keypair.decapsulate(&ct).unwrap();
assert_eq!(ss_sender, ss_receiver);
assert_eq!(ss_sender.len(), mlkem768::SS_SIZE);
}
#[test]
fn test_mlkem1024_roundtrip() {
let keypair = MlKem1024KeyPair::generate().unwrap();
let (ss_sender, ct) = mlkem1024_encapsulate(keypair.encapsulation_key()).unwrap();
let ss_receiver = keypair.decapsulate(&ct).unwrap();
assert_eq!(ss_sender, ss_receiver);
assert_eq!(ss_receiver.len(), mlkem1024::SS_SIZE);
}
#[test]
fn test_mlkem768_key_sizes() {
let keypair = MlKem768KeyPair::generate().unwrap();
assert_eq!(keypair.encapsulation_key().len(), mlkem768::EK_SIZE);
assert_eq!(keypair.decapsulation_key().len(), mlkem768::DK_SIZE);
}
#[test]
fn test_mlkem1024_key_sizes() {
let keypair = MlKem1024KeyPair::generate().unwrap();
assert_eq!(keypair.encapsulation_key().len(), mlkem1024::EK_SIZE);
assert_eq!(keypair.decapsulation_key().len(), mlkem1024::DK_SIZE);
}
#[test]
fn test_invalid_ciphertext_size() {
let keypair = MlKem768KeyPair::generate().unwrap();
let bad_ct = vec![0u8; 100]; let result = keypair.decapsulate(&bad_ct);
assert!(result.is_err());
}
#[test]
fn test_invalid_ek_size() {
let bad_ek = vec![0u8; 100]; let result = mlkem768_encapsulate(&bad_ek);
assert!(result.is_err());
}
#[test]
fn test_session_material_derivation() {
let client_random = [0x11u8; 32];
let server_random = [0x22u8; 32];
let shared_secret = [0x33u8; 32];
let material1 = derive_session_material(&shared_secret, &client_random, &server_random);
let material2 = derive_session_material(&shared_secret, &client_random, &server_random);
assert_eq!(material1, material2);
assert_eq!(material1.len(), 32);
let material3 = derive_session_material(&shared_secret, &server_random, &client_random);
assert_ne!(material1, material3);
}
}
#[cfg(kani)]
mod kani_proofs {
use super::*;
#[kani::proof]
fn proof_mlkem768_decapsulate_no_panic() {
let ct: [u8; mlkem768::CT_SIZE] = [0u8; mlkem768::CT_SIZE];
let slice: &[u8] = &ct;
assert_eq!(slice.len(), mlkem768::CT_SIZE);
let result: core::result::Result<[u8; mlkem768::CT_SIZE], _> = slice.try_into();
assert!(result.is_ok());
}
#[kani::proof]
fn proof_mlkem1024_decapsulate_no_panic() {
let ct: [u8; mlkem1024::CT_SIZE] = [0u8; mlkem1024::CT_SIZE];
let slice: &[u8] = &ct;
assert_eq!(slice.len(), mlkem1024::CT_SIZE);
let result: core::result::Result<[u8; mlkem1024::CT_SIZE], _> = slice.try_into();
assert!(result.is_ok());
}
#[kani::proof]
fn proof_mlkem768_encapsulate_no_panic() {
let ek_len: usize = kani::any();
kani::assume(ek_len <= 2048);
let would_reject = ek_len != mlkem768::EK_SIZE;
if !would_reject {
assert_eq!(ek_len, 1184); }
}
#[kani::proof]
fn proof_mlkem1024_encapsulate_no_panic() {
let ek_len: usize = kani::any();
kani::assume(ek_len <= 2048);
let would_reject = ek_len != mlkem1024::EK_SIZE;
if !would_reject {
assert_eq!(ek_len, 1568); }
}
}