crypto/bls12381: use gnark instead of kilic for bls

This commit is contained in:
Marius van der Wijden 2024-02-06 21:37:08 +01:00
parent 00a3896dc5
commit b6c8158885

View file

@ -23,12 +23,15 @@ import (
"fmt" "fmt"
"math/big" "math/big"
"github.com/consensys/gnark-crypto/ecc"
bls12381 "github.com/consensys/gnark-crypto/ecc/bls12-381"
"github.com/consensys/gnark-crypto/ecc/bls12-381/fp"
"github.com/consensys/gnark-crypto/ecc/bls12-381/fr"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/common/math" "github.com/ethereum/go-ethereum/common/math"
"github.com/ethereum/go-ethereum/core/tracing" "github.com/ethereum/go-ethereum/core/tracing"
"github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/crypto"
"github.com/ethereum/go-ethereum/crypto/blake2b" "github.com/ethereum/go-ethereum/crypto/blake2b"
"github.com/ethereum/go-ethereum/crypto/bls12381"
"github.com/ethereum/go-ethereum/crypto/bn256" "github.com/ethereum/go-ethereum/crypto/bn256"
"github.com/ethereum/go-ethereum/crypto/kzg4844" "github.com/ethereum/go-ethereum/crypto/kzg4844"
"github.com/ethereum/go-ethereum/params" "github.com/ethereum/go-ethereum/params"
@ -673,26 +676,23 @@ func (c *bls12381G1Add) Run(input []byte) ([]byte, error) {
return nil, errBLS12381InvalidInputLength return nil, errBLS12381InvalidInputLength
} }
var err error var err error
var p0, p1 *bls12381.PointG1 var p0, p1 *bls12381.G1Affine
// Initialize G1
g := bls12381.NewG1()
// Decode G1 point p_0 // Decode G1 point p_0
if p0, err = decodePointG1(g, input[:128]); err != nil { if p0, err = decodePointG1(input[:128]); err != nil {
return nil, err return nil, err
} }
// Decode G1 point p_1 // Decode G1 point p_1
if p1, err = decodePointG1(g, input[128:]); err != nil { if p1, err = decodePointG1(input[128:]); err != nil {
return nil, err return nil, err
} }
// Compute r = p_0 + p_1 // Compute r = p_0 + p_1
r := g.New() r := new(bls12381.G1Affine)
g.Add(r, p0, p1) r.Add(p0, p1)
// Encode the G1 point result into 128 bytes // Encode the G1 point result into 128 bytes
return encodePointG1(g, r), nil return encodePointG1(r), nil
} }
// bls12381G1Mul implements EIP-2537 G1Mul precompile. // bls12381G1Mul implements EIP-2537 G1Mul precompile.
@ -711,24 +711,21 @@ func (c *bls12381G1Mul) Run(input []byte) ([]byte, error) {
return nil, errBLS12381InvalidInputLength return nil, errBLS12381InvalidInputLength
} }
var err error var err error
var p0 *bls12381.PointG1 var p0 *bls12381.G1Affine
// Initialize G1
g := bls12381.NewG1()
// Decode G1 point // Decode G1 point
if p0, err = decodePointG1(g, input[:128]); err != nil { if p0, err = decodePointG1(input[:128]); err != nil {
return nil, err return nil, err
} }
// Decode scalar value // Decode scalar value
e := new(big.Int).SetBytes(input[128:]) e := new(big.Int).SetBytes(input[128:])
// Compute r = e * p_0 // Compute r = e * p_0
r := g.New() r := new(bls12381.G1Affine)
g.MulScalarBig(r, p0, e) r.ScalarMultiplication(p0, e)
// Encode the G1 point into 128 bytes // Encode the G1 point into 128 bytes
return encodePointG1(g, r), nil return encodePointG1(r), nil
} }
// bls12381G1MultiExp implements EIP-2537 G1MultiExp precompile. // bls12381G1MultiExp implements EIP-2537 G1MultiExp precompile.
@ -761,31 +758,29 @@ func (c *bls12381G1MultiExp) Run(input []byte) ([]byte, error) {
if len(input) == 0 || len(input)%160 != 0 { if len(input) == 0 || len(input)%160 != 0 {
return nil, errBLS12381InvalidInputLength return nil, errBLS12381InvalidInputLength
} }
var err error points := make([]bls12381.G1Affine, k)
points := make([]*bls12381.PointG1, k) scalars := make([]fr.Element, k)
scalars := make([]*big.Int, k)
// Initialize G1
g := bls12381.NewG1()
// Decode point scalar pairs // Decode point scalar pairs
for i := 0; i < k; i++ { for i := 0; i < k; i++ {
off := 160 * i off := 160 * i
t0, t1, t2 := off, off+128, off+160 t0, t1, t2 := off, off+128, off+160
// Decode G1 point // Decode G1 point
if points[i], err = decodePointG1(g, input[t0:t1]); err != nil { p, err := decodePointG1(input[t0:t1])
if err != nil {
return nil, err return nil, err
} }
points[i] = *p
// Decode scalar value // Decode scalar value
scalars[i] = new(big.Int).SetBytes(input[t1:t2]) scalars[i] = *new(fr.Element).SetBytes(input[t1:t2])
} }
// Compute r = e_0 * p_0 + e_1 * p_1 + ... + e_(k-1) * p_(k-1) // Compute r = e_0 * p_0 + e_1 * p_1 + ... + e_(k-1) * p_(k-1)
r := g.New() r := new(bls12381.G1Affine)
g.MultiExpBig(r, points, scalars) r.MultiExp(points, scalars, ecc.MultiExpConfig{})
// Encode the G1 point to 128 bytes // Encode the G1 point to 128 bytes
return encodePointG1(g, r), nil return encodePointG1(r), nil
} }
// bls12381G2Add implements EIP-2537 G2Add precompile. // bls12381G2Add implements EIP-2537 G2Add precompile.
@ -804,26 +799,23 @@ func (c *bls12381G2Add) Run(input []byte) ([]byte, error) {
return nil, errBLS12381InvalidInputLength return nil, errBLS12381InvalidInputLength
} }
var err error var err error
var p0, p1 *bls12381.PointG2 var p0, p1 *bls12381.G2Affine
// Initialize G2
g := bls12381.NewG2()
r := g.New()
// Decode G2 point p_0 // Decode G2 point p_0
if p0, err = decodePointG2(g, input[:256]); err != nil { if p0, err = decodePointG2(input[:256]); err != nil {
return nil, err return nil, err
} }
// Decode G2 point p_1 // Decode G2 point p_1
if p1, err = decodePointG2(g, input[256:]); err != nil { if p1, err = decodePointG2(input[256:]); err != nil {
return nil, err return nil, err
} }
// Compute r = p_0 + p_1 // Compute r = p_0 + p_1
g.Add(r, p0, p1) r := new(bls12381.G2Affine)
r.Add(p0, p1)
// Encode the G2 point into 256 bytes // Encode the G2 point into 256 bytes
return encodePointG2(g, r), nil return encodePointG2(r), nil
} }
// bls12381G2Mul implements EIP-2537 G2Mul precompile. // bls12381G2Mul implements EIP-2537 G2Mul precompile.
@ -842,24 +834,21 @@ func (c *bls12381G2Mul) Run(input []byte) ([]byte, error) {
return nil, errBLS12381InvalidInputLength return nil, errBLS12381InvalidInputLength
} }
var err error var err error
var p0 *bls12381.PointG2 var p0 *bls12381.G2Affine
// Initialize G2
g := bls12381.NewG2()
// Decode G2 point // Decode G2 point
if p0, err = decodePointG2(g, input[:256]); err != nil { if p0, err = decodePointG2(input[:256]); err != nil {
return nil, err return nil, err
} }
// Decode scalar value // Decode scalar value
e := new(big.Int).SetBytes(input[256:]) e := new(big.Int).SetBytes(input[256:])
// Compute r = e * p_0 // Compute r = e * p_0
r := g.New() r := new(bls12381.G2Affine)
g.MulScalarBig(r, p0, e) r.ScalarMultiplication(p0, e)
// Encode the G2 point into 256 bytes // Encode the G2 point into 256 bytes
return encodePointG2(g, r), nil return encodePointG2(r), nil
} }
// bls12381G2MultiExp implements EIP-2537 G2MultiExp precompile. // bls12381G2MultiExp implements EIP-2537 G2MultiExp precompile.
@ -892,31 +881,29 @@ func (c *bls12381G2MultiExp) Run(input []byte) ([]byte, error) {
if len(input) == 0 || len(input)%288 != 0 { if len(input) == 0 || len(input)%288 != 0 {
return nil, errBLS12381InvalidInputLength return nil, errBLS12381InvalidInputLength
} }
var err error points := make([]bls12381.G2Affine, k)
points := make([]*bls12381.PointG2, k) scalars := make([]fr.Element, k)
scalars := make([]*big.Int, k)
// Initialize G2
g := bls12381.NewG2()
// Decode point scalar pairs // Decode point scalar pairs
for i := 0; i < k; i++ { for i := 0; i < k; i++ {
off := 288 * i off := 288 * i
t0, t1, t2 := off, off+256, off+288 t0, t1, t2 := off, off+256, off+288
// Decode G2 point // Decode G2 point
if points[i], err = decodePointG2(g, input[t0:t1]); err != nil { p, err := decodePointG2(input[t0:t1])
if err != nil {
return nil, err return nil, err
} }
points[i] = *p
// Decode scalar value // Decode scalar value
scalars[i] = new(big.Int).SetBytes(input[t1:t2]) scalars[i] = *new(fr.Element).SetBytes(input[t1:t2])
} }
// Compute r = e_0 * p_0 + e_1 * p_1 + ... + e_(k-1) * p_(k-1) // Compute r = e_0 * p_0 + e_1 * p_1 + ... + e_(k-1) * p_(k-1)
r := g.New() r := new(bls12381.G2Affine)
g.MultiExpBig(r, points, scalars) r.MultiExp(points, scalars, ecc.MultiExpConfig{})
// Encode the G2 point to 256 bytes. // Encode the G2 point to 256 bytes.
return encodePointG2(g, r), nil return encodePointG2(r), nil
} }
// bls12381Pairing implements EIP-2537 Pairing precompile. // bls12381Pairing implements EIP-2537 Pairing precompile.
@ -939,9 +926,10 @@ func (c *bls12381Pairing) Run(input []byte) ([]byte, error) {
return nil, errBLS12381InvalidInputLength return nil, errBLS12381InvalidInputLength
} }
// Initialize BLS12-381 pairing engine var (
e := bls12381.NewEngine() p []bls12381.G1Affine
g1, g2 := e.G1, e.G2 q []bls12381.G2Affine
)
// Decode pairs // Decode pairs
for i := 0; i < k; i++ { for i := 0; i < k; i++ {
@ -949,39 +937,39 @@ func (c *bls12381Pairing) Run(input []byte) ([]byte, error) {
t0, t1, t2 := off, off+128, off+384 t0, t1, t2 := off, off+128, off+384
// Decode G1 point // Decode G1 point
p1, err := decodePointG1(g1, input[t0:t1]) p1, err := decodePointG1(input[t0:t1])
if err != nil { if err != nil {
return nil, err return nil, err
} }
// Decode G2 point // Decode G2 point
p2, err := decodePointG2(g2, input[t1:t2]) p2, err := decodePointG2(input[t1:t2])
if err != nil { if err != nil {
return nil, err return nil, err
} }
// 'point is on curve' check already done, // 'point is on curve' check already done,
// Here we need to apply subgroup checks. // Here we need to apply subgroup checks.
if !g1.InCorrectSubgroup(p1) { if !p1.IsInSubGroup() {
return nil, errBLS12381G1PointSubgroup return nil, errBLS12381G1PointSubgroup
} }
if !g2.InCorrectSubgroup(p2) { if !p2.IsInSubGroup() {
return nil, errBLS12381G2PointSubgroup return nil, errBLS12381G2PointSubgroup
} }
p = append(p, *p1)
// Update pairing engine with G1 and G2 points q = append(q, *p2)
e.AddPair(p1, p2)
} }
// Prepare 32 byte output // Prepare 32 byte output
out := make([]byte, 32) out := make([]byte, 32)
// Compute pairing and set the result // Compute pairing and set the result
if e.Check() { ok, err := bls12381.PairingCheck(p, q)
if err == nil && ok {
out[31] = 1 out[31] = 1
} }
return out, nil return out, nil
} }
func decodePointG1(g *bls12381.G1, in []byte) (*bls12381.PointG1, error) { func decodePointG1(in []byte) (*bls12381.G1Affine, error) {
if len(in) != 128 { if len(in) != 128 {
return nil, errors.New("invalid g1 point length") return nil, errors.New("invalid g1 point length")
} }
@ -998,11 +986,14 @@ func decodePointG1(g *bls12381.G1, in []byte) (*bls12381.PointG1, error) {
} }
copy(pointBytes[:48], xBytes) copy(pointBytes[:48], xBytes)
copy(pointBytes[48:], yBytes) copy(pointBytes[48:], yBytes)
return g.FromBytes(pointBytes)
p := new(bls12381.G1Affine)
p.SetBytes(pointBytes)
return p, nil
} }
// decodePointG2 given encoded (x, y) coordinates in 256 bytes returns a valid G2 Point. // decodePointG2 given encoded (x, y) coordinates in 256 bytes returns a valid G2 Point.
func decodePointG2(g *bls12381.G2, in []byte) (*bls12381.PointG2, error) { func decodePointG2(in []byte) (*bls12381.G2Affine, error) {
if len(in) != 256 { if len(in) != 256 {
return nil, errors.New("invalid g2 point length") return nil, errors.New("invalid g2 point length")
} }
@ -1027,7 +1018,9 @@ func decodePointG2(g *bls12381.G2, in []byte) (*bls12381.PointG2, error) {
copy(pointBytes[48:96], x0Bytes) copy(pointBytes[48:96], x0Bytes)
copy(pointBytes[96:144], y1Bytes) copy(pointBytes[96:144], y1Bytes)
copy(pointBytes[144:192], y0Bytes) copy(pointBytes[144:192], y0Bytes)
return g.FromBytes(pointBytes) p := new(bls12381.G2Affine)
p.SetBytes(pointBytes)
return p, nil
} }
// decodeBLS12381FieldElement decodes BLS12-381 elliptic curve field element. // decodeBLS12381FieldElement decodes BLS12-381 elliptic curve field element.
@ -1048,8 +1041,12 @@ func decodeBLS12381FieldElement(in []byte) ([]byte, error) {
} }
// encodePointG1 encodes a point into 128 bytes. // encodePointG1 encodes a point into 128 bytes.
func encodePointG1(g *bls12381.G1, p *bls12381.PointG1) []byte { func encodePointG1(p *bls12381.G1Affine) []byte {
outRaw := g.ToBytes(p) outRaw := p.RawBytes()
if p.IsInfinity() {
// gnark sets a weird byte if the point is at infinity, remove it
outRaw[0] = 0
}
out := make([]byte, 128) out := make([]byte, 128)
// encode x // encode x
copy(out[16:], outRaw[:48]) copy(out[16:], outRaw[:48])
@ -1059,9 +1056,13 @@ func encodePointG1(g *bls12381.G1, p *bls12381.PointG1) []byte {
} }
// encodePointG2 encodes a point into 256 bytes. // encodePointG2 encodes a point into 256 bytes.
func encodePointG2(g *bls12381.G2, p *bls12381.PointG2) []byte { func encodePointG2(p *bls12381.G2Affine) []byte {
// outRaw is 96 bytes // outRaw is 192 bytes
outRaw := g.ToBytes(p) outRaw := p.RawBytes()
if p.IsInfinity() {
// gnark sets a weird byte if the point is at infinity, remove it
outRaw[0] = 0
}
out := make([]byte, 256) out := make([]byte, 256)
// encode x // encode x
copy(out[16:16+48], outRaw[48:96]) copy(out[16:16+48], outRaw[48:96])
@ -1093,18 +1094,16 @@ func (c *bls12381MapG1) Run(input []byte) ([]byte, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
elem := new(fp.Element).SetBytes(fe)
// Initialize G1
g := bls12381.NewG1()
// Compute mapping // Compute mapping
r, err := g.MapToCurve(fe) r := bls12381.MapToG1(*elem)
if err != nil { if err != nil {
return nil, err return nil, err
} }
// Encode the G1 point to 128 bytes // Encode the G1 point to 128 bytes
return encodePointG1(g, r), nil return encodePointG1(&r), nil
} }
// bls12381MapG2 implements EIP-2537 MapG2 precompile. // bls12381MapG2 implements EIP-2537 MapG2 precompile.
@ -1136,17 +1135,21 @@ func (c *bls12381MapG2) Run(input []byte) ([]byte, error) {
} }
copy(fe[:48], c1) copy(fe[:48], c1)
// Initialize G2 // TODO rework once https://github.com/Consensys/gnark-crypto/issues/483 is solved
g := bls12381.NewG2()
s0 := new(big.Int).SetBytes(c0)
s1 := new(big.Int).SetBytes(c1)
elem := new(bls12381.E2).SetString(s0.String(), s1.String())
// Compute mapping // Compute mapping
r, err := g.MapToCurve(fe) r := bls12381.MapToG2(*elem)
if err != nil { if err != nil {
return nil, err return nil, err
} }
// Encode the G2 point to 256 bytes // Encode the G2 point to 256 bytes
return encodePointG2(g, r), nil return encodePointG2(&r), nil
} }
// kzgPointEvaluation implements the EIP-4844 point evaluation precompile. // kzgPointEvaluation implements the EIP-4844 point evaluation precompile.