use ledger_secure_sdk_sys::*;
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum MlKemParam {
MlKem512,
MlKem768,
MlKem1024,
}
impl MlKemParam {
const fn as_c(self) -> MLKEM_param_t {
match self {
MlKemParam::MlKem512 => MLKEM_512,
MlKemParam::MlKem768 => MLKEM_768,
MlKemParam::MlKem1024 => MLKEM_1024,
}
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum MlKemError {
InvalidParameter,
InvalidParameterValue,
InternalError,
}
impl From<u32> for MlKemError {
fn from(code: u32) -> Self {
match code {
CX_INVALID_PARAMETER => MlKemError::InvalidParameter,
CX_INVALID_PARAMETER_VALUE => MlKemError::InvalidParameterValue,
_ => MlKemError::InternalError,
}
}
}
pub const SHARED_SECRET_LEN: usize = MLKEM_SSBYTES as usize;
pub const MLKEM512_PK_LEN: usize = MLKEM512_PUBLICKEYBYTES as usize;
pub const MLKEM512_SK_LEN: usize = MLKEM512_SECRETKEYBYTES as usize;
pub const MLKEM512_CT_LEN: usize = MLKEM512_CIPHERTEXTBYTES as usize;
pub const MLKEM768_PK_LEN: usize = MLKEM768_PUBLICKEYBYTES as usize;
pub const MLKEM768_SK_LEN: usize = MLKEM768_SECRETKEYBYTES as usize;
pub const MLKEM768_CT_LEN: usize = MLKEM768_CIPHERTEXTBYTES as usize;
pub const MLKEM1024_PK_LEN: usize = MLKEM1024_PUBLICKEYBYTES as usize;
pub const MLKEM1024_SK_LEN: usize = MLKEM1024_SECRETKEYBYTES as usize;
pub const MLKEM1024_CT_LEN: usize = MLKEM1024_CIPHERTEXTBYTES as usize;
impl MlKemParam {
pub const fn pk_len(self) -> usize {
match self {
MlKemParam::MlKem512 => MLKEM512_PK_LEN,
MlKemParam::MlKem768 => MLKEM768_PK_LEN,
MlKemParam::MlKem1024 => MLKEM1024_PK_LEN,
}
}
pub const fn sk_len(self) -> usize {
match self {
MlKemParam::MlKem512 => MLKEM512_SK_LEN,
MlKemParam::MlKem768 => MLKEM768_SK_LEN,
MlKemParam::MlKem1024 => MLKEM1024_SK_LEN,
}
}
pub const fn ct_len(self) -> usize {
match self {
MlKemParam::MlKem512 => MLKEM512_CT_LEN,
MlKemParam::MlKem768 => MLKEM768_CT_LEN,
MlKemParam::MlKem1024 => MLKEM1024_CT_LEN,
}
}
}
pub fn keypair(pk: &mut [u8], sk: &mut [u8], param: MlKemParam) -> Result<(), MlKemError> {
let err = unsafe {
MLKEM_crypto_kem_keypair(
pk.as_mut_ptr(),
pk.len(),
sk.as_mut_ptr(),
sk.len(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(())
}
}
pub fn encapsulate(
ct: &mut [u8],
ss: &mut [u8; SHARED_SECRET_LEN],
pk: &[u8],
param: MlKemParam,
) -> Result<(), MlKemError> {
let err = unsafe {
MLKEM_crypto_kem_enc(
ct.as_mut_ptr(),
ct.len(),
ss.as_mut_ptr(),
ss.len(),
pk.as_ptr(),
pk.len(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(())
}
}
pub fn decapsulate(
ss: &mut [u8; SHARED_SECRET_LEN],
ct: &[u8],
sk: &[u8],
param: MlKemParam,
) -> Result<(), MlKemError> {
let err = unsafe {
MLKEM_crypto_kem_dec(
ss.as_mut_ptr(),
ss.len(),
ct.as_ptr(),
ct.len(),
sk.as_ptr(),
sk.len(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::assert_eq_err as assert_eq;
use crate::testing::TestType;
use testmacro::test_item as test;
#[test]
fn test_mlkem512_keygen_encaps_decaps() {
let mut pk = [0u8; MLKEM512_PK_LEN];
let mut sk = [0u8; MLKEM512_SK_LEN];
keypair(&mut pk, &mut sk, MlKemParam::MlKem512).unwrap();
let mut ct = [0u8; MLKEM512_CT_LEN];
let mut ss_enc = [0u8; SHARED_SECRET_LEN];
encapsulate(&mut ct, &mut ss_enc, &pk, MlKemParam::MlKem512).unwrap();
let mut ss_dec = [0u8; SHARED_SECRET_LEN];
decapsulate(&mut ss_dec, &ct, &sk, MlKemParam::MlKem512).unwrap();
assert_eq!(&ss_enc, &ss_dec);
}
#[test]
fn test_mlkem768_keygen_encaps_decaps() {
let mut pk = [0u8; MLKEM768_PK_LEN];
let mut sk = [0u8; MLKEM768_SK_LEN];
keypair(&mut pk, &mut sk, MlKemParam::MlKem768).unwrap();
let mut ct = [0u8; MLKEM768_CT_LEN];
let mut ss_enc = [0u8; SHARED_SECRET_LEN];
encapsulate(&mut ct, &mut ss_enc, &pk, MlKemParam::MlKem768).unwrap();
let mut ss_dec = [0u8; SHARED_SECRET_LEN];
decapsulate(&mut ss_dec, &ct, &sk, MlKemParam::MlKem768).unwrap();
assert_eq!(&ss_enc, &ss_dec);
}
#[test]
fn test_mlkem1024_keygen_encaps_decaps() {
let mut pk = [0u8; MLKEM1024_PK_LEN];
let mut sk = [0u8; MLKEM1024_SK_LEN];
keypair(&mut pk, &mut sk, MlKemParam::MlKem1024).unwrap();
let mut ct = [0u8; MLKEM1024_CT_LEN];
let mut ss_enc = [0u8; SHARED_SECRET_LEN];
encapsulate(&mut ct, &mut ss_enc, &pk, MlKemParam::MlKem1024).unwrap();
let mut ss_dec = [0u8; SHARED_SECRET_LEN];
decapsulate(&mut ss_dec, &ct, &sk, MlKemParam::MlKem1024).unwrap();
assert_eq!(&ss_enc, &ss_dec);
}
}