Skip to main content

pqc_ml_kem/
lib.rs

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        // TODO: Add checks
235        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        // TODO: Add checks
248        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    // TODO: Add checks
293    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}