diff --git a/core/vm/contracts.go b/core/vm/contracts.go index 1d118a6d7c..928b5511ab 100644 --- a/core/vm/contracts.go +++ b/core/vm/contracts.go @@ -396,33 +396,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 } } @@ -433,13 +440,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) @@ -454,18 +465,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 @@ -483,8 +495,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 @@ -494,9 +505,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 @@ -505,8 +517,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 @@ -516,9 +527,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 @@ -529,8 +541,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 @@ -540,9 +551,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 @@ -581,15 +593,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)) } }