use crate::error::CryptError;
use crate::kem::backend::{rand_core_010, KemBackend, KemId};
use crate::kem::types::{KemCiphertext, KemSharedSecret, KemSize, MlKemPublicKey, MlKemSecretKey};
use ml_kem::{
kem::{Decapsulate, Encapsulate, FromSeed, Kem, KeyExport, TryKeyInit},
MlKem1024, MlKem512, MlKem768,
};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Size512;
impl KemSize for Size512 {}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Size768;
impl KemSize for Size768 {}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Size1024;
impl KemSize for Size1024 {}
#[derive(Clone, Copy, Debug, Default)]
pub struct MlKem512Impl;
#[derive(Clone, Copy, Debug, Default)]
pub struct MlKem768Impl;
#[derive(Clone, Copy, Debug, Default)]
pub struct MlKem1024Impl;
macro_rules! impl_ml_kem {
($impl_ty:ty, $kem_ty:ty, $size_ty:ty, $kem_id:expr) => {
impl KemBackend for $impl_ty {
type Size = $size_ty;
type PublicKey = MlKemPublicKey<$size_ty>;
type SecretKey = MlKemSecretKey<$size_ty>;
type Ciphertext = KemCiphertext;
type SharedSecret = KemSharedSecret;
const ID: KemId = $kem_id;
fn keypair(
rng: &mut impl rand_core_010::CryptoRng,
) -> Result<(Self::PublicKey, Self::SecretKey), CryptError> {
let (dk, ek) = <$kem_ty>::generate_keypair_from_rng(rng);
let ek_bytes = ek.to_bytes();
let pk = MlKemPublicKey::from_bytes(ek_bytes.as_slice().to_vec());
let seed = dk.to_seed().ok_or(CryptError::EncapsulationError)?;
let sk = MlKemSecretKey::from_bytes(seed.as_slice().to_vec());
Ok((pk, sk))
}
fn encapsulate(
pk: &Self::PublicKey,
rng: &mut impl rand_core_010::CryptoRng,
) -> Result<(Self::Ciphertext, Self::SharedSecret), CryptError> {
type EK = <$kem_ty as Kem>::EncapsulationKey;
let ek =
EK::new_from_slice(pk.as_ref()).map_err(|_| CryptError::InvalidKemPublicKey)?;
let (ct, ss) = ek.encapsulate_with_rng(rng);
Ok((
KemCiphertext::from_bytes(ct.as_slice().to_vec()),
KemSharedSecret::from_bytes(ss.as_slice().to_vec()),
))
}
fn decapsulate(
sk: &Self::SecretKey,
ct: &Self::Ciphertext,
) -> Result<Self::SharedSecret, CryptError> {
let seed_arr = <ml_kem::kem::Seed<$kem_ty>>::try_from(sk.as_ref())
.map_err(|_| CryptError::InvalidKemSecretKey)?;
let dk = <$kem_ty>::from_seed(&seed_arr).0;
let ct_arr = <ml_kem::kem::Ciphertext<$kem_ty>>::try_from(ct.as_ref())
.map_err(|_| CryptError::InvalidKemCiphertext)?;
let ss = dk.decapsulate(&ct_arr);
Ok(KemSharedSecret::from_bytes(ss.as_slice().to_vec()))
}
}
};
}
impl_ml_kem!(MlKem512Impl, MlKem512, Size512, KemId::MlKem512);
impl_ml_kem!(MlKem768Impl, MlKem768, Size768, KemId::MlKem768);
impl_ml_kem!(MlKem1024Impl, MlKem1024, Size1024, KemId::MlKem1024);