1mod module;
2mod ring;
3
4use std::vec;
5
6use module::Module;
7use rand::rngs::OsRng;
8use rand::RngCore;
9use sha3::{
10 digest::{ExtendableOutput, Update, XofReader},
11 Digest, Sha3_256, Sha3_512, Shake128, Shake256,
12};
13
14use ring::Ring;
15
16pub enum Type {
17 MlKem512,
18 MlKem768,
19 MlKem1024,
20}
21
22pub struct MLKem {
23 k: u8,
24 eta_1: u8,
25 eta_2: u8,
26 du: u8,
27 dv: u8,
28}
29
30impl MLKem {
31 pub const fn new(type_of: Type) -> Self {
32 match type_of {
33 Type::MlKem512 => MLKem {
34 k: 2,
35 eta_1: 3,
36 eta_2: 2,
37 du: 10,
38 dv: 4,
39 },
40 Type::MlKem768 => MLKem {
41 k: 3,
42 eta_1: 2,
43 eta_2: 2,
44 du: 10,
45 dv: 4,
46 },
47 Type::MlKem1024 => MLKem {
48 k: 4,
49 eta_1: 2,
50 eta_2: 2,
51 du: 11,
52 dv: 5,
53 },
54 }
55 }
56
57 pub fn keygen(&self) -> (Vec<u8>, Vec<u8>) {
58 let d = Self::random_bytes(32);
59 let z = Self::random_bytes(32);
60
61 let (ek, dk) = self._keygen_internal(&d, &z);
62
63 (ek, dk)
64 }
65
66 pub fn encaps(&self, ek: &[u8]) -> (Vec<u8>, Vec<u8>) {
67 let m = Self::random_bytes(32);
68 let (k, c) = self._encaps_internal(ek, &m);
69 (k, c)
70 }
71
72 pub fn decaps(&self, dk: &[u8], c: &[u8]) -> Vec<u8> {
73 self._decaps_internal(dk, c).unwrap()
74 }
75
76 fn _decaps_internal(&self, dk: &[u8], c: &[u8]) -> Result<Vec<u8>, String> {
77 if c.len() != 32 * (self.du * self.k + self.dv) as usize {
78 return Err(String::from("ciphertext type check failed"));
79 }
80 if dk.len() != (768_usize * self.k as usize + 96) {
81 return Err(String::from("decapsulation key type check failed"));
82 }
83
84 let dk_pke = &dk[0..(384_usize * self.k as usize)];
85 let ek_pke = &dk[(384_usize * self.k as usize)..(768_usize * self.k as usize + 32)];
86 let h = &dk[(768_usize * self.k as usize + 32)..(768_usize * self.k as usize + 64)];
87 let z = &dk[(768_usize * self.k as usize + 64)..];
88
89 if Self::_h(ek_pke) != h {
90 return Err(String::from("hash check failed"));
91 }
92
93 let m_prime = self._k_pke_decrypt(dk_pke, c);
94
95 let pre_image = [m_prime.clone(), h.to_vec()].concat();
96 let (k_prime, r_prime) = Self::_g(&pre_image);
97 let pre_image = [z, c].concat();
98 let k_bar = Self::_j(&pre_image);
99
100 let c_prime = self._k_pke_encrypt(ek_pke, &m_prime, &r_prime).unwrap();
101
102 Ok(select_bytes(&k_bar, &k_prime, c == c_prime))
103 }
104
105 fn _encaps_internal(&self, ek: &[u8], m: &[u8]) -> (Vec<u8>, Vec<u8>) {
106 let pre_image = [m, &Self::_h(ek)].concat();
107 let (k, r) = Self::_g(&pre_image);
108 let c = self._k_pke_encrypt(ek, m, &r).unwrap();
109 (k, c)
110 }
111
112 fn _k_pke_encrypt(&self, ek_pke: &[u8], m: &[u8], r: &[u8]) -> Result<Vec<u8>, String> {
113 if ek_pke.len() != 384 * (self.k as usize) + 32 {
114 return Err(String::from(
115 "Type check failed, ek_pke has the wrong length",
116 ));
117 }
118 let t_hat_bytes = &ek_pke[..ek_pke.len() - 32];
119 let rho = &ek_pke[ek_pke.len() - 32..];
120 let t_hat = Module::decode_vector(t_hat_bytes, self.k as usize, 12, true)?;
121
122 if t_hat.encode(12) != t_hat_bytes {
123 return Err(String::from(
124 "Modulus check failed, t_hat does not encode correctly",
125 ));
126 }
127 let a_hat_t = self._generate_matrix_from_seed(rho, true);
128 let n = 0;
129 let (y, n) = self._generate_error_vector(r, self.eta_1, n);
130 let (e_1, n) = self._generate_error_vector(r, self.eta_2, n);
131 let (e_2, _) = self._generate_polynomial(r, self.eta_2, n);
132
133 let y_hat = y.to_ntt();
134
135 let u = &((a_hat_t.mat_mul(&y_hat)?).from_ntt()) + &e_1;
136
137 let mu = Ring::decode(m, 1, false)?.decompress(1);
138
139 let v = &(t_hat.dot(&y_hat)?.from_ntt()) + &(&e_2 + &mu);
140
141 let c_1 = u.compress(self.du).encode(self.du as usize);
142 let c_2 = v.compress(self.dv).encode(self.dv as usize);
143
144 Ok([c_1, c_2].concat())
145 }
146
147 fn _k_pke_decrypt(&self, dk_pke: &[u8], c: &[u8]) -> Vec<u8> {
148 let n = self.k as usize * self.du as usize * 32;
149 let c_1 = &c[..n];
150 let c_2 = &c[n..];
151 let u = Module::decode_vector(c_1, self.k as usize, self.du as usize, false)
152 .unwrap()
153 .decompress(self.du);
154 let v = Ring::decode(c_2, self.dv as usize, false)
155 .unwrap()
156 .decompress(self.dv);
157 let s_hat = Module::decode_vector(dk_pke, self.k as usize, 12, true).unwrap();
158
159 let u_hat = u.to_ntt();
160 let w = &v - &(s_hat.dot(&u_hat).unwrap()).from_ntt();
161
162 w.compress(1).encode(1)
163 }
164
165 fn _keygen_internal(&self, d: &[u8], z: &[u8]) -> (Vec<u8>, Vec<u8>) {
166 let (ek_pke, dk_pke) = self._k_pke_keygen(d);
167
168 let ek = ek_pke;
169 let dk = [dk_pke, ek.clone(), Self::_h(&ek), z.to_vec()].concat();
170
171 (ek, dk)
172 }
173
174 fn _k_pke_keygen(&self, d: &[u8]) -> (Vec<u8>, Vec<u8>) {
175 let pre_image: Vec<u8> = [d, &[self.k]].concat();
176
177 let (rho, sigma) = Self::_g(&pre_image);
178
179 let a_hat = self._generate_matrix_from_seed(&rho, false);
180
181 let n = 0;
182
183 let (s, n) = self._generate_error_vector(&sigma, self.eta_1, n);
184
185 let (e, _) = self._generate_error_vector(&sigma, self.eta_1, n);
186
187 let s_hat = s.to_ntt();
188
189 let e_hat = e.to_ntt();
190
191 let sa_hat = a_hat.mat_mul(&s_hat).unwrap();
192
193 let t_hat = &sa_hat + &e_hat;
194
195 let ek_pke = [t_hat.encode(12), rho].concat();
196
197 let dk_pke = s_hat.encode(12);
198
199 (ek_pke, dk_pke)
200 }
201
202 fn random_bytes(length: usize) -> Vec<u8> {
203 let mut bytes = vec![0u8; length];
204 OsRng.fill_bytes(&mut bytes);
205 bytes
206 }
207
208 fn _g(s: &[u8]) -> (Vec<u8>, Vec<u8>) {
209 let mut hasher = Sha3_512::new();
210 Update::update(&mut hasher, s);
211 let result = hasher.finalize();
212 (result[..32].to_vec(), result[32..].to_vec())
213 }
214
215 fn _h(s: &[u8]) -> Vec<u8> {
216 let mut hasher = Sha3_256::new();
217 Update::update(&mut hasher, s);
218 let result = hasher.finalize();
219 result.to_vec()
220 }
221
222 fn _j(s: &[u8]) -> Vec<u8> {
223 let mut hasher = Shake256::default();
224 hasher.update(s);
225
226 let mut reader = hasher.finalize_xof();
227 let mut buf = [0u8; 32];
228 reader.read(&mut buf);
229
230 buf.to_vec()
231 }
232
233 fn _xof(b: &[u8], i: u8, j: u8) -> Vec<u8> {
234 let mut hasher = Shake128::default();
236 let pre_image: Vec<u8> = [b, &[i], &[j]].concat();
237 hasher.update(&pre_image);
238
239 let mut reader = hasher.finalize_xof();
240 let mut buf = [0u8; 840];
241 reader.read(&mut buf);
242
243 buf.to_vec()
244 }
245
246 fn _prf(eta: u8, s: &[u8], b: u8) -> Vec<u8> {
247 let mut hasher = Shake256::default();
249 let pre_image: Vec<u8> = [s, &[b]].concat();
250 hasher.update(&pre_image);
251
252 let mut reader = hasher.finalize_xof();
253 let mut buf: Vec<u8> = vec![0u8; (eta * 64).into()];
254 reader.read(&mut buf);
255
256 buf.to_vec()
257 }
258
259 fn _generate_matrix_from_seed(&self, rho: &[u8], transpose: bool) -> Module {
260 let k: usize = self.k.into();
261 let mut a_data = vec![vec![Ring::default(); k]; k];
262 for i in 0..k {
263 for j in 0..k {
264 let xof_bytes = Self::_xof(rho, j.try_into().unwrap(), i.try_into().unwrap());
265 a_data[i][j] = Ring::ntt_sample(&xof_bytes);
266 }
267 }
268 Module::new(&a_data, transpose)
269 }
270
271 fn _generate_error_vector(&self, sigma: &[u8], eta: u8, n: u8) -> (Module, u8) {
272 let k: usize = self.k.into();
273 let mut elements = vec![Ring::default(); k];
274 let mut n = n;
275 for i in 0..k {
276 let prf_output = Self::_prf(eta, sigma, n);
277 elements[i] = Ring::cbd(&prf_output, eta, false).unwrap();
278 n += 1;
279 }
280 let data = vec![elements];
281 (Module::new(&data, true), n)
282 }
283
284 fn _generate_polynomial(&self, sigma: &[u8], eta: u8, n: u8) -> (Ring, u8) {
285 let prf_output = Self::_prf(eta, sigma, n);
286 let p = Ring::cbd(&prf_output, eta, false).unwrap();
287 (p, n + 1)
288 }
289}
290
291fn select_bytes(a: &[u8], b: &[u8], cond: bool) -> Vec<u8> {
292 let mut out = vec![0_u8; a.len()];
294 let cw = if !cond { 0 } else { 255 };
295 for i in 0..(a.len()) {
296 out[i] = a[i] ^ (cw & (a[i] ^ b[i]))
297 }
298 out
299}
300
301pub const ML_KEM_512: MLKem = MLKem::new(Type::MlKem512);
302pub const ML_KEM_768: MLKem = MLKem::new(Type::MlKem768);
303pub const ML_KEM_1024: MLKem = MLKem::new(Type::MlKem1024);
304
305#[cfg(test)]
306mod tests {
307 use super::*;
308 use serde_json::Value;
309 use std::fs;
310
311 fn keygen_kat(_type: Type, index: usize) {
312 let data =
313 fs::read_to_string("assets/ML-KEM-keyGen-FIPS203/internalProjection.json").unwrap();
314 let json: Value = serde_json::from_str(&data).unwrap();
315 let tests = json["testGroups"][index]["tests"].as_array().unwrap();
316 let ml_kem = MLKem::new(_type);
317 for value in tests.iter() {
318 let z = &value["z"];
319 let d = &value["d"];
320 let ek = &value["ek"];
321 let dk = &value["dk"];
322
323 let z_as_bytes = hex::decode(z.as_str().unwrap()).unwrap();
324 let d_as_bytes = hex::decode(d.as_str().unwrap()).unwrap();
325
326 let (actual_ek, actual_dk) = ml_kem._keygen_internal(&d_as_bytes, &z_as_bytes);
327
328 let ek_as_bytes = hex::decode(ek.as_str().unwrap()).unwrap();
329 let dk_as_bytes = hex::decode(dk.as_str().unwrap()).unwrap();
330
331 assert_eq!(actual_ek, ek_as_bytes);
332 assert_eq!(actual_dk, dk_as_bytes);
333 }
334 }
335
336 fn encaps_kat(_type: Type, index: usize) {
337 let data =
338 fs::read_to_string("assets/ML-KEM-encapDecap-FIPS203/internalProjection.json").unwrap();
339 let json: Value = serde_json::from_str(&data).unwrap();
340 let tests = json["testGroups"][index]["tests"].as_array().unwrap();
341 let ml_kem = MLKem::new(_type);
342 for value in tests.iter() {
343 let c = &value["c"];
344 let k = &value["k"];
345 let m = &value["m"];
346 let ek = &value["ek"];
347 let dk = &value["dk"];
348
349 let ek_as_bytes = hex::decode(ek.as_str().unwrap()).unwrap();
350 let m_as_bytes = hex::decode(m.as_str().unwrap()).unwrap();
351
352 let (actual_k, actual_c) = ml_kem._encaps_internal(&ek_as_bytes, &m_as_bytes);
353
354 let k_as_bytes = hex::decode(k.as_str().unwrap()).unwrap();
355 let c_as_bytes = hex::decode(c.as_str().unwrap()).unwrap();
356
357 assert_eq!(actual_k, k_as_bytes);
358 assert_eq!(actual_c, c_as_bytes);
359
360 let dk_as_bytes = hex::decode(dk.as_str().unwrap()).unwrap();
361
362 let k_prime = ml_kem.decaps(&dk_as_bytes, &c_as_bytes);
363 assert_eq!(k_prime, k_as_bytes);
364 }
365 }
366
367 fn decaps_kat(_type: Type, index: usize) {
368 let data =
369 fs::read_to_string("assets/ML-KEM-encapDecap-FIPS203/internalProjection.json").unwrap();
370 let json: Value = serde_json::from_str(&data).unwrap();
371 let kat_data = json["testGroups"][3 + index]["tests"].as_array().unwrap();
372 let dk = json["testGroups"][3 + index]["dk"].as_str().unwrap();
373 let dk_as_bytes = hex::decode(dk).unwrap();
374 let ml_kem = MLKem::new(_type);
375 for value in kat_data.iter() {
376 let c = &value["c"];
377 let c_as_bytes = hex::decode(c.as_str().unwrap()).unwrap();
378 let k = &value["k"];
379 let k_as_bytes = hex::decode(k.as_str().unwrap()).unwrap();
380 let k = ml_kem.decaps(&dk_as_bytes, &c_as_bytes);
381 assert_eq!(k, k_as_bytes)
382 }
383 }
384
385 #[test]
386 fn test_keygen_using_kat() {
387 keygen_kat(Type::MlKem512, 0);
388 keygen_kat(Type::MlKem768, 1);
389 keygen_kat(Type::MlKem1024, 2);
390 }
391
392 #[test]
393 fn test_encaps_using_kat() {
394 encaps_kat(Type::MlKem512, 0);
395 encaps_kat(Type::MlKem768, 1);
396 encaps_kat(Type::MlKem1024, 2);
397 }
398
399 #[test]
400 fn test_decaps_using_kat() {
401 decaps_kat(Type::MlKem512, 0);
402 decaps_kat(Type::MlKem768, 1);
403 decaps_kat(Type::MlKem1024, 2);
404 }
405}