diff --git a/gas-report b/gas-report index 0d13dfa..8d8c89d 100644 --- a/gas-report +++ b/gas-report @@ -1,7 +1,7 @@ -Poseidon2Huff_BN254:hash_1: 14845 gas -Poseidon2Huff_BN254:hash_2: 14845 gas -Poseidon2Huff_BN254:hash_3: 14845 gas -Poseidon2Yul_BN254:fallback: 19607 gas +Poseidon2Huff_BN254:hash_1: 14869 gas +Poseidon2Huff_BN254:hash_2: 14869 gas +Poseidon2Huff_BN254:hash_3: 14869 gas +Poseidon2Yul_BN254:fallback: 19639 gas Poseidon2_BN254:hash_1: 217702 gas Poseidon2_BN254:hash_2: 218222 gas Poseidon2_BN254:hash_3: 218801 gas diff --git a/gas-report.sh b/gas-report.sh index 4f47e8a..2fed5ca 100755 --- a/gas-report.sh +++ b/gas-report.sh @@ -9,6 +9,6 @@ OUTPUT=$(forge test --match-contract Poseidon2Test -vvvv 2>&1 | \ grep -E "^[[:space:]]+├─ \[[0-9]+\] (Poseidon2_BN254|Poseidon2Yul_BN254|Poseidon2Huff_BN254)::(hash_1|hash_2|hash_3|fallback)" | \ sed 's/.*├─ \[\([0-9]*\)\] \(.*\)::\([a-z_0-9]*\).*/\2:\3: \1 gas/' | \ - sort -u) + LC_ALL=C sort -u) echo "$OUTPUT" | tee gas-report diff --git a/generate-yul.ts b/generate-yul.ts index d514f69..ec2fd17 100644 --- a/generate-yul.ts +++ b/generate-yul.ts @@ -44,10 +44,10 @@ function poseidon2_core(s0: string, s1: string, s2: string, s3: string) { return ` let PRIME := 0x30644e72e131a029b85045b68181585d2833e84879b9709143e1f593f0000001 - let state0 := ${s0} - let state1 := ${s1} - let state2 := ${s2} - let state3 := ${s3} + let state0 := mod(${s0}, PRIME) + let state1 := mod(${s1}, PRIME) + let state2 := mod(${s2}, PRIME) + let state3 := mod(${s3}, PRIME) ${poseidon2_rounds()} `; diff --git a/src/bn254/huff/Utils.huff b/src/bn254/huff/Utils.huff index e3ebf65..d3384d2 100644 --- a/src/bn254/huff/Utils.huff +++ b/src/bn254/huff/Utils.huff @@ -27,7 +27,7 @@ /// @notice Absorb calldata inputs into sponge state for interface-compatible hashing /// @dev Supports IPoseidon2 interface: hash_1(uint256), hash_2(uint256,uint256), hash_3(uint256,uint256,uint256) /// @dev Calldata layout: [4-byte selector][32-byte arg0][32-byte arg1][32-byte arg2] -/// @dev User-supplied dirty calldata (>= PRIME) is acceptable; will be cleaned by S-box +/// @dev Inputs are cleaned via mod PRIME to ensure proper field element values #define macro ABSORB_CALLDATA() = takes (1) returns (5) { // takes: [PRIME] @@ -37,10 +37,16 @@ 0x5 shr // [(calldatasize - 4) / 32, PRIME] 0x40 shl // [iv, PRIME] - // Load inputs from calldata - 0x44 calldataload // [input2, iv, PRIME] - 0x24 calldataload // [input1, input2, iv, PRIME] - 0x04 calldataload // [input0, input1, input2, iv, PRIME] + // Load inputs from calldata and apply mod PRIME + dup2 // [PRIME, iv, PRIME] + 0x44 calldataload // [input2, PRIME, iv, PRIME] + mod // [input2 % PRIME, iv, PRIME] + dup3 // [PRIME, input2 % PRIME, iv, PRIME] + 0x24 calldataload // [input1, PRIME, input2 % PRIME, iv, PRIME] + mod // [input1 % PRIME, input2 % PRIME, iv, PRIME] + dup4 // [PRIME, input1 % PRIME, input2 % PRIME, iv, PRIME] + 0x04 calldataload // [input0, PRIME, input1 % PRIME, input2 % PRIME, iv, PRIME] + mod // [input0 % PRIME, input1 % PRIME, input2 % PRIME, iv, PRIME] // returns: [state0, state1, state2, state3, PRIME] } diff --git a/src/bn254/yul/LibPoseidon2Yul.sol b/src/bn254/yul/LibPoseidon2Yul.sol index 3fb651b..02c89aa 100644 --- a/src/bn254/yul/LibPoseidon2Yul.sol +++ b/src/bn254/yul/LibPoseidon2Yul.sol @@ -23,10 +23,10 @@ library LibPoseidon2Yul { assembly { let PRIME := 0x30644e72e131a029b85045b68181585d2833e84879b9709143e1f593f0000001 - let state0 := s0 - let state1 := s1 - let state2 := s2 - let state3 := s3 + let state0 := mod(s0, PRIME) + let state1 := mod(s1, PRIME) + let state2 := mod(s2, PRIME) + let state3 := mod(s3, PRIME) // Apply 1st linear layer diff --git a/test/Poseidon2.t.sol b/test/Poseidon2.t.sol index d9a1e6a..fd0ce8e 100644 --- a/test/Poseidon2.t.sol +++ b/test/Poseidon2.t.sol @@ -199,6 +199,93 @@ contract Poseidon2Test is Test { ); } + // ============================================================ + // Input sanitization tests + // ============================================================ + + function test_input_sanitization() public view { + // Test that inputs >= PRIME are handled correctly + uint256 PRIME = Field.PRIME; + + // Test input = PRIME (should be treated as 0) + { + uint256 result = poseidon2Yul.hash_1(PRIME); + uint256 expected = poseidon2Yul.hash_1(0); + assertEq(result, expected, "Yul: PRIME should equal 0"); + } + { + uint256 result = poseidon2Huff.hash_1(PRIME); + uint256 expected = poseidon2Huff.hash_1(0); + assertEq(result, expected, "Huff: PRIME should equal 0"); + } + + // Test input = PRIME + 5 (should be treated as 5) + { + uint256 result = poseidon2Yul.hash_1(PRIME + 5); + uint256 expected = poseidon2Yul.hash_1(5); + assertEq(result, expected, "Yul: PRIME+5 should equal 5"); + } + { + uint256 result = poseidon2Huff.hash_1(PRIME + 5); + uint256 expected = poseidon2Huff.hash_1(5); + assertEq(result, expected, "Huff: PRIME+5 should equal 5"); + } + + // Test input = type(uint256).max (should be treated as type(uint256).max % PRIME) + { + uint256 maxModPrime = type(uint256).max % PRIME; + uint256 result = poseidon2Yul.hash_1(type(uint256).max); + uint256 expected = poseidon2Yul.hash_1(maxModPrime); + assertEq(result, expected, "Yul: uint256.max should equal max % PRIME"); + } + { + uint256 maxModPrime = type(uint256).max % PRIME; + uint256 result = poseidon2Huff.hash_1(type(uint256).max); + uint256 expected = poseidon2Huff.hash_1(maxModPrime); + assertEq(result, expected, "Huff: uint256.max should equal max % PRIME"); + } + + // Test hash_2 with mixed inputs + { + uint256 result = poseidon2Yul.hash_2(PRIME + 3, PRIME + 7); + uint256 expected = poseidon2Yul.hash_2(3, 7); + assertEq(result, expected, "Yul: hash_2 sanitization"); + } + { + uint256 result = poseidon2Huff.hash_2(PRIME + 3, PRIME + 7); + uint256 expected = poseidon2Huff.hash_2(3, 7); + assertEq(result, expected, "Huff: hash_2 sanitization"); + } + + // Test hash_2 with type(uint256).max + { + uint256 maxModPrime = type(uint256).max % PRIME; + uint256 result = poseidon2Yul.hash_2(type(uint256).max, type(uint256).max); + uint256 expected = poseidon2Yul.hash_2(maxModPrime, maxModPrime); + assertEq(result, expected, "Yul: hash_2 with uint256.max"); + } + { + uint256 maxModPrime = type(uint256).max % PRIME; + uint256 result = poseidon2Huff.hash_2(type(uint256).max, type(uint256).max); + uint256 expected = poseidon2Huff.hash_2(maxModPrime, maxModPrime); + assertEq(result, expected, "Huff: hash_2 with uint256.max"); + } + + // Test hash_3 with type(uint256).max + { + uint256 maxModPrime = type(uint256).max % PRIME; + uint256 result = poseidon2Yul.hash_3(type(uint256).max, type(uint256).max, type(uint256).max); + uint256 expected = poseidon2Yul.hash_3(maxModPrime, maxModPrime, maxModPrime); + assertEq(result, expected, "Yul: hash_3 with uint256.max"); + } + { + uint256 maxModPrime = type(uint256).max % PRIME; + uint256 result = poseidon2Huff.hash_3(type(uint256).max, type(uint256).max, type(uint256).max); + uint256 expected = poseidon2Huff.hash_3(maxModPrime, maxModPrime, maxModPrime); + assertEq(result, expected, "Huff: hash_3 with uint256.max"); + } + } + // ============================================================ // Variable length hashing tests // ============================================================