1use crate::{
4 arithmetic::{Matrix, NttMatrix, RingElem},
5 consts::{
6 DOMSEP_KGEXPAND, DOMSEP_PKHASH, MAX_L, MAX_T, MODULUS_P_BITS, MODULUS_Q_BITS, RING_DEG,
7 },
8 sample::{gen_matrix_from_seed, gen_secret_from_seed},
9 ser::deserialize_generic,
10 turboshake256_hash,
11};
12
13use turboshake::CTurboShake256;
14use turboshake::digest::{ExtendableOutput, Update, XofReader};
15use zeroize::{Zeroize, ZeroizeOnDrop};
16
17const H1_VAL: u16 = 1 << (MODULUS_Q_BITS - MODULUS_P_BITS - 1);
18
19const PK_VEC_ELEM_BYTES: usize = MODULUS_P_BITS * RING_DEG / 8;
22
23#[derive(Zeroize, ZeroizeOnDrop)]
29pub(crate) struct PkeSecretKey<const L: usize>(NttMatrix<L, 1>);
30
31#[derive(Clone)]
38pub struct PkePublicKey<const L: usize> {
39 matrix_seed: [u8; 32],
41 mat_a_ntt: NttMatrix<L, L>,
47 vec_bytes: [[u8; PK_VEC_ELEM_BYTES]; L],
51 vec_ntt: NttMatrix<L, 1>,
53}
54
55impl<const L: usize> PkePublicKey<L> {
56 pub const SERIALIZED_LEN: usize = 32 + L * MODULUS_P_BITS * RING_DEG / 8;
57
58 #[allow(clippy::needless_range_loop)]
61 pub(crate) fn serialize(&self, out_buf: &mut [u8]) {
62 assert_eq!(out_buf.len(), Self::SERIALIZED_LEN);
63
64 for i in 0..L {
67 let start = i * PK_VEC_ELEM_BYTES;
68 out_buf[start..start + PK_VEC_ELEM_BYTES].copy_from_slice(&self.vec_bytes[i]);
69 }
70 out_buf[L * PK_VEC_ELEM_BYTES..].copy_from_slice(&self.matrix_seed);
71 }
72
73 #[allow(clippy::unwrap_used, clippy::needless_range_loop)]
75 pub(crate) fn from_bytes(bytes: &[u8]) -> Self {
76 assert_eq!(bytes.len(), Self::SERIALIZED_LEN);
77
78 let (vec_slice, seed) = bytes.split_at(Self::SERIALIZED_LEN - 32);
79 let matrix_seed: [u8; 32] = seed.try_into().unwrap(); let vec = Matrix::deserialize_10(vec_slice);
84 let vec_ntt = NttMatrix::from_uniform_matrix(&vec);
85
86 let mut vec_bytes = [[0u8; PK_VEC_ELEM_BYTES]; L];
89 for i in 0..L {
90 let start = i * PK_VEC_ELEM_BYTES;
91 vec_bytes[i].copy_from_slice(&vec_slice[start..start + PK_VEC_ELEM_BYTES]);
92 }
93
94 let mat_a = gen_matrix_from_seed::<L>(&matrix_seed);
95 let mat_a_ntt = NttMatrix::from_uniform_matrix(&mat_a);
96 Self {
97 matrix_seed,
98 mat_a_ntt,
99 vec_bytes,
100 vec_ntt,
101 }
102 }
103
104 pub(crate) fn hash(&self) -> [u8; 32] {
106 let mut buf = [0u8; max_pke_pubkey_serialized_len()];
108 let pk_slice = &mut buf[..PkePublicKey::<L>::SERIALIZED_LEN];
109 self.serialize(pk_slice);
110 turboshake256_hash::<DOMSEP_PKHASH>(pk_slice, &[])
111 }
112}
113
114pub(crate) const fn max_pke_pubkey_serialized_len() -> usize {
116 32 + MAX_L * MODULUS_P_BITS * RING_DEG / 8
117}
118
119pub const fn max_ciphertext_len() -> usize {
122 MAX_T * RING_DEG / 8 + MAX_L * MODULUS_P_BITS * RING_DEG / 8
124}
125
126pub const fn ciphertext_len<const L: usize, const T: usize>() -> usize {
129 L * MODULUS_P_BITS * RING_DEG / 8 + T * RING_DEG / 8
131}
132
133#[allow(clippy::needless_range_loop)]
142pub(crate) fn expand_decap_key<const L: usize, const MU: usize>(
143 sk: &[u8; 32],
144) -> (PkeSecretKey<L>, [u8; 32], PkePublicKey<L>, [u8; 32]) {
145 let mut mat_seed = [0u8; 32];
147 let mut secret_seed = [0u8; 32];
148 let mut z = [0u8; 32];
149
150 let mut xof = {
151 let mut hasher = CTurboShake256::<DOMSEP_KGEXPAND>::default();
152 hasher.update(sk);
153 hasher.update(&[L as u8]);
154 hasher.finalize_xof()
155 };
156 xof.read(&mut mat_seed);
157 xof.read(&mut secret_seed);
158 xof.read(&mut z);
159
160 let mat_a = gen_matrix_from_seed::<L>(&mat_seed);
161 let vec_s = gen_secret_from_seed::<L, MU>(&secret_seed);
162 let mat_a_ntt = NttMatrix::from_uniform_matrix(&mat_a);
163 let vec_s_ntt = NttMatrix::from_secret_matrix(&vec_s);
164
165 let b = {
167 let mut prod = mat_a_ntt.mul_transpose(&vec_s_ntt);
168 prod.wrapping_add_to_all(H1_VAL);
169 prod.shift_right(MODULUS_Q_BITS - MODULUS_P_BITS);
170 prod
171 };
172
173 let vec_ntt = NttMatrix::from_uniform_matrix(&b);
175
176 let mut vec_bytes = [[0u8; PK_VEC_ELEM_BYTES]; L];
178 for i in 0..L {
179 b.0[i][0].serialize(&mut vec_bytes[i], MODULUS_P_BITS);
180 }
181
182 let pk = PkePublicKey {
183 matrix_seed: mat_seed,
184 mat_a_ntt,
185 vec_bytes,
186 vec_ntt,
187 };
188 let pkh = pk.hash();
189
190 (PkeSecretKey(vec_s_ntt), z, pk, pkh)
191}
192
193pub(crate) fn decrypt<const L: usize, const T: usize>(
196 sk: &PkeSecretKey<L>,
197 ciphertext: &[u8],
198) -> [u8; 32] {
199 assert_eq!(ciphertext.len(), ciphertext_len::<L, T>());
200 let (bprime_bytes, c_bytes) = ciphertext.split_at(L * MODULUS_P_BITS * RING_DEG / 8);
202
203 let bprime: Matrix<L, 1> = Matrix::deserialize_10(bprime_bytes);
204 let bprime_ntt = NttMatrix::from_uniform_matrix(&bprime);
205
206 let mut c = RingElem::deserialize(c_bytes, T);
207 c.shift_left(MODULUS_P_BITS - T);
208
209 let v = bprime_ntt.mul_transpose(&sk.0);
210 let v = v.0[0][0];
211
212 let mut mprime = &v - &c;
214 let h2_val = (1 << (MODULUS_P_BITS - 2)) - (1 << (MODULUS_P_BITS - T - 1))
215 + (1 << (MODULUS_Q_BITS - MODULUS_P_BITS - 1));
216 mprime.wrapping_add_to_all(h2_val);
217 mprime.shift_right(MODULUS_P_BITS - 1);
218
219 let mut m = [0u8; 32];
220 mprime.serialize(&mut m, 1);
221 m
222}
223
224pub(crate) fn encrypt_deterministic<const L: usize, const MU: usize, const T: usize>(
227 pk: &PkePublicKey<L>,
228 msg: &[u8; 32],
229 randomness: &[u8; 32],
230 out_buf: &mut [u8],
231) {
232 assert_eq!(out_buf.len(), ciphertext_len::<L, T>());
233
234 let vec_sprime = gen_secret_from_seed::<L, MU>(randomness);
235 let sprime_ntt = NttMatrix::from_secret_matrix(&vec_sprime);
236
237 let bprime = {
238 let mut prod = pk.mat_a_ntt.mul(&sprime_ntt);
239 prod.wrapping_add_to_all(H1_VAL);
240 prod.shift_right(MODULUS_Q_BITS - MODULUS_P_BITS);
241 prod
242 };
243
244 let vprime: Matrix<1, 1> = pk.vec_ntt.mul_transpose(&sprime_ntt);
245 let vprime = vprime.0[0][0];
246
247 let mut msg_polyn = RingElem(deserialize_generic(msg, 1));
248 msg_polyn.shift_left(MODULUS_P_BITS - 1);
249
250 let mut c = &vprime - &msg_polyn;
252 c.wrapping_add_to_all(H1_VAL);
253 c.shift_right(MODULUS_P_BITS - T);
254
255 let (bprime_buf, c_buf) = out_buf.split_at_mut(L * MODULUS_P_BITS * RING_DEG / 8);
257 bprime.serialize(bprime_buf, MODULUS_P_BITS);
258 c.serialize(c_buf, T);
259}
260
261#[cfg(test)]
262mod test {
263 use super::*;
264 use crate::consts::*;
265
266 use rand::RngCore;
267
268 #[test]
270 fn encryption_correctness() {
271 test_enc_dec::<KOPIS512_L, KOPIS512_T, KOPIS512_MU>();
272 test_enc_dec::<KOPIS768_L, KOPIS768_T, KOPIS768_MU>();
273 test_enc_dec::<KOPIS1024_L, KOPIS1024_T, KOPIS1024_MU>();
274 }
275
276 fn test_enc_dec<const L: usize, const T: usize, const MU: usize>() {
278 let mut rng = rand::rng();
279 let mut backing_buf = [0u8; MAX_T * RING_DEG / 8 + MAX_L * MODULUS_P_BITS * RING_DEG / 8];
280
281 for _ in 0..100 {
282 let mut sk_seed = [0u8; 32];
284 rng.fill_bytes(&mut sk_seed);
285 let (sk, _, pk, _) = expand_decap_key::<L, MU>(&sk_seed);
286
287 let mut enc_seed = [0u8; 32];
288 let mut msg = [0u8; 32];
289 rng.fill_bytes(&mut enc_seed);
290 rng.fill_bytes(&mut msg);
291 let ct_buf = &mut backing_buf[..T * RING_DEG / 8 + L * MODULUS_P_BITS * RING_DEG / 8];
292
293 encrypt_deterministic::<L, MU, T>(&pk, &msg, &enc_seed, ct_buf);
294 let recovered_msg = decrypt::<L, T>(&sk, ct_buf);
295 assert_eq!(msg, recovered_msg);
296 }
297 }
298}