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