swap out big.Int for uint256

This commit is contained in:
Kevaundray Wedderburn 2025-07-10 22:50:45 +01:00
parent 4bc38e05e7
commit 6de5aa5091

View file

@ -37,6 +37,7 @@ import (
"github.com/ethereum/go-ethereum/crypto/kzg4844" "github.com/ethereum/go-ethereum/crypto/kzg4844"
"github.com/ethereum/go-ethereum/crypto/secp256r1" "github.com/ethereum/go-ethereum/crypto/secp256r1"
"github.com/ethereum/go-ethereum/params" "github.com/ethereum/go-ethereum/params"
"github.com/holiman/uint256"
"golang.org/x/crypto/ripemd160" "golang.org/x/crypto/ripemd160"
) )
@ -375,17 +376,15 @@ const (
) )
var ( var (
big1 = big.NewInt(1) // uint256 constants for modexp calculations
big3 = big.NewInt(3) u256_16 = uint256.NewInt(16)
big7 = big.NewInt(7) u256_480 = uint256.NewInt(480)
big20 = big.NewInt(20) u256_199680 = uint256.NewInt(199680)
big32 = big.NewInt(32)
big64 = big.NewInt(64) // divisors as uint256
big96 = big.NewInt(96) u256_byzantiumDivisor = uint256.NewInt(byzantiumDivisor)
big480 = big.NewInt(480) u256_berlinDivisor = uint256.NewInt(berlinDivisor)
big1024 = big.NewInt(1024) u256_osakaDivisor = uint256.NewInt(osakaDivisor)
big3072 = big.NewInt(3072)
big199680 = big.NewInt(199680)
) )
// byzantiumMultComplexity implements bigModexp multComplexity formula, as defined in EIP-198 // byzantiumMultComplexity implements bigModexp multComplexity formula, as defined in EIP-198
@ -396,33 +395,33 @@ var (
// else: return x ** 2 // 16 + 480 * x - 199680 // else: return x ** 2 // 16 + 480 * x - 199680
// //
// where is x is max(length_of_MODULUS, length_of_BASE) // where is x is max(length_of_MODULUS, length_of_BASE)
func byzantiumMultComplexity(x uint64) *big.Int { func byzantiumMultComplexity(x uint64) *uint256.Int {
switch { switch {
case x <= 64: case x <= 64:
return new(big.Int).SetUint64(x * x) return new(uint256.Int).SetUint64(x * x)
case x <= 1024: case x <= 1024:
// x^2 / 4 + 96*x - 3072 // x^2 / 4 + 96*x - 3072
result := x*x/4 + 96*x - 3072 result := x*x/4 + 96*x - 3072
return new(big.Int).SetUint64(result) return new(uint256.Int).SetUint64(result)
default: default:
// For large x, use big.Int arithmetic to avoid overflow // For large x, use uint256 arithmetic to avoid overflow
// x^2 / 16 + 480*x - 199680 // x^2 / 16 + 480*x - 199680
xBig := new(big.Int).SetUint64(x) xUint := new(uint256.Int).SetUint64(x)
// Calculate x^2 // Calculate x^2
xSquared := new(big.Int).Mul(xBig, xBig) xSquared := new(uint256.Int).Mul(xUint, xUint)
// Calculate x^2 / 16 (right shift by 4 bits) // Calculate x^2 / 16 (right shift by 4 bits)
xSquaredDiv16 := new(big.Int).Rsh(xSquared, 4) xSquaredDiv16 := new(uint256.Int).Rsh(xSquared, 4)
// Calculate 480 * x // Calculate 480 * x
x480 := new(big.Int).Mul(big480, xBig) x480 := new(uint256.Int).Mul(xUint, u256_480)
// Calculate 480 * x - 199680 // Calculate 480 * x - 199680
x480Minus199680 := new(big.Int).Sub(x480, big199680) x480Minus199680 := new(uint256.Int).Sub(x480, u256_199680)
// Add the two parts together // Add the two parts together
return new(big.Int).Add(xSquaredDiv16, x480Minus199680) return new(uint256.Int).Add(xSquaredDiv16, x480Minus199680)
} }
} }
@ -433,18 +432,18 @@ func byzantiumMultComplexity(x uint64) *big.Int {
// ceiling(x/8)^2 // ceiling(x/8)^2
// //
// where is x is max(length_of_MODULUS, length_of_BASE) // where is x is max(length_of_MODULUS, length_of_BASE)
func berlinMultComplexity(x uint64) *big.Int { func berlinMultComplexity(x uint64) *uint256.Int {
// TODO: The preceding line is too smart. // TODO: The preceding line is too smart.
// TODO: The issue is that (x+7) / 8 can overflow // TODO: The issue is that (x+7) / 8 can overflow
// TODO: if x > 2^64 - 7 // TODO: if x > 2^64 - 7
ceilDiv8 := (x >> 3) + ((x&7 + 7) >> 3) // safe ceil(x / 8) ceilDiv8 := (x >> 3) + ((x&7 + 7) >> 3) // safe ceil(x / 8)
z := new(big.Int).SetUint64(ceilDiv8) z := new(uint256.Int).SetUint64(ceilDiv8)
return new(big.Int).Mul(z, z) // square without overflow return new(uint256.Int).Mul(z, z) // square without overflow
} }
// Slow Bigint way (benchmark this) // Slow Bigint way (benchmark this)
// func berlinMultComplexity(xInt uint64) *big.Int { // func berlinMultComplexity(xInt uint64) *big.Int {
// x := new(big.Int).SetUint64(xInt) // x := new(big.Int).SetUint64(xInt)
// x = new(big.Int).Add(x, big7) // x + 7 // x = new(big.Int).Add(x, big.NewInt(7)) // x + 7
// x = new(big.Int).Rsh(x, 3) // (x + 7) / 8 // x = new(big.Int).Rsh(x, 3) // (x + 7) / 8
// return new(big.Int).Mul(x, x) // ((x + 7) / 8) ^ 2 // return new(big.Int).Mul(x, x) // ((x + 7) / 8) ^ 2
// } // }
@ -453,21 +452,21 @@ func berlinMultComplexity(x uint64) *big.Int {
// //
// For x <= 32: returns 16 // For x <= 32: returns 16
// For x > 32: returns 2 * ceiling(x/8)^2 // For x > 32: returns 2 * ceiling(x/8)^2
func osakaMultComplexity(x uint64) *big.Int { func osakaMultComplexity(x uint64) *uint256.Int {
if x <= 32 { if x <= 32 {
return big.NewInt(16) return u256_16
} }
// For x > 32, return 2 * berlinMultComplexity(x) // For x > 32, return 2 * berlinMultComplexity(x)
berlinComplexity := berlinMultComplexity(x) berlinComplexity := berlinMultComplexity(x)
return new(big.Int).Lsh(berlinComplexity, 1) // 2 * berlinComplexity return new(uint256.Int).Lsh(berlinComplexity, 1) // 2 * berlinComplexity
} }
// calculateIterationCount calculates the number of iterations for the modexp precompile. // calculateIterationCount calculates the number of iterations for the modexp precompile.
// This is the adjusted exponent length used in gas calculation. // This is the adjusted exponent length used in gas calculation.
func calculateIterationCount(expLen uint64, expHead *big.Int, multiplier uint64) uint64 { func calculateIterationCount(expLen uint64, expHead *uint256.Int, multiplier uint64) uint64 {
var iterationCount uint64 var iterationCount uint64
if expLen <= 32 && expHead.Sign() == 0 { if expLen <= 32 && expHead.IsZero() {
iterationCount = 0 iterationCount = 0
} else if expLen <= 32 { } else if expLen <= 32 {
// For small exponents, use MSB position - 1 // For small exponents, use MSB position - 1
@ -490,7 +489,7 @@ func calculateIterationCount(expLen uint64, expHead *big.Int, multiplier uint64)
} }
// byzantiumGasCalc calculates the gas cost for the modexp precompile using Byzantium rules. // byzantiumGasCalc calculates the gas cost for the modexp precompile using Byzantium rules.
func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 { func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 {
// Calculate max(baseLen, modLen) // Calculate max(baseLen, modLen)
maxLen := baseLen maxLen := baseLen
if modLen > maxLen { if modLen > maxLen {
@ -504,17 +503,18 @@ func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 {
iterationCount := calculateIterationCount(expLen, expHead, byzantiumMultiplier) iterationCount := calculateIterationCount(expLen, expHead, byzantiumMultiplier)
// Calculate gas: (multComplexity * iterationCount) / byzantiumDivisor // Calculate gas: (multComplexity * iterationCount) / byzantiumDivisor
gas := new(big.Int).Mul(multComplexity, new(big.Int).SetUint64(iterationCount)) iterCount := new(uint256.Int).SetUint64(iterationCount)
gas.Div(gas, big.NewInt(byzantiumDivisor)) gas := new(uint256.Int).Mul(multComplexity, iterCount)
gas.Div(gas, u256_byzantiumDivisor)
if gas.BitLen() > 64 { if !gas.IsUint64() {
return math.MaxUint64 return math.MaxUint64
} }
return gas.Uint64() return gas.Uint64()
} }
// berlinGasCalc calculates the gas cost for the modexp precompile using Berlin rules. // berlinGasCalc calculates the gas cost for the modexp precompile using Berlin rules.
func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 { func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 {
// Calculate max(baseLen, modLen) // Calculate max(baseLen, modLen)
maxLen := baseLen maxLen := baseLen
if modLen > maxLen { if modLen > maxLen {
@ -528,10 +528,11 @@ func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 {
iterationCount := calculateIterationCount(expLen, expHead, berlinMultiplier) iterationCount := calculateIterationCount(expLen, expHead, berlinMultiplier)
// Calculate gas: (multComplexity * iterationCount) / berlinDivisor // Calculate gas: (multComplexity * iterationCount) / berlinDivisor
gas := new(big.Int).Mul(multComplexity, new(big.Int).SetUint64(iterationCount)) iterCount := new(uint256.Int).SetUint64(iterationCount)
gas.Div(gas, big.NewInt(berlinDivisor)) gas := new(uint256.Int).Mul(multComplexity, iterCount)
gas.Div(gas, u256_berlinDivisor)
if gas.BitLen() > 64 { if !gas.IsUint64() {
return math.MaxUint64 return math.MaxUint64
} }
@ -544,7 +545,7 @@ func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 {
} }
// osakaGasCalc calculates the gas cost for the modexp precompile using Osaka rules. // osakaGasCalc calculates the gas cost for the modexp precompile using Osaka rules.
func osakaGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 { func osakaGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 {
// Calculate max(baseLen, modLen) // Calculate max(baseLen, modLen)
maxLen := baseLen maxLen := baseLen
if modLen > maxLen { if modLen > maxLen {
@ -558,10 +559,11 @@ func osakaGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 {
iterationCount := calculateIterationCount(expLen, expHead, osakaMultiplier) iterationCount := calculateIterationCount(expLen, expHead, osakaMultiplier)
// Calculate gas: (multComplexity * iterationCount) / osakaDivisor // Calculate gas: (multComplexity * iterationCount) / osakaDivisor
gas := new(big.Int).Mul(multComplexity, new(big.Int).SetUint64(iterationCount)) iterCount := new(uint256.Int).SetUint64(iterationCount)
gas.Div(gas, big.NewInt(osakaDivisor)) gas := new(uint256.Int).Mul(multComplexity, iterCount)
gas.Div(gas, u256_osakaDivisor)
if gas.BitLen() > 64 { if !gas.IsUint64() {
return math.MaxUint64 return math.MaxUint64
} }
@ -576,21 +578,21 @@ func osakaGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 {
// 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 {
// Parse input lengths // Parse input lengths
baseLenBig := new(big.Int).SetBytes(getData(input, 0, 32)) baseLenBig := new(uint256.Int).SetBytes(getData(input, 0, 32))
expLenBig := new(big.Int).SetBytes(getData(input, 32, 32)) expLenBig := new(uint256.Int).SetBytes(getData(input, 32, 32))
modLenBig := new(big.Int).SetBytes(getData(input, 64, 32)) modLenBig := new(uint256.Int).SetBytes(getData(input, 64, 32))
// Convert to uint64, capping at max value // Convert to uint64, capping at max value
baseLen := baseLenBig.Uint64() baseLen := baseLenBig.Uint64()
if baseLenBig.BitLen() > 64 { if !baseLenBig.IsUint64() {
baseLen = math.MaxUint64 baseLen = math.MaxUint64
} }
expLen := expLenBig.Uint64() expLen := expLenBig.Uint64()
if expLenBig.BitLen() > 64 { if !expLenBig.IsUint64() {
expLen = math.MaxUint64 expLen = math.MaxUint64
} }
modLen := modLenBig.Uint64() modLen := modLenBig.Uint64()
if modLenBig.BitLen() > 64 { if !modLenBig.IsUint64() {
modLen = math.MaxUint64 modLen = math.MaxUint64
} }
@ -602,14 +604,15 @@ 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 *uint256.Int
if uint64(len(input)) <= baseLen { if uint64(len(input)) <= baseLen {
expHead = new(big.Int) expHead = new(uint256.Int)
} else { } else {
if expLen > 32 { if expLen > 32 {
expHead = new(big.Int).SetBytes(getData(input, baseLen, 32)) expHead = new(uint256.Int).SetBytes(getData(input, baseLen, 32))
} else { } else {
expHead = new(big.Int).SetBytes(getData(input, baseLen, expLen)) // TODO: Check that if expLen < baseLen, then getData will return an empty slice
expHead = new(uint256.Int).SetBytes(getData(input, baseLen, expLen))
} }
} }