1use super::*;
4use alloy_primitives::U256;
5use solar_data_structures::{index::IndexVec, map::FxHashMap};
6
7#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
9pub enum EvmIrPass {
10 None,
12 StackSchedule,
14 ColdLayout,
16 TerminalDedup,
18}
19
20impl EvmIrPass {
21 #[must_use]
23 pub const fn name(self) -> &'static str {
24 match self {
25 Self::None => "none",
26 Self::StackSchedule => "stack-schedule",
27 Self::ColdLayout => "cold-layout",
28 Self::TerminalDedup => "terminal-dedup",
29 }
30 }
31
32 pub fn run(self, module: &mut EvmIrModule) -> bool {
34 match self {
35 Self::None => false,
36 Self::StackSchedule => super::super::ir_stack_schedule::schedule_stack_ops(module),
37 Self::ColdLayout => move_cold_terminal_blocks(module),
38 Self::TerminalDedup => deduplicate_terminal_blocks(module),
39 }
40 }
41
42 #[must_use]
44 pub fn by_name(name: &str) -> Option<Self> {
45 Some(match name {
46 "none" => Self::None,
47 "stack-schedule" => Self::StackSchedule,
48 "cold-layout" => Self::ColdLayout,
49 "terminal-dedup" => Self::TerminalDedup,
50 _ => return None,
51 })
52 }
53}
54
55pub const EVM_IR_PASSES: &[EvmIrPass] =
57 &[EvmIrPass::None, EvmIrPass::StackSchedule, EvmIrPass::ColdLayout, EvmIrPass::TerminalDedup];
58
59fn move_cold_terminal_blocks(module: &mut EvmIrModule) -> bool {
60 let mut kept = Vec::with_capacity(module.blocks.len());
61 let mut moved = Vec::new();
62
63 for (block_id, block) in module.blocks.iter_enumerated() {
64 if is_movable_cold_terminal_block(module, block_id, block) {
65 moved.push(block_id);
66 } else {
67 kept.push(block_id);
68 }
69 }
70
71 if moved.is_empty() {
72 return false;
73 }
74
75 kept.extend(moved);
76 remap_block_order(module, &kept);
77 true
78}
79
80fn is_movable_cold_terminal_block(
81 module: &EvmIrModule,
82 block_id: EvmIrBlockId,
83 block: &EvmIrBlock,
84) -> bool {
85 if module.entry_block == Some(block_id) || block_id.index() == 0 {
86 return false;
87 }
88 let Some(term) = &block.terminator else {
89 return false;
90 };
91 if block.metadata.hotness != EvmIrBlockHotness::Cold || !is_evm_terminal(&term.kind) {
92 return false;
93 }
94 let previous = EvmIrBlockId::from_usize(block_id.index() - 1);
95 module.blocks[previous].terminator.as_ref().is_some_and(|term| is_layout_barrier(&term.kind))
96}
97
98fn is_layout_barrier(kind: &EvmIrTerminatorKind) -> bool {
99 matches!(kind, EvmIrTerminatorKind::Jump(_)) || is_evm_terminal(kind)
100}
101
102fn deduplicate_terminal_blocks(module: &mut EvmIrModule) -> bool {
103 let mut canonical = Vec::<(TerminalBlockKey, EvmIrBlockId)>::new();
104 let mut changed = false;
105
106 let block_ids: Vec<_> = module.blocks.indices().collect();
107 for block_id in block_ids {
108 let block = &module.blocks[block_id];
109 if !terminal_block_dedup_is_profitable(block) {
110 continue;
111 }
112 let Some(key) = terminal_block_key(block) else { continue };
113 if let Some((_, target)) = canonical.iter().find(|(known, _)| *known == key) {
114 module.blocks[block_id].instructions.clear();
115 module.blocks[block_id].terminator =
116 Some(EvmIrTerminator::new(EvmIrTerminatorKind::Jump(*target)));
117 changed = true;
118 } else {
119 canonical.push((key, block_id));
120 }
121 }
122
123 changed
124}
125
126fn terminal_block_dedup_is_profitable(block: &EvmIrBlock) -> bool {
127 let Some(term) = &block.terminator else { return false };
128 if !is_evm_terminal(&term.kind) {
129 return false;
130 }
131 let current_size = 1
135 + block.instructions.iter().map(estimated_instruction_size).sum::<usize>()
136 + estimated_terminator_size(&term.kind);
137 let replacement_size = 1 + 3 + 1;
138 current_size > replacement_size
139}
140
141fn estimated_instruction_size(inst: &EvmIrInstruction) -> usize {
142 match &inst.kind {
143 EvmIrInstructionKind::Stack(_) => 1,
144 EvmIrInstructionKind::Operation(mnemonic) if mnemonic == "push" => {
145 match inst.operands.as_slice() {
146 [operand] => estimated_push_size(operand),
147 _ => 1,
148 }
149 }
150 EvmIrInstructionKind::Operation(mnemonic) if mnemonic == "push_immutable" => 33,
151 EvmIrInstructionKind::Operation(_) => 1,
152 }
153}
154
155fn estimated_terminator_size(kind: &EvmIrTerminatorKind) -> usize {
156 let operand_pushes = |operands: &[&EvmIrOperand]| {
157 operands.iter().map(|operand| estimated_push_size(operand)).sum::<usize>() + 1
158 };
159 match kind {
160 EvmIrTerminatorKind::Return { offset, size }
161 | EvmIrTerminatorKind::Revert { offset, size } => operand_pushes(&[offset, size]),
162 EvmIrTerminatorKind::SelfDestruct { recipient } => operand_pushes(&[recipient]),
163 EvmIrTerminatorKind::Stop
164 | EvmIrTerminatorKind::Invalid
165 | EvmIrTerminatorKind::RawOpcode(_) => 1,
166 EvmIrTerminatorKind::Fallthrough(_)
167 | EvmIrTerminatorKind::Jump(_)
168 | EvmIrTerminatorKind::Branch { .. }
169 | EvmIrTerminatorKind::Switch { .. } => 0,
170 }
171}
172
173fn estimated_push_size(operand: &EvmIrOperand) -> usize {
174 match operand {
175 EvmIrOperand::Immediate(value) if *value == U256::ZERO => 1,
176 EvmIrOperand::Immediate(value) => value.byte_len() + 1,
177 EvmIrOperand::Block(_) | EvmIrOperand::Symbol(_) => 3,
178 EvmIrOperand::Value(_) => 0,
179 }
180}
181
182fn is_evm_terminal(kind: &EvmIrTerminatorKind) -> bool {
183 matches!(
184 kind,
185 EvmIrTerminatorKind::Return { .. }
186 | EvmIrTerminatorKind::Revert { .. }
187 | EvmIrTerminatorKind::Stop
188 | EvmIrTerminatorKind::Invalid
189 | EvmIrTerminatorKind::SelfDestruct { .. }
190 ) || matches!(kind, EvmIrTerminatorKind::RawOpcode(opcode) if super::super::assembler::op::is_terminal(*opcode))
191}
192
193fn terminal_block_key(block: &EvmIrBlock) -> Option<TerminalBlockKey> {
194 let mut locals = FxHashMap::default();
195 let mut instructions = Vec::with_capacity(block.instructions.len());
196
197 for inst in &block.instructions {
198 let operands =
199 inst.operands.iter().map(|operand| terminal_operand_key(operand, &locals)).collect();
200 let result = inst.result.map(|value| {
201 let index = locals.len();
202 locals.insert(value, index);
203 index
204 });
205 instructions.push(TerminalInstructionKey { result, kind: inst.kind.clone(), operands });
206 }
207
208 let term = block.terminator.as_ref()?;
209 Some(TerminalBlockKey {
210 instructions,
211 terminator: terminal_terminator_key(&term.kind, &locals),
212 })
213}
214
215fn terminal_operand_key(
216 operand: &EvmIrOperand,
217 locals: &FxHashMap<EvmIrValueId, usize>,
218) -> TerminalOperandKey {
219 match operand {
220 EvmIrOperand::Value(value) => locals
221 .get(value)
222 .copied()
223 .map(TerminalOperandKey::LocalValue)
224 .unwrap_or(TerminalOperandKey::ExternalValue(*value)),
225 EvmIrOperand::Immediate(value) => TerminalOperandKey::Immediate(*value),
226 EvmIrOperand::Block(block) => TerminalOperandKey::Block(*block),
227 EvmIrOperand::Symbol(symbol) => TerminalOperandKey::Symbol(symbol.clone()),
228 }
229}
230
231fn terminal_terminator_key(
232 kind: &EvmIrTerminatorKind,
233 locals: &FxHashMap<EvmIrValueId, usize>,
234) -> TerminalTerminatorKey {
235 match kind {
236 EvmIrTerminatorKind::Fallthrough(target) => TerminalTerminatorKey::Fallthrough(*target),
237 EvmIrTerminatorKind::Jump(target) => TerminalTerminatorKey::Jump(*target),
238 EvmIrTerminatorKind::Branch { condition, then_block, else_block } => {
239 TerminalTerminatorKey::Branch {
240 condition: terminal_operand_key(condition, locals),
241 then_block: *then_block,
242 else_block: *else_block,
243 }
244 }
245 EvmIrTerminatorKind::Switch { value, default, cases } => TerminalTerminatorKey::Switch {
246 value: terminal_operand_key(value, locals),
247 default: *default,
248 cases: cases
249 .iter()
250 .map(|(case, target)| (terminal_operand_key(case, locals), *target))
251 .collect(),
252 },
253 EvmIrTerminatorKind::Return { offset, size } => TerminalTerminatorKey::Return {
254 offset: terminal_operand_key(offset, locals),
255 size: terminal_operand_key(size, locals),
256 },
257 EvmIrTerminatorKind::Revert { offset, size } => TerminalTerminatorKey::Revert {
258 offset: terminal_operand_key(offset, locals),
259 size: terminal_operand_key(size, locals),
260 },
261 EvmIrTerminatorKind::Stop => TerminalTerminatorKey::Stop,
262 EvmIrTerminatorKind::Invalid => TerminalTerminatorKey::Invalid,
263 EvmIrTerminatorKind::SelfDestruct { recipient } => TerminalTerminatorKey::SelfDestruct {
264 recipient: terminal_operand_key(recipient, locals),
265 },
266 EvmIrTerminatorKind::RawOpcode(opcode) => TerminalTerminatorKey::RawOpcode(*opcode),
267 }
268}
269
270#[derive(Clone, Debug, PartialEq, Eq)]
271struct TerminalBlockKey {
272 instructions: Vec<TerminalInstructionKey>,
273 terminator: TerminalTerminatorKey,
274}
275
276#[derive(Clone, Debug, PartialEq, Eq)]
277struct TerminalInstructionKey {
278 result: Option<usize>,
279 kind: EvmIrInstructionKind,
280 operands: Vec<TerminalOperandKey>,
281}
282
283#[derive(Clone, Debug, PartialEq, Eq)]
284enum TerminalTerminatorKey {
285 Fallthrough(EvmIrBlockId),
286 Jump(EvmIrBlockId),
287 Branch {
288 condition: TerminalOperandKey,
289 then_block: EvmIrBlockId,
290 else_block: EvmIrBlockId,
291 },
292 Switch {
293 value: TerminalOperandKey,
294 default: EvmIrBlockId,
295 cases: Vec<(TerminalOperandKey, EvmIrBlockId)>,
296 },
297 Return {
298 offset: TerminalOperandKey,
299 size: TerminalOperandKey,
300 },
301 Revert {
302 offset: TerminalOperandKey,
303 size: TerminalOperandKey,
304 },
305 Stop,
306 Invalid,
307 SelfDestruct {
308 recipient: TerminalOperandKey,
309 },
310 RawOpcode(u8),
311}
312
313#[derive(Clone, Debug, PartialEq, Eq)]
314enum TerminalOperandKey {
315 LocalValue(usize),
316 ExternalValue(EvmIrValueId),
317 Immediate(U256),
318 Block(EvmIrBlockId),
319 Symbol(String),
320}
321
322fn remap_block_order(module: &mut EvmIrModule, order: &[EvmIrBlockId]) {
323 debug_assert_eq!(order.len(), module.blocks.len());
324 let mut remap = vec![EvmIrBlockId::from_usize(0); module.blocks.len()];
325 let mut old_blocks: Vec<Option<EvmIrBlock>> =
326 std::mem::take(&mut module.blocks).into_iter().map(Some).collect();
327 let mut blocks = IndexVec::with_capacity(old_blocks.len());
328 for &old_block in order {
329 let block = old_blocks[old_block.index()]
330 .take()
331 .expect("block order must contain each block exactly once");
332 let new_block = blocks.push(block);
333 remap[old_block.index()] = new_block;
334 }
335 debug_assert!(old_blocks.into_iter().all(|block| block.is_none()));
336 module.blocks = blocks;
337 module.entry_block = module.entry_block.map(|block| remap[block.index()]);
338 for block in &mut module.blocks {
339 for inst in &mut block.instructions {
340 for operand in &mut inst.operands {
341 remap_operand_blocks(operand, &remap);
342 }
343 }
344 if let Some(term) = &mut block.terminator {
345 remap_terminator_blocks(&mut term.kind, &remap);
346 }
347 }
348}
349
350fn remap_operand_blocks(operand: &mut EvmIrOperand, remap: &[EvmIrBlockId]) {
351 if let EvmIrOperand::Block(block) = operand {
352 *block = remap[block.index()];
353 }
354}
355
356fn remap_terminator_blocks(kind: &mut EvmIrTerminatorKind, remap: &[EvmIrBlockId]) {
357 visit_terminator_targets_mut(kind, |target| *target = remap[target.index()]);
358}
359
360fn visit_terminator_targets_mut(
361 kind: &mut EvmIrTerminatorKind,
362 mut visit: impl FnMut(&mut EvmIrBlockId),
363) {
364 match kind {
365 EvmIrTerminatorKind::Fallthrough(target) | EvmIrTerminatorKind::Jump(target) => {
366 visit(target)
367 }
368 EvmIrTerminatorKind::Branch { then_block, else_block, .. } => {
369 visit(then_block);
370 visit(else_block);
371 }
372 EvmIrTerminatorKind::Switch { default, cases, .. } => {
373 visit(default);
374 for (_, target) in cases {
375 visit(target);
376 }
377 }
378 EvmIrTerminatorKind::Return { .. }
379 | EvmIrTerminatorKind::Revert { .. }
380 | EvmIrTerminatorKind::Stop
381 | EvmIrTerminatorKind::Invalid
382 | EvmIrTerminatorKind::SelfDestruct { .. }
383 | EvmIrTerminatorKind::RawOpcode(_) => {}
384 }
385}