Skip to main content

solar_codegen/backend/evm/ir/
passes.rs

1//! EVM IR optimization and layout passes.
2
3use super::*;
4use alloy_primitives::U256;
5use solar_data_structures::{index::IndexVec, map::FxHashMap};
6
7/// A named EVM IR pass exposed to `solar evm-opt`.
8#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
9pub enum EvmIrPass {
10    /// No transform; validate and print the module.
11    None,
12    /// Materialize virtual instruction operands with physical stack operations.
13    StackSchedule,
14    /// Move cold terminal blocks after hot fallthrough code when this preserves fallthrough edges.
15    ColdLayout,
16    /// Replace duplicate terminal block bodies with jumps to the first copy when profitable.
17    TerminalDedup,
18}
19
20impl EvmIrPass {
21    /// Stable command-line pass name.
22    #[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    /// Runs this pass on an EVM IR module.
33    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    /// Looks up a pass by command-line name.
43    #[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
55/// All EVM IR passes exposed by `solar evm-opt`.
56pub 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    // A replacement block still needs `JUMPDEST PUSH2(label) JUMP`. Avoid
132    // rewriting tiny revert blocks where size is equal and revert-path gas
133    // would get worse.
134    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}