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+ } ;
47use serde:: Deserialize ;
58use std:: fs:: File ;
69
@@ -33,13 +36,18 @@ struct Source {
3336struct 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+
6885macro_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+
92161mlkem_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