core: reuse the gas calculation bigints in bigmodexp precompile

This commit is contained in:
Guillaume Ballet 2019-01-17 16:38:30 +01:00
parent 81e26d5a48
commit 3c139b8a68

View file

@ -147,7 +147,11 @@ func (c *dataCopy) Run(in []byte) ([]byte, error) {
} }
// bigModExp implements a native big integer exponential modular operation. // bigModExp implements a native big integer exponential modular operation.
type bigModExp struct{} type bigModExp struct {
baseLen *big.Int
expLen *big.Int
modLen *big.Int
}
var ( var (
big1 = big.NewInt(1) big1 = big.NewInt(1)
@ -165,11 +169,10 @@ var (
// RequiredGas returns the gas required to execute the pre-compiled contract. // RequiredGas returns the gas required to execute the pre-compiled contract.
func (c *bigModExp) RequiredGas(input []byte) uint64 { func (c *bigModExp) RequiredGas(input []byte) uint64 {
var ( c.baseLen = new(big.Int).SetBytes(getData(input, 0, 32))
baseLen = new(big.Int).SetBytes(getData(input, 0, 32)) c.expLen = new(big.Int).SetBytes(getData(input, 32, 32))
expLen = new(big.Int).SetBytes(getData(input, 32, 32)) c.modLen = new(big.Int).SetBytes(getData(input, 64, 32))
modLen = new(big.Int).SetBytes(getData(input, 64, 32))
)
if len(input) > 96 { if len(input) > 96 {
input = input[96:] input = input[96:]
} else { } else {
@ -177,13 +180,13 @@ func (c *bigModExp) RequiredGas(input []byte) uint64 {
} }
// Retrieve the head 32 bytes of exp for the adjusted exponent length // Retrieve the head 32 bytes of exp for the adjusted exponent length
var expHead *big.Int var expHead *big.Int
if big.NewInt(int64(len(input))).Cmp(baseLen) <= 0 { if big.NewInt(int64(len(input))).Cmp(c.baseLen) <= 0 {
expHead = new(big.Int) expHead = new(big.Int)
} else { } else {
if expLen.Cmp(big32) > 0 { if c.expLen.Cmp(big32) > 0 {
expHead = new(big.Int).SetBytes(getData(input, baseLen.Uint64(), 32)) expHead = new(big.Int).SetBytes(getData(input, c.baseLen.Uint64(), 32))
} else { } else {
expHead = new(big.Int).SetBytes(getData(input, baseLen.Uint64(), expLen.Uint64())) expHead = new(big.Int).SetBytes(getData(input, c.baseLen.Uint64(), c.expLen.Uint64()))
} }
} }
// Calculate the adjusted exponent length // Calculate the adjusted exponent length
@ -192,14 +195,14 @@ func (c *bigModExp) RequiredGas(input []byte) uint64 {
msb = bitlen - 1 msb = bitlen - 1
} }
adjExpLen := new(big.Int) adjExpLen := new(big.Int)
if expLen.Cmp(big32) > 0 { if c.expLen.Cmp(big32) > 0 {
adjExpLen.Sub(expLen, big32) adjExpLen.Sub(c.expLen, big32)
adjExpLen.Mul(big8, adjExpLen) adjExpLen.Mul(big8, adjExpLen)
} }
adjExpLen.Add(adjExpLen, big.NewInt(int64(msb))) adjExpLen.Add(adjExpLen, big.NewInt(int64(msb)))
// Calculate the gas cost of the operation // Calculate the gas cost of the operation
gas := new(big.Int).Set(math.BigMax(modLen, baseLen)) gas := new(big.Int).Set(math.BigMax(c.modLen, c.baseLen))
switch { switch {
case gas.Cmp(big64) <= 0: case gas.Cmp(big64) <= 0:
gas.Mul(gas, gas) gas.Mul(gas, gas)
@ -224,31 +227,26 @@ func (c *bigModExp) RequiredGas(input []byte) uint64 {
} }
func (c *bigModExp) Run(input []byte) ([]byte, error) { func (c *bigModExp) Run(input []byte) ([]byte, error) {
var (
baseLen = new(big.Int).SetBytes(getData(input, 0, 32)).Uint64()
expLen = new(big.Int).SetBytes(getData(input, 32, 32)).Uint64()
modLen = new(big.Int).SetBytes(getData(input, 64, 32)).Uint64()
)
if len(input) > 96 { if len(input) > 96 {
input = input[96:] input = input[96:]
} else { } else {
input = input[:0] input = input[:0]
} }
// Handle a special case when both the base and mod length is zero // Handle a special case when both the base and mod length is zero
if baseLen == 0 && modLen == 0 { if c.baseLen.Uint64() == 0 && c.modLen.Uint64() == 0 {
return []byte{}, nil return []byte{}, nil
} }
// Retrieve the operands and execute the exponentiation // Retrieve the operands and execute the exponentiation
var ( var (
base = new(big.Int).SetBytes(getData(input, 0, baseLen)) base = new(big.Int).SetBytes(getData(input, 0, c.baseLen.Uint64()))
exp = new(big.Int).SetBytes(getData(input, baseLen, expLen)) exp = new(big.Int).SetBytes(getData(input, c.baseLen.Uint64(), c.expLen.Uint64()))
mod = new(big.Int).SetBytes(getData(input, baseLen+expLen, modLen)) mod = new(big.Int).SetBytes(getData(input, c.baseLen.Uint64()+c.expLen.Uint64(), c.modLen.Uint64()))
) )
if mod.BitLen() == 0 { if mod.BitLen() == 0 {
// Modulo 0 is undefined, return zero // Modulo 0 is undefined, return zero
return common.LeftPadBytes([]byte{}, int(modLen)), nil return common.LeftPadBytes([]byte{}, int(c.modLen.Int64())), nil
} }
return common.LeftPadBytes(base.Exp(base, exp, mod).Bytes(), int(modLen)), nil return common.LeftPadBytes(base.Exp(base, exp, mod).Bytes(), int(c.modLen.Int64())), nil
} }
// newCurvePoint unmarshals a binary blob into a bn256 elliptic curve point, // newCurvePoint unmarshals a binary blob into a bn256 elliptic curve point,