diff --git a/core/vm/contracts.go b/core/vm/contracts.go index 91ff93dff4..cea88b96a0 100644 --- a/core/vm/contracts.go +++ b/core/vm/contracts.go @@ -360,6 +360,20 @@ type bigModExp struct { eip7883 bool } +// Constants for modexp gas calculation +const ( + byzantiumMultiplier = 8 + berlinMultiplier = 8 + osakaMultiplier = 16 + + byzantiumDivisor = 20 + berlinDivisor = 3 + osakaDivisor = 3 + + berlinMinGas = 200 + osakaMinGas = 500 +) + var ( big1 = big.NewInt(1) big3 = big.NewInt(3) @@ -374,7 +388,7 @@ var ( big199680 = big.NewInt(199680) ) -// modexpMultComplexity implements bigModexp multComplexity formula, as defined in EIP-198 +// byzantiumMultComplexity implements bigModexp multComplexity formula, as defined in EIP-198 // // def mult_complexity(x): // if x <= 64: return x ** 2 @@ -382,7 +396,7 @@ var ( // else: return x ** 2 // 16 + 480 * x - 199680 // // where is x is max(length_of_MODULUS, length_of_BASE) -func modexpMultComplexity(x *big.Int) *big.Int { +func byzantiumMultComplexity(x *big.Int) *big.Int { switch { case x.Cmp(big64) <= 0: x.Mul(x, x) // x ** 2 @@ -402,98 +416,194 @@ func modexpMultComplexity(x *big.Int) *big.Int { return x } +// berlinMultComplexity implements the multiplication complexity formula for Berlin. +// +// def mult_complexity(x): +// +// ceiling(x/8)^2 +// +// where is x is max(length_of_MODULUS, length_of_BASE) +func berlinMultComplexity(x *big.Int) *big.Int { + x = new(big.Int).Add(x, big7) // x + 7 + x = new(big.Int).Rsh(x, 3) // (x + 7) / 8 + return new(big.Int).Mul(x, x) // ((x + 7) / 8) ^ 2 +} + +// osakaMultComplexity implements the multiplication complexity formula for Osaka. +// +// For x <= 32: returns 16 +// For x > 32: returns 2 * ceiling(x/8)^2 +func osakaMultComplexity(x *big.Int) *big.Int { + if x.Cmp(big32) <= 0 { + return big.NewInt(16) + } + // For x > 32, return 2 * berlinMultComplexity(x) + berlinComplexity := berlinMultComplexity(x) + return new(big.Int).Lsh(berlinComplexity, 1) // 2 * berlinComplexity +} + +// 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 *big.Int, multiplier uint64) uint64 { + var iterationCount uint64 + + if expLen <= 32 && expHead.Sign() == 0 { + iterationCount = 0 + } else if expLen <= 32 { + // For small exponents, use MSB position - 1 + if bitLen := expHead.BitLen(); bitLen > 0 { + iterationCount = uint64(bitLen - 1) + } + } else { + // For large exponents: (expLen - 32) * multiplier + MSB position - 1 + iterationCount = (expLen - 32) * multiplier + if bitLen := expHead.BitLen(); bitLen > 0 { + iterationCount += uint64(bitLen - 1) + } + } + + // Return at least 1 + if iterationCount == 0 { + return 1 + } + return iterationCount +} + +// byzantiumGasCalc calculates the gas cost for the modexp precompile using Byzantium rules. +func byzantiumGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 { + // Calculate max(baseLen, modLen) + maxLen := baseLen + if modLen > maxLen { + maxLen = modLen + } + + // Calculate multiplication complexity + // Use SetUint64 to avoid int64 overflow + multComplexity := byzantiumMultComplexity(new(big.Int).SetUint64(maxLen)) + + // Calculate iteration count + iterationCount := calculateIterationCount(expLen, expHead, byzantiumMultiplier) + + // Calculate gas: (multComplexity * iterationCount) / byzantiumDivisor + gas := new(big.Int).Mul(multComplexity, new(big.Int).SetUint64(iterationCount)) + gas.Div(gas, big.NewInt(byzantiumDivisor)) + + if gas.BitLen() > 64 { + return math.MaxUint64 + } + return gas.Uint64() +} + +// berlinGasCalc calculates the gas cost for the modexp precompile using Berlin rules. +func berlinGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 { + // Calculate max(baseLen, modLen) + maxLen := baseLen + if modLen > maxLen { + maxLen = modLen + } + + // Calculate multiplication complexity + // Use SetUint64 to avoid int64 overflow + multComplexity := berlinMultComplexity(new(big.Int).SetUint64(maxLen)) + + // Calculate iteration count + iterationCount := calculateIterationCount(expLen, expHead, berlinMultiplier) + + // Calculate gas: (multComplexity * iterationCount) / berlinDivisor + gas := new(big.Int).Mul(multComplexity, new(big.Int).SetUint64(iterationCount)) + gas.Div(gas, big.NewInt(berlinDivisor)) + + if gas.BitLen() > 64 { + return math.MaxUint64 + } + + // Return at least berlinMinGas + gasUint64 := gas.Uint64() + if gasUint64 < berlinMinGas { + return berlinMinGas + } + return gasUint64 +} + +// osakaGasCalc calculates the gas cost for the modexp precompile using Osaka rules. +func osakaGasCalc(baseLen, expLen, modLen uint64, expHead *big.Int) uint64 { + // Calculate max(baseLen, modLen) + maxLen := baseLen + if modLen > maxLen { + maxLen = modLen + } + + // Calculate multiplication complexity + // Use SetUint64 to avoid int64 overflow + multComplexity := osakaMultComplexity(new(big.Int).SetUint64(maxLen)) + + // Calculate iteration count + iterationCount := calculateIterationCount(expLen, expHead, osakaMultiplier) + + // Calculate gas: (multComplexity * iterationCount) / osakaDivisor + gas := new(big.Int).Mul(multComplexity, new(big.Int).SetUint64(iterationCount)) + gas.Div(gas, big.NewInt(osakaDivisor)) + + if gas.BitLen() > 64 { + return math.MaxUint64 + } + + // Return at least osakaMinGas + gasUint64 := gas.Uint64() + if gasUint64 < osakaMinGas { + return osakaMinGas + } + return gasUint64 +} + // RequiredGas returns the gas required to execute the pre-compiled contract. func (c *bigModExp) RequiredGas(input []byte) uint64 { - var ( - baseLen = new(big.Int).SetBytes(getData(input, 0, 32)) - expLen = new(big.Int).SetBytes(getData(input, 32, 32)) - modLen = new(big.Int).SetBytes(getData(input, 64, 32)) - ) + // Parse input lengths + baseLenBig := new(big.Int).SetBytes(getData(input, 0, 32)) + expLenBig := new(big.Int).SetBytes(getData(input, 32, 32)) + modLenBig := new(big.Int).SetBytes(getData(input, 64, 32)) + + // Convert to uint64, capping at max value + baseLen := baseLenBig.Uint64() + if baseLenBig.BitLen() > 64 { + baseLen = math.MaxUint64 + } + expLen := expLenBig.Uint64() + if expLenBig.BitLen() > 64 { + expLen = math.MaxUint64 + } + modLen := modLenBig.Uint64() + if modLenBig.BitLen() > 64 { + modLen = math.MaxUint64 + } + + // Skip the header if len(input) > 96 { input = input[96:] } else { input = input[:0] } + // Retrieve the head 32 bytes of exp for the adjusted exponent length var expHead *big.Int - if big.NewInt(int64(len(input))).Cmp(baseLen) <= 0 { + if uint64(len(input)) <= baseLen { expHead = new(big.Int) } else { - if expLen.Cmp(big32) > 0 { - expHead = new(big.Int).SetBytes(getData(input, baseLen.Uint64(), 32)) + if expLen > 32 { + expHead = new(big.Int).SetBytes(getData(input, baseLen, 32)) } else { - expHead = new(big.Int).SetBytes(getData(input, baseLen.Uint64(), expLen.Uint64())) + expHead = new(big.Int).SetBytes(getData(input, baseLen, expLen)) } } - // Calculate the adjusted exponent length - var msb int - if bitlen := expHead.BitLen(); bitlen > 0 { - msb = bitlen - 1 - } - adjExpLen := new(big.Int) - if expLen.Cmp(big32) > 0 { - adjExpLen.Sub(expLen, big32) - if c.eip7883 { - adjExpLen.Lsh(adjExpLen, 4) - } else { - adjExpLen.Lsh(adjExpLen, 3) - } - } - adjExpLen.Add(adjExpLen, big.NewInt(int64(msb))) - // Calculate the gas cost of the operation - gas := new(big.Int) - if modLen.Cmp(baseLen) < 0 { - gas.Set(baseLen) + + // Choose the appropriate gas calculation based on the EIP flags + if c.eip7883 { + return osakaGasCalc(baseLen, expLen, modLen, expHead) + } else if c.eip2565 { + return berlinGasCalc(baseLen, expLen, modLen, expHead) } else { - gas.Set(modLen) + return byzantiumGasCalc(baseLen, expLen, modLen, expHead) } - - maxLenOver32 := gas.Cmp(big32) > 0 - if c.eip2565 { - // EIP-2565 (Berlin fork) has three changes: - // - // 1. Different multComplexity (inlined here) - // in EIP-2565 (https://eips.ethereum.org/EIPS/eip-2565): - // - // def mult_complexity(x): - // ceiling(x/8)^2 - // - // where is x is max(length_of_MODULUS, length_of_BASE) - gas.Add(gas, big7) - gas.Rsh(gas, 3) - gas.Mul(gas, gas) - - var minPrice uint64 = 200 - if c.eip7883 { - minPrice = 500 - if maxLenOver32 { - gas.Add(gas, gas) - } else { - gas = big.NewInt(16) - } - } - - if adjExpLen.Cmp(big1) > 0 { - gas.Mul(gas, adjExpLen) - } - // 2. Different divisor (`GQUADDIVISOR`) (3) - gas.Div(gas, big3) - if gas.BitLen() > 64 { - return math.MaxUint64 - } - return max(minPrice, gas.Uint64()) - } - - // Pre-Berlin logic. - gas = modexpMultComplexity(gas) - if adjExpLen.Cmp(big1) > 0 { - gas.Mul(gas, adjExpLen) - } - gas.Div(gas, big20) - if gas.BitLen() > 64 { - return math.MaxUint64 - } - return gas.Uint64() } func (c *bigModExp) Run(input []byte) ([]byte, error) {