Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 17 additions & 5 deletions std/math/emulated/custommod.go
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ import (
// binary decomposition, comparison, or hashing), use [Field.ModMulCanonical] or
// reduce with [Field.assertLessThanModulus].
func (f *Field[T]) ModMul(a, b *Element[T], modulus *Element[T]) *Element[T] {
f.checkModulus(modulus)
modulus = f.checkModulus(modulus)
// fast path when either of the inputs is zero then result is always zero
if len(a.Limbs) == 0 || len(b.Limbs) == 0 {
return f.Zero()
Expand Down Expand Up @@ -83,7 +83,7 @@ func (f *Field[T]) ModMulCanonical(a, b *Element[T], modulus *Element[T]) *Eleme
// would silently change the modulus whenever it does not fit T. If reducing it
// is intended, wrap it in [Field.ReduceStrict] explicitly.
func (f *Field[T]) ModAdd(a, b *Element[T], modulus *Element[T]) *Element[T] {
f.checkModulus(modulus)
modulus = f.checkModulus(modulus)
// inlined version of [Field.reduceAndOp] which uses variable-modulus reduction
var nextOverflow uint
var err error
Expand Down Expand Up @@ -145,7 +145,7 @@ func (f *Field[T]) modSub(a, b *Element[T], modulus *Element[T]) *Element[T] {
// would silently change the modulus whenever it does not fit T. If reducing it
// is intended, wrap it in [Field.ReduceStrict] explicitly.
func (f *Field[T]) ModAssertIsEqual(a, b *Element[T], modulus *Element[T]) {
f.checkModulus(modulus)
modulus = f.checkModulus(modulus)
// like fixed modulus AssertIsEqual, but uses current Sub implementation for
// computing the diff
diff := f.modSub(b, a, modulus)
Expand Down Expand Up @@ -176,7 +176,7 @@ func (f *Field[T]) ModAssertIsEqual(a, b *Element[T], modulus *Element[T]) {
// or decompose into bits. Intermediate multiplications use [Field.ModMul]
// without the canonicality assertion, keeping the per-call cost low.
func (f *Field[T]) ModExp(base, exp, modulus *Element[T]) *Element[T] {
f.checkModulus(modulus)
modulus = f.checkModulus(modulus)
// fast path when the base is zero then result is always zero
if len(base.Limbs) == 0 {
return f.Zero()
Expand Down Expand Up @@ -265,7 +265,7 @@ func (f *Field[T]) ModExp(base, exp, modulus *Element[T]) *Element[T] {
// not fit T. See the note on the exported variable-modulus methods.
//
// The method adds no constraints.
func (f *Field[T]) checkModulus(modulus *Element[T]) {
func (f *Field[T]) checkModulus(modulus *Element[T]) *Element[T] {
// populate the limbs in case the modulus was constructed in-circuit with
// [ValueOf]. No-op for a witness element, which is initialized at witness
// parsing time.
Expand All @@ -276,6 +276,18 @@ func (f *Field[T]) checkModulus(modulus *Element[T]) {
if value, isConstant := f.constantValue(modulus); isConstant && value.Cmp(f.fParams.Modulus()) >= 0 {
panic(fmt.Sprintf("variable modulus must be smaller than emulation modulus %s", f.fParams.Modulus()))
}
// the hints of the variable-modulus operations read NbLimbs modulus limbs
// from their inputs, but a constant created with [Field.NewElement] is
// stored on the minimal number of limbs. Pad it with zero limbs.
if nbLimbs := int(f.fParams.NbLimbs()); len(modulus.Limbs) < nbLimbs {
limbs := make([]frontend.Variable, nbLimbs)
copy(limbs, modulus.Limbs)
for i := len(modulus.Limbs); i < nbLimbs; i++ {
limbs[i] = 0
}
modulus = f.newInternalElement(limbs, 0)
}
return modulus
}

// assertLessThanModulus asserts that e < modulus as integers, where modulus is
Expand Down
56 changes: 56 additions & 0 deletions std/math/emulated/custommod_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@ import (

"github.com/consensys/gnark-crypto/ecc"
"github.com/consensys/gnark/frontend"
"github.com/consensys/gnark/frontend/cs/r1cs"
"github.com/consensys/gnark/frontend/cs/scs"
"github.com/consensys/gnark/std/math/emulated/emparams"
"github.com/consensys/gnark/test"
)
Expand Down Expand Up @@ -202,3 +204,57 @@ func TestVariableExpEdgeCases(t *testing.T) {
}, tc.name)
}
}

type constantModulusCircuit[T FieldParams] struct {
A, B Element[T]
Mul, Add, Exp Element[T]
constModulus, constExp *big.Int
}

func (c *constantModulusCircuit[T]) Define(api frontend.API) error {
f, err := NewField[T](api)
if err != nil {
return err
}
modulus := f.NewElement(c.constModulus)
f.AssertIsEqual(f.ModMulCanonical(&c.A, &c.B, modulus), &c.Mul)
f.ModAssertIsEqual(f.ModAdd(&c.A, &c.B, modulus), &c.Add, modulus)
f.AssertIsEqual(f.ModExp(&c.A, f.NewElement(c.constExp), modulus), &c.Exp)
return nil
}

// TestConstantModulus checks the variable-modulus methods with a constant
// modulus which fits in fewer limbs than the emulated parameters use.
func TestConstantModulus(t *testing.T) {
testConstantModulus[emparams.Secp256k1Fp](t)
testConstantModulus[emparams.Mod1e512](t)
}

func testConstantModulus[T FieldParams](t *testing.T) {
assert := test.NewAssert(t)
modulus := big.NewInt(1000003)
exp := big.NewInt(65537)
a, b := big.NewInt(123456), big.NewInt(654321)
mul := new(big.Int).Mod(new(big.Int).Mul(a, b), modulus)
add := new(big.Int).Mod(new(big.Int).Add(a, b), modulus)
expRes := new(big.Int).Exp(a, exp, modulus)
circuit := constantModulusCircuit[T]{constModulus: modulus, constExp: exp}
assignment := constantModulusCircuit[T]{
A: ValueOf[T](a),
B: ValueOf[T](b),
Mul: ValueOf[T](mul),
Add: ValueOf[T](add),
Exp: ValueOf[T](expRes),
}
assert.NoError(test.IsSolved(&circuit, &assignment, ecc.BN254.ScalarField()))
w, err := frontend.NewWitness(&assignment, ecc.BN254.ScalarField())
assert.NoError(err)
ccs, err := frontend.Compile(ecc.BN254.ScalarField(), r1cs.NewBuilder, &circuit)
assert.NoError(err)
_, err = ccs.Solve(w)
assert.NoError(err)
ccs, err = frontend.Compile(ecc.BN254.ScalarField(), scs.NewBuilder, &circuit)
assert.NoError(err)
_, err = ccs.Solve(w)
assert.NoError(err)
}