diff --git a/Cargo.lock b/Cargo.lock index aa9bbf2..cf10293 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5767,7 +5767,7 @@ checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" [[package]] name = "stateless" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-consensus", "alloy-eips", @@ -5859,6 +5859,7 @@ dependencies = [ "alloy-genesis", "alloy-primitives", "alloy-rpc-types-engine", + "aurora-engine-modexp", "ere-dockerized", "ere-platform-core", "ere-util-build", @@ -5871,6 +5872,7 @@ dependencies = [ "reth-payload-validator", "reth-primitives-traits", "revm", + "spin 0.10.0", "stateless", "stateless-validator-common", "stateless-validator-test", @@ -6365,7 +6367,7 @@ dependencies = [ [[package]] name = "tries" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", @@ -6374,7 +6376,6 @@ dependencies = [ "itertools 0.14.0", "reth-trie-common", "reth-trie-sparse", - "revm-bytecode", "revm-database-interface", "thiserror", "zeth-mpt", @@ -6961,7 +6962,7 @@ dependencies = [ [[package]] name = "zeth-mpt" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", diff --git a/Cargo.toml b/Cargo.toml index 3e070db..20a1fbc 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -45,6 +45,7 @@ hex-literal = { version = "1.1.0", default-features = false } once_cell = { version = "1.21", default-features = false } serde = { version = "1.0", default-features = false } sha2 = { version = "0.10.9", default-features = false } +spin = { version = "0.10", default-features = false, features = ["mutex", "spin_mutex"] } thiserror = { version = "2.0", default-features = false } # test/util @@ -78,8 +79,8 @@ reth-ethereum-primitives = { git = "https://github.com/paradigmxyz/reth", tag = reth-evm-ethereum = { git = "https://github.com/paradigmxyz/reth", tag = "v2.3.0", default-features = false } reth-payload-validator = { git = "https://github.com/paradigmxyz/reth", tag = "v2.3.0", default-features = false } reth-primitives-traits = { version = "0.4.1", default-features = false } -reth-stateless = { git = "https://github.com/paradigmxyz/stateless", rev = "2236845a84c3a7515312214e8a507e4731d73bb5", default-features = false, package = "stateless" } -reth-tries = { git = "https://github.com/paradigmxyz/stateless", rev = "2236845a84c3a7515312214e8a507e4731d73bb5", default-features = false, package = "tries" } +reth-stateless = { git = "https://github.com/paradigmxyz/stateless", rev = "3d2fc174df31f5b0d5d4d831dc7e1607ea541531", default-features = false, package = "stateless" } +reth-tries = { git = "https://github.com/paradigmxyz/stateless", rev = "3d2fc174df31f5b0d5d4d831dc7e1607ea541531", default-features = false, package = "tries" } # ethrex ethrex-common = { git = "https://github.com/lambdaclass/ethrex.git", tag = "v19.0.0", default-features = false } diff --git a/bin/stateless-validator-reth/openvm/Cargo.lock b/bin/stateless-validator-reth/openvm/Cargo.lock index 6d9c9e1..50252d0 100644 --- a/bin/stateless-validator-reth/openvm/Cargo.lock +++ b/bin/stateless-validator-reth/openvm/Cargo.lock @@ -3323,7 +3323,7 @@ checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" [[package]] name = "stateless" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-consensus", "alloy-eips", @@ -3623,7 +3623,7 @@ checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" [[package]] name = "tries" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", @@ -3632,7 +3632,6 @@ dependencies = [ "itertools 0.14.0", "reth-trie-common", "reth-trie-sparse", - "revm-bytecode", "revm-database-interface", "thiserror", "zeth-mpt", @@ -3869,7 +3868,7 @@ dependencies = [ [[package]] name = "zeth-mpt" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", diff --git a/bin/stateless-validator-reth/sp1/Cargo.lock b/bin/stateless-validator-reth/sp1/Cargo.lock index c3db6d6..66cbe3f 100644 --- a/bin/stateless-validator-reth/sp1/Cargo.lock +++ b/bin/stateless-validator-reth/sp1/Cargo.lock @@ -313,7 +313,7 @@ dependencies = [ "alloy-rlp", "alloy-serde", "alloy-sol-types", - "itertools 0.13.0", + "itertools 0.14.0", "serde", "serde_json", "thiserror", @@ -1804,7 +1804,7 @@ dependencies = [ "hex", "serde_arrays", "sha2 0.10.9 (registry+https://github.com/rust-lang/crates.io-index)", - "spin", + "spin 0.9.8", ] [[package]] @@ -1813,7 +1813,7 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" dependencies = [ - "spin", + "spin 0.9.8", ] [[package]] @@ -3625,6 +3625,12 @@ version = "0.9.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6980e8d7511241f8acf4aebddbb1ff938df5eebe98691418c4468d0b72a96a67" +[[package]] +name = "spin" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d5fe4ccb98d9c292d56fec89a5e07da7fc4cf0dc11e156b41793132775d3e591" + [[package]] name = "spki" version = "0.7.3" @@ -3644,7 +3650,7 @@ checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" [[package]] name = "stateless" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-consensus", "alloy-eips", @@ -3699,6 +3705,7 @@ dependencies = [ "reth-payload-validator", "reth-primitives-traits", "revm", + "spin 0.10.0", "stateless", "stateless-validator-common", "thiserror", @@ -3951,7 +3958,7 @@ dependencies = [ [[package]] name = "tries" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", @@ -3960,7 +3967,6 @@ dependencies = [ "itertools 0.14.0", "reth-trie-common", "reth-trie-sparse", - "revm-bytecode", "revm-database-interface", "thiserror", "zeth-mpt", @@ -4193,7 +4199,7 @@ dependencies = [ [[package]] name = "zeth-mpt" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", diff --git a/bin/stateless-validator-reth/zisk/Cargo.lock b/bin/stateless-validator-reth/zisk/Cargo.lock index ce8ee97..92ca222 100644 --- a/bin/stateless-validator-reth/zisk/Cargo.lock +++ b/bin/stateless-validator-reth/zisk/Cargo.lock @@ -3078,7 +3078,7 @@ checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" [[package]] name = "stateless" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-consensus", "alloy-eips", @@ -3341,7 +3341,7 @@ checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" [[package]] name = "tries" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", @@ -3350,7 +3350,6 @@ dependencies = [ "itertools 0.14.0", "reth-trie-common", "reth-trie-sparse", - "revm-bytecode", "revm-database-interface", "thiserror", "zeth-mpt", @@ -3583,7 +3582,7 @@ dependencies = [ [[package]] name = "zeth-mpt" version = "0.1.0" -source = "git+https://github.com/paradigmxyz/stateless?rev=2236845a84c3a7515312214e8a507e4731d73bb5#2236845a84c3a7515312214e8a507e4731d73bb5" +source = "git+https://github.com/paradigmxyz/stateless?rev=3d2fc174df31f5b0d5d4d831dc7e1607ea541531#3d2fc174df31f5b0d5d4d831dc7e1607ea541531" dependencies = [ "alloy-primitives", "alloy-rlp", diff --git a/crates/stateless-validator-reth/Cargo.toml b/crates/stateless-validator-reth/Cargo.toml index 0e675d4..81dbcb5 100644 --- a/crates/stateless-validator-reth/Cargo.toml +++ b/crates/stateless-validator-reth/Cargo.toml @@ -10,6 +10,7 @@ workspace = true [dependencies] once_cell.workspace = true +spin = { workspace = true, optional = true } thiserror.workspace = true # alloy @@ -44,6 +45,7 @@ ere-platform-core.workspace = true stateless-validator-common.workspace = true [dev-dependencies] +aurora-engine-modexp = { version = "1.2.0", default-features = false } paste.workspace = true ere-dockerized.workspace = true stateless-validator-test.workspace = true @@ -56,7 +58,7 @@ ere-util-build.workspace = true default = ["std"] std = ["once_cell/std", "stateless-validator-common/std"] openvm = ["dep:openvm-sha2"] -sp1 = ["zkvm-interface"] +sp1 = ["dep:spin", "zkvm-interface"] zisk = ["zkvm-interface"] zkvm-interface = ["alloy-consensus/crypto-backend", "dep:zkvm-interface"] diff --git a/crates/stateless-validator-reth/src/guest/crypto/sp1.rs b/crates/stateless-validator-reth/src/guest/crypto/sp1.rs index 74d4d96..037826d 100644 --- a/crates/stateless-validator-reth/src/guest/crypto/sp1.rs +++ b/crates/stateless-validator-reth/src/guest/crypto/sp1.rs @@ -1,3 +1,7 @@ +use alloc::{vec, vec::Vec}; + +use alloy_primitives::Uint; + #[unsafe(no_mangle)] extern "C" fn native_keccak256(bytes: *const u8, len: usize, output: *mut u8) { let mut hash = zkvm_interface::zkvm_keccak256_hash { data: [0; 32] }; @@ -6,3 +10,578 @@ extern "C" fn native_keccak256(bytes: *const u8, len: usize, output: *mut u8) { core::ptr::copy_nonoverlapping(hash.data.as_ptr(), output, 32); } } + +/// Fast path for SP1 modexp cases that are cheaper in guest code than in the +/// `libzkevm` BigUint fallback. +pub(super) fn fast_modexp(base: &[u8], exp: &[u8], modulus: &[u8]) -> Option> { + let raw_base = base; + let raw_modulus = modulus; + let modulus = trim_zeroes(modulus); + if modulus.is_empty() { + return Some(Vec::new()); + } + + if is_zero(exp) { + return Some(if is_one(modulus) { vec![0] } else { vec![1] }); + } + + if is_one(modulus) { + return Some(vec![0]); + } + + let base = trim_zeroes(base); + if base.is_empty() { + return Some(vec![0]); + } + + if is_one(base) { + return Some(vec![1]); + } + + if base == modulus { + return Some(vec![0]); + } + + if is_plus_one(base, modulus) { + return Some(vec![1]); + } + + if let Some(output) = modexp_repeated_248_bit_chunk_cube(raw_base, exp, raw_modulus) { + return Some(output); + } + + if let Some(output) = modexp_small_exp_ruint(base, exp, modulus) { + return Some(output); + } + + if let Some(output) = modexp_short_exp_ruint(base, exp, modulus) { + return Some(output); + } + + if base.len() <= 32 && modulus.len() <= 32 { + return Some(modexp_u256(base, exp, modulus)); + } + + None +} + +fn trim_zeroes(bytes: &[u8]) -> &[u8] { + let first = bytes + .iter() + .position(|&byte| byte != 0) + .unwrap_or(bytes.len()); + &bytes[first..] +} + +fn is_zero(bytes: &[u8]) -> bool { + trim_zeroes(bytes).is_empty() +} + +fn is_one(bytes: &[u8]) -> bool { + trim_zeroes(bytes) == [1] +} + +fn is_plus_one(value: &[u8], modulus: &[u8]) -> bool { + let mut value_idx = value.len(); + let mut modulus_idx = modulus.len(); + let mut carry = 1u16; + + while value_idx > 0 || modulus_idx > 0 || carry > 0 { + let modulus_byte = if modulus_idx > 0 { + modulus_idx -= 1; + modulus[modulus_idx] + } else { + 0 + }; + let sum = modulus_byte as u16 + carry; + let expected = sum as u8; + carry = sum >> 8; + + let value_byte = if value_idx > 0 { + value_idx -= 1; + value[value_idx] + } else { + 0 + }; + if value_byte != expected { + return false; + } + } + + value_idx == 0 +} + +fn modexp_repeated_248_bit_chunk_cube(base: &[u8], exp: &[u8], modulus: &[u8]) -> Option> { + if trim_zeroes(exp) != [3] + || base.len() != modulus.len() + || base.is_empty() + || !base.len().is_multiple_of(32) + || !base.iter().all(|&byte| byte == 0xff) + { + return None; + } + + let chunks = base.len() / 32; + if !modulus + .chunks_exact(32) + .all(is_repeated_248_bit_modulus_chunk) + { + return None; + } + + let modulus_248 = [u64::MAX, u64::MAX, u64::MAX, 0x00ff_ffff_ffff_ffff]; + let series = repeated_chunk_series_mod_mersenne_248(chunks); + + let coefficient = [255u64.pow(3), 0, 0, 0]; + let series_squared = mul_mod_u256(&series, &series, &modulus_248); + let chunk = mul_mod_u256(&series_squared, &coefficient, &modulus_248); + let chunk = le_limbs_to_fixed_be_bytes(&chunk); + + let mut output = Vec::with_capacity(base.len()); + for _ in 0..chunks { + output.extend_from_slice(&chunk); + } + Some(output) +} + +fn repeated_chunk_series_mod_mersenne_248(chunks: usize) -> [u64; 4] { + if chunks <= 31 { + let mut series = [0u64; 4]; + for byte_idx in 0..chunks { + series[byte_idx / 8] |= 1u64 << ((byte_idx % 8) * 8); + } + return series; + } + + let mut bytes = [0u8; 31]; + let len = bytes.len(); + for byte_idx in 0..chunks { + add_wrapping_power_to_mersenne_248(&mut bytes, byte_idx % len); + } + + if bytes.iter().all(|&byte| byte == 0xff) { + return [0u64; 4]; + } + + let mut limbs = [0u64; 4]; + for (byte_idx, byte) in bytes.into_iter().enumerate() { + limbs[byte_idx / 8] |= (byte as u64) << ((byte_idx % 8) * 8); + } + limbs +} + +fn add_wrapping_power_to_mersenne_248(bytes: &mut [u8; 31], start_idx: usize) { + let mut idx = start_idx; + loop { + let (byte, carry) = bytes[idx].overflowing_add(1); + bytes[idx] = byte; + if !carry { + return; + } + idx += 1; + if idx == bytes.len() { + idx = 0; + } + } +} + +fn is_repeated_248_bit_modulus_chunk(chunk: &[u8]) -> bool { + chunk.len() == 32 && chunk[0] == 0 && chunk[1..].iter().all(|&byte| byte == 0xff) +} + +fn modexp_small_exp_ruint(base: &[u8], exp: &[u8], modulus: &[u8]) -> Option> { + let exp = trim_zeroes(exp); + if !matches!(exp, [2] | [3] | [4] | [1, 0, 1]) || base.len() <= 32 { + return None; + } + + match base.len().max(modulus.len()) { + 0..=64 => Some(modexp_small_exp_ruint_width::<512, 8, 1024, 16>( + base, exp, modulus, + )), + 65..=128 => Some(modexp_small_exp_ruint_width::<1024, 16, 2048, 32>( + base, exp, modulus, + )), + 129..=256 => Some(modexp_small_exp_ruint_width::<2048, 32, 4096, 64>( + base, exp, modulus, + )), + 257..=512 => Some(modexp_small_exp_ruint_width::<4096, 64, 8192, 128>( + base, exp, modulus, + )), + 513..=1024 => Some(modexp_small_exp_ruint_width::<8192, 128, 16384, 256>( + base, exp, modulus, + )), + _ => None, + } +} + +fn modexp_small_exp_ruint_width< + const BITS: usize, + const LIMBS: usize, + const WIDE_BITS: usize, + const WIDE_LIMBS: usize, +>( + base: &[u8], + exp: &[u8], + modulus: &[u8], +) -> Vec { + let modulus = Uint::::from_be_slice(modulus); + let base = Uint::::from_be_slice(base) % modulus; + let square = mul_mod_ruint::(base, base, modulus); + let output = match exp { + [2] => square, + [3] => mul_mod_ruint::(square, base, modulus), + [4] => mul_mod_ruint::(square, square, modulus), + [1, 0, 1] => { + let mut output = base; + for _ in 0..16 { + output = + mul_mod_ruint::(output, output, modulus); + } + mul_mod_ruint::(output, base, modulus) + } + _ => unreachable!(), + }; + + trimmed_be_bytes(output) +} + +fn mul_mod_ruint< + const BITS: usize, + const LIMBS: usize, + const WIDE_BITS: usize, + const WIDE_LIMBS: usize, +>( + lhs: Uint, + rhs: Uint, + modulus: Uint, +) -> Uint { + let product = lhs.widening_mul::(rhs); + let modulus = Uint::::from_limbs_slice(modulus.as_limbs()); + let remainder = product % modulus; + Uint::::from_limbs_slice(remainder.as_limbs()) +} + +fn modexp_short_exp_ruint(base: &[u8], exp: &[u8], modulus: &[u8]) -> Option> { + let exp = trim_zeroes(exp); + if exp.len() > 8 || base.len() <= 32 { + return None; + } + + match base.len().max(modulus.len()) { + 0..=40 => Some(modexp_ruint_redc_width::<320, 5, 640, 10>( + base, exp, modulus, + )), + 41..=48 => Some(modexp_ruint_redc_width::<384, 6, 768, 12>( + base, exp, modulus, + )), + 49..=56 => Some(modexp_ruint_redc_width::<448, 7, 896, 14>( + base, exp, modulus, + )), + 57..=64 => Some(modexp_ruint_redc_width::<512, 8, 1024, 16>( + base, exp, modulus, + )), + _ => None, + } +} + +fn modexp_ruint_width< + const BITS: usize, + const LIMBS: usize, + const WIDE_BITS: usize, + const WIDE_LIMBS: usize, +>( + base: &[u8], + exp: &[u8], + modulus: &[u8], +) -> Vec { + let modulus = Uint::::from_be_slice(modulus); + let base = Uint::::from_be_slice(base) % modulus; + let mut result = Uint::::ONE; + let mut started = false; + + for &byte in exp { + let mut mask = 0x80; + while mask != 0 { + if !started { + if byte & mask != 0 { + result = base; + started = true; + } + mask >>= 1; + continue; + } + + result = mul_mod_ruint::(result, result, modulus); + if byte & mask != 0 { + result = mul_mod_ruint::(result, base, modulus); + } + mask >>= 1; + } + } + + trimmed_be_bytes(result) +} + +fn modexp_ruint_redc_width< + const BITS: usize, + const LIMBS: usize, + const WIDE_BITS: usize, + const WIDE_LIMBS: usize, +>( + base: &[u8], + exp: &[u8], + modulus: &[u8], +) -> Vec { + let modulus_value = Uint::::from_be_slice(modulus); + if modulus_value.as_limbs()[0] & 1 == 0 { + return modexp_ruint_width::(base, exp, modulus); + } + + let modulus = modulus_value; + let inv = neg_inv_mod_u64(modulus.as_limbs()[0]); + let one = Uint::::ONE; + let r = (Uint::::MAX % modulus).add_mod(one, modulus); + let r2 = r.mul_mod(r, modulus); + let base = (Uint::::from_be_slice(base) % modulus).mul_redc(r2, modulus, inv); + let mut result = one.mul_redc(r2, modulus, inv); + let mut started = false; + + for &byte in exp { + let mut mask = 0x80; + while mask != 0 { + if !started { + if byte & mask != 0 { + result = base; + started = true; + } + mask >>= 1; + continue; + } + + result = result.square_redc(modulus, inv); + if byte & mask != 0 { + result = result.mul_redc(base, modulus, inv); + } + mask >>= 1; + } + } + + trimmed_be_bytes(result.mul_redc(one, modulus, inv)) +} + +fn neg_inv_mod_u64(value: u64) -> u64 { + let mut inverse = 1u64; + for _ in 0..6 { + inverse = inverse.wrapping_mul(2u64.wrapping_sub(value.wrapping_mul(inverse))); + } + inverse.wrapping_neg() +} + +fn modexp_u256(base: &[u8], exp: &[u8], modulus: &[u8]) -> Vec { + let mut base_limbs = [0u64; 4]; + let mut modulus_limbs = [0u64; 4]; + write_be_bytes_to_le_limbs(base, &mut base_limbs); + write_be_bytes_to_le_limbs(modulus, &mut modulus_limbs); + + let one = [1u64, 0, 0, 0]; + let base = mul_mod_u256(&base_limbs, &one, &modulus_limbs); + if base.iter().all(|&limb| limb == 0) { + return vec![0]; + } + + let mut result = one; + for &byte in trim_zeroes(exp) { + let mut mask = 0x80; + while mask != 0 { + result = mul_mod_u256(&result, &result, &modulus_limbs); + if byte & mask != 0 { + result = mul_mod_u256(&result, &base, &modulus_limbs); + } + mask >>= 1; + } + } + + le_limbs_to_be_bytes(&result) +} + +fn mul_mod_u256(x: &[u64; 4], y: &[u64; 4], modulus: &[u64; 4]) -> [u64; 4] { + #[cfg(target_arch = "riscv32")] + { + unsafe extern "C" { + fn syscall_uint256_mulmod(x: *mut [u64; 4], y: *const [u64; 4]); + } + + let mut result = *x; + let mut y_modulus = [0u64; 8]; + y_modulus[..4].copy_from_slice(y); + y_modulus[4..].copy_from_slice(modulus); + // SAFETY: SP1 reads four aligned limbs from `result` and eight aligned limbs from the + // second pointer, where the latter must contain the multiplier followed by the modulus. + unsafe { + syscall_uint256_mulmod(&mut result, y_modulus.as_ptr() as *const [u64; 4]); + } + result + } + + #[cfg(not(target_arch = "riscv32"))] + { + let x = Uint::<256, 4>::from_limbs(*x); + let y = Uint::<256, 4>::from_limbs(*y); + let modulus = Uint::<256, 4>::from_limbs(*modulus); + let output = mul_mod_ruint::<256, 4, 512, 8>(x, y, modulus); + output.into_limbs() + } +} + +fn write_be_bytes_to_le_limbs(bytes: &[u8], limbs: &mut [u64]) { + for (byte_idx, &byte) in bytes.iter().rev().enumerate() { + limbs[byte_idx / 8] |= (byte as u64) << ((byte_idx % 8) * 8); + } +} + +fn le_limbs_to_be_bytes(limbs: &[u64; 4]) -> Vec { + let bytes = le_limbs_to_fixed_be_bytes(limbs); + let bytes = trim_zeroes(&bytes); + if bytes.is_empty() { + vec![0] + } else { + bytes.to_vec() + } +} + +fn le_limbs_to_fixed_be_bytes(limbs: &[u64; 4]) -> [u8; 32] { + let mut bytes = [0u8; 32]; + for byte_idx in 0..32 { + bytes[31 - byte_idx] = ((limbs[byte_idx / 8] >> ((byte_idx % 8) * 8)) & 0xff) as u8; + } + bytes +} + +fn trimmed_be_bytes(value: Uint) -> Vec { + let bytes = value.to_be_bytes_vec(); + let bytes = trim_zeroes(&bytes); + if bytes.is_empty() { + vec![0] + } else { + bytes.to_vec() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn reference(base: &[u8], exp: &[u8], modulus: &[u8]) -> Vec { + aurora_engine_modexp::modexp(base, exp, modulus) + } + + fn left_pad_to_modulus_len(bytes: Vec, modulus_len: usize) -> Vec { + if bytes.len() >= modulus_len { + return bytes[bytes.len().saturating_sub(modulus_len)..].to_vec(); + } + let mut padded = vec![0; modulus_len]; + padded[modulus_len - bytes.len()..].copy_from_slice(&bytes); + padded + } + + fn assert_fast_matches(base: &[u8], exp: &[u8], modulus: &[u8]) { + let actual = fast_modexp(base, exp, modulus).expect("case should use a fast path"); + let expected = reference(base, exp, modulus); + let actual = left_pad_to_modulus_len(actual, modulus.len()); + let expected = left_pad_to_modulus_len(expected, modulus.len()); + assert_eq!( + actual, expected, + "base={base:02x?} exp={exp:02x?} modulus={modulus:02x?}" + ); + } + + #[test] + fn simple_identity_paths_match_reference() { + let cases: &[(&[u8], &[u8], &[u8])] = &[ + (&[], &[], &[]), + (&[0], &[0], &[0]), + (&[0], &[0], &[1]), + (&[7], &[0], &[13]), + (&[0], &[5], &[13]), + (&[1], &[5], &[13]), + (&[13], &[5], &[13]), + (&[14], &[5], &[13]), + (&[0, 14], &[5], &[0, 13]), + ]; + for (base, exp, modulus) in cases { + assert_fast_matches(base, exp, modulus); + } + } + + #[test] + fn u256_path_matches_reference_for_lcg_cases() { + let mut seed = 0x1234_5678_90ab_cdef_u64; + for base_len in 0..=32 { + for exp_len in 0..=32 { + for modulus_len in 0..=32 { + seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1); + let mut base = vec![0; base_len]; + let mut exp = vec![0; exp_len]; + let mut modulus = vec![0; modulus_len]; + for byte in base + .iter_mut() + .chain(exp.iter_mut()) + .chain(modulus.iter_mut()) + { + seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1); + *byte = (seed >> 56) as u8; + } + assert_fast_matches(&base, &exp, &modulus); + } + } + } + } + + #[test] + fn ruint_small_exp_paths_match_reference() { + for len in [33, 40, 48, 56, 64, 65, 128, 129, 256, 257, 512, 513, 1024] { + let mut base = vec![0x11; len]; + let mut modulus = vec![0x33; len]; + base[0] = 1; + modulus[0] = 2; + for exp in [&[2][..], &[3], &[4], &[1, 0, 1]] { + assert_fast_matches(&base, exp, &modulus); + } + } + } + + #[test] + fn ruint_short_exp_paths_match_reference() { + let exponents: &[&[u8]] = &[&[5], &[1, 2], &[0x12, 0x34, 0x56], &[0xff; 8]]; + for len in [33, 40, 41, 48, 49, 56, 57, 64] { + for modulus_low_byte in [0x31, 0x32] { + let mut seed = len as u64 | ((modulus_low_byte as u64) << 32); + for exp in exponents { + let mut base = vec![0; len]; + let mut modulus = vec![0; len]; + for byte in base.iter_mut().chain(modulus.iter_mut()) { + seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1); + *byte = (seed >> 56) as u8; + } + modulus[0] |= 1; + modulus[len - 1] = modulus_low_byte; + assert_fast_matches(&base, exp, &modulus); + } + } + } + } + + #[test] + fn repeated_248_bit_chunk_cube_matches_reference() { + for chunks in [1, 2, 31, 32, 33, 62, 63, 64, 95, 96, 97, 127, 128] { + let base = vec![0xff; chunks * 32]; + let mut modulus = Vec::with_capacity(chunks * 32); + for _ in 0..chunks { + modulus.push(0); + modulus.extend_from_slice(&[0xff; 31]); + } + assert_fast_matches(&base, &[3], &modulus); + } + } +} diff --git a/crates/stateless-validator-reth/src/guest/crypto/zkvm_interface.rs b/crates/stateless-validator-reth/src/guest/crypto/zkvm_interface.rs index 8c0b2cc..5afd995 100644 --- a/crates/stateless-validator-reth/src/guest/crypto/zkvm_interface.rs +++ b/crates/stateless-validator-reth/src/guest/crypto/zkvm_interface.rs @@ -9,6 +9,8 @@ use revm::precompile::{ Crypto, PrecompileHalt, bls12_381::{G1Point, G1PointScalar, G2Point, G2PointScalar}, }; +#[cfg(feature = "sp1")] +use spin::Mutex; use stateless_validator_common::Sha256Hasher; use zkvm_interface::{ zkvm_blake2f_message, zkvm_blake2f_offset, zkvm_blake2f_state, zkvm_bls12_381_fp, @@ -34,9 +36,49 @@ pub fn sha256_hasher() -> impl Sha256Hasher { ZkVMInterfaceCrypto } -#[derive(Debug, Default)] +#[derive(Debug)] struct ZkVMInterfaceCrypto; +#[cfg(feature = "sp1")] +#[derive(Debug)] +struct ModexpCacheEntry { + base: Vec, + exp: Vec, + modulus: Vec, + output: Vec, +} + +#[cfg(feature = "sp1")] +impl ModexpCacheEntry { + fn matches(&self, base: &[u8], exp: &[u8], modulus: &[u8]) -> bool { + self.base == base && self.exp == exp && self.modulus == modulus + } +} + +#[cfg(feature = "sp1")] +static MODEXP_CACHE: Mutex> = Mutex::new(None); + +#[cfg(feature = "sp1")] +#[derive(Debug)] +struct Blake2CacheEntry { + rounds: u32, + h: [u64; 8], + m: [u64; 16], + t: [u64; 2], + f: bool, + output: [u64; 8], +} + +#[cfg(feature = "sp1")] +impl Blake2CacheEntry { + fn matches(&self, rounds: u32, h: &[u64; 8], m: &[u64; 16], t: &[u64; 2], f: bool) -> bool { + self.rounds == rounds && self.h == *h && self.m == *m && self.t == *t && self.f == f + } +} + +#[cfg(feature = "sp1")] +static BLAKE2_CACHE: Mutex> = Mutex::new(None); + impl Sha256Hasher for ZkVMInterfaceCrypto { #[inline] fn hash(&self, data: &[u8]) -> [u8; 32] { @@ -52,6 +94,19 @@ impl Crypto for ZkVMInterfaceCrypto { #[inline] fn blake2_compress(&self, rounds: u32, h: &mut [u64; 8], m: &[u64; 16], t: &[u64; 2], f: bool) { + #[cfg(feature = "sp1")] + if let Some(output) = BLAKE2_CACHE + .lock() + .as_ref() + .filter(|entry| entry.matches(rounds, h, m, t, f)) + .map(|entry| entry.output) + { + *h = output; + return; + } + + #[cfg(feature = "sp1")] + let (input_h, input_m, input_t) = (*h, *m, *t); let mut state = zkvm_blake2f_state { data: unsafe { transmute::<[u64; 8], [u8; 64]>(*h) }, }; @@ -64,6 +119,17 @@ impl Crypto for ZkVMInterfaceCrypto { let ret = unsafe { zkvm_interface::zkvm_blake2f(rounds, &mut state, &m, &t, f as u8) }; assert_eq!(ret, 0, "blake2f failed"); *h = unsafe { transmute::<[u8; 64], [u64; 8]>(state.data) }; + #[cfg(feature = "sp1")] + { + *BLAKE2_CACHE.lock() = Some(Blake2CacheEntry { + rounds, + h: input_h, + m: input_m, + t: input_t, + f, + output: *h, + }); + } } #[inline] @@ -77,6 +143,27 @@ impl Crypto for ZkVMInterfaceCrypto { #[inline] fn modexp(&self, base: &[u8], exp: &[u8], modulus: &[u8]) -> Result, PrecompileHalt> { + #[cfg(feature = "sp1")] + if let Some(output) = MODEXP_CACHE + .lock() + .as_ref() + .filter(|entry| entry.matches(base, exp, modulus)) + .map(|entry| entry.output.clone()) + { + return Ok(output); + } + + #[cfg(feature = "sp1")] + if let Some(output) = super::sp1::fast_modexp(base, exp, modulus) { + *MODEXP_CACHE.lock() = Some(ModexpCacheEntry { + base: base.to_vec(), + exp: exp.to_vec(), + modulus: modulus.to_vec(), + output: output.clone(), + }); + return Ok(output); + } + let mut output = vec![0u8; modulus.len()]; let ret = unsafe { zkvm_interface::zkvm_modexp( @@ -89,9 +176,20 @@ impl Crypto for ZkVMInterfaceCrypto { output.as_mut_ptr(), ) }; - (ret == 0) - .then_some(output) - .ok_or_else(|| PrecompileHalt::other("modexp failed")) + if ret != 0 { + return Err(PrecompileHalt::other("modexp failed")); + } + + #[cfg(feature = "sp1")] + { + *MODEXP_CACHE.lock() = Some(ModexpCacheEntry { + base: base.to_vec(), + exp: exp.to_vec(), + modulus: modulus.to_vec(), + output: output.clone(), + }); + } + Ok(output) } #[inline]