package composite import ( "crypto/rand" "io" "math" "math/big" "git.smesh.lol/gnarl-hamadryad/crypto/ring" ) func Demo16() NTRUParam { return NTRUParam{ Ring: ring.Params{N: 16, Q: 7681, RootOfUnity: 5235, MontR: 1 << 16, QInv: 7679}, Sigma: math.Sqrt(16.0) * 2.0, Tail: 13, } } func DefaultParam() NTRUParam { return Demo16() } type NTRUParam struct{ Ring ring.Params; Sigma float64; Tail int } type NTRUPublicKey struct{ H *ring.Poly; P NTRUParam } type NTRUPrivateKey struct{ F, G, f, g *ring.Poly; Sigma float64; PK *NTRUPublicKey } type NTRUSignature struct{ Salt [16]byte; S2 *ring.Poly } func NTRUKeyGen(p NTRUParam) (*NTRUPublicKey, *NTRUPrivateKey) { return NTRUKeyGenFrom(p, rand.Reader) } func NTRUKeyGenFrom(p NTRUParam, rng io.Reader) (*NTRUPublicKey, *NTRUPrivateKey) { rp := p.Ring; sigma := p.Sigma gs := ring.NewGaussianSamplerFrom(sigma, rng) var f, g, h *ring.Poly for { f = gs.SamplePoly(rp); g = gs.SamplePoly(rp) fInv := polyInverseModQ(f, rp) if fInv == nil { continue } h = ring.Mul(g, fInv) break } F, G := computeFGAdj(f, g, rp) pk := &NTRUPublicKey{H: h, P: p} sk := &NTRUPrivateKey{F: F, G: G, f: f, g: g, Sigma: sigma, PK: pk} return pk, sk } // computeFGAdj solves f*G - g*F = q using the adjugate approach: // 1. Compute M_f, its determinant D, and first adjugate column A[:,0]. // 2. G_init = round(q * A[:,0] / D) — exact rational, rounded to integer. // 3. Newton-correct the residual r = f*G - q (bounded by ||f|| per coeff). // 4. Babai-reduce (F,G) against (f,g). func computeFGAdj(f, g *ring.Poly, rp ring.Params) (*ring.Poly, *ring.Poly) { n := rp.N; q := int64(rp.Q) mf := buildRingHMatrix(f, rp) // Compute D = det(M_f) and adjCol0 = first column of adjugate (both big.Int). D, adjCol0 := determinantAndAdjCol0(mf) if D.Sign() == 0 { return ring.New(rp), ring.New(rp) } fInt := polyToInt64(f, q) gInt := polyToInt64(g, q) Gint := make([]int64, n) Fint := make([]int64, n) bigQ := big.NewInt(q) for i := 0; i < n; i++ { num := new(big.Int).Mul(bigQ, new(big.Int).Set(adjCol0[i])) qq := new(big.Int).Quo(num, D) rem := new(big.Int).Rem(num, D) twiceRem := new(big.Int).Mul(rem, big.NewInt(2)) twiceRem.Abs(twiceRem) DAbs := new(big.Int).Abs(D) if twiceRem.Cmp(DAbs) >= 0 { if num.Sign() >= 0 { qq.Add(qq, big.NewInt(1)) } else { qq.Sub(qq, big.NewInt(1)) } } Gint[i] = qq.Int64() } // Newton: correct residual. for iter := 0; iter < 4; iter++ { fG := mulPolyInt(fInt, Gint, n) gF := mulPolyInt(gInt, Fint, n) allZero := true for i := 0; i < n; i++ { r := fG[i] - gF[i]; if i == 0 { r -= q } if r != 0 { allZero = false; break } } if allZero { break } // Approx correction: δF=0, solve f*δG = -r over Z_q. rhs := make([]int64, n) for i := 0; i < n; i++ { r := fG[i] - gF[i]; if i == 0 { r -= q } rhs[i] = (-r) % q; if rhs[i] < 0 { rhs[i] += q } } dG := solveModQ(mf, rhs, n, q) if dG == nil { break } for i := 0; i < n; i++ { cv := int64(dG[i]); if cv > q/2 { cv -= q } Gint[i] += cv } } // Babai reduction. for iter := 0; iter < 20; iter++ { ns := vecDot(fInt, fInt) + vecDot(gInt, gInt) if ns == 0 { break } dp := vecDot(Fint, fInt) + vecDot(Gint, gInt) k := int(math.Round(float64(dp) / float64(ns))) if k == 0 { break } for i := 0; i < n; i++ { Fint[i] -= int64(k) * fInt[i]; Gint[i] -= int64(k) * gInt[i] } } F := ring.New(rp); G := ring.New(rp) for i := 0; i < n; i++ { F.Coeffs[i] = modU32(Fint[i], rp.Q); G.Coeffs[i] = modU32(Gint[i], rp.Q) } return F, G } // determinantAndAdjCol0 computes det(M) and the first column of adj(M) for a // square n×n matrix M. Uses Bareiss fraction-free Gaussian elimination to // compute D. Then solves M·x = D·e_0 to get adjCol0 = D·x (the scaled adjugate). func determinantAndAdjCol0(M [][]int64) (*big.Int, []*big.Int) { n := len(M) if n == 0 { return nil, nil } // Bareiss: maintain matrix in-place, prev = previous pivot (initially 1). A := make([][]*big.Int, n) for i := 0; i < n; i++ { A[i] = make([]*big.Int, n) for j := 0; j < n; j++ { A[i][j] = big.NewInt(M[i][j]) } } prev := big.NewInt(1) for k := 0; k < n-1; k++ { // If A[k][k] == 0, find non-zero pivot below and swap. if A[k][k].Sign() == 0 { piv := k + 1 for piv < n && A[piv][k].Sign() == 0 { piv++ } if piv == n { return big.NewInt(0), nil } A[k], A[piv] = A[piv], A[k] } for i := k + 1; i < n; i++ { for j := k + 1; j < n; j++ { // A[i][j] = (A[k][k]*A[i][j] - A[i][k]*A[k][j]) / prev t1 := new(big.Int).Mul(A[k][k], A[i][j]) t2 := new(big.Int).Mul(A[i][k], A[k][j]) A[i][j].Sub(t1, t2) A[i][j].Div(A[i][j], prev) } A[i][k].SetInt64(0) } prev = A[k][k] } // D = A[n-1][n-1]. D := new(big.Int).Set(A[n-1][n-1]) // Solve M · x = D · e_0 for x. Then adjCol0 = D · x. // x = M^{-1} · (D · e_0) = D · (first column of M^{-1}). // Equivalent: solve M · y = e_0 for y = first column of M^{-1}, then adjCol0 = D · y. // Use augmented matrix [M | e_0] solved over Q via big.Rat. rows := make([][]*big.Rat, n) for i := 0; i < n; i++ { rows[i] = make([]*big.Rat, n+1) for j := 0; j < n; j++ { rows[i][j] = new(big.Rat).SetInt64(M[i][j]) } rows[i][n] = new(big.Rat) } rows[0][n].SetInt64(1) for col := 0; col < n; col++ { p := col for r := col; r < n; r++ { if rows[r][col].Sign() != 0 { p = r; break } } if p != col { rows[col], rows[p] = rows[p], rows[col] } piv := new(big.Rat).Set(rows[col][col]) for j := col; j <= n; j++ { rows[col][j].Quo(rows[col][j], piv) } for row := 0; row < n; row++ { if row == col { continue } fc := new(big.Rat).Set(rows[row][col]) if fc.Sign() == 0 { continue } for j := col; j <= n; j++ { tmp := new(big.Rat).Mul(fc, rows[col][j]) rows[row][j].Sub(rows[row][j], tmp) } } } adjCol0 := make([]*big.Int, n) for i := 0; i < n; i++ { // y_i = first column of M^{-1} = rows[i][n] (rational). // adjCol0[i] = D * y_i (integer, since D*M^{-1} = adj(M) is integer). num := new(big.Int).Mul(D, rows[i][n].Num()) adjCol0[i] = new(big.Int).Quo(num, rows[i][n].Denom()) } return D, adjCol0 } // solveModQ solves M*x = rhs mod q via Gaussian elimination. func solveModQ(mf [][]int64, rhs []int64, n int, q int64) []int64 { aug := make([][]int64, n) for i := 0; i < n; i++ { aug[i] = make([]int64, n+1) for j := 0; j < n; j++ { aug[i][j] = modI(mf[i][j], q) } aug[i][n] = modI(rhs[i], q) } for col := 0; col < n; col++ { p := col for r := col; r < n; r++ { if aug[r][col]%q != 0 { p = r; break } } if aug[p][col]%q == 0 { return nil } if p != col { aug[col], aug[p] = aug[p], aug[col] } piv := modI(aug[col][col], q); pi := modInverse(piv, q) for j := col; j <= n; j++ { aug[col][j] = modI(aug[col][j]*pi, q) } for row := 0; row < n; row++ { if row == col { continue } fc := modI(aug[row][col], q); if fc == 0 { continue } for j := col; j <= n; j++ { aug[row][j] = modI(aug[row][j]-fc*aug[col][j], q) } } } res := make([]int64, n) for i := 0; i < n; i++ { res[i] = aug[i][n] } return res } func modI(x, q int64) int64 { x = x % q; if x < 0 { x += q }; return x } func polyInverseModQ(f *ring.Poly, rp ring.Params) *ring.Poly { n := rp.N; q := rp.Q mf := buildRingHMatrix(f, rp) aug := make([][]int64, n) for i := 0; i < n; i++ { aug[i] = make([]int64, n+1) for j := 0; j < n; j++ { aug[i][j] = mf[i][j] } } aug[0][n] = 1 for col := 0; col < n; col++ { p := col for r := col; r < n; r++ { if aug[r][col] != 0 { p = r; break } } if aug[p][col] == 0 { return nil } if p != col { aug[col], aug[p] = aug[p], aug[col] } pv := aug[col][col]; pi := modInverse(pv, int64(q)) for j := col; j <= n; j++ { aug[col][j] = modI(aug[col][j]*pi, int64(q)) } for row := 0; row < n; row++ { if row == col { continue } fc := aug[row][col]; if fc == 0 { continue } for j := col; j <= n; j++ { aug[row][j] = modI(aug[row][j]-fc*aug[col][j], int64(q)) } } } r := ring.New(rp) for i := 0; i < n; i++ { r.Coeffs[i] = uint32(aug[i][n]) } return r } func modInverse(a, mod int64) int64 { a %= mod; if a < 0 { a += mod } or, r, os, s := a, mod, int64(1), int64(0) for r != 0 { qq := or / r; or, r = r, or-qq*r; os, s = s, os-qq*s } if or != 1 { return 0 } if os < 0 { os += mod } return os } func buildRingHMatrix(h *ring.Poly, rp ring.Params) [][]int64 { n := rp.N H := make([][]int64, n) for k := 0; k < n; k++ { H[k] = make([]int64, n) for j := 0; j <= k; j++ { H[k][j] = int64(h.Coeffs[k-j]) } for j := k + 1; j < n; j++ { H[k][j] = -int64(h.Coeffs[k-j+n]) } } return H } func modU32(x int64, q uint32) uint32 { for x < 0 { x += int64(q) }; return uint32(uint64(x) % uint64(q)) } func polyToInt64(p *ring.Poly, q int64) []int64 { o := make([]int64, len(p.Coeffs)); h := q / 2 for i, c := range p.Coeffs { cv := int64(c); if cv > h { cv -= q }; o[i] = cv } return o } func mulPolyInt(a, b []int64, n int) []int64 { o := make([]int64, n) for i := 0; i < n; i++ { for j := 0; j < n; j++ { k := i + j; if k < n { o[k] += a[i]*b[j] } else { o[k-n] -= a[i]*b[j] } }} return o } func vecDot(a, b []int64) int64 { var s int64; for i := range a { s += a[i]*b[i] }; return s } func ffSampling(f, g, F, G *ring.Poly, target *ring.Poly, sigma float64, rng io.Reader, rp ring.Params) (*ring.Poly, *ring.Poly) { b1T := g; b1B := ring.Neg(f); b2T := ring.Neg(G); b2B := F tT := ring.New(rp); tB := target.Clone() n1 := float64(polyDot(b1T,b1T)+polyDot(b1B,b1B)) mu := float64(polyDot(b2T,b1T)+polyDot(b2B,b1B)) / n1 b2sT := polySub(b2T, polyScale(b1T, mu)) b2sB := polySub(b2B, polyScale(b1B, mu)) n2 := float64(polyDot(b2sT,b2sT)+polyDot(b2sB,b2sB)) if n1 < 1 { n1 = 1 }; if n2 < 1 { n2 = 1 } c2 := (float64(polyDot(tT,b2T)+polyDot(tB,b2B))-mu*float64(polyDot(tT,b1T)+polyDot(tB,b1B)))/n2 z2 := ring.NewGaussianSamplerFrom(sigma/math.Sqrt(n2), rng).SampleZ(c2) tT = polySub(tT, polyScale(b2T, float64(z2))); tB = polySub(tB, polyScale(b2B, float64(z2))) c1 := float64(polyDot(tT,b1T)+polyDot(tB,b1B))/n1 z1 := ring.NewGaussianSamplerFrom(sigma/math.Sqrt(n1), rng).SampleZ(c1) tT = polySub(tT, polyScale(b1T, float64(z1))); tB = polySub(tB, polyScale(b1B, float64(z1))) return tT, tB } func polyDot(a, b *ring.Poly) int64 { q := a.Params().Q; h := int64(q/2); var s int64 for i := range a.Coeffs { ai := int64(a.Coeffs[i]); bi := int64(b.Coeffs[i]) if ai > h { ai -= int64(q) }; if bi > h { bi -= int64(q) }; s += ai*bi } return s } func polyScale(a *ring.Poly, s float64) *ring.Poly { r := ring.New(a.Params()); si := int64(math.Round(math.Abs(s))) if si == 0 { return r } q := int64(a.Params().Q); h := q/2 for i, c := range a.Coeffs { cv := int64(c); if cv > h { cv -= q } cv *= si; if s < 0 { cv = -cv }; cv %= q; if cv < 0 { cv += q }; r.Coeffs[i] = uint32(cv) } return r } func polySub(a, b *ring.Poly) *ring.Poly { return ring.Sub(a, b) } func NTRUSign(sk *NTRUPrivateKey, msg []byte) *NTRUSignature { return NTRUSignFrom(sk, msg, rand.Reader) } func NTRUSignFrom(sk *NTRUPrivateKey, msg []byte, rng io.Reader) *NTRUSignature { if rng == nil { rng = rand.Reader } var salt [16]byte; io.ReadFull(rng, salt[:]) t := hashToPoly(salt[:], sk.PK.H, msg, sk.PK.P.Ring) _, s2 := ffSampling(sk.f, sk.g, sk.F, sk.G, t, sk.Sigma, rng, sk.PK.P.Ring) return &NTRUSignature{Salt: salt, S2: s2} } func NTRUVerify(pk *NTRUPublicKey, msg []byte, sig *NTRUSignature) bool { t := hashToPoly(sig.Salt[:], pk.H, msg, pk.P.Ring) hs2 := ring.Mul(pk.H, sig.S2); s1 := ring.Sub(t, hs2) b := uint32(pk.P.Sigma*(1.5+float64(pk.P.Tail))) return ring.Norm(s1) <= b && ring.Norm(sig.S2) <= b } func hashToPoly(salt []byte, h *ring.Poly, msg []byte, rp ring.Params) *ring.Poly { in := make([]byte, 0, len(salt)+len(msg)+128) in = append(in, []byte("comp-ntru-v1")...) in = append(in, salt...); in = append(in, ring.Serialize(h)...); in = append(in, msg...) var st uint64 = 14695981039346656037 c := ring.New(rp) for i := 0; i < rp.N; i++ { for _, b := range in { st ^= uint64(b); st *= 1099511628211 } st ^= uint64(i); st *= 1099511628211; c.Coeffs[i] = uint32(st % uint64(rp.Q)) } return c } func (pk *NTRUPublicKey) MarshalBinary() []byte { return ring.Serialize(pk.H) } func UnmarshalNTRUPK(d []byte, p NTRUParam) (*NTRUPublicKey, error) { h := ring.Deserialize(p.Ring, d); if h == nil { return nil, errShort } return &NTRUPublicKey{H: h, P: p}, nil } func (sig *NTRUSignature) MarshalBinary() []byte { b := make([]byte, 16+len(ring.Serialize(sig.S2))) copy(b[:16], sig.Salt[:]); copy(b[16:], ring.Serialize(sig.S2)); return b } func UnmarshalNTRUSig(d []byte, rp ring.Params) (*NTRUSignature, error) { bits := 0; for v := rp.Q-1; v > 0; v >>= 1 { bits++ } sb := (rp.N*bits+7)/8 if len(d) < 16+sb { return nil, errShort } sig := &NTRUSignature{}; copy(sig.Salt[:], d[:16]) sig.S2 = ring.Deserialize(rp, d[16:16+sb]) if sig.S2 == nil { return nil, errShort }; return sig, nil } func SigBytes(p NTRUParam) int { bits := 0; for v := p.Ring.Q-1; v > 0; v >>= 1 { bits++ }; return 16+(p.Ring.N*bits+7)/8 } var errShort = errStr("composite: data too short") type errStr string func (e errStr) Error() string { return string(e) }