mirror of
https://github.com/luxfi/math.git
synced 2026-07-27 03:38:49 +00:00
ntt, poly, rns, sample
Eight pure-Go reference packages that own the canonical semantics of
every cryptographic-math primitive Lux protocols share. Backends
(AVX2 / NEON / cgo+C++ / CUDA / Metal / WGSL) plug in behind the
substrate; KATs gate byte-equality across them.
Packages added:
* params/ — ModulusID, NTTParamID, FHEParamID, PulsarParamID,
HashSuiteID, BackendID; KATHeader required-fields
schema. Stable string IDs; renaming is a breaking
change.
* backend/ — Policy enum (PureGo / NativeCPU / GPUPreferred /
GPURequired); Resolve(policy, registered) returns the
BackendID per LP-107's fallback chain. GPURequired
returns ErrUnavailable when no GPU backend is
registered.
* codec/ — Bounded readers (closes lattice issues #2 + #4 DoS
class permanently). Limits + LimitError + Reader.
ReadUint{16,32,64}Slice all reject the 70T-element
attack input before allocating; depth + frame-bytes
caps; iterative (no recursion). Regression test
TestReadUint64Slice_RejectsHugeLength encodes the
original lattice issue #4 input.
* modarith/ — Modulus type with QInv, R2, Barrett constants
(computed via math/big once per modulus). AddMod /
SubMod / MulMod (Div64 reference path). MontMulMod /
ToMontgomery / FromMontgomery cross-checked against
MulMod across 100 random Pulsar-Q pairs.
ReductionMode enum byte-equal to lattice/types.
* ntt/ — Service + Backend interface; pure-Go backend
delegates to lattice/v7/ring.SubRing.NTT/INTT (the
canonical Lattigo-derived Montgomery NTT body — no
re-implementation, just dispatch). Round-trip +
batch + determinism tests on Pulsar N=256.
* poly/ — Add / Sub / ScalarMul / PointwiseMul (NTT-domain) /
negacyclic Mul (NTT round-trip). Composes modarith +
ntt; no lower-level reach.
* rns/ — Basis(Moduli, Name) for FHE RNS chains. Single +
two-prime construction validated; rejects even moduli.
Phase 3 will add basis-extension and modulus-switching.
* sample/ — Uniform (rejection sampling, mask-and-retry); Ternary
(density + sign byte); CenteredBinomial (popcount
difference); DiscreteGaussianRejection (6-sigma cutoff).
All take an io.Reader so callers can KAT-replay.
Architecture invariants enforced in package docs:
- Go is the canonical semantic reference.
- Backend selection MUST NOT alter transcript bytes (consensus
paths default to PolicyPureGo / PolicyNativeCPU).
- No unbounded codec readers anywhere near wire formats.
- No re-implementation: where lattice/v7/ring already owns the
canonical body, we delegate, not fork.
- One ID space; renaming is a breaking change.
Test posture: 8/8 packages green via `GOWORK=off go test ./...`.
Phases 3-7 (queued):
3. lattice consumes math
4. pulsar consumes lattice + math
5. fhe consumes math
6. luxcpp/crypto/math native backend mirror
7. cross-runtime KAT release gate
See SUBSTRATE.md for the full posture and migration plan; full LP at
~/work/lux/lps/LP-107-lux-math-substrate.md.
104 lines
3.0 KiB
Go
104 lines
3.0 KiB
Go
// Copyright (c) 2026 Lux Industries Inc.
|
|
// SPDX-License-Identifier: BSD-3-Clause
|
|
|
|
// Package poly provides polynomial-arithmetic primitives over R_q =
|
|
// Z_q[X] / (X^N + 1) used by every Lux lattice protocol.
|
|
//
|
|
// LP-107 §"Polynomial and RNS operations" — the canonical motivation.
|
|
// Add, sub, scalar-mul, NTT-domain mul; converted to /from NTT domain
|
|
// via package ntt.
|
|
//
|
|
// Phase 2 (this file): pure-Go reference implementation. Body uses
|
|
// luxfi/math/modarith for the field arithmetic and luxfi/math/ntt
|
|
// for the transform; no re-implementation, just composition.
|
|
package poly
|
|
|
|
import (
|
|
"fmt"
|
|
|
|
"github.com/luxfi/math/modarith"
|
|
"github.com/luxfi/math/ntt"
|
|
)
|
|
|
|
// Add returns dst = a + b (mod q). All inputs must have length N.
|
|
func Add(dst, a, b []uint64, q uint64) error {
|
|
if len(dst) != len(a) || len(a) != len(b) {
|
|
return fmt.Errorf("poly.Add: length mismatch dst=%d a=%d b=%d",
|
|
len(dst), len(a), len(b))
|
|
}
|
|
for i := range a {
|
|
dst[i] = modarith.AddMod(a[i], b[i], q)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Sub returns dst = a - b (mod q).
|
|
func Sub(dst, a, b []uint64, q uint64) error {
|
|
if len(dst) != len(a) || len(a) != len(b) {
|
|
return fmt.Errorf("poly.Sub: length mismatch dst=%d a=%d b=%d",
|
|
len(dst), len(a), len(b))
|
|
}
|
|
for i := range a {
|
|
dst[i] = modarith.SubMod(a[i], b[i], q)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// ScalarMul returns dst = a * scalar (mod q).
|
|
func ScalarMul(dst, a []uint64, scalar, q uint64) error {
|
|
if len(dst) != len(a) {
|
|
return fmt.Errorf("poly.ScalarMul: length mismatch dst=%d a=%d",
|
|
len(dst), len(a))
|
|
}
|
|
for i := range a {
|
|
dst[i] = modarith.MulMod(a[i], scalar, q)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// PointwiseMul returns dst = a * b (pointwise, NTT domain) (mod q).
|
|
// Inputs must already be in NTT domain; output is also in NTT domain.
|
|
// Use ntt.Service.Inverse to bring back to coefficient domain.
|
|
func PointwiseMul(dst, a, b []uint64, q uint64) error {
|
|
if len(dst) != len(a) || len(a) != len(b) {
|
|
return fmt.Errorf("poly.PointwiseMul: length mismatch dst=%d a=%d b=%d",
|
|
len(dst), len(a), len(b))
|
|
}
|
|
for i := range a {
|
|
dst[i] = modarith.MulMod(a[i], b[i], q)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Mul computes the negacyclic polynomial product dst = a * b (mod q,
|
|
// mod X^N + 1) via NTT round-trip: NTT(a), NTT(b), pointwise mul,
|
|
// inverse NTT. dst, a, b must each have length p.N. Inputs are in
|
|
// coefficient domain; output is in coefficient domain.
|
|
//
|
|
// dst MAY alias a or b.
|
|
func Mul(dst, a, b []uint64, svc *ntt.Service) error {
|
|
N := int(svc.Params().N)
|
|
q := svc.Params().Q
|
|
if len(dst) != N || len(a) != N || len(b) != N {
|
|
return fmt.Errorf("poly.Mul: length mismatch dst=%d a=%d b=%d N=%d",
|
|
len(dst), len(a), len(b), N)
|
|
}
|
|
aN := make([]uint64, N)
|
|
bN := make([]uint64, N)
|
|
copy(aN, a)
|
|
copy(bN, b)
|
|
if err := svc.Forward(aN, 1); err != nil {
|
|
return fmt.Errorf("poly.Mul: NTT(a): %w", err)
|
|
}
|
|
if err := svc.Forward(bN, 1); err != nil {
|
|
return fmt.Errorf("poly.Mul: NTT(b): %w", err)
|
|
}
|
|
if err := PointwiseMul(dst, aN, bN, q); err != nil {
|
|
return err
|
|
}
|
|
if err := svc.Inverse(dst, 1); err != nil {
|
|
return fmt.Errorf("poly.Mul: INTT(dst): %w", err)
|
|
}
|
|
return nil
|
|
}
|