use crate::{
Deserializable, HpkeError, Serializable,
aead::{Aead, AesGcm128, AesGcm256, ChaCha20Poly1305, ExportOnlyAead},
kdf::{
HkdfSha256, HkdfSha384, HkdfSha512, Kdf as KdfTrait, KdfShake128, KdfShake256,
KdfTurboShake128, KdfTurboShake256,
},
kem::{
DhP256HkdfSha256, DhP384HkdfSha384, DhP521HkdfSha512, Kem as KemTrait, MlKem768,
MlKem768P256, MlKem1024, MlKem1024P384, SharedSecret, X25519HkdfSha256, XWing,
},
op_mode::{OpModeR, PskBundle},
setup::setup_receiver,
};
use std::{fs::File, string::String, vec::Vec};
use ml_kem::KeyExport;
use serde::{Deserialize, Deserializer, de::Error as SError};
pub(crate) trait TestableKem: KemTrait {
type EphemeralKey: Deserializable;
#[doc(hidden)]
fn encap_with_eph(
pk_recip: &Self::PublicKey,
sender_id_keypair: Option<(&Self::PrivateKey, &Self::PublicKey)>,
sk_eph: Self::EphemeralKey,
) -> Result<(SharedSecret<Self>, Self::EncappedKey), HpkeError>;
#[doc(hidden)]
fn encap_det(
pk_recip: &Self::PublicKey,
sender_id_keypair: Option<(&Self::PrivateKey, &Self::PublicKey)>,
randomness: &[u8],
) -> Result<(SharedSecret<Self>, Self::EncappedKey), HpkeError>;
}
macro_rules! assert_serializable_eq {
($a:expr, $b:expr, $args:tt) => {
assert_eq!($a.to_bytes(), $b.to_bytes(), $args)
};
}
fn bytes_from_hex<'de, D>(deserializer: D) -> Result<Vec<u8>, D::Error>
where
D: Deserializer<'de>,
{
let mut hex_str = String::deserialize(deserializer)?;
if hex_str.len() % 2 == 1 {
hex_str.insert(0, '0');
}
hex::decode(hex_str).map_err(|e| SError::custom(format!("{:?}", e)))
}
fn bytes_from_hex_opt<'de, D>(deserializer: D) -> Result<Option<Vec<u8>>, D::Error>
where
D: Deserializer<'de>,
{
bytes_from_hex(deserializer).map(Some)
}
#[derive(Clone, serde::Deserialize, Debug)]
struct MainTestVector {
mode: u8,
kem_id: u16,
kdf_id: u16,
aead_id: u16,
#[serde(deserialize_with = "bytes_from_hex")]
info: Vec<u8>,
#[serde(rename = "ikmR", deserialize_with = "bytes_from_hex")]
ikm_recip: Vec<u8>,
#[serde(default, rename = "ikmS", deserialize_with = "bytes_from_hex_opt")]
ikm_sender: Option<Vec<u8>>,
#[serde(rename = "ikmE", deserialize_with = "bytes_from_hex")]
ikm_eph: Vec<u8>,
#[serde(rename = "skRm", deserialize_with = "bytes_from_hex")]
sk_recip: Vec<u8>,
#[serde(default, rename = "skSm", deserialize_with = "bytes_from_hex_opt")]
sk_sender: Option<Vec<u8>>,
#[serde(default, rename = "skEm", deserialize_with = "bytes_from_hex_opt")]
sk_eph: Option<Vec<u8>>,
#[serde(default, deserialize_with = "bytes_from_hex_opt")]
psk: Option<Vec<u8>>,
#[serde(default, rename = "psk_id", deserialize_with = "bytes_from_hex_opt")]
psk_id: Option<Vec<u8>>,
#[serde(rename = "pkRm", deserialize_with = "bytes_from_hex")]
pk_recip: Vec<u8>,
#[serde(default, rename = "pkSm", deserialize_with = "bytes_from_hex_opt")]
pk_sender: Option<Vec<u8>>,
#[serde(default, rename = "pkEm", deserialize_with = "bytes_from_hex_opt")]
_pk_eph: Option<Vec<u8>>,
#[serde(rename = "enc", deserialize_with = "bytes_from_hex")]
encapped_key: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
shared_secret: Vec<u8>,
#[serde(
default,
rename = "key_schedule_context",
deserialize_with = "bytes_from_hex_opt"
)]
_hpke_context: Option<Vec<u8>>,
#[serde(default, rename = "secret", deserialize_with = "bytes_from_hex_opt")]
_key_schedule_secret: Option<Vec<u8>>,
#[serde(rename = "key", deserialize_with = "bytes_from_hex")]
_aead_key: Vec<u8>,
#[serde(rename = "base_nonce", deserialize_with = "bytes_from_hex")]
_aead_base_nonce: Vec<u8>,
#[serde(rename = "exporter_secret", deserialize_with = "bytes_from_hex")]
_exporter_secret: Vec<u8>,
encryptions: Vec<EncryptionTestVector>,
exports: Vec<ExporterTestVector>,
}
#[derive(Clone, serde::Deserialize, Debug)]
struct EncryptionTestVector {
#[serde(rename = "pt", deserialize_with = "bytes_from_hex")]
plaintext: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
aad: Vec<u8>,
#[serde(rename = "nonce", deserialize_with = "bytes_from_hex")]
_nonce: Vec<u8>,
#[serde(rename = "ct", deserialize_with = "bytes_from_hex")]
ciphertext: Vec<u8>,
}
#[derive(Clone, serde::Deserialize, Debug)]
struct ExporterTestVector {
#[serde(rename = "exporter_context", deserialize_with = "bytes_from_hex")]
export_ctx: Vec<u8>,
#[serde(rename = "L")]
export_len: usize,
#[serde(rename = "exported_value", deserialize_with = "bytes_from_hex")]
export_val: Vec<u8>,
}
fn deser_keypair<Kem: KemTrait>(
sk_bytes: &[u8],
pk_bytes: &[u8],
) -> (Kem::PrivateKey, Kem::PublicKey) {
let sk = <Kem as KemTrait>::PrivateKey::from_bytes(sk_bytes).unwrap();
let pk = <Kem as KemTrait>::PublicKey::from_bytes(pk_bytes).unwrap();
(sk, pk)
}
fn make_op_mode_r<'a, Kem: KemTrait>(
mode_id: u8,
pk: Option<Kem::PublicKey>,
psk: Option<&'a [u8]>,
psk_id: Option<&'a [u8]>,
) -> OpModeR<'a, Kem> {
let bundle = psk.map(|bytes| PskBundle::new(bytes, psk_id.unwrap()).unwrap());
match mode_id {
0 => OpModeR::Base,
1 => OpModeR::Psk(bundle.unwrap()),
2 => OpModeR::Auth(pk.unwrap()),
3 => OpModeR::AuthPsk(pk.unwrap(), bundle.unwrap()),
_ => panic!("Invalid mode ID: {}", mode_id),
}
}
fn test_case<A: Aead, Kdf: KdfTrait, Kem: TestableKem>(tv: MainTestVector) {
let recip_keypair = deser_keypair::<Kem>(&tv.sk_recip, &tv.pk_recip);
let sender_keypair = {
let pk_sender = &tv.pk_sender.as_ref();
tv.sk_sender
.as_ref()
.map(|sk| deser_keypair::<Kem>(sk, pk_sender.unwrap()))
};
{
let derived_kp = Kem::derive_keypair(&tv.ikm_recip);
assert_serializable_eq!(recip_keypair.0, derived_kp.0, "sk recip doesn't match");
assert_serializable_eq!(recip_keypair.1, derived_kp.1, "pk recip doesn't match");
}
if let Some(kp) = sender_keypair.as_ref() {
let derived_kp = Kem::derive_keypair(&tv.ikm_sender.unwrap());
assert_serializable_eq!(kp.0, derived_kp.0, "sk sender doesn't match");
assert_serializable_eq!(kp.1, derived_kp.1, "pk sender doesn't match");
}
let (sk_recip, pk_recip) = recip_keypair;
let sender_keypair = sender_keypair.as_ref().map(|(sk, pk)| (sk, pk)); let (shared_secret, encapped_key) =
Kem::encap_det(&pk_recip, sender_keypair, tv.ikm_eph.as_slice()).expect("encap failed");
if let Some(sk_eph) = tv
.sk_eph
.map(|b| Kem::EphemeralKey::from_bytes(&b).unwrap())
{
let (other_shared_secret, other_encapped_key) =
Kem::encap_with_eph(&pk_recip, sender_keypair, sk_eph).expect("encap failed");
assert!(
shared_secret.0 == other_shared_secret.0,
"ikm shared secret doesn't match sk_eph shared secret"
);
assert_serializable_eq!(
encapped_key,
other_encapped_key,
"ikm encapped key doesn't match sk_eph encapped key"
);
}
assert_eq!(
shared_secret.0.as_slice(),
tv.shared_secret.as_slice(),
"shared_secret doesn't match"
);
{
let provided_encapped_key =
<Kem as KemTrait>::EncappedKey::from_bytes(&tv.encapped_key).unwrap();
assert_serializable_eq!(
encapped_key,
provided_encapped_key,
"encapped keys don't match"
);
}
let mode = make_op_mode_r(
tv.mode,
sender_keypair.map(|(_, pk)| pk.clone()),
tv.psk.as_deref(),
tv.psk_id.as_deref(),
);
let mut aead_ctx = setup_receiver::<A, Kdf, Kem>(&mode, &sk_recip, &encapped_key, &tv.info)
.expect("setup_receiver failed");
for enc_packet in tv.encryptions {
let EncryptionTestVector {
aad,
ciphertext,
plaintext,
..
} = enc_packet;
let decrypted = aead_ctx.open(&ciphertext, &aad).expect("open failed");
assert_eq!(decrypted, plaintext, "plaintexts don't match");
}
for export in tv.exports {
let mut exported_val = vec![0u8; export.export_len];
aead_ctx
.export(&export.export_ctx, &mut exported_val)
.unwrap();
assert_eq!(exported_val, export.export_val, "export values don't match");
}
}
macro_rules! dispatch_testcase {
($tv:ident, ($( $aead_ty:ty ),*), ($( $kdf_ty:ty ),*), ($( $kem_ty:ty ),*)) => {
dispatch_testcase!(@tup1 $tv, ($( $aead_ty ),*), ($( $kdf_ty ),*), ($( $kem_ty ),*))
};
(@tup1 $tv:ident, ($( $aead_ty:ty ),*), $kdf_tup:tt, $kem_tup:tt) => {
$(
dispatch_testcase!(@tup2 $tv, $aead_ty, $kdf_tup, $kem_tup);
)*
};
(@tup2 $tv:ident, $aead_ty:ty, ($( $kdf_ty:ty ),*), $kem_tup:tt) => {
$(
dispatch_testcase!(@tup3 $tv, $aead_ty, $kdf_ty, $kem_tup);
)*
};
(@tup3 $tv:ident, $aead_ty:ty, $kdf_ty:ty, ($( $kem_ty:ty ),*)) => {
$(
dispatch_testcase!(@base $tv, $aead_ty, $kdf_ty, $kem_ty);
)*
};
(@base $tv:ident, $aead_ty:ty, $kdf_ty:ty, $kem_ty:ty) => {
if let (<$aead_ty>::AEAD_ID, <$kdf_ty>::KDF_ID, <$kem_ty>::KEM_ID) =
($tv.aead_id, $tv.kdf_id, $tv.kem_id)
{
println!(
"Running test case on {}, {}, {}",
stringify!($aead_ty),
stringify!($kdf_ty),
stringify!($kem_ty)
);
let tv = $tv.clone();
test_case::<$aead_ty, $kdf_ty, $kem_ty>(tv);
continue;
}
};
}
#[test]
fn classical_pq_and_hybrid() {
let ref_tvs: Vec<MainTestVector> = {
let file = File::open("test-vectors/origrfc-5f503c5.json").unwrap();
serde_json::from_reader(file).unwrap()
};
let pq_tvs: Vec<MainTestVector> = {
let file = File::open("test-vectors/pq-6433c8f.json").unwrap();
serde_json::from_reader(file).unwrap()
};
for tv in ref_tvs.into_iter().chain(pq_tvs.into_iter()) {
if ![
X25519HkdfSha256::KEM_ID,
DhP256HkdfSha256::KEM_ID,
DhP384HkdfSha384::KEM_ID,
DhP521HkdfSha512::KEM_ID,
XWing::KEM_ID,
MlKem768P256::KEM_ID,
MlKem1024P384::KEM_ID,
MlKem768::KEM_ID,
MlKem1024::KEM_ID,
]
.contains(&tv.kem_id)
{
continue;
}
dispatch_testcase!(
tv,
(AesGcm128, AesGcm256, ChaCha20Poly1305, ExportOnlyAead),
(
HkdfSha256,
HkdfSha384,
HkdfSha512,
KdfShake128,
KdfShake256,
KdfTurboShake128,
KdfTurboShake256
),
(
X25519HkdfSha256,
DhP256HkdfSha256,
DhP384HkdfSha384,
DhP521HkdfSha512,
XWing,
MlKem768P256,
MlKem1024P384,
MlKem768,
MlKem1024
)
);
panic!(
"Unrecognized (AEAD ID, KDF ID, KEM ID) combo: ({:x}, {:x}, {:x})",
tv.aead_id, tv.kdf_id, tv.kem_id
);
}
}
#[derive(Clone, serde::Deserialize, Debug)]
struct HybridTestVector {
#[serde(deserialize_with = "bytes_from_hex")]
randomness: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
encapsulation_key: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
decapsulation_key: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
decapsulation_key_pq: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
decapsulation_key_t: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
ciphertext: Vec<u8>,
#[serde(deserialize_with = "bytes_from_hex")]
shared_secret: Vec<u8>,
}
#[derive(Clone, serde::Deserialize, Debug)]
struct HybridTestVectors {
mlkem768_p256: Vec<HybridTestVector>,
mlkem768_x25519: Vec<HybridTestVector>,
mlkem1024_p384: Vec<HybridTestVector>,
}
macro_rules! test_hybrid_vector {
($kem:ty, $tv:ident, unpack_dk) => {
test_hybrid_vector!($kem, $tv);
let dk = <$kem as KemTrait>::PrivateKey::from_bytes(&$tv.decapsulation_key).unwrap();
assert_eq!(
dk.dk_t.to_bytes().as_slice(),
$tv.decapsulation_key_t.as_slice()
);
assert_eq!(
dk.dk_pq.to_bytes().as_slice(),
$tv.decapsulation_key_pq.as_slice()
);
};
($kem:ty, $tv:ident) => {
let dk = <$kem as KemTrait>::PrivateKey::from_bytes(&$tv.decapsulation_key).unwrap();
let ek = <$kem as KemTrait>::PublicKey::from_bytes(&$tv.encapsulation_key).unwrap();
assert_eq!(<$kem>::sk_to_pk(&dk), ek);
let (ss, ct) = <$kem>::encap_det(&ek, None, &$tv.randomness).unwrap();
assert_eq!(ss.0.as_slice(), $tv.shared_secret.as_slice());
assert_eq!(ct.to_bytes().as_slice(), $tv.ciphertext);
let enc = <$kem as KemTrait>::EncappedKey::from_bytes(&$tv.ciphertext).unwrap();
let ss = <$kem>::decap(&dk, None, &enc).unwrap();
assert_eq!(ss.0.as_slice(), $tv.shared_secret.as_slice());
};
}
#[test]
fn hybrid() {
let hybrid_tvs: HybridTestVectors = {
let file = File::open("test-vectors/hybrid-defafa2.json").unwrap();
serde_json::from_reader(file).unwrap()
};
for tv in hybrid_tvs.mlkem768_p256 {
test_hybrid_vector!(MlKem768P256, tv, unpack_dk);
}
for tv in hybrid_tvs.mlkem1024_p384 {
test_hybrid_vector!(MlKem1024P384, tv, unpack_dk);
}
for tv in hybrid_tvs.mlkem768_x25519 {
test_hybrid_vector!(XWing, tv);
}
}