1use alloy_contract::{CallDecoder, EthCall};
4use alloy_network::Network;
5use alloy_primitives::{Address, B256, Bytes, U256, map::HashMap};
6use alloy_rpc_types::{
7 BlockOverrides,
8 state::{StateOverride, StateOverridesBuilder},
9};
10use clap::Args;
11use eyre::Result;
12use regex::Regex;
13use std::{str::FromStr, sync::LazyLock};
14
15static OVERRIDE_PATTERN: LazyLock<Regex> =
18 LazyLock::new(|| Regex::new(r"^([^:]+):([^:]+):([^:]+)$").unwrap());
19
20#[derive(Args, Clone, Debug, Default)]
22pub struct CallOverrideOpts {
23 #[arg(long = "override-balance", value_name = "ADDRESS:BALANCE", value_delimiter = ',')]
26 pub balance_overrides: Option<Vec<String>>,
27
28 #[arg(long = "override-nonce", value_name = "ADDRESS:NONCE", value_delimiter = ',')]
31 pub nonce_overrides: Option<Vec<String>>,
32
33 #[arg(long = "override-code", value_name = "ADDRESS:CODE", value_delimiter = ',')]
36 pub code_overrides: Option<Vec<String>>,
37
38 #[arg(long = "override-state", value_name = "ADDRESS:SLOT:VALUE", value_delimiter = ',')]
41 pub state_overrides: Option<Vec<String>>,
42
43 #[arg(long = "override-state-diff", value_name = "ADDRESS:SLOT:VALUE", value_delimiter = ',')]
46 pub state_diff_overrides: Option<Vec<String>>,
47
48 #[arg(long = "block.time", value_name = "TIME")]
50 pub block_time: Option<u64>,
51
52 #[arg(long = "block.number", value_name = "NUMBER")]
54 pub block_number: Option<u64>,
55}
56
57impl CallOverrideOpts {
58 pub const fn is_empty(&self) -> bool {
60 !self.has_state_overrides() && self.block_time.is_none() && self.block_number.is_none()
61 }
62
63 const fn has_state_overrides(&self) -> bool {
64 self.balance_overrides.is_some()
65 || self.nonce_overrides.is_some()
66 || self.code_overrides.is_some()
67 || self.state_overrides.is_some()
68 || self.state_diff_overrides.is_some()
69 }
70
71 pub fn apply<'a, D, N>(&self, mut call: EthCall<'a, D, N>) -> Result<EthCall<'a, D, N>>
73 where
74 D: CallDecoder,
75 N: Network,
76 {
77 if let Some(state_overrides) = self.get_state_overrides()? {
78 call = call.overrides(state_overrides);
79 }
80 if let Some(block_overrides) = self.get_block_overrides()? {
81 call = call.with_block_overrides(block_overrides);
82 }
83 Ok(call)
84 }
85
86 pub fn get_state_overrides(&self) -> Result<Option<StateOverride>> {
88 if !self.has_state_overrides() {
90 return Ok(None);
91 }
92
93 let mut builder = StateOverridesBuilder::default();
94 for (addr, balance) in address_value_overrides(&self.balance_overrides)? {
95 builder = builder.with_balance(addr.parse()?, balance.parse()?);
96 }
97 for (addr, nonce) in address_value_overrides(&self.nonce_overrides)? {
98 builder = builder.with_nonce(addr.parse()?, nonce.parse()?);
99 }
100 for (addr, code) in address_value_overrides(&self.code_overrides)? {
101 builder = builder.with_code(addr.parse()?, Bytes::from_str(code)?);
102 }
103 for (addr, entries) in address_slot_value_overrides(&self.state_overrides)? {
104 builder = builder.with_state(addr, entries);
105 }
106 for (addr, entries) in address_slot_value_overrides(&self.state_diff_overrides)? {
107 builder = builder.with_state_diff(addr, entries)
108 }
109 Ok(Some(builder.build()))
110 }
111
112 pub fn get_block_overrides(&self) -> Result<Option<BlockOverrides>> {
114 let mut overrides = BlockOverrides::default();
115 if let Some(number) = self.block_number {
116 overrides = overrides.with_number(U256::from(number));
117 }
118 if let Some(time) = self.block_time {
119 overrides = overrides.with_time(time);
120 }
121 Ok((!overrides.is_empty()).then_some(overrides))
122 }
123}
124
125fn address_value_overrides(overrides: &Option<Vec<String>>) -> Result<Vec<(&str, &str)>> {
127 overrides
128 .iter()
129 .flatten()
130 .map(|s| {
131 s.split_once(':')
132 .ok_or_else(|| eyre::eyre!("Invalid override {s}. Expected <address>:<value>"))
133 })
134 .collect()
135}
136
137fn address_slot_value_overrides(
139 overrides: &Option<Vec<String>>,
140) -> Result<HashMap<Address, HashMap<B256, B256>>> {
141 let mut parsed = HashMap::<Address, HashMap<B256, B256>>::default();
142 for s in overrides.iter().flatten() {
143 let captures = OVERRIDE_PATTERN.captures(s).ok_or_else(|| {
144 eyre::eyre!("Invalid override {s}. Expected <address>:<slot>:<value>")
145 })?;
146 let (slot, value): (U256, U256) = (captures[2].parse()?, captures[3].parse()?);
147 parsed.entry(captures[1].parse()?).or_default().insert(slot.into(), value.into());
148 }
149 Ok(parsed)
150}
151
152#[cfg(test)]
153mod tests {
154 use super::*;
155 use alloy_primitives::{address, b256};
156 use clap::Parser;
157
158 #[derive(Debug, Parser)]
159 struct TestArgs {
160 #[command(flatten)]
161 overrides: CallOverrideOpts,
162 }
163
164 #[test]
165 fn test_get_state_overrides() {
166 let args = TestArgs::parse_from([
167 "foundry-cli",
168 "--override-balance",
169 "0x0000000000000000000000000000000000000001:2",
170 "--override-nonce",
171 "0x0000000000000000000000000000000000000001:3",
172 "--override-code",
173 "0x0000000000000000000000000000000000000001:0x04",
174 "--override-state",
175 "0x0000000000000000000000000000000000000001:5:6",
176 "--override-state-diff",
177 "0x0000000000000000000000000000000000000001:7:8",
178 ]);
179 let overrides = args.overrides.get_state_overrides().unwrap().unwrap();
180 let address = address!("0x0000000000000000000000000000000000000001");
181 let account = overrides.get(&address).unwrap();
182
183 assert_eq!(account.balance, Some(U256::from(2)));
184 assert_eq!(account.nonce, Some(3));
185 assert_eq!(account.code, Some(Bytes::from([0x04])));
186 assert_eq!(
187 account
188 .state
189 .as_ref()
190 .unwrap()
191 .get(&b256!("0x0000000000000000000000000000000000000000000000000000000000000005")),
192 Some(&b256!("0x0000000000000000000000000000000000000000000000000000000000000006"))
193 );
194 assert_eq!(
195 account
196 .state_diff
197 .as_ref()
198 .unwrap()
199 .get(&b256!("0x0000000000000000000000000000000000000000000000000000000000000007")),
200 Some(&b256!("0x0000000000000000000000000000000000000000000000000000000000000008"))
201 );
202 }
203
204 #[test]
205 fn test_get_state_overrides_empty() {
206 let args = TestArgs::parse_from([""]);
207 assert_eq!(args.overrides.get_state_overrides().unwrap(), None);
208 }
209
210 #[test]
211 fn test_invalid_overrides() {
212 let args = TestArgs::parse_from(["foundry-cli", "--override-balance", "invalid_value"]);
213 assert_eq!(
214 args.overrides.get_state_overrides().unwrap_err().to_string(),
215 "Invalid override invalid_value. Expected <address>:<value>"
216 );
217 let args = TestArgs::parse_from(["foundry-cli", "--override-state", "invalid_value"]);
218 assert_eq!(
219 args.overrides.get_state_overrides().unwrap_err().to_string(),
220 "Invalid override invalid_value. Expected <address>:<slot>:<value>"
221 );
222 }
223
224 #[test]
225 fn test_get_block_overrides() {
226 let args =
227 TestArgs::parse_from(["foundry-cli", "--block.number", "1", "--block.time", "2"]);
228 let overrides = args.overrides.get_block_overrides().unwrap().unwrap();
229 assert_eq!(overrides.number, Some(U256::from(1)));
230 assert_eq!(overrides.time, Some(2));
231
232 let args = TestArgs::parse_from([""]);
233 assert_eq!(args.overrides.get_block_overrides().unwrap(), None);
234 }
235}