1
2
3
4
5 package mldsa
6
7 import (
8 "crypto/internal/fips140/sha3"
9 "encoding/hex"
10 "fmt"
11 "math/big"
12 "testing"
13 )
14
15 type interestingValue struct {
16 v uint32
17 m fieldElement
18 }
19
20
21
22
23 func interestingValues() []interestingValue {
24 if testing.Short() {
25 return []interestingValue{{v: q - 1, m: minusOne}}
26 }
27 var values []interestingValue
28 for _, v := range []uint32{
29 0,
30 1,
31 2,
32 3,
33 q - 3,
34 q - 2,
35 q - 1,
36 q / 2,
37 (q + 1) / 2,
38 } {
39 m, _ := fieldToMontgomery(v)
40 values = append(values, interestingValue{v: v, m: m})
41
42 values = append(values, interestingValue{
43 v: fieldFromMontgomery(fieldElement(v)), m: fieldElement(v)})
44 }
45 return values
46 }
47
48 func TestToFromMontgomery(t *testing.T) {
49 for a := range uint32(q) {
50 m, err := fieldToMontgomery(a)
51 if err != nil {
52 t.Fatalf("fieldToMontgomery(%d) returned error: %v", a, err)
53 }
54 exp := fieldElement((uint64(a) * R) % q)
55 if m != exp {
56 t.Fatalf("fieldToMontgomery(%d) = %d, expected %d", a, m, exp)
57 }
58 got := fieldFromMontgomery(m)
59 if got != a {
60 t.Fatalf("fieldFromMontgomery(fieldToMontgomery(%d)) = %d, expected %d", a, got, a)
61 }
62 }
63 }
64
65 func TestFieldAdd(t *testing.T) {
66 t.Parallel()
67 for _, a := range interestingValues() {
68 for b := range fieldElement(q) {
69 got := fieldAdd(a.m, b)
70 exp := (a.m + b) % q
71 if got != exp {
72 t.Fatalf("%d + %d = %d, expected %d", a, b, got, exp)
73 }
74 }
75 }
76 }
77
78 func TestFieldSub(t *testing.T) {
79 t.Parallel()
80 for _, a := range interestingValues() {
81 for b := range fieldElement(q) {
82 got := fieldSub(a.m, b)
83 exp := (a.m + q - b) % q
84 if got != exp {
85 t.Fatalf("%d - %d = %d, expected %d", a, b, got, exp)
86 }
87 }
88 }
89 }
90
91 func TestFieldSubToMontgomery(t *testing.T) {
92 t.Parallel()
93 for _, a := range interestingValues() {
94 for b := range uint32(q) {
95 got := fieldSubToMontgomery(a.v, b)
96 diff := (a.v + q - b) % q
97 exp := fieldElement((uint64(diff) * R) % q)
98 if got != exp {
99 t.Fatalf("fieldSubToMontgomery(%d, %d) = %d, expected %d", a.v, b, got, exp)
100 }
101 }
102 }
103 }
104
105 func TestFieldReduceOnce(t *testing.T) {
106 t.Parallel()
107 for a := range uint32(2 * q) {
108 got := fieldReduceOnce(a)
109 var exp uint32
110 if a < q {
111 exp = a
112 } else {
113 exp = a - q
114 }
115 if uint32(got) != exp {
116 t.Fatalf("fieldReduceOnce(%d) = %d, expected %d", a, got, exp)
117 }
118 }
119 }
120
121 func TestFieldMul(t *testing.T) {
122 t.Parallel()
123 for _, a := range interestingValues() {
124 for b := range fieldElement(q) {
125 got := fieldFromMontgomery(fieldMontgomeryMul(a.m, b))
126 exp := uint32((uint64(a.v) * uint64(fieldFromMontgomery(b))) % q)
127 if got != exp {
128 t.Fatalf("%d * %d = %d, expected %d", a, b, got, exp)
129 }
130 }
131 }
132 }
133
134 func TestFieldToMontgomeryOverflow(t *testing.T) {
135
136 inputs := []uint32{
137 q,
138 q + 1,
139 q + 2,
140 1<<23 - 1,
141 1 << 23,
142 q + 1<<23,
143 q + 1<<31,
144 ^uint32(0),
145 }
146 for _, in := range inputs {
147 if _, err := fieldToMontgomery(in); err == nil {
148 t.Fatalf("fieldToMontgomery(%d) did not return an error", in)
149 }
150 }
151 }
152
153 func TestFieldMulSub(t *testing.T) {
154 for _, a := range interestingValues() {
155 for _, b := range interestingValues() {
156 for _, c := range interestingValues() {
157 got := fieldFromMontgomery(fieldMontgomeryMulSub(a.m, b.m, c.m))
158 exp := uint32((uint64(a.v) * (uint64(b.v) + q - uint64(c.v))) % q)
159 if got != exp {
160 t.Fatalf("%d * (%d - %d) = %d, expected %d", a.v, b.v, c.v, got, exp)
161 }
162 }
163 }
164 }
165 }
166
167 func TestFieldAddMul(t *testing.T) {
168 for _, a := range interestingValues() {
169 for _, b := range interestingValues() {
170 for _, c := range interestingValues() {
171 for _, d := range interestingValues() {
172 got := fieldFromMontgomery(fieldMontgomeryAddMul(a.m, b.m, c.m, d.m))
173 exp := uint32((uint64(a.v)*uint64(b.v) + uint64(c.v)*uint64(d.v)) % q)
174 if got != exp {
175 t.Fatalf("%d + %d * %d = %d, expected %d", a.v, b.v, c.v, got, exp)
176 }
177 }
178 }
179 }
180 }
181 }
182
183 func BitRev8(n uint8) uint8 {
184 var r uint8
185 r |= n >> 7 & 0b0000_0001
186 r |= n >> 5 & 0b0000_0010
187 r |= n >> 3 & 0b0000_0100
188 r |= n >> 1 & 0b0000_1000
189 r |= n << 1 & 0b0001_0000
190 r |= n << 3 & 0b0010_0000
191 r |= n << 5 & 0b0100_0000
192 r |= n << 7 & 0b1000_0000
193 return r
194 }
195
196 func CenteredMod(x, m uint32) int32 {
197 x = x % m
198 if x > m/2 {
199 return int32(x) - int32(m)
200 }
201 return int32(x)
202 }
203
204 func reduceModQ(x int32) uint32 {
205 x %= q
206 if x < 0 {
207 return uint32(x + q)
208 }
209 return uint32(x)
210 }
211
212 func TestCenteredMod(t *testing.T) {
213 for x := range uint32(q * 2) {
214 got := CenteredMod(uint32(x), q)
215 if reduceModQ(got) != (x % q) {
216 t.Fatalf("CenteredMod(%d) = %d, which is not congruent to %d mod %d", x, got, x, q)
217 }
218 }
219
220 for x := range uint32(q) {
221 r, _ := fieldToMontgomery(x)
222 got := fieldCenteredMod(r)
223 exp := CenteredMod(x, q)
224 if got != exp {
225 t.Fatalf("fieldCenteredMod(%d) = %d, expected %d", x, got, exp)
226 }
227 }
228 }
229
230 func TestInfinityNorm(t *testing.T) {
231 for x := range uint32(q) {
232 r, _ := fieldToMontgomery(x)
233 got := fieldInfinityNorm(r)
234 exp := CenteredMod(x, q)
235 if exp < 0 {
236 exp = -exp
237 }
238 if got != uint32(exp) {
239 t.Fatalf("fieldInfinityNorm(%d) = %d, expected %d", x, got, exp)
240 }
241 }
242 }
243
244 func TestConstants(t *testing.T) {
245 if fieldFromMontgomery(one) != 1 {
246 t.Errorf("one constant incorrect")
247 }
248 if fieldFromMontgomery(minusOne) != q-1 {
249 t.Errorf("minusOne constant incorrect")
250 }
251 if fieldInfinityNorm(one) != 1 {
252 t.Errorf("one infinity norm incorrect")
253 }
254 if fieldInfinityNorm(minusOne) != 1 {
255 t.Errorf("minusOne infinity norm incorrect")
256 }
257
258 if PublicKeySize44 != pubKeySize(params44) {
259 t.Errorf("PublicKeySize44 constant incorrect")
260 }
261 if PublicKeySize65 != pubKeySize(params65) {
262 t.Errorf("PublicKeySize65 constant incorrect")
263 }
264 if PublicKeySize87 != pubKeySize(params87) {
265 t.Errorf("PublicKeySize87 constant incorrect")
266 }
267 if SignatureSize44 != sigSize(params44) {
268 t.Errorf("SignatureSize44 constant incorrect")
269 }
270 if SignatureSize65 != sigSize(params65) {
271 t.Errorf("SignatureSize65 constant incorrect")
272 }
273 if SignatureSize87 != sigSize(params87) {
274 t.Errorf("SignatureSize87 constant incorrect")
275 }
276 }
277
278 func TestPower2Round(t *testing.T) {
279 t.Parallel()
280 for x := range uint32(q) {
281 rr, _ := fieldToMontgomery(x)
282 t1, t0 := power2Round(rr)
283
284 hi, err := fieldToMontgomery(uint32(t1) << 13)
285 if err != nil {
286 t.Fatalf("power2Round(%d): failed to convert high part to Montgomery: %v", x, err)
287 }
288 if r := fieldFromMontgomery(fieldAdd(hi, t0)); r != x {
289 t.Fatalf("power2Round(%d) = (%d, %d), which reconstructs to %d, expected %d", x, t1, t0, r, x)
290 }
291 }
292 }
293
294 func SpecDecompose(rr fieldElement, p parameters) (R1 uint32, R0 int32) {
295 r := fieldFromMontgomery(rr)
296 if (q-1)%p.γ2 != 0 {
297 panic("mldsa: internal error: unsupported denγ2")
298 }
299 γ2 := (q - 1) / uint32(p.γ2)
300 r0 := CenteredMod(r, 2*γ2)
301 diff := int32(r) - r0
302 if diff == q-1 {
303 r0 = r0 - 1
304 return 0, r0
305 } else {
306 if diff < 0 || uint32(diff)%γ2 != 0 {
307 panic("mldsa: internal error: invalid decomposition")
308 }
309 r1 := uint32(diff) / (2 * γ2)
310 return r1, r0
311 }
312 }
313
314 func TestDecompose(t *testing.T) {
315 t.Run("ML-DSA-44", func(t *testing.T) {
316 testDecompose(t, params44)
317 })
318 t.Run("ML-DSA-65,87", func(t *testing.T) {
319 testDecompose(t, params65)
320 })
321 }
322
323 func testDecompose(t *testing.T, p parameters) {
324 t.Parallel()
325 for x := range uint32(q) {
326 rr, _ := fieldToMontgomery(x)
327 r1, r0 := SpecDecompose(rr, p)
328
329
330
331 γ2 := (q - 1) / uint32(p.γ2)
332 reconstructed := reduceModQ(int32(r1*2*γ2) + r0)
333 if reconstructed != x {
334 t.Fatalf("SpecDecompose(%d) = (%d, %d), which reconstructs to %d, expected %d", x, r1, r0, reconstructed, x)
335 }
336
337 var gotR1 byte
338 var gotR0 int32
339 switch p.γ2 {
340 case 88:
341 gotR1, gotR0 = decompose88(rr)
342 if gotR1 > 43 {
343 t.Fatalf("decompose88(%d) returned r1 = %d, which is out of range", x, gotR1)
344 }
345 case 32:
346 gotR1, gotR0 = decompose32(rr)
347 if gotR1 > 15 {
348 t.Fatalf("decompose32(%d) returned r1 = %d, which is out of range", x, gotR1)
349 }
350 default:
351 t.Fatalf("unsupported denγ2: %d", p.γ2)
352 }
353 if uint32(gotR1) != r1 {
354 t.Fatalf("highBits(%d) = %d, expected %d", x, gotR1, r1)
355 }
356 if gotR0 != r0 {
357 t.Fatalf("lowBits(%d) = %d, expected %d", x, gotR0, r0)
358 }
359 }
360 }
361
362 func TestZetas(t *testing.T) {
363 ζ := big.NewInt(1753)
364 q := big.NewInt(q)
365 for k, zeta := range zetas {
366
367 exp := new(big.Int).Exp(ζ, big.NewInt(int64(BitRev8(uint8(k)))), q)
368 got := fieldFromMontgomery(zeta)
369 if big.NewInt(int64(got)).Cmp(exp) != 0 {
370 t.Errorf("zetas[%d] = %v, expected %v", k, got, exp)
371 }
372 }
373 }
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396 func TestAccumulated(t *testing.T) {
397 if testing.Short() {
398 t.Skip("skipping accumulated test in short mode")
399 }
400
401 o := sha3.NewShake128()
402 for x := range uint32(q) {
403 r, _ := fieldToMontgomery(x)
404 fmt.Fprintf(o, "%d\n", fieldCenteredMod(r))
405 fmt.Fprintf(o, "%d\n", fieldInfinityNorm(r))
406
407 hi, lo := power2Round(r)
408 fmt.Fprintf(o, "%d\n", hi)
409 fmt.Fprintf(o, "%d\n", fieldFromMontgomery(lo))
410
411 r1, r0 := decompose88(r)
412 if r1x := highBits88(fieldFromMontgomery(r)); r1x != r1 {
413 t.Fatalf("highBits88(%d) = %d, expected %d", x, r1x, r1)
414 }
415 if r1h0 := useHint88(r, 0); r1h0 != r1 {
416 t.Fatalf("useHint88(%d, 0) = %d, expected %d", x, r1h0, r1)
417 }
418
419 fmt.Fprintf(o, "%d\n", r1)
420 fmt.Fprintf(o, "%d\n", useHint88(r, 1))
421 fmt.Fprintf(o, "%d\n", r0)
422 fmt.Fprintf(o, "%d\n", constantTimeAbs(r0))
423
424 r1, r0 = decompose32(r)
425 if r1x := highBits32(fieldFromMontgomery(r)); r1x != r1 {
426 t.Fatalf("highBits32(%d) = %d, expected %d", x, r1x, r1)
427 }
428 if r1h0 := useHint32(r, 0); r1h0 != r1 {
429 t.Fatalf("useHint32(%d, 0) = %d, expected %d", x, r1h0, r1)
430 }
431
432 fmt.Fprintf(o, "%d\n", r1)
433 fmt.Fprintf(o, "%d\n", useHint32(r, 1))
434 fmt.Fprintf(o, "%d\n", r0)
435 fmt.Fprintf(o, "%d\n", constantTimeAbs(r0))
436 }
437
438
439
440 expected := "f930663417278156ab05d940294a77210a809c924d8ab63ec72f4526247602c7"
441 if got := hex.EncodeToString(o.Sum(nil)); got != expected {
442 t.Errorf("got %s, expected %s", got, expected)
443 }
444 }
445
View as plain text