mirror of
https://github.com/ethereum/go-ethereum.git
synced 2026-08-18 09:53:48 +00:00
crypto/bn256/cloudflare: enable curve mul lattice optimization
This commit is contained in:
parent
f0c7ee3eb9
commit
41d1314ff3
4 changed files with 171 additions and 6 deletions
|
|
@ -14,6 +14,8 @@
|
||||||
// You should have received a copy of the GNU Lesser General Public License
|
// You should have received a copy of the GNU Lesser General Public License
|
||||||
// along with the go-ethereum library. If not, see <http://www.gnu.org/licenses/>.
|
// along with the go-ethereum library. If not, see <http://www.gnu.org/licenses/>.
|
||||||
|
|
||||||
|
// +build gofuzz
|
||||||
|
|
||||||
package bn256
|
package bn256
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
|
@ -39,6 +41,8 @@ func FuzzAdd(data []byte) int {
|
||||||
|
|
||||||
if (errc == nil) != (errg == nil) {
|
if (errc == nil) != (errg == nil) {
|
||||||
panic("parse mismatch")
|
panic("parse mismatch")
|
||||||
|
} else if errc != nil {
|
||||||
|
return 0
|
||||||
}
|
}
|
||||||
// Ensure both libs can parse the second curve point
|
// Ensure both libs can parse the second curve point
|
||||||
yc := new(cloudflare.G1)
|
yc := new(cloudflare.G1)
|
||||||
|
|
@ -49,6 +53,8 @@ func FuzzAdd(data []byte) int {
|
||||||
|
|
||||||
if (errc == nil) != (errg == nil) {
|
if (errc == nil) != (errg == nil) {
|
||||||
panic("parse mismatch")
|
panic("parse mismatch")
|
||||||
|
} else if errc != nil {
|
||||||
|
return 0
|
||||||
}
|
}
|
||||||
// Add the two points and ensure they result in the same output
|
// Add the two points and ensure they result in the same output
|
||||||
rc := new(cloudflare.G1)
|
rc := new(cloudflare.G1)
|
||||||
|
|
@ -79,6 +85,8 @@ func FuzzMul(data []byte) int {
|
||||||
|
|
||||||
if (errc == nil) != (errg == nil) {
|
if (errc == nil) != (errg == nil) {
|
||||||
panic("parse mismatch")
|
panic("parse mismatch")
|
||||||
|
} else if errc != nil {
|
||||||
|
return 0
|
||||||
}
|
}
|
||||||
// Add the two points and ensure they result in the same output
|
// Add the two points and ensure they result in the same output
|
||||||
rc := new(cloudflare.G1)
|
rc := new(cloudflare.G1)
|
||||||
|
|
@ -107,6 +115,8 @@ func FuzzPair(data []byte) int {
|
||||||
|
|
||||||
if (errc == nil) != (errg == nil) {
|
if (errc == nil) != (errg == nil) {
|
||||||
panic("parse mismatch")
|
panic("parse mismatch")
|
||||||
|
} else if errc != nil {
|
||||||
|
return 0
|
||||||
}
|
}
|
||||||
// Ensure both libs can parse the twist point
|
// Ensure both libs can parse the twist point
|
||||||
tc := new(cloudflare.G2)
|
tc := new(cloudflare.G2)
|
||||||
|
|
@ -117,6 +127,8 @@ func FuzzPair(data []byte) int {
|
||||||
|
|
||||||
if (errc == nil) != (errg == nil) {
|
if (errc == nil) != (errg == nil) {
|
||||||
panic("parse mismatch")
|
panic("parse mismatch")
|
||||||
|
} else if errc != nil {
|
||||||
|
return 0
|
||||||
}
|
}
|
||||||
// Pair the two points and ensure thet result in the same output
|
// Pair the two points and ensure thet result in the same output
|
||||||
if cloudflare.PairingCheck([]*cloudflare.G1{pc}, []*cloudflare.G2{tc}) != google.PairingCheck([]*google.G1{pg}, []*google.G2{tg}) {
|
if cloudflare.PairingCheck([]*cloudflare.G1{pc}, []*cloudflare.G2{tc}) != google.PairingCheck([]*google.G1{pg}, []*google.G2{tg}) {
|
||||||
|
|
|
||||||
|
|
@ -183,15 +183,24 @@ func (c *curvePoint) Double(a *curvePoint) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *curvePoint) Mul(a *curvePoint, scalar *big.Int) {
|
func (c *curvePoint) Mul(a *curvePoint, scalar *big.Int) {
|
||||||
sum, t := &curvePoint{}, &curvePoint{}
|
precomp := [1 << 2]*curvePoint{nil, {}, {}, {}}
|
||||||
sum.SetInfinity()
|
precomp[1].Set(a)
|
||||||
|
precomp[2].Set(a)
|
||||||
|
gfpMul(&precomp[2].x, &precomp[2].x, xiTo2PSquaredMinus2Over3)
|
||||||
|
precomp[3].Add(precomp[1], precomp[2])
|
||||||
|
|
||||||
for i := scalar.BitLen(); i >= 0; i-- {
|
multiScalar := curveLattice.Multi(scalar)
|
||||||
|
|
||||||
|
sum := &curvePoint{}
|
||||||
|
sum.SetInfinity()
|
||||||
|
t := &curvePoint{}
|
||||||
|
|
||||||
|
for i := len(multiScalar) - 1; i >= 0; i-- {
|
||||||
t.Double(sum)
|
t.Double(sum)
|
||||||
if scalar.Bit(i) != 0 {
|
if multiScalar[i] == 0 {
|
||||||
sum.Add(t, a)
|
|
||||||
} else {
|
|
||||||
sum.Set(t)
|
sum.Set(t)
|
||||||
|
} else {
|
||||||
|
sum.Add(t, precomp[multiScalar[i]])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
c.Set(sum)
|
c.Set(sum)
|
||||||
|
|
|
||||||
115
crypto/bn256/cloudflare/lattice.go
Normal file
115
crypto/bn256/cloudflare/lattice.go
Normal file
|
|
@ -0,0 +1,115 @@
|
||||||
|
package bn256
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math/big"
|
||||||
|
)
|
||||||
|
|
||||||
|
var half = new(big.Int).Rsh(Order, 1)
|
||||||
|
|
||||||
|
var curveLattice = &lattice{
|
||||||
|
vectors: [][]*big.Int{
|
||||||
|
{bigFromBase10("147946756881789319000765030803803410728"), bigFromBase10("147946756881789319010696353538189108491")},
|
||||||
|
{bigFromBase10("147946756881789319020627676272574806254"), bigFromBase10("-147946756881789318990833708069417712965")},
|
||||||
|
},
|
||||||
|
inverse: []*big.Int{
|
||||||
|
bigFromBase10("147946756881789318990833708069417712965"),
|
||||||
|
bigFromBase10("147946756881789319010696353538189108491"),
|
||||||
|
},
|
||||||
|
det: bigFromBase10("43776485743678550444492811490514550177096728800832068687396408373151616991234"),
|
||||||
|
}
|
||||||
|
|
||||||
|
var targetLattice = &lattice{
|
||||||
|
vectors: [][]*big.Int{
|
||||||
|
{bigFromBase10("9931322734385697761"), bigFromBase10("9931322734385697761"), bigFromBase10("9931322734385697763"), bigFromBase10("9931322734385697764")},
|
||||||
|
{bigFromBase10("4965661367192848881"), bigFromBase10("4965661367192848881"), bigFromBase10("4965661367192848882"), bigFromBase10("-9931322734385697762")},
|
||||||
|
{bigFromBase10("-9931322734385697762"), bigFromBase10("-4965661367192848881"), bigFromBase10("4965661367192848881"), bigFromBase10("-4965661367192848882")},
|
||||||
|
{bigFromBase10("9931322734385697763"), bigFromBase10("-4965661367192848881"), bigFromBase10("-4965661367192848881"), bigFromBase10("-4965661367192848881")},
|
||||||
|
},
|
||||||
|
inverse: []*big.Int{
|
||||||
|
bigFromBase10("734653495049373973658254490726798021314063399421879442165"),
|
||||||
|
bigFromBase10("147946756881789319000765030803803410728"),
|
||||||
|
bigFromBase10("-147946756881789319005730692170996259609"),
|
||||||
|
bigFromBase10("1469306990098747947464455738335385361643788813749140841702"),
|
||||||
|
},
|
||||||
|
det: new(big.Int).Set(Order),
|
||||||
|
}
|
||||||
|
|
||||||
|
type lattice struct {
|
||||||
|
vectors [][]*big.Int
|
||||||
|
inverse []*big.Int
|
||||||
|
det *big.Int
|
||||||
|
}
|
||||||
|
|
||||||
|
// decompose takes a scalar mod Order as input and finds a short, positive decomposition of it wrt to the lattice basis.
|
||||||
|
func (l *lattice) decompose(k *big.Int) []*big.Int {
|
||||||
|
n := len(l.inverse)
|
||||||
|
|
||||||
|
// Calculate closest vector in lattice to <k,0,0,...> with Babai's rounding.
|
||||||
|
c := make([]*big.Int, n)
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
c[i] = new(big.Int).Mul(k, l.inverse[i])
|
||||||
|
round(c[i], l.det)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Transform vectors according to c and subtract <k,0,0,...>.
|
||||||
|
out := make([]*big.Int, n)
|
||||||
|
temp := new(big.Int)
|
||||||
|
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
out[i] = new(big.Int)
|
||||||
|
|
||||||
|
for j := 0; j < n; j++ {
|
||||||
|
temp.Mul(c[j], l.vectors[j][i])
|
||||||
|
out[i].Add(out[i], temp)
|
||||||
|
}
|
||||||
|
|
||||||
|
out[i].Neg(out[i])
|
||||||
|
out[i].Add(out[i], l.vectors[0][i]).Add(out[i], l.vectors[0][i])
|
||||||
|
}
|
||||||
|
out[0].Add(out[0], k)
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *lattice) Precompute(add func(i, j uint)) {
|
||||||
|
n := uint(len(l.vectors))
|
||||||
|
total := uint(1) << n
|
||||||
|
|
||||||
|
for i := uint(0); i < n; i++ {
|
||||||
|
for j := uint(0); j < total; j++ {
|
||||||
|
if (j>>i)&1 == 1 {
|
||||||
|
add(i, j)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *lattice) Multi(scalar *big.Int) []uint8 {
|
||||||
|
decomp := l.decompose(scalar)
|
||||||
|
|
||||||
|
maxLen := 0
|
||||||
|
for _, x := range decomp {
|
||||||
|
if x.BitLen() > maxLen {
|
||||||
|
maxLen = x.BitLen()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
out := make([]uint8, maxLen)
|
||||||
|
for j, x := range decomp {
|
||||||
|
for i := 0; i < maxLen; i++ {
|
||||||
|
out[i] += uint8(x.Bit(i)) << uint(j)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// round sets num to num/denom rounded to the nearest integer.
|
||||||
|
func round(num, denom *big.Int) {
|
||||||
|
r := new(big.Int)
|
||||||
|
num.DivMod(num, denom, r)
|
||||||
|
|
||||||
|
if r.Cmp(half) == 1 {
|
||||||
|
num.Add(num, big.NewInt(1))
|
||||||
|
}
|
||||||
|
}
|
||||||
29
crypto/bn256/cloudflare/lattice_test.go
Normal file
29
crypto/bn256/cloudflare/lattice_test.go
Normal file
|
|
@ -0,0 +1,29 @@
|
||||||
|
package bn256
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLatticeReduceCurve(t *testing.T) {
|
||||||
|
k, _ := rand.Int(rand.Reader, Order)
|
||||||
|
ks := curveLattice.decompose(k)
|
||||||
|
|
||||||
|
if ks[0].BitLen() > 130 || ks[1].BitLen() > 130 {
|
||||||
|
t.Fatal("reduction too large")
|
||||||
|
} else if ks[0].Sign() < 0 || ks[1].Sign() < 0 {
|
||||||
|
t.Fatal("reduction must be positive")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLatticeReduceTarget(t *testing.T) {
|
||||||
|
k, _ := rand.Int(rand.Reader, Order)
|
||||||
|
ks := targetLattice.decompose(k)
|
||||||
|
|
||||||
|
if ks[0].BitLen() > 66 || ks[1].BitLen() > 66 || ks[2].BitLen() > 66 || ks[3].BitLen() > 66 {
|
||||||
|
t.Fatal("reduction too large")
|
||||||
|
} else if ks[0].Sign() < 0 || ks[1].Sign() < 0 || ks[2].Sign() < 0 || ks[3].Sign() < 0 {
|
||||||
|
t.Fatal("reduction must be positive")
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Reference in a new issue