1use crate::{Cheatcode, CheatsCtxt, Result, Vm::*, evm::journaled_account};
2use alloy_primitives::Address;
3use foundry_evm_core::evm::FoundryEvmNetwork;
4
5#[derive(Clone, Copy, Debug, Default)]
7pub struct Prank {
8 pub prank_caller: Address,
10 pub prank_origin: Address,
12 pub new_caller: Address,
14 pub new_origin: Option<Address>,
16 pub depth: usize,
18 pub single_call: bool,
20 pub delegate_call: bool,
22 pub used: bool,
24}
25
26impl Prank {
27 pub const fn new(
29 prank_caller: Address,
30 prank_origin: Address,
31 new_caller: Address,
32 new_origin: Option<Address>,
33 depth: usize,
34 single_call: bool,
35 delegate_call: bool,
36 ) -> Self {
37 Self {
38 prank_caller,
39 prank_origin,
40 new_caller,
41 new_origin,
42 depth,
43 single_call,
44 delegate_call,
45 used: false,
46 }
47 }
48
49 pub const fn first_time_applied(&self) -> Option<Self> {
52 if self.used { None } else { Some(Self { used: true, ..*self }) }
53 }
54
55 pub(crate) fn changes_for(&self, depth: usize, caller: Address) -> Option<PrankChanges> {
57 if depth < self.depth || caller != self.prank_caller {
58 return None;
59 }
60
61 let new_caller = (depth == self.depth).then_some(self.new_caller);
63
64 let applied = new_caller.is_some() || self.new_origin.is_some();
66
67 Some(PrankChanges {
68 caller: new_caller,
69 origin: self.new_origin,
70 used: if applied { self.first_time_applied() } else { None },
71 })
72 }
73}
74
75#[derive(Clone, Copy, Debug)]
77pub(crate) struct PrankChanges {
78 pub(crate) caller: Option<Address>,
80 pub(crate) origin: Option<Address>,
82 pub(crate) used: Option<Prank>,
84}
85
86impl Cheatcode for prank_0Call {
87 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
88 let Self { msgSender } = self;
89 prank(ccx, msgSender, None, true, false)
90 }
91}
92
93impl Cheatcode for startPrank_0Call {
94 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
95 let Self { msgSender } = self;
96 prank(ccx, msgSender, None, false, false)
97 }
98}
99
100impl Cheatcode for prank_1Call {
101 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
102 let Self { msgSender, txOrigin } = self;
103 prank(ccx, msgSender, Some(txOrigin), true, false)
104 }
105}
106
107impl Cheatcode for startPrank_1Call {
108 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
109 let Self { msgSender, txOrigin } = self;
110 prank(ccx, msgSender, Some(txOrigin), false, false)
111 }
112}
113
114impl Cheatcode for prank_2Call {
115 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
116 let Self { msgSender, delegateCall } = self;
117 prank(ccx, msgSender, None, true, *delegateCall)
118 }
119}
120
121impl Cheatcode for startPrank_2Call {
122 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
123 let Self { msgSender, delegateCall } = self;
124 prank(ccx, msgSender, None, false, *delegateCall)
125 }
126}
127
128impl Cheatcode for prank_3Call {
129 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
130 let Self { msgSender, txOrigin, delegateCall } = self;
131 prank(ccx, msgSender, Some(txOrigin), true, *delegateCall)
132 }
133}
134
135impl Cheatcode for startPrank_3Call {
136 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
137 let Self { msgSender, txOrigin, delegateCall } = self;
138 prank(ccx, msgSender, Some(txOrigin), false, *delegateCall)
139 }
140}
141
142impl Cheatcode for stopPrankCall {
143 fn apply_stateful<FEN: FoundryEvmNetwork>(&self, ccx: &mut CheatsCtxt<'_, '_, FEN>) -> Result {
144 let Self {} = self;
145 ccx.state.pranks.remove(&ccx.depth());
146 Ok(Default::default())
147 }
148}
149
150fn prank<FEN: FoundryEvmNetwork>(
151 ccx: &mut CheatsCtxt<'_, '_, FEN>,
152 new_caller: &Address,
153 new_origin: Option<&Address>,
154 single_call: bool,
155 delegate_call: bool,
156) -> Result {
157 let account = journaled_account(ccx.ecx, *new_caller)?;
161
162 if delegate_call {
164 ensure!(
165 account.info.code.as_ref().is_some_and(|code| !code.is_empty()),
166 "cannot `prank` delegate call from an EOA"
167 );
168 }
169
170 let depth = ccx.depth();
171 if let Some(Prank { used, single_call: current_single_call, .. }) = ccx.state.get_prank(depth) {
172 ensure!(used, "cannot overwrite a prank until it is applied at least once");
173 ensure!(
176 single_call == *current_single_call,
177 "cannot override an ongoing prank with a single vm.prank; \
178 use vm.startPrank to override the current prank"
179 );
180 }
181
182 let prank = Prank::new(
183 ccx.caller,
184 ccx.tx_caller(),
185 *new_caller,
186 new_origin.copied(),
187 depth,
188 single_call,
189 delegate_call,
190 );
191
192 ensure!(
193 ccx.state.broadcast.is_none(),
194 "cannot `prank` for a broadcasted transaction; \
195 pass the desired `tx.origin` into the `broadcast` cheatcode call"
196 );
197
198 ccx.state.pranks.insert(prank.depth, prank);
199 Ok(Default::default())
200}
201
202#[cfg(test)]
203mod tests {
204 use super::*;
205
206 const CALLER: Address = Address::repeat_byte(0x01);
207 const SENDER: Address = Address::repeat_byte(0x02);
208 const ORIGIN: Address = Address::repeat_byte(0x03);
209
210 #[test]
211 fn changes_for_applies_the_sender_only_at_the_prank_depth() {
212 let prank = Prank::new(CALLER, Address::ZERO, SENDER, None, 2, false, false);
213 assert!(prank.changes_for(1, CALLER).is_none());
214 assert!(prank.changes_for(2, ORIGIN).is_none());
215
216 let at_depth = prank.changes_for(2, CALLER).unwrap();
217 assert_eq!((at_depth.caller, at_depth.origin), (Some(SENDER), None));
218 let used = at_depth.used.unwrap();
219 assert!(used.used);
220 assert!(used.changes_for(2, CALLER).unwrap().used.is_none());
221
222 let deeper = prank.changes_for(3, CALLER).unwrap();
223 assert_eq!((deeper.caller, deeper.origin), (None, None));
224 assert!(deeper.used.is_none());
225
226 let with_origin = Prank::new(CALLER, Address::ZERO, SENDER, Some(ORIGIN), 2, false, false);
227 let deeper = with_origin.changes_for(3, CALLER).unwrap();
228 assert_eq!((deeper.caller, deeper.origin), (None, Some(ORIGIN)));
229 assert!(deeper.used.is_some());
230 }
231}