From 5eb27425ea17d1466148b42b44106b0869f5ee67 Mon Sep 17 00:00:00 2001 From: Kevaundray Wedderburn Date: Fri, 11 Jul 2025 11:12:05 +0100 Subject: [PATCH] cleanup Co-authored-by: Felix Lange --- core/vm/contracts.go | 92 ++++++++++++++++++++++++-------------------- 1 file changed, 51 insertions(+), 41 deletions(-) diff --git a/core/vm/contracts.go b/core/vm/contracts.go index a817cfd3bd..19abca8a63 100644 --- a/core/vm/contracts.go +++ b/core/vm/contracts.go @@ -395,33 +395,40 @@ var ( // else: return x ** 2 // 16 + 480 * x - 199680 // // where is x is max(length_of_MODULUS, length_of_BASE) -func byzantiumMultComplexity(x uint64) *uint256.Int { +func byzantiumMultComplexity(x uint64) uint256.Int { switch { case x <= 64: - return new(uint256.Int).SetUint64(x * x) + var z uint256.Int + z.SetUint64(x * x) + return z case x <= 1024: // x^2 / 4 + 96*x - 3072 result := x*x/4 + 96*x - 3072 - return new(uint256.Int).SetUint64(result) + var z uint256.Int + z.SetUint64(result) + return z default: // For large x, use uint256 arithmetic to avoid overflow // x^2 / 16 + 480*x - 199680 - xUint := new(uint256.Int).SetUint64(x) + var xUint, xSquared, result uint256.Int + xUint.SetUint64(x) // Calculate x^2 - xSquared := new(uint256.Int).Mul(xUint, xUint) + xSquared.Mul(&xUint, &xUint) // Calculate x^2 / 16 (right shift by 4 bits) - xSquaredDiv16 := new(uint256.Int).Rsh(xSquared, 4) + result.Rsh(&xSquared, 4) // Calculate 480 * x - x480 := new(uint256.Int).Mul(xUint, u256_480) + var x480 uint256.Int + x480.Mul(&xUint, u256_480) // Calculate 480 * x - 199680 - x480Minus199680 := new(uint256.Int).Sub(x480, u256_199680) + x480.Sub(&x480, u256_199680) // Add the two parts together - return new(uint256.Int).Add(xSquaredDiv16, x480Minus199680) + result.Add(&result, &x480) + return result } } @@ -432,13 +439,17 @@ func byzantiumMultComplexity(x uint64) *uint256.Int { // ceiling(x/8)^2 // // where is x is max(length_of_MODULUS, length_of_BASE) -func berlinMultComplexity(x uint64) *uint256.Int { +func berlinMultComplexity(x uint64) uint256.Int { // TODO: The preceding line is too smart. // TODO: The issue is that (x+7) / 8 can overflow // TODO: if x > 2^64 - 7 + // Safe ceiling division by 8 to avoid overflow ceilDiv8 := (x >> 3) + ((x&7 + 7) >> 3) // safe ceil(x / 8) - z := new(uint256.Int).SetUint64(ceilDiv8) - return new(uint256.Int).Mul(z, z) // square without overflow + + var z uint256.Int + z.SetUint64(ceilDiv8) + z.Mul(&z, &z) + return z } // Slow Bigint way (benchmark this) @@ -453,18 +464,19 @@ func berlinMultComplexity(x uint64) *uint256.Int { // // For x <= 32: returns 16 // For x > 32: returns 2 * ceiling(x/8)^2 -func osakaMultComplexity(x uint64) *uint256.Int { +func osakaMultComplexity(x uint64) uint256.Int { if x <= 32 { - return u256_16 + return *u256_16 } // For x > 32, return 2 * berlinMultComplexity(x) - berlinComplexity := berlinMultComplexity(x) - return new(uint256.Int).Lsh(berlinComplexity, 1) // 2 * berlinComplexity + result := berlinMultComplexity(x) + result.Lsh(&result, 1) // 2 * berlinComplexity + return result } // calculateIterationCount calculates the number of iterations for the modexp precompile. // This is the adjusted exponent length used in gas calculation. -func calculateIterationCount(expLen uint64, expHead *uint256.Int, multiplier uint64) uint64 { +func calculateIterationCount(expLen uint64, expHead uint256.Int, multiplier uint64) uint64 { var iterationCount uint64 // For large exponents (expLen > 32), add (expLen - 32) * multiplier @@ -482,8 +494,7 @@ func calculateIterationCount(expLen uint64, expHead *uint256.Int, multiplier uin } // byzantiumGasCalc calculates the gas cost for the modexp precompile using Byzantium rules. -func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 { - // Calculate max(baseLen, modLen) +func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead uint256.Int) uint64 { maxLen := max(baseLen, modLen) // Calculate multiplication complexity @@ -493,9 +504,10 @@ func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint iterationCount := calculateIterationCount(expLen, expHead, byzantiumMultiplier) // Calculate gas: (multComplexity * iterationCount) / byzantiumDivisor - iterCount := new(uint256.Int).SetUint64(iterationCount) - gas := new(uint256.Int).Mul(multComplexity, iterCount) - gas.Div(gas, u256_byzantiumDivisor) + var gas uint256.Int + gas.SetUint64(iterationCount) + gas.Mul(&multComplexity, &gas) + gas.Div(&gas, u256_byzantiumDivisor) if !gas.IsUint64() { return math.MaxUint64 @@ -504,8 +516,7 @@ func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint } // berlinGasCalc calculates the gas cost for the modexp precompile using Berlin rules. -func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 { - // Calculate max(baseLen, modLen) +func berlinGasCalc(baseLen, expLen, modLen uint64, expHead uint256.Int) uint64 { maxLen := max(baseLen, modLen) // Calculate multiplication complexity @@ -515,9 +526,10 @@ func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 iterationCount := calculateIterationCount(expLen, expHead, berlinMultiplier) // Calculate gas: (multComplexity * iterationCount) / berlinDivisor - iterCount := new(uint256.Int).SetUint64(iterationCount) - gas := new(uint256.Int).Mul(multComplexity, iterCount) - gas.Div(gas, u256_berlinDivisor) + var gas uint256.Int + gas.SetUint64(iterationCount) + gas.Mul(&multComplexity, &gas) + gas.Div(&gas, u256_berlinDivisor) if !gas.IsUint64() { return math.MaxUint64 @@ -528,8 +540,7 @@ func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 } // osakaGasCalc calculates the gas cost for the modexp precompile using Osaka rules. -func osakaGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 { - // Calculate max(baseLen, modLen) +func osakaGasCalc(baseLen, expLen, modLen uint64, expHead uint256.Int) uint64 { maxLen := max(baseLen, modLen) // Calculate multiplication complexity @@ -539,9 +550,10 @@ func osakaGasCalc(baseLen, expLen, modLen uint64, expHead *uint256.Int) uint64 { iterationCount := calculateIterationCount(expLen, expHead, osakaMultiplier) // Calculate gas: (multComplexity * iterationCount) / osakaDivisor - iterCount := new(uint256.Int).SetUint64(iterationCount) - gas := new(uint256.Int).Mul(multComplexity, iterCount) - gas.Div(gas, u256_osakaDivisor) + var gas uint256.Int + gas.SetUint64(iterationCount) + gas.Mul(&multComplexity, &gas) + gas.Div(&gas, u256_osakaDivisor) if !gas.IsUint64() { return math.MaxUint64 @@ -580,15 +592,13 @@ func (c *bigModExp) RequiredGas(input []byte) uint64 { } // Retrieve the head 32 bytes of exp for the adjusted exponent length - var expHead *uint256.Int - if uint64(len(input)) <= baseLen { - expHead = new(uint256.Int) - } else { + var expHead uint256.Int + if uint64(len(input)) > baseLen { if expLen > 32 { - expHead = new(uint256.Int).SetBytes(getData(input, baseLen, 32)) + expHead.SetBytes(getData(input, baseLen, 32)) } else { // TODO: Check that if expLen < baseLen, then getData will return an empty slice - expHead = new(uint256.Int).SetBytes(getData(input, baseLen, expLen)) + expHead.SetBytes(getData(input, baseLen, expLen)) } } @@ -604,9 +614,9 @@ func (c *bigModExp) RequiredGas(input []byte) uint64 { 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() + baseLen = new(uint256.Int).SetBytes(getData(input, 0, 32)).Uint64() + expLen = new(uint256.Int).SetBytes(getData(input, 32, 32)).Uint64() + modLen = new(uint256.Int).SetBytes(getData(input, 64, 32)).Uint64() ) if len(input) > 96 { input = input[96:]