Skip to content

Commit d8c7434

Browse files
authored
ml-kem: add Wycheproof mlkem_*_encaps_test (#214)
Tests decoding of encapsulation keys, and that they are able to generate the correct ciphertext and shared secret via the `EncapsulationKey::encapsulate_deterministic` API (when the `hazmat` feature is enabled)
1 parent b2fc552 commit d8c7434

2 files changed

Lines changed: 95 additions & 10 deletions

File tree

ml-kem/Cargo.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,10 +19,10 @@ exclude = ["tests/key-gen.rs", "tests/key-gen.json", "tests/encap-decap.rs", "te
1919
alloc = ["pkcs8?/alloc"]
2020

2121
getrandom = ["kem/getrandom"]
22+
hazmat = []
2223
pem = ["pkcs8/pem"]
2324
pkcs8 = ["dep:const-oid", "dep:pkcs8"]
2425
zeroize = ["module-lattice/zeroize", "dep:zeroize"]
25-
hazmat = []
2626

2727
[dependencies]
2828
array = { package = "hybrid-array", version = "0.4.4", features = ["extra-sizes", "subtle"] }

ml-kem/tests/wycheproof.rs

Lines changed: 94 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,9 @@
11
//! Test against the Wycheproof test vectors.
22
3-
use ml_kem::{EncodedSizeUser, KemCore, MlKem512, MlKem768, MlKem1024, kem::KeyExport};
3+
use ml_kem::{
4+
EncodedSizeUser, KemCore, MlKem512, MlKem768, MlKem1024,
5+
kem::{KeyExport, TryKeyInit},
6+
};
47
use serde::Deserialize;
58
use std::fs::File;
69

@@ -33,13 +36,18 @@ struct Source {
3336
struct Test {
3437
#[serde(rename(deserialize = "tcId"))]
3538
id: usize,
36-
comment: String,
37-
#[serde(with = "hex::serde")]
38-
seed: Vec<u8>,
39+
comment: Option<String>,
40+
seed: Option<String>,
3941
#[serde(default, with = "hex::serde")]
4042
ek: Vec<u8>,
41-
#[serde(with = "hex::serde")]
42-
dk: Vec<u8>,
43+
dk: Option<String>,
44+
#[cfg(feature = "hazmat")]
45+
m: Option<String>,
46+
#[cfg(feature = "hazmat")]
47+
c: Option<String>,
48+
#[cfg(feature = "hazmat")]
49+
#[serde(default, rename(deserialize = "K"))]
50+
k: Option<String>,
4351
result: ExpectedResult,
4452
}
4553

@@ -65,6 +73,15 @@ macro_rules! load_json_file {
6573
}};
6674
}
6775

76+
fn decode_optional_hex(opt: &Option<String>, field: &str) -> Vec<u8> {
77+
match opt {
78+
Some(h) => {
79+
hex::decode(h).unwrap_or_else(|e| panic!("invalid hex for field '{field}': {e}"))
80+
}
81+
None => panic!("missing field: {field}"),
82+
}
83+
}
84+
6885
macro_rules! mlkem_keygen_seed_test {
6986
($name:ident, $json_file:expr, $kem:ident) => {
7087
#[test]
@@ -78,17 +95,69 @@ macro_rules! mlkem_keygen_seed_test {
7895
);
7996

8097
for test in &group.tests {
81-
println!("Test #{}: {} ({:?})", test.id, &test.comment, &test.result);
98+
println!(
99+
"Test #{}: {} ({:?})",
100+
test.id,
101+
test.comment.as_ref().unwrap(),
102+
&test.result
103+
);
104+
let test_seed = decode_optional_hex(&test.seed, "seed");
105+
let test_dk = decode_optional_hex(&test.dk, "dk");
82106

83-
let (dk, ek) = $kem::from_seed(test.seed.as_slice().try_into().unwrap());
84-
assert_eq!(test.dk.as_slice(), dk.to_encoded_bytes().as_slice());
107+
let (dk, ek) = $kem::from_seed(test_seed.as_slice().try_into().unwrap());
108+
assert_eq!(test_dk.as_slice(), dk.to_encoded_bytes().as_slice());
85109
assert_eq!(test.ek.as_slice(), ek.to_bytes().as_slice());
86110
}
87111
}
88112
}
89113
};
90114
}
91115

116+
macro_rules! mlkem_encaps_test {
117+
($name:ident, $json_file:expr, $kem_module:ident) => {
118+
#[test]
119+
fn $name() {
120+
let tests = load_json_file!($json_file);
121+
122+
for group in tests.groups {
123+
println!(
124+
"Parameter set: {} ({} v{})\n",
125+
&group.parameter_set, &group.source.name, &group.source.version
126+
);
127+
128+
for test in &group.tests {
129+
println!("Test #{} ({:?})", test.id, &test.result);
130+
131+
use ml_kem::$kem_module::EncapsulationKey;
132+
let ek_result = EncapsulationKey::new_from_slice(&test.ek);
133+
134+
#[cfg_attr(not(feature = "hazmat"), allow(unused_variables))]
135+
let ek = match test.result {
136+
ExpectedResult::Valid => ek_result.expect("should be valid"),
137+
ExpectedResult::Invalid => {
138+
assert!(ek_result.is_err());
139+
continue;
140+
}
141+
other => todo!("{:?}", other),
142+
};
143+
144+
#[cfg(feature = "hazmat")]
145+
{
146+
let test_m = decode_optional_hex(&test.m, "m");
147+
let test_m = test_m.as_slice().try_into().unwrap();
148+
let (c, k) = ek.encapsulate_deterministic(test_m);
149+
150+
let test_c = decode_optional_hex(&test.c, "c");
151+
let test_k = decode_optional_hex(&test.k, "K");
152+
assert_eq!(test_c.as_slice(), c.as_slice());
153+
assert_eq!(test_k.as_slice(), k.as_slice());
154+
}
155+
}
156+
}
157+
}
158+
};
159+
}
160+
92161
mlkem_keygen_seed_test!(
93162
mlkem_512_keygen_seed_test,
94163
"mlkem_512_keygen_seed_test.json",
@@ -104,3 +173,19 @@ mlkem_keygen_seed_test!(
104173
"mlkem_1024_keygen_seed_test.json",
105174
MlKem1024
106175
);
176+
177+
mlkem_encaps_test!(
178+
mlkem_512_encaps_test,
179+
"mlkem_512_encaps_test.json",
180+
ml_kem_512
181+
);
182+
mlkem_encaps_test!(
183+
mlkem_768_encaps_test,
184+
"mlkem_768_encaps_test.json",
185+
ml_kem_768
186+
);
187+
mlkem_encaps_test!(
188+
mlkem_1024_encaps_test,
189+
"mlkem_1024_encaps_test.json",
190+
ml_kem_1024
191+
);

0 commit comments

Comments
 (0)