Source file src/crypto/internal/fips140/mldsa/field_test.go

     1  // Copyright 2025 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     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  // q is large enough that we can't exhaustively test all q × q inputs, so when
    21  // we have two inputs  we test [0, q) on one side and a set of interesting
    22  // values on the other side.
    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  		// Also test values that have an interesting Montgomery representation.
    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  	// fieldToMontgomery should reject inputs ≥ q.
   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  		// Check that SpecDecompose is correct.
   330  		// r ≡ r1 * (2 * γ2) + r0 mod q
   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  		// ζ^BitRev₈(k) mod q
   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  // TestAccumulated computes the hash of the following 12 values, as ASCII
   376  // decimals with an optional leading - sign and separated by newlines, for all
   377  // elements r in ℤq from 0 to q-1:
   378  //
   379  //   - r mod± q
   380  //   - ‖r‖∞ = |r mod± q|
   381  //   - r1, r0 = Power2Round(r)
   382  //
   383  // For ML-DSA-44 (γ₂ = (q - 1) / 88):
   384  //   - HighBits(r) = UseHint(0, r)
   385  //   - UseHint(1, r)
   386  //   - LowBits(r)
   387  //   - ‖LowBits(r)‖∞ = |LowBits(r)|
   388  //
   389  // For ML-DSA-65 and ML-DSA-87 (γ₂ = (q - 1) / 32):
   390  //   - HighBits(r) = UseHint(0, r)
   391  //   - UseHint(1, r)
   392  //   - LowBits(r)
   393  //   - ‖LowBits(r)‖∞ = |LowBits(r)|
   394  //
   395  // Note that HighBits(r), LowBits(r) = Decompose(r).
   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  	// The expected value is documented at https://c2sp.org/CCTV/ML-DSA, and
   439  	// tested against https://github.com/FiloSottile/mldsa-py.
   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