1use constant_time_eq::constant_time_eq;
4#[cfg(feature = "zeroize")]
5use zeroize::{Zeroize, ZeroizeOnDrop};
6
7use crate::{
8 Xof,
9 sha3::{Shake128, Shake256},
10};
11
12pub const ML_DSA_65_PUBLIC_KEY_SIZE: usize = 1952;
13pub const ML_DSA_65_SIGNATURE_SIZE: usize = 3309;
14pub const ML_DSA_65_SEED_SIZE: usize = 32;
15pub const ML_DSA_65_CONTEXT_MAX_LEN: usize = 255;
16
17const Q: u32 = 8380417;
18const N: usize = 256;
19const D: u32 = 13;
20const ONE: u32 = 4193792;
21const MINUS_ONE: u32 = 4186625;
22const RR: u32 = 2365951;
23const QINV: u32 = 4236238847;
24const N_INV: u32 = 16382;
25const GAMMA1: u32 = 1 << 19;
26const GAMMA2: u32 = (Q - 1) / 32;
27const BETA: u32 = 196;
28const TAU: usize = 49;
29const LAMBDA_OVER_4: usize = 48;
30const POLYZ_BYTES: usize = (19 + 1) * N / 8;
31const K: usize = 6;
32const L: usize = 5;
33const OMEGA: usize = 55;
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
36pub enum MlDsaError {
37 ContextTooLong,
38 InvalidSignature,
39 InvalidPublicKey,
40 InvalidSignatureLength,
41}
42
43#[cfg(feature = "alloc")]
44impl core::fmt::Display for MlDsaError {
45 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
46 match self {
47 MlDsaError::ContextTooLong => write!(f, "context length exceeds 255 bytes"),
48 MlDsaError::InvalidSignature => write!(f, "signature is not valid"),
49 MlDsaError::InvalidPublicKey => write!(f, "public key is not valid"),
50 MlDsaError::InvalidSignatureLength => write!(f, "signature length is not valid"),
51 }
52 }
53}
54
55type FieldElement = u32;
56
57fn field_to_montgomery(a: u32) -> FieldElement {
58 debug_assert!(a < Q);
59 field_montgomery_mul(a, RR)
60}
61
62fn field_from_montgomery(a: FieldElement) -> u32 {
63 field_montgomery_reduce(a as u64)
64}
65
66fn field_montgomery_reduce(x: u64) -> u32 {
67 let t = (x as u32).wrapping_mul(QINV);
68 let u = (x + (t as u64) * (Q as u64)) >> 32;
69 field_reduce_once(u as u32)
70}
71
72fn field_montgomery_mul(a: FieldElement, b: FieldElement) -> FieldElement {
73 field_montgomery_reduce(a as u64 * b as u64)
74}
75
76fn field_reduce_once(x: u32) -> FieldElement {
77 let t = x.wrapping_sub(Q);
78 let mask = ((t as i32) >> 31) as u32;
79 t.wrapping_add(Q & mask)
80}
81
82fn field_add(a: FieldElement, b: FieldElement) -> FieldElement {
83 field_reduce_once(a.wrapping_add(b))
84}
85
86fn field_sub(a: FieldElement, b: FieldElement) -> FieldElement {
87 field_reduce_once(a.wrapping_sub(b).wrapping_add(Q))
88}
89
90fn field_sub_to_montgomery(a: u32, b: u32) -> FieldElement {
91 let x = a.wrapping_sub(b).wrapping_add(Q);
92 field_montgomery_mul(x, RR)
93}
94
95fn field_infinity_norm(r: FieldElement) -> u32 {
96 let x = field_from_montgomery(r);
97 let q_minus_x = Q - x;
98 let half_q = Q / 2;
99 let mask = ((half_q.wrapping_sub(x)) as i32 >> 31) as u32;
100 (mask & q_minus_x) | (!mask & x)
101}
102
103fn field_centered_mod(r: FieldElement) -> i32 {
104 let x = field_from_montgomery(r);
105 let x = x as i32;
106 let half_q = (Q / 2) as i32;
107 let mask = ((half_q - x) >> 31) as i32;
108 (mask & (x - Q as i32)) | (!mask & x)
109}
110
111fn power2round(r: FieldElement) -> (u16, FieldElement) {
112 let rr = field_from_montgomery(r);
113 let r1 = (rr + (1 << 12) - 1) >> 13;
114 let r0 = field_sub_to_montgomery(rr, r1 << 13);
115 (r1 as u16, r0)
116}
117
118fn highbits32(x: u32) -> u8 {
119 let r1 = (x + 127) >> 7;
120 let r1 = (r1 * 1025 + (1 << 21)) >> 22;
121 (r1 & 0b1111) as u8
122}
123
124fn decompose32(r: FieldElement) -> (u8, i32) {
125 let x = field_from_montgomery(r) as i32;
126 let r1 = highbits32(x as u32);
127 let r0 = x - (r1 as i32) * 2 * (Q as i32 - 1) / 32;
128 let half_q = (Q / 2) as i32;
129 let mask = ((half_q - r0) >> 31) as i32;
130 let r0 = (mask & (r0 - Q as i32)) | (!mask & r0);
131 (r1, r0)
132}
133
134fn make_hint32(ct0: FieldElement, w: FieldElement, cs2: FieldElement) -> u8 {
135 let r_plus_z = field_sub(w, cs2);
136 let v1 = highbits32(field_from_montgomery(r_plus_z));
137 let r = field_add(r_plus_z, ct0);
138 let r1 = highbits32(field_from_montgomery(r));
139 (v1 ^ r1) as u8 & 1u8
140}
141
142fn use_hint32(r: FieldElement, hint: u8) -> u8 {
143 let (r1, r0) = decompose32(r);
144 if hint == 0 {
145 return r1;
146 }
147 let r0_gt_0 = !(r0.wrapping_sub(1) >> 31) as u8;
148 let r1_plus = r1.wrapping_add(1) & 0x0F;
149 let r1_minus = r1.wrapping_sub(1) & 0x0F;
150 (r0_gt_0 & r1_plus) | ((!r0_gt_0) & r1_minus)
151}
152
153#[derive(Clone, Debug, PartialEq, Eq)]
154#[cfg_attr(feature = "zeroize", derive(Zeroize, ZeroizeOnDrop))]
155struct Poly {
156 coeffs: [FieldElement; N],
157}
158
159impl Default for Poly {
160 fn default() -> Self {
161 Self {
162 coeffs: [0u32; N],
163 }
164 }
165}
166
167#[derive(Clone, Debug, PartialEq, Eq)]
168#[cfg_attr(feature = "zeroize", derive(Zeroize, ZeroizeOnDrop))]
169struct NttPoly {
170 coeffs: [FieldElement; N],
171}
172
173impl Default for NttPoly {
174 fn default() -> Self {
175 Self {
176 coeffs: [0u32; N],
177 }
178 }
179}
180
181fn poly_add(a: &Poly, b: &Poly) -> Poly {
182 let mut r = Poly::default();
183 for i in 0..N {
184 r.coeffs[i] = field_add(a.coeffs[i], b.coeffs[i]);
185 }
186 r
187}
188
189fn poly_sub(a: &Poly, b: &Poly) -> Poly {
190 let mut r = Poly::default();
191 for i in 0..N {
192 r.coeffs[i] = field_sub(a.coeffs[i], b.coeffs[i]);
193 }
194 r
195}
196
197fn ntt_add(a: &NttPoly, b: &NttPoly) -> NttPoly {
198 let mut r = NttPoly::default();
199 for i in 0..N {
200 r.coeffs[i] = field_add(a.coeffs[i], b.coeffs[i]);
201 }
202 r
203}
204
205fn ntt_sub(a: &NttPoly, b: &NttPoly) -> NttPoly {
206 let mut r = NttPoly::default();
207 for i in 0..N {
208 r.coeffs[i] = field_sub(a.coeffs[i], b.coeffs[i]);
209 }
210 r
211}
212
213fn ntt_mul(a: &NttPoly, b: &NttPoly) -> NttPoly {
214 let mut r = NttPoly::default();
215 for i in 0..N {
216 r.coeffs[i] = field_montgomery_mul(a.coeffs[i], b.coeffs[i]);
217 }
218 r
219}
220
221const ZETAS: [FieldElement; 256] = [
222 4193792, 25847, 5771523, 7861508, 237124, 7602457, 7504169, 466468, 1826347, 2353451, 8021166, 6288512, 3119733,
223 5495562, 3111497, 2680103, 2725464, 1024112, 7300517, 3585928, 7830929, 7260833, 2619752, 6271868, 6262231,
224 4520680, 6980856, 5102745, 1757237, 8360995, 4010497, 280005, 2706023, 95776, 3077325, 3530437, 6718724, 4788269,
225 5842901, 3915439, 4519302, 5336701, 3574422, 5512770, 3539968, 8079950, 2348700, 7841118, 6681150, 6736599,
226 3505694, 4558682, 3507263, 6239768, 6779997, 3699596, 811944, 531354, 954230, 3881043, 3900724, 5823537, 2071892,
227 5582638, 4450022, 6851714, 4702672, 5339162, 6927966, 3475950, 2176455, 6795196, 7122806, 1939314, 4296819,
228 7380215, 5190273, 5223087, 4747489, 126922, 3412210, 7396998, 2147896, 2715295, 5412772, 4686924, 7969390, 5903370,
229 7709315, 7151892, 8357436, 7072248, 7998430, 1349076, 1852771, 6949987, 5037034, 264944, 508951, 3097992, 44288,
230 7280319, 904516, 3958618, 4656075, 8371839, 1653064, 5130689, 2389356, 8169440, 759969, 7063561, 189548, 4827145,
231 3159746, 6529015, 5971092, 8202977, 1315589, 1341330, 1285669, 6795489, 7567685, 6940675, 5361315, 4499357,
232 4751448, 3839961, 2091667, 3407706, 2316500, 3817976, 5037939, 2244091, 5933984, 4817955, 266997, 2434439, 7144689,
233 3513181, 4860065, 4621053, 7183191, 5187039, 900702, 1859098, 909542, 819034, 495491, 6767243, 8337157, 7857917,
234 7725090, 5257975, 2031748, 3207046, 4823422, 7855319, 7611795, 4784579, 342297, 286988, 5942594, 4108315, 3437287,
235 5038140, 1735879, 203044, 2842341, 2691481, 5790267, 1265009, 4055324, 1247620, 2486353, 1595974, 4613401, 1250494,
236 2635921, 4832145, 5386378, 1869119, 1903435, 7329447, 7047359, 1237275, 5062207, 6950192, 7929317, 1312455,
237 3306115, 6417775, 7100756, 1917081, 5834105, 7005614, 1500165, 777191, 2235880, 3406031, 7838005, 5548557, 6709241,
238 6533464, 5796124, 4656147, 594136, 4603424, 6366809, 2432395, 2454455, 8215696, 1957272, 3369112, 185531, 7173032,
239 5196991, 162844, 1616392, 3014001, 810149, 1652634, 4686184, 6581310, 5341501, 3523897, 3866901, 269760, 2213111,
240 7404533, 1717735, 472078, 7953734, 1723600, 6577327, 1910376, 6712985, 7276084, 8119771, 4546524, 5441381, 6144432,
241 7959518, 6094090, 183443, 7403526, 1612842, 4834730, 7826001, 3919660, 8332111, 7018208, 3937738, 1400424, 7534263,
242 1976782,
243];
244
245fn ntt(f: &Poly) -> NttPoly {
246 let mut f = NttPoly {
247 coeffs: f.coeffs,
248 };
249 let mut m: usize = 0;
250
251 let mut len: usize = 128;
252 while len >= 8 {
253 let mut start: usize = 0;
254 while start < N {
255 m += 1;
256 let zeta = ZETAS[m];
257 let mid = start + len;
258 for j in (start..mid).step_by(2) {
259 let t = field_montgomery_mul(zeta, f.coeffs[j + len]);
260 f.coeffs[j + len] = field_sub(f.coeffs[j], t);
261 f.coeffs[j] = field_add(f.coeffs[j], t);
262 let t = field_montgomery_mul(zeta, f.coeffs[j + len + 1]);
263 f.coeffs[j + len + 1] = field_sub(f.coeffs[j + 1], t);
264 f.coeffs[j + 1] = field_add(f.coeffs[j + 1], t);
265 }
266 start += 2 * len;
267 }
268 len /= 2;
269 }
270
271 let mut start: usize = 0;
272 while start < N {
273 m += 1;
274 let zeta = ZETAS[m];
275 let t = field_montgomery_mul(zeta, f.coeffs[start + 4]);
276 f.coeffs[start + 4] = field_sub(f.coeffs[start], t);
277 f.coeffs[start] = field_add(f.coeffs[start], t);
278 let t = field_montgomery_mul(zeta, f.coeffs[start + 5]);
279 f.coeffs[start + 5] = field_sub(f.coeffs[start + 1], t);
280 f.coeffs[start + 1] = field_add(f.coeffs[start + 1], t);
281 let t = field_montgomery_mul(zeta, f.coeffs[start + 6]);
282 f.coeffs[start + 6] = field_sub(f.coeffs[start + 2], t);
283 f.coeffs[start + 2] = field_add(f.coeffs[start + 2], t);
284 let t = field_montgomery_mul(zeta, f.coeffs[start + 7]);
285 f.coeffs[start + 7] = field_sub(f.coeffs[start + 3], t);
286 f.coeffs[start + 3] = field_add(f.coeffs[start + 3], t);
287 start += 8;
288 }
289
290 start = 0;
291 while start < N {
292 m += 1;
293 let zeta = ZETAS[m];
294 let t = field_montgomery_mul(zeta, f.coeffs[start + 2]);
295 f.coeffs[start + 2] = field_sub(f.coeffs[start], t);
296 f.coeffs[start] = field_add(f.coeffs[start], t);
297 let t = field_montgomery_mul(zeta, f.coeffs[start + 3]);
298 f.coeffs[start + 3] = field_sub(f.coeffs[start + 1], t);
299 f.coeffs[start + 1] = field_add(f.coeffs[start + 1], t);
300 start += 4;
301 }
302
303 start = 0;
304 while start < N {
305 m += 1;
306 let zeta = ZETAS[m];
307 let t = field_montgomery_mul(zeta, f.coeffs[start + 1]);
308 f.coeffs[start + 1] = field_sub(f.coeffs[start], t);
309 f.coeffs[start] = field_add(f.coeffs[start], t);
310 start += 2;
311 }
312
313 f
314}
315
316fn invntt(f: &NttPoly) -> Poly {
317 let mut f = NttPoly {
318 coeffs: f.coeffs,
319 };
320 let mut m: usize = 255;
321
322 let mut start: usize = 0;
323 while start < N {
324 let zeta = ZETAS[m];
325 m -= 1;
326 let t = f.coeffs[start];
327 f.coeffs[start] = field_add(t, f.coeffs[start + 1]);
328 f.coeffs[start + 1] = field_montgomery_mul(zeta, field_sub(f.coeffs[start + 1], t));
329 start += 2;
330 }
331
332 start = 0;
333 while start < N {
334 let zeta = ZETAS[m];
335 m -= 1;
336 let t = f.coeffs[start];
337 f.coeffs[start] = field_add(t, f.coeffs[start + 2]);
338 f.coeffs[start + 2] = field_montgomery_mul(zeta, field_sub(f.coeffs[start + 2], t));
339 let t = f.coeffs[start + 1];
340 f.coeffs[start + 1] = field_add(t, f.coeffs[start + 3]);
341 f.coeffs[start + 3] = field_montgomery_mul(zeta, field_sub(f.coeffs[start + 3], t));
342 start += 4;
343 }
344
345 start = 0;
346 while start < N {
347 let zeta = ZETAS[m];
348 m -= 1;
349 let t = f.coeffs[start];
350 f.coeffs[start] = field_add(t, f.coeffs[start + 4]);
351 f.coeffs[start + 4] = field_montgomery_mul(zeta, field_sub(f.coeffs[start + 4], t));
352 let t = f.coeffs[start + 1];
353 f.coeffs[start + 1] = field_add(t, f.coeffs[start + 5]);
354 f.coeffs[start + 5] = field_montgomery_mul(zeta, field_sub(f.coeffs[start + 5], t));
355 let t = f.coeffs[start + 2];
356 f.coeffs[start + 2] = field_add(t, f.coeffs[start + 6]);
357 f.coeffs[start + 6] = field_montgomery_mul(zeta, field_sub(f.coeffs[start + 6], t));
358 let t = f.coeffs[start + 3];
359 f.coeffs[start + 3] = field_add(t, f.coeffs[start + 7]);
360 f.coeffs[start + 7] = field_montgomery_mul(zeta, field_sub(f.coeffs[start + 7], t));
361 start += 8;
362 }
363
364 let mut len: usize = 8;
365 while len < N {
366 let mut start: usize = 0;
367 while start < N {
368 let zeta = ZETAS[m];
369 m -= 1;
370 let mid = start + len;
371 for j in (start..mid).step_by(2) {
372 let t = f.coeffs[j];
373 f.coeffs[j] = field_add(t, f.coeffs[j + len]);
374 let diff = field_sub(f.coeffs[j + len], t);
375 f.coeffs[j + len] = field_montgomery_mul(zeta, diff);
376 let t = f.coeffs[j + 1];
377 f.coeffs[j + 1] = field_add(t, f.coeffs[j + len + 1]);
378 let diff = field_sub(f.coeffs[j + len + 1], t);
379 f.coeffs[j + len + 1] = field_montgomery_mul(zeta, diff);
380 }
381 start += 2 * len;
382 }
383 len *= 2;
384 }
385
386 let mut r = Poly::default();
387 for i in 0..N {
388 r.coeffs[i] = field_montgomery_mul(f.coeffs[i], N_INV);
389 }
390 r
391}
392
393fn sample_ntt(rho: &[u8; 32], s: u8, r: u8) -> NttPoly {
394 let mut shake = Shake128::new();
395 shake.absorb(rho);
396 shake.absorb(&[s, r]);
397
398 let mut a = NttPoly::default();
399 let mut j: usize = 0;
400 let mut buf = [0u8; 168];
401 let mut off: usize = 168;
402
403 loop {
404 if off >= 168 {
405 shake.squeeze(&mut buf);
406 off = 0;
407 }
408 let v = (buf[off] as u32) | ((buf[off + 1] as u32) << 8) | ((buf[off + 2] as u32) << 16);
409 off += 3;
410 let v = v & 0x7FFFFF;
411 if v < Q {
412 a.coeffs[j] = field_to_montgomery(v);
413 j += 1;
414 if j >= N {
415 break;
416 }
417 }
418 }
419 a
420}
421
422fn sample_bounded_poly(rho: &[u8], r: u8) -> Poly {
423 let mut shake = Shake256::new();
424 shake.absorb(rho);
425 shake.absorb(&[r, 0]);
426
427 let mut a = Poly::default();
428 let mut j: usize = 0;
429 let mut buf = [0u8; 136];
430 let mut off: usize = 136;
431
432 loop {
433 if off >= 136 {
434 shake.squeeze(&mut buf);
435 off = 0;
436 }
437 let z0 = buf[off] & 0x0F;
438 let z1 = buf[off] >> 4;
439 off += 1;
440
441 if z0 <= 8 {
442 a.coeffs[j] = field_sub_to_montgomery(4, z0 as u32);
443 j += 1;
444 if j >= N {
445 break;
446 }
447 }
448 if z1 <= 8 {
449 a.coeffs[j] = field_sub_to_montgomery(4, z1 as u32);
450 j += 1;
451 if j >= N {
452 break;
453 }
454 }
455 }
456 a
457}
458
459fn sample_in_ball(rho: &[u8]) -> Poly {
460 let mut shake = Shake256::new();
461 shake.absorb(rho);
462 let mut s = [0u8; 8];
463 shake.squeeze(&mut s);
464
465 let mut c = Poly::default();
466 let mut signs: u64 = u64::from_le_bytes(s);
467
468 for i in (N - TAU)..N {
469 let mut jb = [0u8; 1];
470 loop {
471 shake.squeeze(&mut jb);
472 if jb[0] as usize <= i {
473 break;
474 }
475 }
476 let j = jb[0] as usize;
477 c.coeffs[i] = c.coeffs[j];
478 if (signs & 1) == 0 {
479 c.coeffs[j] = ONE;
480 } else {
481 c.coeffs[j] = MINUS_ONE;
482 }
483 signs >>= 1;
484 }
485
486 c
487}
488
489fn expand_mask(nonce: &[u8; 64], kappa: usize) -> Poly {
490 let mut shake = Shake256::new();
491 shake.absorb(nonce);
492 shake.absorb(&(kappa as u16).to_le_bytes());
493
494 let b = 1u32 << 19;
495 let mask20 = (1u32 << 20) - 1;
496 let mut buf = [0u8; POLYZ_BYTES];
497 shake.squeeze(&mut buf);
498 let mut r = Poly::default();
499 let mut p = &buf[..];
500 for i in (0..N).step_by(2) {
501 let w0 = (p[0] as u32) | ((p[1] as u32) << 8) | ((p[2] as u32) << 16);
502 r.coeffs[i] = field_sub_to_montgomery(b, w0 & mask20);
503 let w1 = ((p[2] as u32) >> 4) | ((p[3] as u32) << 4) | ((p[4] as u32) << 12);
504 r.coeffs[i + 1] = field_sub_to_montgomery(b, w1 & mask20);
505 p = &p[5..];
506 }
507 r
508}
509
510fn highbits_vec(w: &Poly) -> [u8; N] {
511 let mut r = [0u8; N];
512 for i in 0..N {
513 r[i] = highbits32(field_from_montgomery(w.coeffs[i]));
514 }
515 r
516}
517
518fn make_hint_vec(ct0: &Poly, w: &Poly, cs2: &Poly) -> ([u8; N], usize) {
519 let mut h = [0u8; N];
520 let mut count = 0usize;
521 for i in 0..N {
522 h[i] = make_hint32(ct0.coeffs[i], w.coeffs[i], cs2.coeffs[i]);
523 count += h[i] as usize;
524 }
525 (h, count)
526}
527
528fn use_hint_vec(r: &Poly, h: &[u8; N]) -> [u8; N] {
529 let mut w = [0u8; N];
530 for i in 0..N {
531 w[i] = use_hint32(r.coeffs[i], h[i]);
532 }
533 w
534}
535
536fn coefficients_exceed_bound(w: &Poly, bound: u32) -> bool {
537 for i in 0..N {
538 if field_infinity_norm(w.coeffs[i]) >= bound {
539 return true;
540 }
541 }
542 false
543}
544
545fn lowbits_exceed_bound(w: &Poly, bound: u32) -> bool {
546 for i in 0..N {
547 let (_, r0) = decompose32(w.coeffs[i]);
548 let abs_r0 = (r0 ^ (r0 >> 31)).wrapping_sub(r0 >> 31) as u32;
549 if abs_r0 >= bound {
550 return true;
551 }
552 }
553 false
554}
555
556fn pk_encode(rho: &[u8; 32], t1: &[[u16; N]; K]) -> [u8; ML_DSA_65_PUBLIC_KEY_SIZE] {
557 let mut pk = [0u8; ML_DSA_65_PUBLIC_KEY_SIZE];
558 pk[..32].copy_from_slice(rho);
559 let mut pos = 32;
560
561 for w in t1.iter() {
562 for i in (0..N).step_by(4) {
563 let c0 = w[i] as u32;
564 let c1 = w[i + 1] as u32;
565 let c2 = w[i + 2] as u32;
566 let c3 = w[i + 3] as u32;
567 pk[pos] = (c0 & 0xFF) as u8;
568 pk[pos + 1] = ((c0 >> 8) | (c1 << 2)) as u8;
569 pk[pos + 2] = ((c1 >> 6) | (c2 << 4)) as u8;
570 pk[pos + 3] = ((c2 >> 4) | (c3 << 6)) as u8;
571 pk[pos + 4] = (c3 >> 2) as u8;
572 pos += 5;
573 }
574 }
575 pk
576}
577
578fn pk_decode(pk: &[u8; ML_DSA_65_PUBLIC_KEY_SIZE]) -> Result<([u8; 32], [[u16; N]; K]), MlDsaError> {
579 let mut rho = [0u8; 32];
580 rho.copy_from_slice(&pk[..32]);
581 let mut t1 = [[0u16; N]; K];
582 let mut pos = 32;
583
584 for r in 0..K {
585 for i in (0..N).step_by(4) {
586 let b0 = pk[pos] as u16;
587 let b1 = pk[pos + 1] as u16;
588 let b2 = pk[pos + 2] as u16;
589 let b3 = pk[pos + 3] as u16;
590 let b4 = pk[pos + 4] as u16;
591 t1[r][i] = b0 | ((b1 & 0b0000_0011) << 8);
592 t1[r][i + 1] = (b1 >> 2) | ((b2 & 0b0000_1111) << 6);
593 t1[r][i + 2] = (b2 >> 4) | ((b3 & 0b0011_1111) << 4);
594 t1[r][i + 3] = (b3 >> 6) | ((b4 & 0b1111_1111) << 2);
595 pos += 5;
596 }
597 }
598 Ok((rho, t1))
599}
600
601fn bitpack_20(z: &Poly) -> [u8; POLYZ_BYTES] {
602 let b = 1u32 << 19;
603 let mut out = [0u8; POLYZ_BYTES];
604 let mut q = 0usize;
605
606 for i in (0..N).step_by(2) {
607 let w0 = (b as i32 - field_centered_mod(z.coeffs[i])) as u32;
608 out[q] = w0 as u8;
609 out[q + 1] = (w0 >> 8) as u8;
610 out[q + 2] = (w0 >> 16) as u8;
611 let w1 = (b as i32 - field_centered_mod(z.coeffs[i + 1])) as u32;
612 out[q + 2] |= ((w1 & 0x0F) << 4) as u8;
613 out[q + 3] = (w1 >> 4) as u8;
614 out[q + 4] = (w1 >> 12) as u8;
615 q += 5;
616 }
617 out
618}
619
620fn bitunpack_20(v: &[u8]) -> Poly {
621 let b = 1u32 << 19;
622 let mask20 = (1u32 << 20) - 1;
623 let mut r = Poly::default();
624 let mut p = v;
625
626 for i in (0..N).step_by(2) {
627 let w0 = (p[0] as u32) | ((p[1] as u32) << 8) | ((p[2] as u32) << 16);
628 r.coeffs[i] = field_sub_to_montgomery(b, w0 & mask20);
629 let w1 = ((p[2] as u32) >> 4) | ((p[3] as u32) << 4) | ((p[4] as u32) << 12);
630 r.coeffs[i + 1] = field_sub_to_montgomery(b, w1 & mask20);
631 p = &p[5..];
632 }
633 r
634}
635
636fn hint_encode(h: &[[u8; N]; K]) -> [u8; OMEGA + K] {
637 let mut sig = [0u8; OMEGA + K];
638 let mut idx: u8 = 0;
639
640 for i in 0..K {
641 for j in 0..N {
642 if h[i][j] != 0 {
643 sig[idx as usize] = j as u8;
644 idx += 1;
645 }
646 }
647 sig[OMEGA + i] = idx;
648 }
649 sig
650}
651
652fn hint_decode(sig: &[u8; OMEGA + K]) -> Result<[[u8; N]; K], MlDsaError> {
653 let mut h = [[0u8; N]; K];
654 let mut idx: u8 = 0;
655
656 for i in 0..K {
657 let limit = sig[OMEGA + i];
658 if limit < idx || limit > OMEGA as u8 {
659 return Err(MlDsaError::InvalidSignature);
660 }
661 let poly_start = idx;
663 while idx < limit {
664 let j = sig[idx as usize];
665 if idx > poly_start && sig[(idx - 1) as usize] >= j {
667 return Err(MlDsaError::InvalidSignature);
668 }
669 if j as usize >= N {
670 return Err(MlDsaError::InvalidSignature);
671 }
672 h[i][j as usize] = 1;
673 idx += 1;
674 }
675 }
676 for k in idx as usize..OMEGA {
677 if sig[k] != 0 {
678 return Err(MlDsaError::InvalidSignature);
679 }
680 }
681 Ok(h)
682}
683
684fn sig_encode(ch: &[u8; LAMBDA_OVER_4], z: &[Poly; L], h: &[[u8; N]; K]) -> [u8; ML_DSA_65_SIGNATURE_SIZE] {
685 let mut sig = [0u8; ML_DSA_65_SIGNATURE_SIZE];
686 sig[..LAMBDA_OVER_4].copy_from_slice(ch);
687
688 let mut pos = LAMBDA_OVER_4;
689 for i in 0..L {
690 let packed = bitpack_20(&z[i]);
691 sig[pos..pos + POLYZ_BYTES].copy_from_slice(&packed);
692 pos += POLYZ_BYTES;
693 }
694
695 let hint_sig = hint_encode(h);
696 sig[pos..].copy_from_slice(&hint_sig);
697 sig
698}
699
700fn sig_decode(sig: &[u8]) -> Result<([u8; LAMBDA_OVER_4], [Poly; L], [[u8; N]; K]), MlDsaError> {
701 if sig.len() != ML_DSA_65_SIGNATURE_SIZE {
702 return Err(MlDsaError::InvalidSignatureLength);
703 }
704 let mut ch = [0u8; LAMBDA_OVER_4];
705 ch.copy_from_slice(&sig[..LAMBDA_OVER_4]);
706
707 let mut z: [Poly; L] = Default::default();
708 let mut pos = LAMBDA_OVER_4;
709 for i in 0..L {
710 z[i] = bitunpack_20(&sig[pos..pos + POLYZ_BYTES]);
711 pos += POLYZ_BYTES;
712 }
713
714 let mut hint_bytes = [0u8; OMEGA + K];
715 hint_bytes.copy_from_slice(&sig[pos..]);
716 let h = hint_decode(&hint_bytes)?;
717
718 Ok((ch, z, h))
719}
720
721fn w1_encode_bytes(w1: &[[u8; N]; K]) -> [u8; K * N / 2] {
722 let mut buf = [0u8; K * N / 2];
723 let mut pos = 0;
724 for w in w1.iter() {
725 for i in (0..N).step_by(2) {
726 buf[pos] = w[i] | (w[i + 1] << 4);
727 pos += 1;
728 }
729 }
730 buf
731}
732
733fn compute_matrix_a(rho: &[u8; 32]) -> [[NttPoly; L]; K] {
734 let mut a: [[NttPoly; L]; K] = Default::default();
735 for r in 0..K {
736 for s in 0..L {
737 a[r][s] = sample_ntt(rho, s as u8, r as u8);
738 }
739 }
740 a
741}
742
743fn compute_pubkey_hash(pk: &[u8; ML_DSA_65_PUBLIC_KEY_SIZE]) -> [u8; 64] {
744 let mut shake = Shake256::new();
745 shake.absorb(pk);
746 let mut tr = [0u8; 64];
747 shake.squeeze(&mut tr);
748 tr
749}
750
751fn compute_message_hash(tr: &[u8; 64], message: &[u8], ctx: &[u8]) -> Result<[u8; 64], MlDsaError> {
752 if ctx.len() > 255 {
753 return Err(MlDsaError::ContextTooLong);
754 }
755 let mut shake = Shake256::new();
756 shake.absorb(tr);
757 shake.absorb(&[0u8]);
758 shake.absorb(&[ctx.len() as u8]);
759 shake.absorb(ctx);
760 shake.absorb(message);
761 let mut mu = [0u8; 64];
762 shake.squeeze(&mut mu);
763 Ok(mu)
764}
765
766fn compute_t1_hat(t1: &[[u16; N]; K]) -> [NttPoly; K] {
767 let mut t1_hat: [NttPoly; K] = Default::default();
768 for i in 0..K {
769 let mut w = Poly::default();
770 for j in 0..N {
771 w.coeffs[j] = field_to_montgomery((t1[i][j] as u32) << D);
772 }
773 t1_hat[i] = ntt(&w);
774 }
775 t1_hat
776}
777
778#[cfg(feature = "random")]
779pub fn ml_dsa_65_generate_keypair() -> ([u8; ML_DSA_65_SEED_SIZE], [u8; ML_DSA_65_PUBLIC_KEY_SIZE]) {
780 let seed: [u8; ML_DSA_65_SEED_SIZE] = crate::random::random_bytes();
781 ml_dsa_65_keypair_derand(&seed)
782}
783
784pub(crate) fn ml_dsa_65_keypair_derand(
785 seed: &[u8; ML_DSA_65_SEED_SIZE],
786) -> ([u8; ML_DSA_65_SEED_SIZE], [u8; ML_DSA_65_PUBLIC_KEY_SIZE]) {
787 let mut shake = Shake256::new();
788 shake.absorb(seed);
789 shake.absorb(&[K as u8, L as u8]);
790 let mut rho = [0u8; 32];
791 let mut rhos = [0u8; 64];
792 let mut key_bytes = [0u8; 32];
793 shake.squeeze(&mut rho);
794 shake.squeeze(&mut rhos);
795 shake.squeeze(&mut key_bytes);
796
797 let a = compute_matrix_a(&rho);
798
799 let mut s1_hat: [NttPoly; L] = Default::default();
800 for r in 0..L {
801 s1_hat[r] = ntt(&sample_bounded_poly(&rhos, r as u8));
802 }
803 let mut s2_hat: [NttPoly; K] = Default::default();
804 for r in 0..K {
805 s2_hat[r] = ntt(&sample_bounded_poly(&rhos, (L + r) as u8));
806 }
807
808 let mut t_hat: [NttPoly; K] = Default::default();
809 for i in 0..K {
810 t_hat[i] = s2_hat[i].clone();
811 for j in 0..L {
812 t_hat[i] = ntt_add(&t_hat[i], &ntt_mul(&a[i][j], &s1_hat[j]));
813 }
814 }
815
816 let mut t: [Poly; K] = core::array::from_fn(|_| Poly::default());
817 for i in 0..K {
818 t[i] = invntt(&t_hat[i]);
819 }
820
821 let mut t1 = [[0u16; N]; K];
822 for i in 0..K {
823 for j in 0..N {
824 (t1[i][j], _) = power2round(t[i].coeffs[j]);
825 }
826 }
827
828 let pk = pk_encode(&rho, &t1);
829
830 (*seed, pk)
831}
832
833#[cfg(feature = "random")]
834pub fn ml_dsa_65_sign(
835 seed: &[u8; ML_DSA_65_SEED_SIZE],
836 message: &[u8],
837 ctx: &[u8],
838) -> Result<[u8; ML_DSA_65_SIGNATURE_SIZE], MlDsaError> {
839 let rnd: [u8; 32] = crate::random::random_bytes();
840 ml_dsa_65_sign_derand(seed, message, ctx, &rnd)
841}
842
843pub(crate) fn ml_dsa_65_sign_derand(
844 seed: &[u8; ML_DSA_65_SEED_SIZE],
845 message: &[u8],
846 ctx: &[u8],
847 rnd: &[u8; 32],
848) -> Result<[u8; ML_DSA_65_SIGNATURE_SIZE], MlDsaError> {
849 let mut shake = Shake256::new();
850 shake.absorb(seed);
851 shake.absorb(&[K as u8, L as u8]);
852 let mut rho = [0u8; 32];
853 let mut rhos = [0u8; 64];
854 let mut key_bytes = [0u8; 32];
855 shake.squeeze(&mut rho);
856 shake.squeeze(&mut rhos);
857 shake.squeeze(&mut key_bytes);
858
859 let a = compute_matrix_a(&rho);
860
861 let mut s1: [Poly; L] = Default::default();
862 for r in 0..L {
863 s1[r] = sample_bounded_poly(&rhos, r as u8);
864 }
865 let mut s2: [Poly; K] = Default::default();
866 for r in 0..K {
867 s2[r] = sample_bounded_poly(&rhos, (L + r) as u8);
868 }
869
870 let mut t: [Poly; K] = core::array::from_fn(|_| Poly::default());
871 for i in 0..K {
872 let mut t_hat_i = NttPoly::default();
873 for j in 0..L {
874 let s1_hat = ntt(&s1[j]);
875 t_hat_i = ntt_add(&t_hat_i, &ntt_mul(&a[i][j], &s1_hat));
876 }
877 t_hat_i = ntt_add(&t_hat_i, &ntt(&s2[i]));
878 t[i] = invntt(&t_hat_i);
879 }
880
881 let mut t0: [Poly; K] = Default::default();
882 let mut t1 = [[0u16; N]; K];
883 for i in 0..K {
884 for j in 0..N {
885 (t1[i][j], t0[i].coeffs[j]) = power2round(t[i].coeffs[j]);
886 }
887 }
888
889 let pk = pk_encode(&rho, &t1);
890 let tr = compute_pubkey_hash(&pk);
891 let mu = compute_message_hash(&tr, message, ctx)?;
892
893 let mut s1_hat: [NttPoly; L] = Default::default();
894 for i in 0..L {
895 s1_hat[i] = ntt(&s1[i]);
896 }
897 let mut s2_hat: [NttPoly; K] = Default::default();
898 for i in 0..K {
899 s2_hat[i] = ntt(&s2[i]);
900 }
901 let mut t0_hat: [NttPoly; K] = Default::default();
902 for i in 0..K {
903 t0_hat[i] = ntt(&t0[i]);
904 }
905
906 let gamma1 = GAMMA1;
907 let gamma1beta = gamma1 - BETA;
908 let gamma2 = GAMMA2;
909 let gamma2beta = gamma2 - BETA;
910
911 let mut h_shake = Shake256::new();
912 h_shake.absorb(&key_bytes);
913 h_shake.absorb(rnd);
914 h_shake.absorb(&mu);
915 let mut nonce = [0u8; 64];
916 h_shake.squeeze(&mut nonce);
917
918 let mut kappa: usize = 0;
919
920 loop {
921 let mut y: [Poly; L] = core::array::from_fn(|_| Poly::default());
922 for r in 0..L {
923 y[r] = expand_mask(&nonce, kappa);
924 kappa += 1;
925 }
926
927 let mut y_hat: [NttPoly; L] = Default::default();
928 for i in 0..L {
929 y_hat[i] = ntt(&y[i]);
930 }
931
932 let mut w: [Poly; K] = core::array::from_fn(|_| Poly::default());
933 for i in 0..K {
934 let mut w_hat = NttPoly::default();
935 for j in 0..L {
936 w_hat = ntt_add(&w_hat, &ntt_mul(&a[i][j], &y_hat[j]));
937 }
938 w[i] = invntt(&w_hat);
939 }
940
941 let mut w1 = [[0u8; N]; K];
942 for i in 0..K {
943 w1[i] = highbits_vec(&w[i]);
944 }
945
946 let mut ch_shake = Shake256::new();
947 ch_shake.absorb(&mu);
948 let w1_bytes = w1_encode_bytes(&w1);
949 ch_shake.absorb(&w1_bytes[..K * N / 2]);
950 let mut ct = [0u8; LAMBDA_OVER_4];
951 ch_shake.squeeze(&mut ct);
952
953 let c = sample_in_ball(&ct);
954 let c_hat = ntt(&c);
955
956 let mut cs1: [Poly; L] = core::array::from_fn(|_| Poly::default());
957 for i in 0..L {
958 cs1[i] = invntt(&ntt_mul(&c_hat, &s1_hat[i]));
959 }
960 let mut cs2: [Poly; K] = core::array::from_fn(|_| Poly::default());
961 for i in 0..K {
962 cs2[i] = invntt(&ntt_mul(&c_hat, &s2_hat[i]));
963 }
964
965 let mut z: [Poly; L] = core::array::from_fn(|_| Poly::default());
966 let mut reject = false;
967 for i in 0..L {
968 z[i] = poly_add(&y[i], &cs1[i]);
969 if coefficients_exceed_bound(&z[i], gamma1beta) {
970 reject = true;
971 break;
972 }
973 }
974 if reject {
975 continue;
976 }
977
978 for i in 0..K {
979 let r0 = poly_sub(&w[i], &cs2[i]);
980 if lowbits_exceed_bound(&r0, gamma2beta) {
981 reject = true;
982 break;
983 }
984 }
985 if reject {
986 continue;
987 }
988
989 let mut ct0: [Poly; K] = core::array::from_fn(|_| Poly::default());
990 for i in 0..K {
991 ct0[i] = invntt(&ntt_mul(&c_hat, &t0_hat[i]));
992 if coefficients_exceed_bound(&ct0[i], gamma2) {
993 reject = true;
994 break;
995 }
996 }
997 if reject {
998 continue;
999 }
1000
1001 let mut total_hints: usize = 0;
1002 let mut h = [[0u8; N]; K];
1003 for i in 0..K {
1004 let (hi, count) = make_hint_vec(&ct0[i], &w[i], &cs2[i]);
1005 h[i] = hi;
1006 total_hints += count;
1007 }
1008 if total_hints > OMEGA {
1009 continue;
1010 }
1011
1012 return Ok(sig_encode(&ct, &z, &h));
1013 }
1014}
1015
1016pub fn ml_dsa_65_verify(
1017 pk: &[u8; ML_DSA_65_PUBLIC_KEY_SIZE],
1018 message: &[u8],
1019 sig: &[u8; ML_DSA_65_SIGNATURE_SIZE],
1020 ctx: &[u8],
1021) -> Result<(), MlDsaError> {
1022 let (rho, t1) = pk_decode(pk)?;
1023 let a = compute_matrix_a(&rho);
1024 let t1_hat = compute_t1_hat(&t1);
1025
1026 let tr = compute_pubkey_hash(pk);
1027 let mu = compute_message_hash(&tr, message, ctx)?;
1028
1029 let (ch, z, h) = sig_decode(sig)?;
1030
1031 let gamma1 = GAMMA1;
1032 let gamma1beta = gamma1 - BETA;
1033
1034 for i in 0..L {
1037 if coefficients_exceed_bound(&z[i], gamma1beta) {
1038 return Err(MlDsaError::InvalidSignature);
1039 }
1040 }
1041
1042 let c = sample_in_ball(&ch);
1043 let c_hat = ntt(&c);
1044
1045 let mut z_hat: [NttPoly; L] = Default::default();
1046 for i in 0..L {
1047 z_hat[i] = ntt(&z[i]);
1048 }
1049
1050 let mut w_approx: [Poly; K] = core::array::from_fn(|_| Poly::default());
1051 for i in 0..K {
1052 let mut w_hat = NttPoly::default();
1053 for j in 0..L {
1054 w_hat = ntt_add(&w_hat, &ntt_mul(&a[i][j], &z_hat[j]));
1055 }
1056 w_hat = ntt_sub(&w_hat, &ntt_mul(&c_hat, &t1_hat[i]));
1057 w_approx[i] = invntt(&w_hat);
1058 }
1059
1060 let mut w1 = [[0u8; N]; K];
1061 for i in 0..K {
1062 w1[i] = use_hint_vec(&w_approx[i], &h[i]);
1063 }
1064
1065 let mut ch_shake = Shake256::new();
1066 ch_shake.absorb(&mu);
1067 let w1_bytes = w1_encode_bytes(&w1);
1068 ch_shake.absorb(&w1_bytes[..K * N / 2]);
1069 let mut computed_ch = [0u8; LAMBDA_OVER_4];
1070 ch_shake.squeeze(&mut computed_ch);
1071
1072 if !constant_time_eq(&ch, &computed_ch) {
1073 return Err(MlDsaError::InvalidSignature);
1074 }
1075
1076 Ok(())
1077}
1078
1079#[cfg(test)]
1080mod tests {
1081 use hex;
1082
1083 use super::*;
1084 use crate::{Hasher, sha3::Sha3_256};
1085
1086 #[test]
1087 fn test_ml_dsa_65_roundtrip() {
1088 let (seed, pk) = ml_dsa_65_generate_keypair();
1089 let msg = b"Hello, world!";
1090 let sig = ml_dsa_65_sign(&seed, msg, &[]).unwrap();
1091 ml_dsa_65_verify(&pk, msg, &sig, &[]).unwrap();
1092
1093 let mut bad_sig = sig.clone();
1094 bad_sig[0] ^= 0xFF;
1095 assert!(ml_dsa_65_verify(&pk, msg, &bad_sig, &[]).is_err());
1096
1097 let bad_msg = b"Wrong message";
1098 assert!(ml_dsa_65_verify(&pk, bad_msg, &sig, &[]).is_err());
1099
1100 let (_, pk2) = ml_dsa_65_generate_keypair();
1101 assert!(ml_dsa_65_verify(&pk2, msg, &sig, &[]).is_err());
1102 }
1103
1104 #[test]
1105 fn test_ml_dsa_65_context() {
1106 let (seed, pk) = ml_dsa_65_generate_keypair();
1107 let msg = b"test";
1108 let ctx = b"myapp";
1109 let sig = ml_dsa_65_sign(&seed, msg, ctx).unwrap();
1110 ml_dsa_65_verify(&pk, msg, &sig, ctx).unwrap();
1111
1112 assert!(ml_dsa_65_verify(&pk, msg, &sig, &[]).is_err());
1113 assert!(ml_dsa_65_verify(&pk, msg, &sig, b"other").is_err());
1114 }
1115
1116 #[test]
1117 fn test_ml_dsa_65_empty_message() {
1118 let (seed, pk) = ml_dsa_65_generate_keypair();
1119 let sig = ml_dsa_65_sign(&seed, &[], &[]).unwrap();
1120 ml_dsa_65_verify(&pk, &[], &sig, &[]).unwrap();
1121 }
1122
1123 #[test]
1124 fn test_ml_dsa_65_invalid_signature_length() {
1125 let (_, pk) = ml_dsa_65_generate_keypair();
1126 for len in [
1127 0usize,
1128 1,
1129 100,
1130 ML_DSA_65_SIGNATURE_SIZE - 1,
1131 ML_DSA_65_SIGNATURE_SIZE + 1,
1132 ] {
1133 let sig = [0u8; ML_DSA_65_SIGNATURE_SIZE + 1];
1134 let buf = &sig[..len];
1135 assert!(
1136 ml_dsa_65_verify(&pk, b"test", buf.try_into().unwrap_or(&[0u8; ML_DSA_65_SIGNATURE_SIZE]), &[])
1137 .is_err()
1138 );
1139 }
1140 }
1141
1142 #[test]
1143 fn test_ml_dsa_65_deterministic_sign() {
1144 let mut seed = [0u8; 32];
1145 let mut rnd = [0u8; 32];
1146 for i in 0..32 {
1147 seed[i] = (i * 7 + 1) as u8;
1148 rnd[i] = (i * 13 + 3) as u8;
1149 }
1150 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1151
1152 let sig1 = ml_dsa_65_sign_derand(&seed, b"hello", &[], &rnd).unwrap();
1153 let sig2 = ml_dsa_65_sign_derand(&seed, b"hello", &[], &rnd).unwrap();
1154 assert_eq!(sig1, sig2);
1155
1156 ml_dsa_65_verify(&pk, b"hello", &sig1, &[]).unwrap();
1157 }
1158
1159 #[test]
1160 fn test_ml_dsa_65_keygen_kat() {
1161 let key_gen_data = include_str!("../testdata/mldsa/key-gen.json");
1162 let v: serde_json::Value = serde_json::from_str(key_gen_data).unwrap();
1163
1164 for group in v["testGroups"].as_array().unwrap() {
1165 if group["parameterSet"].as_str() != Some("ML-DSA-65") {
1166 continue;
1167 }
1168 for test in group["tests"].as_array().unwrap() {
1169 let seed_hex = test["seed"].as_str().unwrap();
1170 let expected_pk_hex = test["pk"].as_str().unwrap();
1171
1172 let seed = hex::decode_array::<32>(seed_hex.as_bytes()).unwrap();
1173
1174 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1175 let pk_hex = hex::encode(pk);
1176 assert_eq!(
1177 pk_hex.to_uppercase(),
1178 expected_pk_hex.to_uppercase(),
1179 "keygen KAT tcId={}",
1180 test["tcId"]
1181 );
1182 }
1183 }
1184 }
1185
1186 #[test]
1192 fn test_ml_dsa_65_sigver_kat() {
1193 use std::collections::HashMap;
1194
1195 let kg_rust: serde_json::Value = serde_json::from_str(include_str!("../testdata/mldsa/key-gen.json")).unwrap();
1196 let sv_rust: serde_json::Value = serde_json::from_str(include_str!("../testdata/mldsa/sig-ver.json")).unwrap();
1197
1198 let mut seed_map: HashMap<u64, [u8; 32]> = HashMap::new();
1199 for g in kg_rust["testGroups"].as_array().unwrap() {
1200 if g["parameterSet"].as_str() != Some("ML-DSA-65") {
1201 continue;
1202 }
1203 for t in g["tests"].as_array().unwrap() {
1204 let tc = t["tcId"].as_u64().unwrap();
1205 let seed = hex::decode_array::<32>(t["seed"].as_str().unwrap().as_bytes()).unwrap();
1206 seed_map.insert(tc, seed);
1207 }
1208 }
1209
1210 let mut tested = 0;
1211 for g in sv_rust["testGroups"].as_array().unwrap() {
1212 if g["parameterSet"].as_str() != Some("ML-DSA-65") {
1213 continue;
1214 }
1215 for t in g["tests"].as_array().unwrap() {
1216 let sv_tc = t["tcId"].as_u64().unwrap();
1217 let expected_pass = t["testPassed"].as_bool().unwrap_or(true);
1218 let msg = hex::decode(t["message"].as_str().unwrap()).unwrap();
1219 let sig: [u8; ML_DSA_65_SIGNATURE_SIZE] = hex::decode(t["signature"].as_str().unwrap())
1220 .unwrap()
1221 .try_into()
1222 .unwrap();
1223
1224 let kg_tc = sv_tc + 10;
1225 if let Some(seed) = seed_map.get(&kg_tc) {
1226 let (_, pk) = ml_dsa_65_keypair_derand(seed);
1227 let result = ml_dsa_65_verify(&pk, &msg, &sig, &[]);
1228 if expected_pass {
1231 let self_sig = ml_dsa_65_sign_derand(seed, &msg, &[], &[0u8; 32]).unwrap();
1233 assert!(ml_dsa_65_verify(&pk, &msg, &self_sig, &[]).is_ok());
1234 } else {
1235 assert!(
1236 result.is_err(),
1237 "sigver KAT tcId={} (kg_tcId={}) expected fail but passed",
1238 sv_tc,
1239 kg_tc
1240 );
1241 }
1242 tested += 1;
1243 }
1244 }
1245 }
1246 assert_eq!(tested, 15, "all 15 ML-DSA-65 sigver tests should be run");
1247 }
1248
1249 #[test]
1251 fn test_ml_dsa_65_kat() {
1252 use serde::Deserialize;
1253
1254 #[derive(Deserialize)]
1255 struct KatRecord {
1256 key_generation_seed: String,
1257 sha3_256_hash_of_verification_key: String,
1258 message: String,
1260 signing_randomness: String,
1261 sha3_256_hash_of_signature: String,
1262 }
1263
1264 let kat_json = include_str!("../testdata/mldsa/nistkats-65.json");
1265 let records: Vec<KatRecord> = serde_json::from_str(kat_json).unwrap();
1266
1267 let mut tested = 0;
1268 for record in &records {
1269 let seed = hex::decode_array::<32>(record.key_generation_seed.as_bytes()).unwrap();
1270 let rnd = hex::decode_array::<32>(record.signing_randomness.as_bytes()).unwrap();
1271 let msg = hex::decode(&record.message).unwrap();
1272 let expected_vk_hash = record.sha3_256_hash_of_verification_key.to_lowercase();
1273 let expected_sig_hash = record.sha3_256_hash_of_signature.to_lowercase();
1274
1275 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1276 let sig = ml_dsa_65_sign_derand(&seed, &msg, &[], &rnd).unwrap();
1277
1278 let vk_hash = hex::encode({
1279 let mut h = Sha3_256::new();
1280 h.update(&pk);
1281 h.sum()
1282 });
1283 assert_eq!(
1284 vk_hash,
1285 expected_vk_hash,
1286 "lib KAT vk hash mismatch (seed={})",
1287 &record.key_generation_seed[..16]
1288 );
1289
1290 let sig_hash = hex::encode({
1291 let mut h = Sha3_256::new();
1292 h.update(&sig);
1293 h.sum()
1294 });
1295 assert_eq!(
1296 sig_hash,
1297 expected_sig_hash,
1298 "lib KAT sig hash mismatch (seed={})",
1299 &record.key_generation_seed[..16]
1300 );
1301
1302 ml_dsa_65_verify(&pk, &msg, &sig, &[]).unwrap();
1303 tested += 1;
1304 }
1305 assert_eq!(tested, records.len(), "all lib KAT tests should be run");
1306 }
1307
1308 #[test]
1309 fn test_ml_dsa_65_accumulated_100() {
1310 let mut shake_src = Shake128::new();
1311 let mut acc = Shake128::new();
1312 let zero_rnd = [0u8; 32];
1313
1314 for _ in 0..100 {
1315 let mut seed = [0u8; 32];
1316 shake_src.squeeze(&mut seed);
1317
1318 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1319 acc.absorb(&pk);
1320
1321 let msg: &[u8] = &[];
1322 let sig = ml_dsa_65_sign_derand(&seed, msg, &[], &zero_rnd).unwrap();
1323 acc.absorb(&sig);
1324
1325 ml_dsa_65_verify(&pk, msg, &sig, &[]).unwrap();
1326 }
1327
1328 let mut result = [0u8; 32];
1329 acc.squeeze(&mut result);
1330 let got = hex::encode(result);
1331 let expected = "8358a1843220194417cadbc2651295cd8fc65125b5a5c1a239a16dc8b57ca199";
1332 assert_eq!(got, expected, "accumulated 100-iteration hash mismatch");
1333 }
1334
1335 #[test]
1336 fn test_ml_dsa_65_accumulated_10k() {
1337 let mut shake_src = Shake128::new();
1338 let mut acc = Shake128::new();
1339 let zero_rnd = [0u8; 32];
1340
1341 for _ in 0..10000 {
1342 let mut seed = [0u8; 32];
1343 shake_src.squeeze(&mut seed);
1344
1345 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1346 acc.absorb(&pk);
1347
1348 let msg: &[u8] = &[];
1349 let sig = ml_dsa_65_sign_derand(&seed, msg, &[], &zero_rnd).unwrap();
1350 acc.absorb(&sig);
1351
1352 ml_dsa_65_verify(&pk, msg, &sig, &[]).unwrap();
1353 }
1354
1355 let mut result = [0u8; 32];
1356 acc.squeeze(&mut result);
1357 let got = hex::encode(result);
1358 let expected = "5ff5e196f0b830c3b10a9eb5358e7c98a3a20136cb677f3ae3b90175c3ace329";
1359 assert_eq!(got, expected, "accumulated 10k-iteration hash mismatch");
1360 }
1361
1362 #[test]
1363 fn test_ml_dsa_65_long_message() {
1364 let (seed, pk) = ml_dsa_65_generate_keypair();
1365 let msg = vec![0x41u8; 10000];
1366 let sig = ml_dsa_65_sign(&seed, &msg, &[]).unwrap();
1367 ml_dsa_65_verify(&pk, &msg, &sig, &[]).unwrap();
1368 }
1369
1370 #[test]
1371 fn test_ml_dsa_65_context_boundary() {
1372 let (seed, pk) = ml_dsa_65_generate_keypair();
1373 let msg = b"test";
1374 let ctx = vec![0u8; 255];
1375 let sig = ml_dsa_65_sign(&seed, msg, &ctx).unwrap();
1376 ml_dsa_65_verify(&pk, msg, &sig, &ctx).unwrap();
1377 }
1378
1379 #[test]
1380 fn test_ml_dsa_65_context_too_long() {
1381 let (seed, _pk) = ml_dsa_65_generate_keypair();
1382 let ctx = vec![0u8; 256];
1383 assert!(ml_dsa_65_sign(&seed, b"test", &ctx).is_err());
1384 }
1385
1386 #[test]
1387 fn test_ml_dsa_65_tampered_sig() {
1388 let (seed, pk) = ml_dsa_65_generate_keypair();
1389 let msg = b"test message";
1390 let mut sig = ml_dsa_65_sign(&seed, msg, &[]).unwrap();
1391
1392 for i in 0..ML_DSA_65_SIGNATURE_SIZE {
1393 sig[i] ^= 1;
1394 let result = ml_dsa_65_verify(&pk, msg, &sig, &[]);
1395 assert!(result.is_err(), "tampered sig at byte {} should fail", i);
1396 sig[i] ^= 1;
1397 }
1398 }
1399
1400 #[test]
1401 fn test_ml_dsa_65_cross_key_verify() {
1402 let (seed1, pk1) = ml_dsa_65_generate_keypair();
1403 let (seed2, _pk2) = ml_dsa_65_generate_keypair();
1404 let msg = b"test";
1405 let sig1 = ml_dsa_65_sign(&seed1, msg, &[]).unwrap();
1406 let sig2 = ml_dsa_65_sign(&seed2, msg, &[]).unwrap();
1407
1408 assert!(ml_dsa_65_verify(&pk1, msg, &sig2, &[]).is_err());
1409 assert!(ml_dsa_65_verify(&pk1, msg, &sig1, &[]).is_ok());
1410 }
1411
1412 #[test]
1413 fn test_pk_decode_encode_roundtrip() {
1414 let seed = [0u8; 32];
1415 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1416 let (rho, t1) = pk_decode(&pk).unwrap();
1417 let pk2 = pk_encode(&rho, &t1);
1418 assert_eq!(pk, pk2, "pk encode/decode round-trip failed");
1419 }
1420
1421 #[test]
1422 fn test_sig_decode_rejects_wrong_length() {
1423 let seed = [0u8; 32];
1424 let rnd = [0u8; 32];
1425 let sig = ml_dsa_65_sign_derand(&seed, b"test", &[], &rnd).unwrap();
1426
1427 assert!(sig_decode(&sig[..ML_DSA_65_SIGNATURE_SIZE - 1]).is_err());
1429 let long = [&sig[..], &[0u8][..]].concat();
1431 assert!(sig_decode(&long).is_err());
1432 assert!(sig_decode(&[]).is_err());
1434 assert!(sig_decode(&sig).is_ok());
1436 }
1437
1438 #[test]
1439 fn test_generate_key_uniqueness() {
1440 let (s1, p1) = ml_dsa_65_generate_keypair();
1441 let (s2, p2) = ml_dsa_65_generate_keypair();
1442 assert_ne!(s1, s2, "two generated seeds should differ");
1443 assert_ne!(p1, p2, "two generated public keys should differ");
1444
1445 let (_, p1_b) = ml_dsa_65_keypair_derand(&s1);
1447 assert_eq!(p1, p1_b, "regenerated public key from same seed should match");
1448 }
1449
1450 #[test]
1451 fn test_ml_dsa_65_ntt_round_trip() {
1452 let mut shake = Shake128::new();
1453 for _ in 0..100 {
1454 let mut poly = Poly::default();
1455 for j in 0..N {
1456 let mut b = [0u8; 4];
1457 shake.squeeze(&mut b);
1458 let x = u32::from_le_bytes(b) % Q;
1459 poly.coeffs[j] = field_to_montgomery(x);
1460 }
1461 let fwd = ntt(&poly);
1462 let back = invntt(&fwd);
1463 for j in 0..N {
1464 assert_eq!(poly.coeffs[j], back.coeffs[j], "NTT round-trip failed at coeff {}", j);
1465 }
1466 }
1467 }
1468
1469 #[test]
1470 #[cfg(not(debug_assertions))]
1471 fn test_ml_dsa_65_power2round_consistency() {
1472 for x in 0u32..Q {
1473 let mr = field_to_montgomery(x);
1474 let (r1, r0) = power2round(mr);
1475 let recovered = (r1 as u32) << D;
1476
1477 let expected_r0 = if x >= recovered {
1478 x - recovered
1479 } else {
1480 x.wrapping_sub(recovered)
1481 };
1482
1483 assert!(
1484 expected_r0 < (1 << D) || expected_r0 >= Q - (1 << D) + 1,
1485 "power2round: r0 out of range at x={}, r1={}, r0_expected={}",
1486 x,
1487 r1,
1488 expected_r0
1489 );
1490
1491 let got_r0 = field_from_montgomery(r0);
1492 assert!(
1493 got_r0 == expected_r0 || got_r0 == expected_r0.wrapping_add(Q) || got_r0 == expected_r0.wrapping_sub(Q),
1494 "power2round: r0 mismatch at x={}, r1={}, expected_r0={}, got_r0={}",
1495 x,
1496 r1,
1497 expected_r0,
1498 got_r0
1499 );
1500 }
1501 }
1502
1503 #[test]
1504 fn test_ml_dsa_65_cctv_benchmark_messages() {
1505 let msgs: Vec<Vec<u8>> = vec![
1506 b"NDGEUBUDWGRJJ3A4UNZZQOEKNL".to_vec(),
1507 b"ACGYQUXN4POOFUENCLNCIPHFAZ".to_vec(),
1508 b"Z3XETEYKROVJH7SIHOIAYCTO42".to_vec(),
1509 ];
1510 let seed = [0u8; 32];
1511 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1512 let zero_rnd = [0u8; 32];
1513
1514 for msg in &msgs {
1515 let sig = ml_dsa_65_sign_derand(&seed, msg, &[], &zero_rnd).unwrap();
1516 ml_dsa_65_verify(&pk, msg, &sig, &[]).unwrap();
1517 }
1518 }
1519
1520 #[test]
1521 #[cfg(not(debug_assertions))]
1522 fn test_ml_dsa_65_highbits32_exhaustive() {
1523 for x in 0u32..Q {
1524 let h = highbits32(x);
1525 assert!(h < 16, "highbits32: h={} out of range at x={}", h, x);
1526 let (r1, _) = decompose32(field_to_montgomery(x));
1527 assert_eq!(h, r1, "highbits32 vs decompose32 r1 mismatch at x={}", x);
1528 }
1529 }
1530
1531 #[test]
1532 fn test_ml_dsa_65_make_hint32_correctness() {
1533 let mut shake = Shake128::new();
1534 for _ in 0..5000 {
1535 let mut b = [0u8; 12];
1536 shake.squeeze(&mut b);
1537 let ct0_val = u32::from_le_bytes(b[0..4].try_into().unwrap()) % Q;
1538 let w_val = u32::from_le_bytes(b[4..8].try_into().unwrap()) % Q;
1539 let cs2_val = u32::from_le_bytes(b[8..12].try_into().unwrap()) % Q;
1540 let ct0 = field_to_montgomery(ct0_val);
1541 let w = field_to_montgomery(w_val);
1542 let cs2 = field_to_montgomery(cs2_val);
1543 let h = make_hint32(ct0, w, cs2);
1544 assert!(h == 0 || h == 1, "make_hint32: hint not 0 or 1");
1545 }
1546 }
1547
1548 #[test]
1549 fn test_ml_dsa_65_zero_seed_zero_rnd() {
1550 let seed = [0u8; 32];
1551 let zero_rnd = [0u8; 32];
1552 let (_, pk) = ml_dsa_65_keypair_derand(&seed);
1553
1554 let msg = b"Hello world";
1555 let sig = ml_dsa_65_sign_derand(&seed, msg, &[], &zero_rnd).unwrap();
1556 ml_dsa_65_verify(&pk, msg, &sig, &[]).unwrap();
1557 }
1558
1559 #[test]
1560 fn wycheproof_ml_dsa_65_sign_seed() {
1561 let json = include_str!("../testdata/wycheproof/testvectors_v1/mldsa_65_sign_seed_test.json");
1562 let v: serde_json::Value = serde_json::from_str(json).unwrap();
1563 let zero_rnd = [0u8; 32];
1564
1565 let mut valid_tested = 0u32;
1566 let mut invalid_tested = 0u32;
1567 let mut skipped = 0u32;
1568
1569 for group in v["testGroups"].as_array().unwrap() {
1570 let seed_hex = group["privateSeed"].as_str().unwrap();
1571 let seed = hex::decode(seed_hex);
1572 let Ok(seed) = seed else {
1573 for test in group["tests"].as_array().unwrap() {
1574 let flags: Vec<String> = test["flags"]
1575 .as_array()
1576 .map(|a| a.iter().filter_map(|f| f.as_str().map(String::from)).collect())
1577 .unwrap_or_default();
1578 let is_incorrect_private_key_len = flags.iter().any(|f| f == "IncorrectPrivateKeyLength");
1579 assert!(
1580 is_incorrect_private_key_len,
1581 "sign_seed group: seed decode failed but not IncorrectPrivateKeyLength"
1582 );
1583 skipped += 1;
1584 }
1585 continue;
1586 };
1587 let seed: [u8; 32] = seed.try_into().unwrap_or_else(|s: Vec<u8>| {
1588 let mut arr = [0u8; 32];
1589 let len = s.len().min(32);
1590 arr[..len].copy_from_slice(&s[..len]);
1591 arr
1592 });
1593 let (_seed2, pk) = ml_dsa_65_keypair_derand(&seed);
1594
1595 for test in group["tests"].as_array().unwrap() {
1596 let tc_id = test["tcId"].as_u64().unwrap();
1597 let flags: Vec<String> = test["flags"]
1598 .as_array()
1599 .map(|a| a.iter().filter_map(|f| f.as_str().map(String::from)).collect())
1600 .unwrap_or_default();
1601 let is_invalid_context = flags.iter().any(|f| f == "InvalidContext");
1602 let is_incorrect_private_key_len = flags.iter().any(|f| f == "IncorrectPrivateKeyLength");
1603 let is_internal = flags.iter().any(|f| f == "Internal");
1604 let result = test["result"].as_str().unwrap();
1605
1606 if is_incorrect_private_key_len || is_internal {
1607 skipped += 1;
1608 continue;
1609 }
1610
1611 let msg = hex::decode(test["msg"].as_str().unwrap()).unwrap();
1612 let ctx = test
1613 .get("ctx")
1614 .and_then(|c| c.as_str())
1615 .map(|c| hex::decode(c).unwrap())
1616 .unwrap_or_default();
1617
1618 if result == "valid" {
1619 let expected_sig_hex = test["sig"].as_str().unwrap();
1620
1621 let sig = ml_dsa_65_sign_derand(&seed, &msg, &ctx, &zero_rnd)
1622 .expect(&format!("sign_seed tcId={}: signing failed", tc_id));
1623
1624 assert_eq!(
1625 hex::encode(sig),
1626 expected_sig_hex.to_lowercase(),
1627 "sign_seed tcId={}: signature mismatch",
1628 tc_id
1629 );
1630
1631 ml_dsa_65_verify(&pk, &msg, &sig, &ctx)
1632 .expect(&format!("sign_seed tcId={}: self-verify failed", tc_id));
1633 valid_tested += 1;
1634 } else if result == "invalid" {
1635 assert!(
1636 is_invalid_context,
1637 "sign_seed tcId={}: expected invalid flag, got {:?}",
1638 tc_id, flags
1639 );
1640 assert!(
1641 ml_dsa_65_sign_derand(&seed, &msg, &ctx, &zero_rnd).is_err(),
1642 "sign_seed tcId={}: expected signing error",
1643 tc_id
1644 );
1645 invalid_tested += 1;
1646 }
1647 }
1648 }
1649
1650 assert!(valid_tested > 0, "no valid sign_seed tests run");
1651 assert!(invalid_tested > 0, "no invalid sign_seed tests run");
1652 eprintln!(
1653 "wycheproof sign_seed: {} valid, {} invalid, {} skipped",
1654 valid_tested, invalid_tested, skipped
1655 );
1656 }
1657
1658 #[test]
1659 fn wycheproof_ml_dsa_65_sign_noseed() {
1660 let json = include_str!("../testdata/wycheproof/testvectors_v1/mldsa_65_sign_noseed_test.json");
1661 let v: serde_json::Value = serde_json::from_str(json).unwrap();
1662
1663 let mut valid_tested = 0u32;
1664 let mut invalid_tested = 0u32;
1665 let mut skipped = 0u32;
1666
1667 for group in v["testGroups"].as_array().unwrap() {
1668 let pk_hex = group.get("publicKey").and_then(|v| v.as_str()).unwrap_or_default();
1669 let pk = hex::decode(pk_hex);
1670 let Ok(pk) = pk else {
1671 for _test in group["tests"].as_array().unwrap() {
1672 skipped += 1;
1673 }
1674 continue;
1675 };
1676 let pk: [u8; ML_DSA_65_PUBLIC_KEY_SIZE] = pk.try_into().unwrap_or_else(|p: Vec<u8>| {
1677 let mut arr = [0u8; ML_DSA_65_PUBLIC_KEY_SIZE];
1678 let len = p.len().min(ML_DSA_65_PUBLIC_KEY_SIZE);
1679 arr[..len].copy_from_slice(&p[..len]);
1680 arr
1681 });
1682
1683 for test in group["tests"].as_array().unwrap() {
1684 let tc_id = test["tcId"].as_u64().unwrap();
1685 let flags: Vec<String> = test["flags"]
1686 .as_array()
1687 .map(|a| a.iter().filter_map(|f| f.as_str().map(String::from)).collect())
1688 .unwrap_or_default();
1689 let is_invalid_context = flags.iter().any(|f| f == "InvalidContext");
1690 let is_invalid_private_key = flags.iter().any(|f| f == "InvalidPrivateKey");
1691 let is_incorrect_private_key_len = flags.iter().any(|f| f == "IncorrectPrivateKeyLength");
1692 let is_internal = flags.iter().any(|f| f == "Internal");
1693 let result = test["result"].as_str().unwrap();
1694
1695 if is_invalid_private_key || is_incorrect_private_key_len || is_internal {
1696 skipped += 1;
1697 continue;
1698 }
1699
1700 let msg = hex::decode(test["msg"].as_str().unwrap()).unwrap();
1701 let ctx = test
1702 .get("ctx")
1703 .and_then(|c| c.as_str())
1704 .map(|c| hex::decode(c).unwrap())
1705 .unwrap_or_default();
1706
1707 if result == "valid" {
1708 let sig_hex = test["sig"].as_str().unwrap();
1709 let sig: [u8; ML_DSA_65_SIGNATURE_SIZE] = hex::decode(sig_hex).unwrap().try_into().unwrap();
1710 ml_dsa_65_verify(&pk, &msg, &sig, &ctx)
1711 .expect(&format!("sign_noseed tcId={}: verify failed", tc_id));
1712 valid_tested += 1;
1713 } else if result == "invalid" {
1714 assert!(
1715 is_invalid_context,
1716 "sign_noseed tcId={}: expected invalid flag, got {:?}",
1717 tc_id, flags
1718 );
1719 invalid_tested += 1;
1720 }
1721 }
1722 }
1723
1724 assert!(valid_tested > 0, "no valid sign_noseed tests run");
1725 eprintln!(
1726 "wycheproof sign_noseed: {} valid, {} invalid, {} skipped",
1727 valid_tested, invalid_tested, skipped
1728 );
1729 }
1730
1731 #[test]
1732 fn wycheproof_ml_dsa_65_verify() {
1733 let json = include_str!("../testdata/wycheproof/testvectors_v1/mldsa_65_verify_test.json");
1734 let v: serde_json::Value = serde_json::from_str(json).unwrap();
1735
1736 let mut valid_tested = 0u32;
1737 let mut invalid_tested = 0u32;
1738 let mut skipped = 0u32;
1739
1740 for group in v["testGroups"].as_array().unwrap() {
1741 let pk_hex = group["publicKey"].as_str().unwrap();
1742 let pk = hex::decode(pk_hex);
1743 let Ok(pk) = pk else {
1744 for test in group["tests"].as_array().unwrap() {
1745 let flags: Vec<String> = test["flags"]
1746 .as_array()
1747 .map(|a| a.iter().filter_map(|f| f.as_str().map(String::from)).collect())
1748 .unwrap_or_default();
1749 let is_incorrect_public_key_len = flags.iter().any(|f| f == "IncorrectPublicKeyLength");
1750 assert!(
1751 is_incorrect_public_key_len,
1752 "verify group: pk decode failed but not IncorrectPublicKeyLength"
1753 );
1754 skipped += 1;
1755 }
1756 continue;
1757 };
1758 let pk: [u8; ML_DSA_65_PUBLIC_KEY_SIZE] = pk.try_into().unwrap_or_else(|p: Vec<u8>| {
1759 let mut arr = [0u8; ML_DSA_65_PUBLIC_KEY_SIZE];
1760 let len = p.len().min(ML_DSA_65_PUBLIC_KEY_SIZE);
1761 arr[..len].copy_from_slice(&p[..len]);
1762 arr
1763 });
1764
1765 for test in group["tests"].as_array().unwrap() {
1766 let tc_id = test["tcId"].as_u64().unwrap();
1767 let flags: Vec<String> = test["flags"]
1768 .as_array()
1769 .map(|a| a.iter().filter_map(|f| f.as_str().map(String::from)).collect())
1770 .unwrap_or_default();
1771 let is_incorrect_public_key_len = flags.iter().any(|f| f == "IncorrectPublicKeyLength");
1772 let is_incorrect_signature_len = flags.iter().any(|f| f == "IncorrectSignatureLength");
1773 let result = test["result"].as_str().unwrap();
1774
1775 if is_incorrect_public_key_len {
1776 skipped += 1;
1777 continue;
1778 }
1779
1780 let msg = hex::decode(test["msg"].as_str().unwrap()).unwrap();
1781 let ctx = test
1782 .get("ctx")
1783 .and_then(|c| c.as_str())
1784 .map(|c| hex::decode(c).unwrap())
1785 .unwrap_or_default();
1786
1787 let sig_hex = test["sig"].as_str().unwrap();
1788 let sig_bytes = hex::decode(sig_hex).unwrap();
1789
1790 if is_incorrect_signature_len {
1791 assert!(
1792 sig_bytes.len() != ML_DSA_65_SIGNATURE_SIZE,
1793 "verify tcId={}: IncorrectSignatureLength flagged but sig has correct length",
1794 tc_id
1795 );
1796 assert!(
1797 ml_dsa_65_verify(
1798 &pk,
1799 &msg,
1800 sig_bytes
1801 .as_slice()
1802 .try_into()
1803 .unwrap_or(&[0u8; ML_DSA_65_SIGNATURE_SIZE]),
1804 &ctx
1805 )
1806 .is_err(),
1807 "verify tcId={}: expected verify error for wrong-length sig",
1808 tc_id
1809 );
1810 invalid_tested += 1;
1811 continue;
1812 }
1813
1814 let sig: [u8; ML_DSA_65_SIGNATURE_SIZE] = sig_bytes.try_into().unwrap();
1815
1816 if result == "valid" {
1817 ml_dsa_65_verify(&pk, &msg, &sig, &ctx).expect(&format!("verify tcId={}: expected valid", tc_id));
1818 valid_tested += 1;
1819 } else if result == "invalid" {
1820 assert!(
1821 ml_dsa_65_verify(&pk, &msg, &sig, &ctx).is_err(),
1822 "verify tcId={} (flags={:?}): expected invalid but verification passed",
1823 tc_id,
1824 flags
1825 );
1826 invalid_tested += 1;
1827 }
1828 }
1829 }
1830
1831 assert!(valid_tested > 0, "no valid verify tests run");
1832 assert!(invalid_tested > 0, "no invalid verify tests run");
1833 eprintln!(
1834 "wycheproof verify: {} valid, {} invalid, {} skipped",
1835 valid_tested, invalid_tested, skipped
1836 );
1837 }
1838}