Source file src/crypto/internal/fips140/mldsa/mldsa.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  	"bytes"
     9  	"crypto/internal/fips140"
    10  	"crypto/internal/fips140/drbg"
    11  	"crypto/internal/fips140/sha3"
    12  	"crypto/internal/fips140/subtle"
    13  	"crypto/internal/fips140deps/byteorder"
    14  	"errors"
    15  )
    16  
    17  type parameters struct {
    18  	k, l int // dimensions of A
    19  	η    int // bound for secret coefficients
    20  	γ1   int // log₂(γ₁), where [-γ₁+1, γ₁] is the bound of y
    21  	γ2   int // denominator of γ₂ = (q - 1) / γ2
    22  	λ    int // collison strength
    23  	τ    int // number of non-zero coefficients in challenge
    24  	ω    int // max number of hints in MakeHint
    25  }
    26  
    27  var (
    28  	params44 = parameters{k: 4, l: 4, η: 2, γ1: 17, γ2: 88, λ: 128, τ: 39, ω: 80}
    29  	params65 = parameters{k: 6, l: 5, η: 4, γ1: 19, γ2: 32, λ: 192, τ: 49, ω: 55}
    30  	params87 = parameters{k: 8, l: 7, η: 2, γ1: 19, γ2: 32, λ: 256, τ: 60, ω: 75}
    31  )
    32  
    33  func pubKeySize(p parameters) int {
    34  	// ρ + k × n × 10-bit coefficients of t₁
    35  	return 32 + p.k*n*10/8
    36  }
    37  
    38  func sigSize(p parameters) int {
    39  	// challenge + l × n × (γ₁+1)-bit coefficients of z + hint
    40  	return (p.λ / 4) + p.l*n*(p.γ1+1)/8 + p.ω + p.k
    41  }
    42  
    43  const (
    44  	PrivateKeySize = 32
    45  
    46  	PublicKeySize44 = 32 + 4*n*10/8
    47  	PublicKeySize65 = 32 + 6*n*10/8
    48  	PublicKeySize87 = 32 + 8*n*10/8
    49  
    50  	SignatureSize44 = 128/4 + 4*n*(17+1)/8 + 80 + 4
    51  	SignatureSize65 = 192/4 + 5*n*(19+1)/8 + 55 + 6
    52  	SignatureSize87 = 256/4 + 7*n*(19+1)/8 + 75 + 8
    53  )
    54  
    55  const maxK, maxL, maxλ, maxγ1 = 8, 7, 256, 19
    56  const maxPubKeySize = PublicKeySize87
    57  
    58  type PrivateKey struct {
    59  	seed [32]byte
    60  	pub  PublicKey
    61  	a    [maxK * maxL]nttElement
    62  	t1   [maxK]nttElement // NTT(t₁ ⋅ 2ᵈ)
    63  	s1   [maxL]nttElement
    64  	s2   [maxK]nttElement
    65  	t0   [maxK]nttElement
    66  	k    [32]byte
    67  }
    68  
    69  func (priv *PrivateKey) Equal(x *PrivateKey) bool {
    70  	return priv.pub.p == x.pub.p && subtle.ConstantTimeCompare(priv.seed[:], x.seed[:]) == 1
    71  }
    72  
    73  func (priv *PrivateKey) Bytes() []byte {
    74  	seed := priv.seed
    75  	return seed[:]
    76  }
    77  
    78  func (priv *PrivateKey) PublicKey() *PublicKey {
    79  	// Note that this is likely to keep the entire PrivateKey reachable for
    80  	// the lifetime of the PublicKey, which may be undesirable.
    81  	return &priv.pub
    82  }
    83  
    84  type PublicKey struct {
    85  	raw [maxPubKeySize]byte
    86  	p   parameters
    87  	tr  [64]byte // public key hash
    88  }
    89  
    90  func (pub *PublicKey) Equal(x *PublicKey) bool {
    91  	size := pubKeySize(pub.p)
    92  	return pub.p == x.p && subtle.ConstantTimeCompare(pub.raw[:size], x.raw[:size]) == 1
    93  }
    94  
    95  func (pub *PublicKey) Bytes() []byte {
    96  	size := pubKeySize(pub.p)
    97  	return bytes.Clone(pub.raw[:size])
    98  }
    99  
   100  func (pub *PublicKey) Parameters() string {
   101  	switch pub.p {
   102  	case params44:
   103  		return "ML-DSA-44"
   104  	case params65:
   105  		return "ML-DSA-65"
   106  	case params87:
   107  		return "ML-DSA-87"
   108  	default:
   109  		panic("mldsa: internal error: unknown parameters")
   110  	}
   111  }
   112  
   113  func GenerateKey44() *PrivateKey {
   114  	fipsSelfTest()
   115  	fips140.RecordApproved()
   116  	var seed [32]byte
   117  	drbg.Read(seed[:])
   118  	priv := newPrivateKey(&seed, params44)
   119  	fipsPCT(priv)
   120  	return priv
   121  }
   122  
   123  func GenerateKey65() *PrivateKey {
   124  	fipsSelfTest()
   125  	fips140.RecordApproved()
   126  	var seed [32]byte
   127  	drbg.Read(seed[:])
   128  	priv := newPrivateKey(&seed, params65)
   129  	fipsPCT(priv)
   130  	return priv
   131  }
   132  
   133  func GenerateKey87() *PrivateKey {
   134  	fipsSelfTest()
   135  	fips140.RecordApproved()
   136  	var seed [32]byte
   137  	drbg.Read(seed[:])
   138  	priv := newPrivateKey(&seed, params87)
   139  	fipsPCT(priv)
   140  	return priv
   141  }
   142  
   143  var errInvalidSeedLength = errors.New("mldsa: invalid seed length")
   144  
   145  func NewPrivateKey44(seed []byte) (*PrivateKey, error) {
   146  	fipsSelfTest()
   147  	fips140.RecordApproved()
   148  	if len(seed) != 32 {
   149  		return nil, errInvalidSeedLength
   150  	}
   151  	return newPrivateKey((*[32]byte)(seed), params44), nil
   152  }
   153  
   154  func NewPrivateKey65(seed []byte) (*PrivateKey, error) {
   155  	fipsSelfTest()
   156  	fips140.RecordApproved()
   157  	if len(seed) != 32 {
   158  		return nil, errInvalidSeedLength
   159  	}
   160  	return newPrivateKey((*[32]byte)(seed), params65), nil
   161  }
   162  
   163  func NewPrivateKey87(seed []byte) (*PrivateKey, error) {
   164  	fipsSelfTest()
   165  	fips140.RecordApproved()
   166  	if len(seed) != 32 {
   167  		return nil, errInvalidSeedLength
   168  	}
   169  	return newPrivateKey((*[32]byte)(seed), params87), nil
   170  }
   171  
   172  func newPrivateKey(seed *[32]byte, p parameters) *PrivateKey {
   173  	k, l := p.k, p.l
   174  
   175  	priv := &PrivateKey{pub: PublicKey{p: p}}
   176  	priv.seed = *seed
   177  
   178  	ξ := sha3.NewShake256()
   179  	ξ.Write(seed[:])
   180  	ξ.Write([]byte{byte(k), byte(l)})
   181  	ρ, ρs := make([]byte, 32), make([]byte, 64)
   182  	ξ.Read(ρ)
   183  	ξ.Read(ρs)
   184  	ξ.Read(priv.k[:])
   185  
   186  	A := priv.a[:k*l]
   187  	computeMatrixA(A, ρ, p)
   188  
   189  	s1 := priv.s1[:l]
   190  	for r := range l {
   191  		s1[r] = ntt(sampleBoundedPoly(ρs, byte(r), p))
   192  	}
   193  	s2 := priv.s2[:k]
   194  	for r := range k {
   195  		s2[r] = ntt(sampleBoundedPoly(ρs, byte(l+r), p))
   196  	}
   197  
   198  	// ˆt = Â ∘ ŝ₁ + ŝ₂
   199  	tHat := make([]nttElement, k, maxK)
   200  	for i := range tHat {
   201  		tHat[i] = s2[i]
   202  		for j := range s1 {
   203  			tHat[i] = polyAdd(tHat[i], nttMul(A[i*l+j], s1[j]))
   204  		}
   205  	}
   206  	// t = NTT⁻¹(ˆt)
   207  	t := make([]ringElement, k, maxK)
   208  	for i := range tHat {
   209  		t[i] = inverseNTT(tHat[i])
   210  	}
   211  	// (t₁, _) = Power2Round(t)
   212  	// (_, ˆt₀) = NTT(Power2Round(t))
   213  	t1, t0 := make([][n]uint16, k, maxK), priv.t0[:k]
   214  	for i := range t {
   215  		var w ringElement
   216  		for j := range t[i] {
   217  			t1[i][j], w[j] = power2Round(t[i][j])
   218  		}
   219  		t0[i] = ntt(w)
   220  	}
   221  
   222  	pk := pkEncode(priv.pub.raw[:0], ρ, t1, p)
   223  	priv.pub.tr = computePublicKeyHash(pk)
   224  	computeT1Hat(priv.t1[:k], t1) // NTT(t₁ ⋅ 2ᵈ)
   225  
   226  	return priv
   227  }
   228  
   229  func computeMatrixA(A []nttElement, ρ []byte, p parameters) {
   230  	k, l := p.k, p.l
   231  	for r := range k {
   232  		for s := range l {
   233  			A[r*l+s] = sampleNTT(ρ, byte(s), byte(r))
   234  		}
   235  	}
   236  }
   237  
   238  func computePublicKeyHash(pk []byte) [64]byte {
   239  	H := sha3.NewShake256()
   240  	H.Write(pk)
   241  	var tr [64]byte
   242  	H.Read(tr[:])
   243  	return tr
   244  }
   245  
   246  func computeT1Hat(t1Hat []nttElement, t1 [][n]uint16) {
   247  	for i := range t1 {
   248  		var w ringElement
   249  		for j := range t1[i] {
   250  			// t₁ <= 2¹⁰ - 1
   251  			// t₁ ⋅ 2ᵈ <= 2ᵈ(2¹⁰ - 1) = 2²³ - 2¹³ < q = 2²³ - 2¹³ + 1
   252  			z, _ := fieldToMontgomery(uint32(t1[i][j]) << 13)
   253  			w[j] = z
   254  		}
   255  		t1Hat[i] = ntt(w)
   256  	}
   257  }
   258  
   259  func pkEncode(buf []byte, ρ []byte, t1 [][n]uint16, p parameters) []byte {
   260  	pk := append(buf, ρ...)
   261  	for _, w := range t1[:p.k] {
   262  		// Encode four at a time into 4 * 10 bits = 5 bytes.
   263  		for i := 0; i < n; i += 4 {
   264  			c0 := w[i]
   265  			c1 := w[i+1]
   266  			c2 := w[i+2]
   267  			c3 := w[i+3]
   268  			b0 := byte(c0 >> 0)
   269  			b1 := byte((c0 >> 8) | (c1 << 2))
   270  			b2 := byte((c1 >> 6) | (c2 << 4))
   271  			b3 := byte((c2 >> 4) | (c3 << 6))
   272  			b4 := byte(c3 >> 2)
   273  			pk = append(pk, b0, b1, b2, b3, b4)
   274  		}
   275  	}
   276  	return pk
   277  }
   278  
   279  func pkDecode(pk []byte, t1 [][n]uint16, p parameters) (ρ []byte, err error) {
   280  	if len(pk) != pubKeySize(p) {
   281  		return nil, errInvalidPublicKeyLength
   282  	}
   283  	ρ, pk = pk[:32], pk[32:]
   284  	for r := range t1 {
   285  		// Decode four at a time from 4 * 10 bits = 5 bytes.
   286  		for i := 0; i < n; i += 4 {
   287  			b0, b1, b2, b3, b4 := pk[0], pk[1], pk[2], pk[3], pk[4]
   288  			t1[r][i+0] = uint16(b0>>0) | uint16(b1&0b0000_0011)<<8
   289  			t1[r][i+1] = uint16(b1>>2) | uint16(b2&0b0000_1111)<<6
   290  			t1[r][i+2] = uint16(b2>>4) | uint16(b3&0b0011_1111)<<4
   291  			t1[r][i+3] = uint16(b3>>6) | uint16(b4&0b1111_1111)<<2
   292  			pk = pk[5:]
   293  		}
   294  	}
   295  	return ρ, nil
   296  }
   297  
   298  var errInvalidPublicKeyLength = errors.New("mldsa: invalid public key length")
   299  
   300  func NewPublicKey44(pk []byte) (*PublicKey, error) {
   301  	return newPublicKey(&PublicKey{}, pk, params44)
   302  }
   303  
   304  func NewPublicKey65(pk []byte) (*PublicKey, error) {
   305  	return newPublicKey(&PublicKey{}, pk, params65)
   306  }
   307  
   308  func NewPublicKey87(pk []byte) (*PublicKey, error) {
   309  	return newPublicKey(&PublicKey{}, pk, params87)
   310  }
   311  
   312  func newPublicKey(pub *PublicKey, pk []byte, p parameters) (*PublicKey, error) {
   313  	if len(pk) != pubKeySize(p) {
   314  		return nil, errInvalidPublicKeyLength
   315  	}
   316  
   317  	// We don't precompute A and t1Hat here, because they would make the
   318  	// PublicKey over 68KB. Unlike private keys, public keys are often used to
   319  	// verify a signature only once, so precomputation doesn't help as often,
   320  	// but they can stay around in memory, for example as part of a TLS
   321  	// connection's PeerCertificates, so their size is more of a concern.
   322  	// Instead, we compute A and t1Hat on demand in Verify.
   323  
   324  	pub.p = p
   325  	copy(pub.raw[:], pk)
   326  	pub.tr = computePublicKeyHash(pk)
   327  
   328  	return pub, nil
   329  }
   330  
   331  var (
   332  	errContextTooLong    = errors.New("mldsa: context too long")
   333  	errMessageHashLength = errors.New("mldsa: invalid message hash length")
   334  	errRandomLength      = errors.New("mldsa: invalid random length")
   335  )
   336  
   337  func Sign(priv *PrivateKey, msg []byte, context string) ([]byte, error) {
   338  	fipsSelfTest()
   339  	fips140.RecordApproved()
   340  	var random [32]byte
   341  	drbg.Read(random[:])
   342  	μ, err := computeMessageHash(priv.pub.tr[:], msg, context)
   343  	if err != nil {
   344  		return nil, err
   345  	}
   346  	return signInternal(priv, &μ, &random), nil
   347  }
   348  
   349  func SignDeterministic(priv *PrivateKey, msg []byte, context string) ([]byte, error) {
   350  	fipsSelfTest()
   351  	fips140.RecordApproved()
   352  	var random [32]byte
   353  	μ, err := computeMessageHash(priv.pub.tr[:], msg, context)
   354  	if err != nil {
   355  		return nil, err
   356  	}
   357  	return signInternal(priv, &μ, &random), nil
   358  }
   359  
   360  func TestingOnlySignWithRandom(priv *PrivateKey, msg []byte, context string, random []byte) ([]byte, error) {
   361  	fipsSelfTest()
   362  	fips140.RecordApproved()
   363  	μ, err := computeMessageHash(priv.pub.tr[:], msg, context)
   364  	if err != nil {
   365  		return nil, err
   366  	}
   367  	if len(random) != 32 {
   368  		return nil, errRandomLength
   369  	}
   370  	return signInternal(priv, &μ, (*[32]byte)(random)), nil
   371  }
   372  
   373  func SignExternalMu(priv *PrivateKey, μ []byte) ([]byte, error) {
   374  	fipsSelfTest()
   375  	fips140.RecordApproved()
   376  	var random [32]byte
   377  	drbg.Read(random[:])
   378  	if len(μ) != 64 {
   379  		return nil, errMessageHashLength
   380  	}
   381  	return signInternal(priv, (*[64]byte)(μ), &random), nil
   382  }
   383  
   384  func SignExternalMuDeterministic(priv *PrivateKey, μ []byte) ([]byte, error) {
   385  	fipsSelfTest()
   386  	fips140.RecordApproved()
   387  	var random [32]byte
   388  	if len(μ) != 64 {
   389  		return nil, errMessageHashLength
   390  	}
   391  	return signInternal(priv, (*[64]byte)(μ), &random), nil
   392  }
   393  
   394  func TestingOnlySignExternalMuWithRandom(priv *PrivateKey, μ []byte, random []byte) ([]byte, error) {
   395  	fipsSelfTest()
   396  	fips140.RecordApproved()
   397  	if len(μ) != 64 {
   398  		return nil, errMessageHashLength
   399  	}
   400  	if len(random) != 32 {
   401  		return nil, errRandomLength
   402  	}
   403  	return signInternal(priv, (*[64]byte)(μ), (*[32]byte)(random)), nil
   404  }
   405  
   406  func computeMessageHash(tr []byte, msg []byte, context string) ([64]byte, error) {
   407  	if len(context) > 255 {
   408  		return [64]byte{}, errContextTooLong
   409  	}
   410  	H := sha3.NewShake256()
   411  	H.Write(tr)
   412  	H.Write([]byte{0}) // ML-DSA / HashML-DSA domain separator
   413  	H.Write([]byte{byte(len(context))})
   414  	H.Write([]byte(context))
   415  	H.Write(msg)
   416  	var μ [64]byte
   417  	H.Read(μ[:])
   418  	return μ, nil
   419  }
   420  
   421  func signInternal(priv *PrivateKey, μ *[64]byte, random *[32]byte) []byte {
   422  	p, k, l := priv.pub.p, priv.pub.p.k, priv.pub.p.l
   423  	A, s1, s2, t0 := priv.a[:k*l], priv.s1[:l], priv.s2[:k], priv.t0[:k]
   424  
   425  	β := p.τ * p.η
   426  	γ1 := uint32(1 << p.γ1)
   427  	γ1β := γ1 - uint32(β)
   428  	γ2 := (q - 1) / uint32(p.γ2)
   429  	γ2β := γ2 - uint32(β)
   430  
   431  	H := sha3.NewShake256()
   432  	H.Write(priv.k[:])
   433  	H.Write(random[:])
   434  	H.Write(μ[:])
   435  	nonce := make([]byte, 64)
   436  	H.Read(nonce)
   437  
   438  	κ := 0
   439  sign:
   440  	for {
   441  		// Main rejection sampling loop. Note that leaking rejected signatures
   442  		// leaks information about the private key. However, as explained in
   443  		// https://pq-crystals.org/dilithium/data/dilithium-specification-round3.pdf
   444  		// Section 5.5, we are free to leak rejected ch values, as well as which
   445  		// check causes the rejection and which coefficient failed the check
   446  		// (but not the value or sign of the coefficient).
   447  
   448  		y := make([]ringElement, l, maxL)
   449  		for r := range y {
   450  			counter := make([]byte, 2)
   451  			byteorder.LEPutUint16(counter, uint16(κ))
   452  			κ++
   453  
   454  			H.Reset()
   455  			H.Write(nonce)
   456  			H.Write(counter)
   457  			v := make([]byte, (p.γ1+1)*n/8, (maxγ1+1)*n/8)
   458  			H.Read(v)
   459  
   460  			y[r] = bitUnpack(v, p)
   461  		}
   462  
   463  		// w = NTT⁻¹(Â ∘ NTT(y))
   464  		yHat := make([]nttElement, l, maxL)
   465  		for i := range y {
   466  			yHat[i] = ntt(y[i])
   467  		}
   468  		w := make([]ringElement, k, maxK)
   469  		for i := range w {
   470  			var wHat nttElement
   471  			for j := range l {
   472  				wHat = polyAdd(wHat, nttMul(A[i*l+j], yHat[j]))
   473  			}
   474  			w[i] = inverseNTT(wHat)
   475  		}
   476  
   477  		H.Reset()
   478  		H.Write(μ[:])
   479  		for i := range w {
   480  			w1Encode(H, highBits(w[i], p), p)
   481  		}
   482  		ch := make([]byte, p.λ/4, maxλ/4)
   483  		H.Read(ch)
   484  
   485  		// sampleInBall is not constant time, but see comment above about
   486  		// leaking rejected ch values being acceptable.
   487  		c := ntt(sampleInBall(ch, p))
   488  
   489  		cs1 := make([]ringElement, l, maxL)
   490  		for i := range cs1 {
   491  			cs1[i] = inverseNTT(nttMul(c, s1[i]))
   492  		}
   493  		cs2 := make([]ringElement, k, maxK)
   494  		for i := range cs2 {
   495  			cs2[i] = inverseNTT(nttMul(c, s2[i]))
   496  		}
   497  
   498  		z := make([]ringElement, l, maxL)
   499  		for i := range y {
   500  			z[i] = polyAdd(y[i], cs1[i])
   501  
   502  			// Reject if ||z||∞ ≥ γ1 − β
   503  			if coefficientsExceedBound(z[i], γ1β) {
   504  				if testingOnlyRejectionReason != nil {
   505  					testingOnlyRejectionReason("z")
   506  				}
   507  				continue sign
   508  			}
   509  		}
   510  
   511  		for i := range w {
   512  			r0 := polySub(w[i], cs2[i])
   513  
   514  			// Reject if ||LowBits(r0)||∞ ≥ γ2 − β
   515  			if lowBitsExceedBound(r0, γ2β, p) {
   516  				if testingOnlyRejectionReason != nil {
   517  					testingOnlyRejectionReason("r0")
   518  				}
   519  				continue sign
   520  			}
   521  		}
   522  
   523  		ct0 := make([]ringElement, k, maxK)
   524  		for i := range ct0 {
   525  			ct0[i] = inverseNTT(nttMul(c, t0[i]))
   526  
   527  			// Reject if ||ct0||∞ ≥ γ2
   528  			if coefficientsExceedBound(ct0[i], γ2) {
   529  				if testingOnlyRejectionReason != nil {
   530  					testingOnlyRejectionReason("ct0")
   531  				}
   532  				continue sign
   533  			}
   534  		}
   535  
   536  		count1s := 0
   537  		h := make([][n]byte, k, maxK)
   538  		for i := range w {
   539  			var count int
   540  			h[i], count = makeHint(ct0[i], w[i], cs2[i], p)
   541  			count1s += count
   542  		}
   543  		// Reject if number of hints > ω
   544  		if count1s > p.ω {
   545  			if testingOnlyRejectionReason != nil {
   546  				testingOnlyRejectionReason("h")
   547  			}
   548  			continue sign
   549  		}
   550  
   551  		return sigEncode(ch, z, h, p)
   552  	}
   553  }
   554  
   555  // testingOnlyRejectionReason is set in tests, to ensure that all rejection
   556  // paths are covered. If not nil, it is called with a string describing the
   557  // reason for rejection: "z", "r0", "ct0", or "h".
   558  var testingOnlyRejectionReason func(reason string)
   559  
   560  // w1Encode implements w1Encode from FIPS 204, writing directly into H.
   561  func w1Encode(H *sha3.SHAKE, w [n]byte, p parameters) {
   562  	switch p.γ2 {
   563  	case 32:
   564  		// Coefficients are <= (q − 1)/(2γ2) − 1 = 15, four bits each.
   565  		buf := make([]byte, 4*n/8)
   566  		for i := 0; i < n; i += 2 {
   567  			b0 := w[i]
   568  			b1 := w[i+1]
   569  			buf[i/2] = b0 | b1<<4
   570  		}
   571  		H.Write(buf)
   572  	case 88:
   573  		// Coefficients are <= (q − 1)/(2γ2) − 1 = 43, six bits each.
   574  		buf := make([]byte, 6*n/8)
   575  		for i := 0; i < n; i += 4 {
   576  			b0 := w[i]
   577  			b1 := w[i+1]
   578  			b2 := w[i+2]
   579  			b3 := w[i+3]
   580  			buf[3*i/4+0] = (b0 >> 0) | (b1 << 6)
   581  			buf[3*i/4+1] = (b1 >> 2) | (b2 << 4)
   582  			buf[3*i/4+2] = (b2 >> 4) | (b3 << 2)
   583  		}
   584  		H.Write(buf)
   585  	default:
   586  		panic("mldsa: internal error: unsupported γ2")
   587  	}
   588  }
   589  
   590  func coefficientsExceedBound(w ringElement, bound uint32) bool {
   591  	// If this function appears in profiles, it might be possible to deduplicate
   592  	// the work of fieldFromMontgomery inside fieldInfinityNorm with the
   593  	// subsequent encoding of w.
   594  	for i := range w {
   595  		if fieldInfinityNorm(w[i]) >= bound {
   596  			return true
   597  		}
   598  	}
   599  	return false
   600  }
   601  
   602  func lowBitsExceedBound(w ringElement, bound uint32, p parameters) bool {
   603  	switch p.γ2 {
   604  	case 32:
   605  		for i := range w {
   606  			_, r0 := decompose32(w[i])
   607  			if constantTimeAbs(r0) >= bound {
   608  				return true
   609  			}
   610  		}
   611  	case 88:
   612  		for i := range w {
   613  			_, r0 := decompose88(w[i])
   614  			if constantTimeAbs(r0) >= bound {
   615  				return true
   616  			}
   617  		}
   618  	default:
   619  		panic("mldsa: internal error: unsupported γ2")
   620  	}
   621  	return false
   622  }
   623  
   624  var (
   625  	errInvalidSignatureLength           = errors.New("mldsa: invalid signature length")
   626  	errInvalidSignatureCoeffBounds      = errors.New("mldsa: invalid signature")
   627  	errInvalidSignatureChallenge        = errors.New("mldsa: invalid signature")
   628  	errInvalidSignatureHintLimits       = errors.New("mldsa: invalid signature encoding")
   629  	errInvalidSignatureHintIndexOrder   = errors.New("mldsa: invalid signature encoding")
   630  	errInvalidSignatureHintExtraIndices = errors.New("mldsa: invalid signature encoding")
   631  )
   632  
   633  func Verify(pub *PublicKey, msg, sig []byte, context string) error {
   634  	fipsSelfTest()
   635  	fips140.RecordApproved()
   636  	μ, err := computeMessageHash(pub.tr[:], msg, context)
   637  	if err != nil {
   638  		return err
   639  	}
   640  	return verifyInternal(pub, &μ, sig)
   641  }
   642  
   643  func VerifyExternalMu(pub *PublicKey, μ []byte, sig []byte) error {
   644  	fipsSelfTest()
   645  	fips140.RecordApproved()
   646  	if len(μ) != 64 {
   647  		return errMessageHashLength
   648  	}
   649  	return verifyInternal(pub, (*[64]byte)(μ), sig)
   650  }
   651  
   652  func verifyInternal(pub *PublicKey, μ *[64]byte, sig []byte) error {
   653  	p, k, l := pub.p, pub.p.k, pub.p.l
   654  
   655  	β := p.τ * p.η
   656  	γ1 := uint32(1 << p.γ1)
   657  	γ1β := γ1 - uint32(β)
   658  
   659  	t1 := make([][n]uint16, k, maxK)
   660  	ρ, err := pkDecode(pub.raw[:pubKeySize(pub.p)], t1, p)
   661  	if err != nil {
   662  		return err
   663  	}
   664  	A := make([]nttElement, k*l, maxK*maxL)
   665  	computeMatrixA(A, ρ, p)
   666  	t1Hat := make([]nttElement, k, maxK)
   667  	computeT1Hat(t1Hat, t1) // NTT(t₁ ⋅ 2ᵈ)
   668  
   669  	z := make([]ringElement, l, maxL)
   670  	h := make([][n]byte, k, maxK)
   671  	ch, err := sigDecode(sig, z, h, p)
   672  	if err != nil {
   673  		return err
   674  	}
   675  
   676  	c := ntt(sampleInBall(ch, p))
   677  
   678  	// w = Â ∘ NTT(z) − NTT(c) ∘ NTT(t₁ ⋅ 2ᵈ)
   679  	zHat := make([]nttElement, l, maxL)
   680  	for i := range zHat {
   681  		zHat[i] = ntt(z[i])
   682  	}
   683  	w := make([]ringElement, k, maxK)
   684  	for i := range w {
   685  		var wHat nttElement
   686  		for j := range l {
   687  			wHat = polyAdd(wHat, nttMul(A[i*l+j], zHat[j]))
   688  		}
   689  		wHat = polySub(wHat, nttMul(c, t1Hat[i]))
   690  		w[i] = inverseNTT(wHat)
   691  	}
   692  
   693  	// Use hints h to compute w₁ from w(approx).
   694  	w1 := make([][n]byte, k, maxK)
   695  	for i := range w {
   696  		w1[i] = useHint(w[i], h[i], p)
   697  	}
   698  
   699  	H := sha3.NewShake256()
   700  	H.Write(μ[:])
   701  	for i := range w {
   702  		w1Encode(H, w1[i], p)
   703  	}
   704  	computedCH := make([]byte, p.λ/4, maxλ/4)
   705  	H.Read(computedCH)
   706  
   707  	for i := range z {
   708  		if coefficientsExceedBound(z[i], γ1β) {
   709  			return errInvalidSignatureCoeffBounds
   710  		}
   711  	}
   712  
   713  	if !bytes.Equal(ch, computedCH) {
   714  		return errInvalidSignatureChallenge
   715  	}
   716  
   717  	return nil
   718  }
   719  
   720  func sigEncode(ch []byte, z []ringElement, h [][n]byte, p parameters) []byte {
   721  	sig := make([]byte, 0, sigSize(p))
   722  	sig = append(sig, ch...)
   723  	for i := range z {
   724  		sig = bitPack(sig, z[i], p)
   725  	}
   726  	sig = hintEncode(sig, h, p)
   727  	return sig
   728  }
   729  
   730  func sigDecode(sig []byte, z []ringElement, h [][n]byte, p parameters) (ch []byte, err error) {
   731  	if len(sig) != sigSize(p) {
   732  		return nil, errInvalidSignatureLength
   733  	}
   734  	ch, sig = sig[:p.λ/4], sig[p.λ/4:]
   735  	for i := range z {
   736  		length := (p.γ1 + 1) * n / 8
   737  		z[i] = bitUnpack(sig[:length], p)
   738  		sig = sig[length:]
   739  	}
   740  	if err := hintDecode(sig, h, p); err != nil {
   741  		return nil, err
   742  	}
   743  	return ch, nil
   744  }
   745  
   746  func hintEncode(buf []byte, h [][n]byte, p parameters) []byte {
   747  	ω, k := p.ω, p.k
   748  	out, y := sliceForAppend(buf, ω+k)
   749  	var idx byte
   750  	for i := range k {
   751  		for j := range n {
   752  			if h[i][j] != 0 {
   753  				y[idx] = byte(j)
   754  				idx++
   755  			}
   756  		}
   757  		y[ω+i] = idx
   758  	}
   759  	return out
   760  }
   761  
   762  func hintDecode(y []byte, h [][n]byte, p parameters) error {
   763  	ω, k := p.ω, p.k
   764  	if len(y) != ω+k {
   765  		return errors.New("mldsa: internal error: invalid signature hint length")
   766  	}
   767  	var idx byte
   768  	for i := range k {
   769  		limit := y[ω+i]
   770  		if limit < idx || limit > byte(ω) {
   771  			return errInvalidSignatureHintLimits
   772  		}
   773  		first := idx
   774  		for idx < limit {
   775  			if idx > first && y[idx-1] >= y[idx] {
   776  				return errInvalidSignatureHintIndexOrder
   777  			}
   778  			h[i][y[idx]] = 1
   779  			idx++
   780  		}
   781  	}
   782  	for i := idx; i < byte(ω); i++ {
   783  		if y[i] != 0 {
   784  			return errInvalidSignatureHintExtraIndices
   785  		}
   786  	}
   787  	return nil
   788  }
   789  

View as plain text