ntru.go raw

   1  package composite
   2  
   3  import (
   4  	"crypto/rand"
   5  	"io"
   6  	"math"
   7  	"math/big"
   8  
   9  	"git.smesh.lol/gnarl-hamadryad/crypto/ring"
  10  )
  11  
  12  func Demo16() NTRUParam {
  13  	return NTRUParam{
  14  		Ring: ring.Params{N: 16, Q: 7681, RootOfUnity: 5235, MontR: 1 << 16, QInv: 7679},
  15  		Sigma: math.Sqrt(16.0) * 2.0, Tail: 13,
  16  	}
  17  }
  18  func DefaultParam() NTRUParam { return Demo16() }
  19  
  20  type NTRUParam struct{ Ring ring.Params; Sigma float64; Tail int }
  21  type NTRUPublicKey struct{ H *ring.Poly; P NTRUParam }
  22  type NTRUPrivateKey struct{ F, G, f, g *ring.Poly; Sigma float64; PK *NTRUPublicKey }
  23  type NTRUSignature struct{ Salt [16]byte; S2 *ring.Poly }
  24  
  25  func NTRUKeyGen(p NTRUParam) (*NTRUPublicKey, *NTRUPrivateKey) {
  26  	return NTRUKeyGenFrom(p, rand.Reader)
  27  }
  28  
  29  func NTRUKeyGenFrom(p NTRUParam, rng io.Reader) (*NTRUPublicKey, *NTRUPrivateKey) {
  30  	rp := p.Ring; sigma := p.Sigma
  31  	gs := ring.NewGaussianSamplerFrom(sigma, rng)
  32  
  33  	var f, g, h *ring.Poly
  34  	for {
  35  		f = gs.SamplePoly(rp); g = gs.SamplePoly(rp)
  36  		fInv := polyInverseModQ(f, rp)
  37  		if fInv == nil { continue }
  38  		h = ring.Mul(g, fInv)
  39  		break
  40  	}
  41  
  42  	F, G := computeFGAdj(f, g, rp)
  43  
  44  	pk := &NTRUPublicKey{H: h, P: p}
  45  	sk := &NTRUPrivateKey{F: F, G: G, f: f, g: g, Sigma: sigma, PK: pk}
  46  	return pk, sk
  47  }
  48  
  49  // computeFGAdj solves f*G - g*F = q using the adjugate approach:
  50  // 1. Compute M_f, its determinant D, and first adjugate column A[:,0].
  51  // 2. G_init = round(q * A[:,0] / D) — exact rational, rounded to integer.
  52  // 3. Newton-correct the residual r = f*G - q (bounded by ||f|| per coeff).
  53  // 4. Babai-reduce (F,G) against (f,g).
  54  func computeFGAdj(f, g *ring.Poly, rp ring.Params) (*ring.Poly, *ring.Poly) {
  55  	n := rp.N; q := int64(rp.Q)
  56  	mf := buildRingHMatrix(f, rp)
  57  
  58  	// Compute D = det(M_f) and adjCol0 = first column of adjugate (both big.Int).
  59  	D, adjCol0 := determinantAndAdjCol0(mf)
  60  	if D.Sign() == 0 { return ring.New(rp), ring.New(rp) }
  61  
  62  	fInt := polyToInt64(f, q)
  63  	gInt := polyToInt64(g, q)
  64  	Gint := make([]int64, n)
  65  	Fint := make([]int64, n)
  66  
  67  	bigQ := big.NewInt(q)
  68  	for i := 0; i < n; i++ {
  69  		num := new(big.Int).Mul(bigQ, new(big.Int).Set(adjCol0[i]))
  70  		qq := new(big.Int).Quo(num, D)
  71  		rem := new(big.Int).Rem(num, D)
  72  		twiceRem := new(big.Int).Mul(rem, big.NewInt(2))
  73  		twiceRem.Abs(twiceRem)
  74  		DAbs := new(big.Int).Abs(D)
  75  		if twiceRem.Cmp(DAbs) >= 0 {
  76  			if num.Sign() >= 0 { qq.Add(qq, big.NewInt(1)) } else { qq.Sub(qq, big.NewInt(1)) }
  77  		}
  78  		Gint[i] = qq.Int64()
  79  	}
  80  
  81  	// Newton: correct residual.
  82  	for iter := 0; iter < 4; iter++ {
  83  		fG := mulPolyInt(fInt, Gint, n)
  84  		gF := mulPolyInt(gInt, Fint, n)
  85  		allZero := true
  86  		for i := 0; i < n; i++ {
  87  			r := fG[i] - gF[i]; if i == 0 { r -= q }
  88  			if r != 0 { allZero = false; break }
  89  		}
  90  		if allZero { break }
  91  
  92  		// Approx correction: δF=0, solve f*δG = -r over Z_q.
  93  		rhs := make([]int64, n)
  94  		for i := 0; i < n; i++ {
  95  			r := fG[i] - gF[i]; if i == 0 { r -= q }
  96  			rhs[i] = (-r) % q; if rhs[i] < 0 { rhs[i] += q }
  97  		}
  98  		dG := solveModQ(mf, rhs, n, q)
  99  		if dG == nil { break }
 100  		for i := 0; i < n; i++ {
 101  			cv := int64(dG[i]); if cv > q/2 { cv -= q }
 102  			Gint[i] += cv
 103  		}
 104  	}
 105  
 106  	// Babai reduction.
 107  	for iter := 0; iter < 20; iter++ {
 108  		ns := vecDot(fInt, fInt) + vecDot(gInt, gInt)
 109  		if ns == 0 { break }
 110  		dp := vecDot(Fint, fInt) + vecDot(Gint, gInt)
 111  		k := int(math.Round(float64(dp) / float64(ns)))
 112  		if k == 0 { break }
 113  		for i := 0; i < n; i++ { Fint[i] -= int64(k) * fInt[i]; Gint[i] -= int64(k) * gInt[i] }
 114  	}
 115  
 116  	F := ring.New(rp); G := ring.New(rp)
 117  	for i := 0; i < n; i++ { F.Coeffs[i] = modU32(Fint[i], rp.Q); G.Coeffs[i] = modU32(Gint[i], rp.Q) }
 118  	return F, G
 119  }
 120  
 121  // determinantAndAdjCol0 computes det(M) and the first column of adj(M) for a
 122  // square n×n matrix M. Uses Bareiss fraction-free Gaussian elimination to
 123  // compute D. Then solves M·x = D·e_0 to get adjCol0 = D·x (the scaled adjugate).
 124  func determinantAndAdjCol0(M [][]int64) (*big.Int, []*big.Int) {
 125  	n := len(M)
 126  	if n == 0 { return nil, nil }
 127  
 128  	// Bareiss: maintain matrix in-place, prev = previous pivot (initially 1).
 129  	A := make([][]*big.Int, n)
 130  	for i := 0; i < n; i++ {
 131  		A[i] = make([]*big.Int, n)
 132  		for j := 0; j < n; j++ {
 133  			A[i][j] = big.NewInt(M[i][j])
 134  		}
 135  	}
 136  
 137  	prev := big.NewInt(1)
 138  	for k := 0; k < n-1; k++ {
 139  		// If A[k][k] == 0, find non-zero pivot below and swap.
 140  		if A[k][k].Sign() == 0 {
 141  			piv := k + 1
 142  			for piv < n && A[piv][k].Sign() == 0 { piv++ }
 143  			if piv == n { return big.NewInt(0), nil }
 144  			A[k], A[piv] = A[piv], A[k]
 145  		}
 146  
 147  		for i := k + 1; i < n; i++ {
 148  			for j := k + 1; j < n; j++ {
 149  				// A[i][j] = (A[k][k]*A[i][j] - A[i][k]*A[k][j]) / prev
 150  				t1 := new(big.Int).Mul(A[k][k], A[i][j])
 151  				t2 := new(big.Int).Mul(A[i][k], A[k][j])
 152  				A[i][j].Sub(t1, t2)
 153  				A[i][j].Div(A[i][j], prev)
 154  			}
 155  			A[i][k].SetInt64(0)
 156  		}
 157  		prev = A[k][k]
 158  	}
 159  
 160  	// D = A[n-1][n-1].
 161  	D := new(big.Int).Set(A[n-1][n-1])
 162  
 163  	// Solve M · x = D · e_0 for x. Then adjCol0 = D · x.
 164  	// x = M^{-1} · (D · e_0) = D · (first column of M^{-1}).
 165  	// Equivalent: solve M · y = e_0 for y = first column of M^{-1}, then adjCol0 = D · y.
 166  	// Use augmented matrix [M | e_0] solved over Q via big.Rat.
 167  
 168  	rows := make([][]*big.Rat, n)
 169  	for i := 0; i < n; i++ {
 170  		rows[i] = make([]*big.Rat, n+1)
 171  		for j := 0; j < n; j++ {
 172  			rows[i][j] = new(big.Rat).SetInt64(M[i][j])
 173  		}
 174  		rows[i][n] = new(big.Rat)
 175  	}
 176  	rows[0][n].SetInt64(1)
 177  
 178  	for col := 0; col < n; col++ {
 179  		p := col
 180  		for r := col; r < n; r++ { if rows[r][col].Sign() != 0 { p = r; break } }
 181  		if p != col { rows[col], rows[p] = rows[p], rows[col] }
 182  		piv := new(big.Rat).Set(rows[col][col])
 183  		for j := col; j <= n; j++ { rows[col][j].Quo(rows[col][j], piv) }
 184  		for row := 0; row < n; row++ {
 185  			if row == col { continue }
 186  			fc := new(big.Rat).Set(rows[row][col])
 187  			if fc.Sign() == 0 { continue }
 188  			for j := col; j <= n; j++ {
 189  				tmp := new(big.Rat).Mul(fc, rows[col][j])
 190  				rows[row][j].Sub(rows[row][j], tmp)
 191  			}
 192  		}
 193  	}
 194  
 195  	adjCol0 := make([]*big.Int, n)
 196  	for i := 0; i < n; i++ {
 197  		// y_i = first column of M^{-1} = rows[i][n] (rational).
 198  		// adjCol0[i] = D * y_i (integer, since D*M^{-1} = adj(M) is integer).
 199  		num := new(big.Int).Mul(D, rows[i][n].Num())
 200  		adjCol0[i] = new(big.Int).Quo(num, rows[i][n].Denom())
 201  	}
 202  	return D, adjCol0
 203  }
 204  
 205  // solveModQ solves M*x = rhs mod q via Gaussian elimination.
 206  func solveModQ(mf [][]int64, rhs []int64, n int, q int64) []int64 {
 207  	aug := make([][]int64, n)
 208  	for i := 0; i < n; i++ {
 209  		aug[i] = make([]int64, n+1)
 210  		for j := 0; j < n; j++ { aug[i][j] = modI(mf[i][j], q) }
 211  		aug[i][n] = modI(rhs[i], q)
 212  	}
 213  	for col := 0; col < n; col++ {
 214  		p := col
 215  		for r := col; r < n; r++ { if aug[r][col]%q != 0 { p = r; break } }
 216  		if aug[p][col]%q == 0 { return nil }
 217  		if p != col { aug[col], aug[p] = aug[p], aug[col] }
 218  		piv := modI(aug[col][col], q); pi := modInverse(piv, q)
 219  		for j := col; j <= n; j++ { aug[col][j] = modI(aug[col][j]*pi, q) }
 220  		for row := 0; row < n; row++ {
 221  			if row == col { continue }
 222  			fc := modI(aug[row][col], q); if fc == 0 { continue }
 223  			for j := col; j <= n; j++ { aug[row][j] = modI(aug[row][j]-fc*aug[col][j], q) }
 224  		}
 225  	}
 226  	res := make([]int64, n)
 227  	for i := 0; i < n; i++ { res[i] = aug[i][n] }
 228  	return res
 229  }
 230  
 231  func modI(x, q int64) int64 { x = x % q; if x < 0 { x += q }; return x }
 232  
 233  func polyInverseModQ(f *ring.Poly, rp ring.Params) *ring.Poly {
 234  	n := rp.N; q := rp.Q
 235  	mf := buildRingHMatrix(f, rp)
 236  	aug := make([][]int64, n)
 237  	for i := 0; i < n; i++ {
 238  		aug[i] = make([]int64, n+1)
 239  		for j := 0; j < n; j++ { aug[i][j] = mf[i][j] }
 240  	}
 241  	aug[0][n] = 1
 242  	for col := 0; col < n; col++ {
 243  		p := col
 244  		for r := col; r < n; r++ { if aug[r][col] != 0 { p = r; break } }
 245  		if aug[p][col] == 0 { return nil }
 246  		if p != col { aug[col], aug[p] = aug[p], aug[col] }
 247  		pv := aug[col][col]; pi := modInverse(pv, int64(q))
 248  		for j := col; j <= n; j++ { aug[col][j] = modI(aug[col][j]*pi, int64(q)) }
 249  		for row := 0; row < n; row++ {
 250  			if row == col { continue }
 251  			fc := aug[row][col]; if fc == 0 { continue }
 252  			for j := col; j <= n; j++ { aug[row][j] = modI(aug[row][j]-fc*aug[col][j], int64(q)) }
 253  		}
 254  	}
 255  	r := ring.New(rp)
 256  	for i := 0; i < n; i++ { r.Coeffs[i] = uint32(aug[i][n]) }
 257  	return r
 258  }
 259  
 260  func modInverse(a, mod int64) int64 {
 261  	a %= mod; if a < 0 { a += mod }
 262  	or, r, os, s := a, mod, int64(1), int64(0)
 263  	for r != 0 { qq := or / r; or, r = r, or-qq*r; os, s = s, os-qq*s }
 264  	if or != 1 { return 0 }
 265  	if os < 0 { os += mod }
 266  	return os
 267  }
 268  
 269  func buildRingHMatrix(h *ring.Poly, rp ring.Params) [][]int64 {
 270  	n := rp.N
 271  	H := make([][]int64, n)
 272  	for k := 0; k < n; k++ {
 273  		H[k] = make([]int64, n)
 274  		for j := 0; j <= k; j++ { H[k][j] = int64(h.Coeffs[k-j]) }
 275  		for j := k + 1; j < n; j++ { H[k][j] = -int64(h.Coeffs[k-j+n]) }
 276  	}
 277  	return H
 278  }
 279  
 280  func modU32(x int64, q uint32) uint32 { for x < 0 { x += int64(q) }; return uint32(uint64(x) % uint64(q)) }
 281  func polyToInt64(p *ring.Poly, q int64) []int64 {
 282  	o := make([]int64, len(p.Coeffs)); h := q / 2
 283  	for i, c := range p.Coeffs { cv := int64(c); if cv > h { cv -= q }; o[i] = cv }
 284  	return o
 285  }
 286  func mulPolyInt(a, b []int64, n int) []int64 {
 287  	o := make([]int64, n)
 288  	for i := 0; i < n; i++ { for j := 0; j < n; j++ {
 289  		k := i + j; if k < n { o[k] += a[i]*b[j] } else { o[k-n] -= a[i]*b[j] }
 290  	}}
 291  	return o
 292  }
 293  func vecDot(a, b []int64) int64 { var s int64; for i := range a { s += a[i]*b[i] }; return s }
 294  
 295  func ffSampling(f, g, F, G *ring.Poly, target *ring.Poly, sigma float64, rng io.Reader, rp ring.Params) (*ring.Poly, *ring.Poly) {
 296  	b1T := g; b1B := ring.Neg(f); b2T := ring.Neg(G); b2B := F
 297  	tT := ring.New(rp); tB := target.Clone()
 298  	n1 := float64(polyDot(b1T,b1T)+polyDot(b1B,b1B))
 299  	mu := float64(polyDot(b2T,b1T)+polyDot(b2B,b1B)) / n1
 300  	b2sT := polySub(b2T, polyScale(b1T, mu))
 301  	b2sB := polySub(b2B, polyScale(b1B, mu))
 302  	n2 := float64(polyDot(b2sT,b2sT)+polyDot(b2sB,b2sB))
 303  	if n1 < 1 { n1 = 1 }; if n2 < 1 { n2 = 1 }
 304  	c2 := (float64(polyDot(tT,b2T)+polyDot(tB,b2B))-mu*float64(polyDot(tT,b1T)+polyDot(tB,b1B)))/n2
 305  	z2 := ring.NewGaussianSamplerFrom(sigma/math.Sqrt(n2), rng).SampleZ(c2)
 306  	tT = polySub(tT, polyScale(b2T, float64(z2))); tB = polySub(tB, polyScale(b2B, float64(z2)))
 307  	c1 := float64(polyDot(tT,b1T)+polyDot(tB,b1B))/n1
 308  	z1 := ring.NewGaussianSamplerFrom(sigma/math.Sqrt(n1), rng).SampleZ(c1)
 309  	tT = polySub(tT, polyScale(b1T, float64(z1))); tB = polySub(tB, polyScale(b1B, float64(z1)))
 310  	return tT, tB
 311  }
 312  
 313  func polyDot(a, b *ring.Poly) int64 {
 314  	q := a.Params().Q; h := int64(q/2); var s int64
 315  	for i := range a.Coeffs {
 316  		ai := int64(a.Coeffs[i]); bi := int64(b.Coeffs[i])
 317  		if ai > h { ai -= int64(q) }; if bi > h { bi -= int64(q) }; s += ai*bi
 318  	}
 319  	return s
 320  }
 321  func polyScale(a *ring.Poly, s float64) *ring.Poly {
 322  	r := ring.New(a.Params()); si := int64(math.Round(math.Abs(s)))
 323  	if si == 0 { return r }
 324  	q := int64(a.Params().Q); h := q/2
 325  	for i, c := range a.Coeffs {
 326  		cv := int64(c); if cv > h { cv -= q }
 327  		cv *= si; if s < 0 { cv = -cv }; cv %= q; if cv < 0 { cv += q }; r.Coeffs[i] = uint32(cv)
 328  	}
 329  	return r
 330  }
 331  func polySub(a, b *ring.Poly) *ring.Poly { return ring.Sub(a, b) }
 332  
 333  func NTRUSign(sk *NTRUPrivateKey, msg []byte) *NTRUSignature {
 334  	return NTRUSignFrom(sk, msg, rand.Reader)
 335  }
 336  func NTRUSignFrom(sk *NTRUPrivateKey, msg []byte, rng io.Reader) *NTRUSignature {
 337  	if rng == nil { rng = rand.Reader }
 338  	var salt [16]byte; io.ReadFull(rng, salt[:])
 339  	t := hashToPoly(salt[:], sk.PK.H, msg, sk.PK.P.Ring)
 340  	_, s2 := ffSampling(sk.f, sk.g, sk.F, sk.G, t, sk.Sigma, rng, sk.PK.P.Ring)
 341  	return &NTRUSignature{Salt: salt, S2: s2}
 342  }
 343  func NTRUVerify(pk *NTRUPublicKey, msg []byte, sig *NTRUSignature) bool {
 344  	t := hashToPoly(sig.Salt[:], pk.H, msg, pk.P.Ring)
 345  	hs2 := ring.Mul(pk.H, sig.S2); s1 := ring.Sub(t, hs2)
 346  	b := uint32(pk.P.Sigma*(1.5+float64(pk.P.Tail)))
 347  	return ring.Norm(s1) <= b && ring.Norm(sig.S2) <= b
 348  }
 349  func hashToPoly(salt []byte, h *ring.Poly, msg []byte, rp ring.Params) *ring.Poly {
 350  	in := make([]byte, 0, len(salt)+len(msg)+128)
 351  	in = append(in, []byte("comp-ntru-v1")...)
 352  	in = append(in, salt...); in = append(in, ring.Serialize(h)...); in = append(in, msg...)
 353  	var st uint64 = 14695981039346656037
 354  	c := ring.New(rp)
 355  	for i := 0; i < rp.N; i++ {
 356  		for _, b := range in { st ^= uint64(b); st *= 1099511628211 }
 357  		st ^= uint64(i); st *= 1099511628211; c.Coeffs[i] = uint32(st % uint64(rp.Q))
 358  	}
 359  	return c
 360  }
 361  
 362  func (pk *NTRUPublicKey) MarshalBinary() []byte { return ring.Serialize(pk.H) }
 363  func UnmarshalNTRUPK(d []byte, p NTRUParam) (*NTRUPublicKey, error) {
 364  	h := ring.Deserialize(p.Ring, d); if h == nil { return nil, errShort }
 365  	return &NTRUPublicKey{H: h, P: p}, nil
 366  }
 367  func (sig *NTRUSignature) MarshalBinary() []byte {
 368  	b := make([]byte, 16+len(ring.Serialize(sig.S2)))
 369  	copy(b[:16], sig.Salt[:]); copy(b[16:], ring.Serialize(sig.S2)); return b
 370  }
 371  func UnmarshalNTRUSig(d []byte, rp ring.Params) (*NTRUSignature, error) {
 372  	bits := 0; for v := rp.Q-1; v > 0; v >>= 1 { bits++ }
 373  	sb := (rp.N*bits+7)/8
 374  	if len(d) < 16+sb { return nil, errShort }
 375  	sig := &NTRUSignature{}; copy(sig.Salt[:], d[:16])
 376  	sig.S2 = ring.Deserialize(rp, d[16:16+sb])
 377  	if sig.S2 == nil { return nil, errShort }; return sig, nil
 378  }
 379  func SigBytes(p NTRUParam) int {
 380  	bits := 0; for v := p.Ring.Q-1; v > 0; v >>= 1 { bits++ }; return 16+(p.Ring.N*bits+7)/8
 381  }
 382  var errShort = errStr("composite: data too short")
 383  type errStr string
 384  func (e errStr) Error() string { return string(e) }
 385