Skip to main content

solar_codegen/transform/
memory_dse.rs

1//! Local dead memory optimization.
2//!
3//! This pass removes full-word `mstore` instructions that are overwritten by a
4//! later full-word `mstore` to the same exact address within the same basic
5//! block, before any operation can observe memory or gas. It also forwards
6//! same-block `mload` instructions from the latest exact-address `mstore` when
7//! no intervening operation can mutate memory.
8
9use crate::{
10    analysis::CfgInfo,
11    mir::{
12        BlockId, Function, Immediate, InstId, InstKind, Terminator, Value, ValueId,
13        utils as mir_utils,
14    },
15    pass::FunctionPass,
16};
17use alloy_primitives::{U256, keccak256};
18use solar_data_structures::map::{FxHashMap, FxHashSet};
19
20/// Local dead memory optimization pass.
21#[derive(Debug, Default)]
22pub struct MemoryStoreEliminator {
23    /// Number of memory instructions eliminated.
24    pub eliminated_count: usize,
25}
26
27/// Function pass for local dead memory-store elimination.
28pub struct MemoryDsePass;
29
30impl FunctionPass for MemoryDsePass {
31    fn name(&self) -> &str {
32        "memory-dse"
33    }
34
35    fn run_on_function(&mut self, func: &mut Function) -> bool {
36        MemoryStoreEliminator::new().run_to_fixpoint(func) != 0
37    }
38}
39
40#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
41enum MemAddrKey {
42    Const(u64),
43    BaseOffset { base: ValueId, offset: u64 },
44}
45
46#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
47struct ImmutableCopyKey {
48    len: u64,
49    offset: u64,
50}
51
52#[derive(Clone, Copy, Debug)]
53struct CachedImmutableCopy {
54    block: BlockId,
55    index: usize,
56    value: ValueId,
57}
58
59impl MemoryStoreEliminator {
60    /// Creates a new memory optimization pass.
61    pub fn new() -> Self {
62        Self::default()
63    }
64
65    /// Runs local memory optimization on a function.
66    pub fn run(&mut self, func: &mut Function) -> usize {
67        self.eliminated_count = 0;
68
69        self.reuse_redundant_immutable_copies(func);
70        self.remove_unused_internal_frame_stores(func);
71
72        let block_ids: Vec<BlockId> = func.blocks.indices().collect();
73        for block_id in block_ids {
74            self.process_block(func, block_id);
75        }
76        self.remove_cross_block_overwrites(func);
77
78        self.eliminated_count
79    }
80
81    /// Runs local memory optimization until no more instructions can be eliminated.
82    pub fn run_to_fixpoint(&mut self, func: &mut Function) -> usize {
83        let mut total = 0;
84        loop {
85            let eliminated = self.run(func);
86            if eliminated == 0 {
87                break;
88            }
89            total += eliminated;
90        }
91        total
92    }
93
94    fn reuse_redundant_immutable_copies(&mut self, func: &mut Function) {
95        let inst_results = func.inst_results();
96        let cfg = CfgInfo::new(func);
97        let mut cached: FxHashMap<ImmutableCopyKey, CachedImmutableCopy> = FxHashMap::default();
98        let mut replacements = FxHashMap::default();
99        let mut dead = FxHashSet::default();
100
101        for (block_id, block) in func.blocks.iter_enumerated() {
102            let insts = block.instructions.clone();
103            for (index, window) in insts.windows(2).enumerate() {
104                let codecopy = window[0];
105                let load = window[1];
106                let InstKind::CodeCopy(dest, src, size) = func.instructions[codecopy].kind else {
107                    continue;
108                };
109                if func.value_u64(size) != Some(32) {
110                    continue;
111                }
112                let Some(key) = Self::immutable_copy_key(func, src) else {
113                    continue;
114                };
115                let InstKind::MLoad(load_addr) = func.instructions[load].kind else {
116                    continue;
117                };
118                if Self::mem_addr_key(func, dest) != Self::mem_addr_key(func, load_addr) {
119                    continue;
120                }
121                let Some(&loaded_value) = inst_results.get(&load) else {
122                    continue;
123                };
124
125                if let Some(cached_copy) = cached.get(&key).copied()
126                    && Self::copy_dominates(cfg.dominators(), cached_copy, block_id, index)
127                {
128                    replacements.insert(loaded_value, cached_copy.value);
129                    func.instructions[codecopy].kind = InstKind::MStore(dest, cached_copy.value);
130                    dead.insert(load);
131                    self.eliminated_count += 1;
132                    continue;
133                }
134
135                cached.insert(
136                    key,
137                    CachedImmutableCopy { block: block_id, index, value: loaded_value },
138                );
139            }
140        }
141
142        if replacements.is_empty() && dead.is_empty() {
143            return;
144        }
145
146        func.replace_uses_canonicalized(&replacements);
147        for block in func.blocks.iter_mut() {
148            block.instructions.retain(|id| !dead.contains(id));
149        }
150    }
151
152    fn remove_unused_internal_frame_stores(&mut self, func: &mut Function) {
153        if Self::has_frame_observer(func) {
154            return;
155        }
156
157        let Some(reads) = Self::internal_frame_read_ranges(func) else {
158            return;
159        };
160        let mut dead = FxHashSet::default();
161
162        for block in func.blocks.iter() {
163            for &inst_id in &block.instructions {
164                let InstKind::MStore(addr, _) = func.instructions[inst_id].kind else {
165                    continue;
166                };
167                let Some(offset) = Self::internal_frame_offset(func, addr) else {
168                    continue;
169                };
170                if !reads.iter().any(|&(read_offset, read_size)| {
171                    mir_utils::ranges_overlap(offset, 32, read_offset, read_size)
172                }) {
173                    dead.insert(inst_id);
174                }
175            }
176        }
177
178        if dead.is_empty() {
179            return;
180        }
181
182        self.eliminated_count += dead.len();
183        for block in func.blocks.iter_mut() {
184            block.instructions.retain(|id| !dead.contains(id));
185        }
186    }
187
188    fn process_block(&mut self, func: &mut Function, block_id: BlockId) {
189        self.fold_constant_keccak(func, block_id);
190        self.forward_loads(func, block_id);
191        self.remove_equal_stores(func, block_id);
192
193        let inst_ids = func.blocks[block_id].instructions.clone();
194        let mut overwritten: FxHashSet<MemAddrKey> = FxHashSet::default();
195        let mut dead: FxHashSet<InstId> = FxHashSet::default();
196
197        for &inst_id in inst_ids.iter().rev() {
198            let inst = &func.instructions[inst_id];
199            match &inst.kind {
200                InstKind::MStore(addr, _) => {
201                    if let Some(key) = Self::mem_addr_key(func, *addr) {
202                        if overwritten.contains(&key) {
203                            dead.insert(inst_id);
204                            self.eliminated_count += 1;
205                        } else {
206                            overwritten.insert(key);
207                        }
208                    } else {
209                        overwritten.clear();
210                    }
211                }
212                InstKind::MLoad(addr) => {
213                    if let Some(key) = Self::mem_addr_key(func, *addr) {
214                        Self::remove_overlapping_set(&mut overwritten, key);
215                    } else {
216                        overwritten.clear();
217                    }
218                }
219                InstKind::CalldataCopy(dest, _, size)
220                | InstKind::CodeCopy(dest, _, size)
221                | InstKind::ReturnDataCopy(dest, _, size) => {
222                    Self::insert_or_clear_full_word_overwritten_range(
223                        func,
224                        &mut overwritten,
225                        *dest,
226                        *size,
227                    );
228                }
229                InstKind::ExtCodeCopy(_, dest, _, size) => {
230                    Self::insert_or_clear_full_word_overwritten_range(
231                        func,
232                        &mut overwritten,
233                        *dest,
234                        *size,
235                    );
236                }
237                kind if Self::is_memory_or_gas_observer(kind) => {
238                    overwritten.clear();
239                }
240                _ => {}
241            }
242        }
243
244        if dead.is_empty() {
245            return;
246        }
247
248        func.blocks[block_id].instructions.retain(|id| !dead.contains(id));
249    }
250
251    fn remove_cross_block_overwrites(&mut self, func: &mut Function) {
252        let mut dead = FxHashSet::default();
253
254        for pred in func.blocks.indices() {
255            let Some(succ) = Self::single_jump_successor(func, pred) else {
256                continue;
257            };
258            if func.blocks[succ].predecessors.as_slice() != [pred] {
259                continue;
260            }
261
262            let Some((store, pred_key)) = self.last_cross_block_store_candidate(func, pred) else {
263                continue;
264            };
265            let Some(succ_key) = self.first_cross_block_overwrite(func, succ) else {
266                continue;
267            };
268            if pred_key == succ_key {
269                dead.insert(store);
270            }
271        }
272
273        if dead.is_empty() {
274            return;
275        }
276
277        self.eliminated_count += dead.len();
278        for block in func.blocks.iter_mut() {
279            block.instructions.retain(|id| !dead.contains(id));
280        }
281    }
282
283    fn single_jump_successor(func: &Function, block: BlockId) -> Option<BlockId> {
284        let Some(Terminator::Jump(target)) = func.blocks[block].terminator.as_ref() else {
285            return None;
286        };
287        Some(*target)
288    }
289
290    fn last_cross_block_store_candidate(
291        &self,
292        func: &Function,
293        block: BlockId,
294    ) -> Option<(InstId, MemAddrKey)> {
295        for &inst_id in func.blocks[block].instructions.iter().rev() {
296            match func.instructions[inst_id].kind {
297                InstKind::MStore(addr, _) => {
298                    let key = Self::mem_addr_key(func, addr)?;
299                    return Some((inst_id, key));
300                }
301                ref kind if Self::cross_block_memory_barrier(kind) => return None,
302                _ => {}
303            }
304        }
305        None
306    }
307
308    fn first_cross_block_overwrite(&self, func: &Function, block: BlockId) -> Option<MemAddrKey> {
309        for &inst_id in &func.blocks[block].instructions {
310            match func.instructions[inst_id].kind {
311                InstKind::MStore(addr, _) => return Self::mem_addr_key(func, addr),
312                ref kind if Self::cross_block_memory_barrier(kind) => return None,
313                _ => {}
314            }
315        }
316        None
317    }
318
319    fn fold_constant_keccak(&mut self, func: &mut Function, block_id: BlockId) {
320        let inst_ids = func.blocks[block_id].instructions.clone();
321        let inst_results = func.inst_results();
322        let mut stored_words: FxHashMap<MemAddrKey, U256> = FxHashMap::default();
323        let mut replacements: FxHashMap<ValueId, ValueId> = FxHashMap::default();
324        let mut dead: FxHashSet<InstId> = FxHashSet::default();
325
326        for &inst_id in &inst_ids {
327            match &func.instructions[inst_id].kind {
328                InstKind::MStore(addr, value) => {
329                    let Some(key) = Self::mem_addr_key(func, *addr) else {
330                        stored_words.clear();
331                        continue;
332                    };
333                    Self::remove_overlapping_map(&mut stored_words, key);
334                    if let Some(value) = func.value_u256(*value) {
335                        stored_words.insert(key, value);
336                    }
337                }
338                InstKind::Keccak256(offset, size) => {
339                    let Some(bytes) =
340                        Self::constant_memory_bytes(func, &stored_words, *offset, *size)
341                    else {
342                        continue;
343                    };
344                    let Some(&result) = inst_results.get(&inst_id) else {
345                        continue;
346                    };
347                    let hash = keccak256(&bytes);
348                    let replacement = func.alloc_value(Value::Immediate(Immediate::uint256(
349                        U256::from_be_bytes(hash.0),
350                    )));
351                    replacements.insert(result, replacement);
352                    dead.insert(inst_id);
353                    self.eliminated_count += 1;
354                }
355                kind if Self::can_mutate_memory(kind) => {
356                    stored_words.clear();
357                }
358                _ => {}
359            }
360        }
361
362        if dead.is_empty() {
363            return;
364        }
365
366        func.replace_uses_canonicalized(&replacements);
367        func.blocks[block_id].instructions.retain(|id| !dead.contains(id));
368    }
369
370    fn remove_equal_stores(&mut self, func: &mut Function, block_id: BlockId) {
371        let inst_ids = func.blocks[block_id].instructions.clone();
372        let mut stored_values: FxHashMap<MemAddrKey, ValueId> = FxHashMap::default();
373        let mut dead: FxHashSet<InstId> = FxHashSet::default();
374
375        for &inst_id in &inst_ids {
376            let inst = &func.instructions[inst_id];
377            match &inst.kind {
378                InstKind::MStore(addr, value) => {
379                    let Some(key) = Self::mem_addr_key(func, *addr) else {
380                        stored_values.clear();
381                        continue;
382                    };
383
384                    if stored_values.get(&key).is_some_and(|&stored| stored == *value) {
385                        dead.insert(inst_id);
386                        self.eliminated_count += 1;
387                        continue;
388                    }
389
390                    Self::remove_overlapping_map(&mut stored_values, key);
391                    stored_values.insert(key, *value);
392                }
393                kind if Self::can_mutate_memory(kind) => {
394                    stored_values.clear();
395                }
396                _ => {}
397            }
398        }
399
400        if dead.is_empty() {
401            return;
402        }
403
404        func.blocks[block_id].instructions.retain(|id| !dead.contains(id));
405    }
406
407    fn forward_loads(&mut self, func: &mut Function, block_id: BlockId) {
408        let inst_ids = func.blocks[block_id].instructions.clone();
409        let inst_results = func.inst_results();
410        let mut stored_values: FxHashMap<MemAddrKey, ValueId> = FxHashMap::default();
411        let mut replacements: FxHashMap<ValueId, ValueId> = FxHashMap::default();
412        let mut dead: FxHashSet<InstId> = FxHashSet::default();
413
414        for &inst_id in &inst_ids {
415            let inst = &func.instructions[inst_id];
416            match &inst.kind {
417                InstKind::MStore(addr, value) => {
418                    if let Some(key) = Self::mem_addr_key(func, *addr) {
419                        if !Self::remove_overlapping_write_range(
420                            func,
421                            &mut stored_values,
422                            *addr,
423                            32,
424                        ) {
425                            stored_values.clear();
426                            continue;
427                        }
428                        stored_values
429                            .insert(key, mir_utils::resolve_replacement(*value, &replacements));
430                    } else {
431                        stored_values.clear();
432                    }
433                }
434                InstKind::MLoad(addr) => {
435                    let Some(key) = Self::mem_addr_key(func, *addr) else {
436                        continue;
437                    };
438                    let Some(&stored_value) = stored_values.get(&key) else {
439                        continue;
440                    };
441                    if let Some(&loaded_value) = inst_results.get(&inst_id) {
442                        replacements.insert(loaded_value, stored_value);
443                        dead.insert(inst_id);
444                    }
445                }
446                InstKind::MStore8(addr, _)
447                    if !Self::remove_overlapping_write_range(
448                        func,
449                        &mut stored_values,
450                        *addr,
451                        1,
452                    ) =>
453                {
454                    stored_values.clear();
455                }
456                InstKind::CalldataCopy(dest, _, size)
457                | InstKind::CodeCopy(dest, _, size)
458                | InstKind::ReturnDataCopy(dest, _, size) => {
459                    let Some(size) = func.value_u64(*size) else {
460                        stored_values.clear();
461                        continue;
462                    };
463                    if !Self::remove_overlapping_write_range(func, &mut stored_values, *dest, size)
464                    {
465                        stored_values.clear();
466                    }
467                }
468                InstKind::ExtCodeCopy(_, dest, _, size) => {
469                    let Some(size) = func.value_u64(*size) else {
470                        stored_values.clear();
471                        continue;
472                    };
473                    if !Self::remove_overlapping_write_range(func, &mut stored_values, *dest, size)
474                    {
475                        stored_values.clear();
476                    }
477                }
478                kind if Self::can_mutate_memory(kind) => {
479                    stored_values.clear();
480                }
481                _ => {}
482            }
483        }
484
485        if dead.is_empty() {
486            return;
487        }
488
489        func.replace_uses_canonicalized(&replacements);
490        self.eliminated_count += dead.len();
491        func.blocks[block_id].instructions.retain(|id| !dead.contains(id));
492    }
493
494    fn mem_addr_key(func: &Function, value: ValueId) -> Option<MemAddrKey> {
495        Self::mem_addr_key_with_depth(func, value, 0)
496    }
497
498    fn mem_addr_key_with_depth(
499        func: &Function,
500        value: ValueId,
501        depth: usize,
502    ) -> Option<MemAddrKey> {
503        if depth > 8 {
504            return Some(MemAddrKey::BaseOffset { base: value, offset: 0 });
505        }
506
507        match &func.values[value] {
508            Value::Immediate(imm) => {
509                let addr = imm.as_u256()?;
510                u64::try_from(addr).ok().map(MemAddrKey::Const)
511            }
512            Value::Inst(inst_id) => match func.instructions[*inst_id].kind {
513                InstKind::Add(a, b) => Self::add_addr_offset(func, a, b, depth)
514                    .or_else(|| Self::add_addr_offset(func, b, a, depth))
515                    .or(Some(MemAddrKey::BaseOffset { base: value, offset: 0 })),
516                _ => Some(MemAddrKey::BaseOffset { base: value, offset: 0 }),
517            },
518            Value::Arg { .. } => Some(MemAddrKey::BaseOffset { base: value, offset: 0 }),
519            Value::Undef(_) | Value::Error(_) => None,
520        }
521    }
522
523    fn add_addr_offset(
524        func: &Function,
525        base: ValueId,
526        offset: ValueId,
527        depth: usize,
528    ) -> Option<MemAddrKey> {
529        let offset = func.value_u64(offset)?;
530        match Self::mem_addr_key_with_depth(func, base, depth + 1)? {
531            MemAddrKey::Const(addr) => addr.checked_add(offset).map(MemAddrKey::Const),
532            MemAddrKey::BaseOffset { base, offset: base_offset } => base_offset
533                .checked_add(offset)
534                .map(|offset| MemAddrKey::BaseOffset { base, offset }),
535        }
536    }
537
538    fn immutable_copy_key(func: &Function, src: ValueId) -> Option<ImmutableCopyKey> {
539        match func.values[src] {
540            Value::Inst(inst_id) => match func.instructions[inst_id].kind {
541                InstKind::Sub(code_size, len) if Self::is_codesize(func, code_size) => {
542                    Some(ImmutableCopyKey { len: func.value_u64(len)?, offset: 0 })
543                }
544                InstKind::Add(base, offset) => {
545                    Self::immutable_copy_key_with_offset(func, base, offset)
546                        .or_else(|| Self::immutable_copy_key_with_offset(func, offset, base))
547                }
548                _ => None,
549            },
550            _ => None,
551        }
552    }
553
554    fn immutable_copy_key_with_offset(
555        func: &Function,
556        base: ValueId,
557        offset: ValueId,
558    ) -> Option<ImmutableCopyKey> {
559        let mut key = Self::immutable_copy_key(func, base)?;
560        key.offset = key.offset.checked_add(func.value_u64(offset)?)?;
561        Some(key)
562    }
563
564    fn is_codesize(func: &Function, value: ValueId) -> bool {
565        matches!(func.values[value], Value::Inst(inst_id) if matches!(func.instructions[inst_id].kind, InstKind::CodeSize))
566    }
567
568    fn copy_dominates(
569        dominators: &crate::analysis::DominatorTree,
570        cached: CachedImmutableCopy,
571        block: BlockId,
572        index: usize,
573    ) -> bool {
574        if cached.block == block {
575            return cached.index < index;
576        }
577        dominators.dominates(cached.block, block)
578    }
579
580    fn constant_memory_bytes(
581        func: &Function,
582        stored_words: &FxHashMap<MemAddrKey, U256>,
583        offset: ValueId,
584        size: ValueId,
585    ) -> Option<Vec<u8>> {
586        let offset = func.value_u64(offset)?;
587        let size = func.value_u64(size)?;
588        if size > 4096 || size % 32 != 0 {
589            return None;
590        }
591
592        let mut bytes = Vec::with_capacity(size as usize);
593        for word_offset in (0..size).step_by(32) {
594            let addr = offset.checked_add(word_offset)?;
595            let word = stored_words.get(&MemAddrKey::Const(addr))?;
596            bytes.extend_from_slice(&word.to_be_bytes::<32>());
597        }
598        Some(bytes)
599    }
600
601    fn overlaps(a: MemAddrKey, b: MemAddrKey) -> bool {
602        match (a, b) {
603            (MemAddrKey::Const(a), MemAddrKey::Const(b)) => mir_utils::ranges_overlap(a, 32, b, 32),
604            (
605                MemAddrKey::BaseOffset { base: a_base, offset: a_offset },
606                MemAddrKey::BaseOffset { base: b_base, offset: b_offset },
607            ) if a_base == b_base => mir_utils::ranges_overlap(a_offset, 32, b_offset, 32),
608            _ => true,
609        }
610    }
611
612    fn remove_overlapping_map<T>(map: &mut FxHashMap<MemAddrKey, T>, key: MemAddrKey) {
613        map.retain(|&stored, _| !Self::overlaps(stored, key));
614    }
615
616    fn remove_overlapping_set(set: &mut FxHashSet<MemAddrKey>, key: MemAddrKey) {
617        set.retain(|&stored| !Self::overlaps(stored, key));
618    }
619
620    fn remove_overlapping_write_range<T>(
621        func: &Function,
622        map: &mut FxHashMap<MemAddrKey, T>,
623        dest: ValueId,
624        size: u64,
625    ) -> bool {
626        let Some(write) = Self::mem_addr_key(func, dest) else {
627            return false;
628        };
629        map.retain(|&stored, _| !Self::ranges_overlap_mem_keys(func, stored, 32, write, size));
630        true
631    }
632
633    fn insert_full_word_overwritten_range(
634        func: &Function,
635        overwritten: &mut FxHashSet<MemAddrKey>,
636        dest: ValueId,
637        size: ValueId,
638    ) -> bool {
639        let Some(size) = func.value_u64(size) else {
640            return false;
641        };
642        if size % 32 != 0 || size > 4096 {
643            return false;
644        }
645
646        let Some(base) = Self::mem_addr_key(func, dest) else {
647            return false;
648        };
649        for offset in (0..size).step_by(32) {
650            let Some(key) = Self::offset_mem_addr_key(base, offset) else {
651                return false;
652            };
653            overwritten.insert(key);
654        }
655        true
656    }
657
658    fn insert_or_clear_full_word_overwritten_range(
659        func: &Function,
660        overwritten: &mut FxHashSet<MemAddrKey>,
661        dest: ValueId,
662        size: ValueId,
663    ) {
664        if !Self::insert_full_word_overwritten_range(func, overwritten, dest, size) {
665            overwritten.clear();
666        }
667    }
668
669    fn offset_mem_addr_key(key: MemAddrKey, add: u64) -> Option<MemAddrKey> {
670        match key {
671            MemAddrKey::Const(offset) => offset.checked_add(add).map(MemAddrKey::Const),
672            MemAddrKey::BaseOffset { base, offset } => {
673                offset.checked_add(add).map(|offset| MemAddrKey::BaseOffset { base, offset })
674            }
675        }
676    }
677
678    fn ranges_overlap_mem_keys(
679        func: &Function,
680        read: MemAddrKey,
681        read_size: u64,
682        write: MemAddrKey,
683        write_size: u64,
684    ) -> bool {
685        if (Self::is_scratch_const(read) && Self::is_fmp_heap_key(func, write))
686            || (Self::is_fmp_heap_key(func, read) && Self::is_scratch_const(write))
687        {
688            return false;
689        }
690
691        match (read, write) {
692            (MemAddrKey::Const(read), MemAddrKey::Const(write)) => {
693                mir_utils::ranges_overlap(read, read_size, write, write_size)
694            }
695            (
696                MemAddrKey::BaseOffset { base: read_base, offset: read },
697                MemAddrKey::BaseOffset { base: write_base, offset: write },
698            ) if read_base == write_base => {
699                mir_utils::ranges_overlap(read, read_size, write, write_size)
700            }
701            _ => true,
702        }
703    }
704
705    fn is_scratch_const(key: MemAddrKey) -> bool {
706        matches!(key, MemAddrKey::Const(offset) if offset < 128)
707    }
708
709    fn is_fmp_heap_key(func: &Function, key: MemAddrKey) -> bool {
710        let MemAddrKey::BaseOffset { base, .. } = key else {
711            return false;
712        };
713        Self::is_fmp_heap_value(func, base, 0)
714    }
715
716    fn is_fmp_heap_value(func: &Function, value: ValueId, depth: usize) -> bool {
717        if depth > 8 {
718            return false;
719        }
720        let Value::Inst(inst_id) = func.values[value] else {
721            return false;
722        };
723        match func.instructions[inst_id].kind {
724            InstKind::MLoad(addr) => func.value_u64(addr) == Some(0x40),
725            InstKind::Add(a, b) => {
726                Self::is_fmp_heap_value(func, a, depth + 1)
727                    || Self::is_fmp_heap_value(func, b, depth + 1)
728            }
729            _ => false,
730        }
731    }
732
733    fn is_memory_or_gas_observer(kind: &InstKind) -> bool {
734        matches!(
735            kind,
736            InstKind::MStore8(_, _)
737                | InstKind::MSize
738                | InstKind::MCopy(_, _, _)
739                | InstKind::CalldataCopy(_, _, _)
740                | InstKind::CodeCopy(_, _, _)
741                | InstKind::ReturnDataCopy(_, _, _)
742                | InstKind::ExtCodeCopy(_, _, _, _)
743                | InstKind::Keccak256(_, _)
744                | InstKind::Call { .. }
745                | InstKind::StaticCall { .. }
746                | InstKind::DelegateCall { .. }
747                | InstKind::InternalCall { .. }
748                | InstKind::Create(_, _, _)
749                | InstKind::Create2(_, _, _, _)
750                | InstKind::Log0(_, _)
751                | InstKind::Log1(_, _, _)
752                | InstKind::Log2(_, _, _, _)
753                | InstKind::Log3(_, _, _, _, _)
754                | InstKind::Log4(_, _, _, _, _, _)
755                | InstKind::Gas
756        )
757    }
758
759    fn has_frame_observer(func: &Function) -> bool {
760        func.blocks.iter().any(|block| {
761            block.instructions.iter().any(|&inst_id| {
762                matches!(
763                    func.instructions[inst_id].kind,
764                    InstKind::Gas | InstKind::MSize | InstKind::InternalCall { .. }
765                )
766            })
767        })
768    }
769
770    fn internal_frame_read_ranges(func: &Function) -> Option<Vec<(u64, u64)>> {
771        let mut reads = Vec::new();
772
773        for block in func.blocks.iter() {
774            for &inst_id in &block.instructions {
775                match func.instructions[inst_id].kind {
776                    InstKind::MLoad(addr) => {
777                        if let Some(offset) = Self::internal_frame_offset(func, addr) {
778                            reads.push((offset, 32));
779                        }
780                    }
781                    InstKind::Keccak256(offset, size)
782                    | InstKind::Log0(offset, size)
783                    | InstKind::Create(_, offset, size) => {
784                        Self::push_frame_read(func, &mut reads, offset, size)?;
785                    }
786                    InstKind::Log1(offset, size, _) => {
787                        Self::push_frame_read(func, &mut reads, offset, size)?;
788                    }
789                    InstKind::Log2(offset, size, _, _)
790                    | InstKind::MCopy(_, offset, size)
791                    | InstKind::Create2(_, offset, size, _) => {
792                        Self::push_frame_read(func, &mut reads, offset, size)?;
793                    }
794                    InstKind::Log3(offset, size, _, _, _) => {
795                        Self::push_frame_read(func, &mut reads, offset, size)?;
796                    }
797                    InstKind::Log4(offset, size, _, _, _, _) => {
798                        Self::push_frame_read(func, &mut reads, offset, size)?;
799                    }
800                    InstKind::Call { args_offset, args_size, .. }
801                    | InstKind::StaticCall { args_offset, args_size, .. }
802                    | InstKind::DelegateCall { args_offset, args_size, .. } => {
803                        Self::push_frame_read(func, &mut reads, args_offset, args_size)?;
804                    }
805                    _ => {}
806                }
807            }
808        }
809
810        for block in func.blocks.iter() {
811            match block.terminator.as_ref() {
812                Some(Terminator::ReturnData { offset, size })
813                | Some(Terminator::Revert { offset, size }) => {
814                    Self::push_frame_read(func, &mut reads, *offset, *size)?;
815                }
816                _ => {}
817            }
818        }
819
820        Some(reads)
821    }
822
823    fn push_frame_read(
824        func: &Function,
825        reads: &mut Vec<(u64, u64)>,
826        offset: ValueId,
827        size: ValueId,
828    ) -> Option<()> {
829        if let Some(frame_offset) = Self::internal_frame_offset(func, offset) {
830            reads.push((frame_offset, func.value_u64(size)?));
831        }
832        Some(())
833    }
834
835    fn internal_frame_offset(func: &Function, value: ValueId) -> Option<u64> {
836        Self::internal_frame_offset_with_depth(func, value, 0)
837    }
838
839    fn internal_frame_offset_with_depth(
840        func: &Function,
841        value: ValueId,
842        depth: usize,
843    ) -> Option<u64> {
844        if depth > 8 {
845            return None;
846        }
847
848        match func.values[value] {
849            Value::Inst(inst_id) => match func.instructions[inst_id].kind {
850                InstKind::InternalFrameAddr(offset) => Some(offset),
851                InstKind::Add(a, b) => Self::internal_frame_add_offset(func, a, b, depth)
852                    .or_else(|| Self::internal_frame_add_offset(func, b, a, depth)),
853                _ => None,
854            },
855            _ => None,
856        }
857    }
858
859    fn internal_frame_add_offset(
860        func: &Function,
861        base: ValueId,
862        offset: ValueId,
863        depth: usize,
864    ) -> Option<u64> {
865        let base = Self::internal_frame_offset_with_depth(func, base, depth + 1)?;
866        base.checked_add(func.value_u64(offset)?)
867    }
868
869    fn can_mutate_memory(kind: &InstKind) -> bool {
870        matches!(
871            kind,
872            InstKind::MStore8(_, _)
873                | InstKind::MCopy(_, _, _)
874                | InstKind::CalldataCopy(_, _, _)
875                | InstKind::CodeCopy(_, _, _)
876                | InstKind::ReturnDataCopy(_, _, _)
877                | InstKind::ExtCodeCopy(_, _, _, _)
878                | InstKind::Call { .. }
879                | InstKind::StaticCall { .. }
880                | InstKind::DelegateCall { .. }
881                | InstKind::InternalCall { .. }
882        )
883    }
884
885    fn cross_block_memory_barrier(kind: &InstKind) -> bool {
886        matches!(kind, InstKind::MLoad(_)) || Self::is_memory_or_gas_observer(kind)
887    }
888}
889
890#[cfg(test)]
891mod tests {
892    use super::*;
893    use crate::mir::{FunctionBuilder, FunctionId};
894    use solar_interface::Ident;
895
896    fn test_func() -> Function {
897        Function::new(Ident::DUMMY)
898    }
899
900    #[test]
901    fn removes_overwritten_store() {
902        let mut func = test_func();
903        let mut builder = FunctionBuilder::new(&mut func);
904        let addr = builder.imm_u64(128);
905        let zero = builder.imm_u64(0);
906        let value = builder.imm_u64(42);
907        builder.mstore(addr, zero);
908        builder.mstore(addr, value);
909        builder.stop();
910
911        let mut pass = MemoryStoreEliminator::new();
912        assert_eq!(pass.run(&mut func), 1);
913        assert_eq!(func.blocks[func.entry_block].instructions.len(), 1);
914    }
915
916    #[test]
917    fn forwards_store_observed_only_by_load() {
918        let mut func = test_func();
919        let mut builder = FunctionBuilder::new(&mut func);
920        let addr = builder.imm_u64(128);
921        let zero = builder.imm_u64(0);
922        let value = builder.imm_u64(42);
923        builder.mstore(addr, zero);
924        let loaded = builder.mload(addr);
925        builder.mstore(addr, value);
926        builder.ret(vec![loaded]);
927
928        let mut pass = MemoryStoreEliminator::new();
929        assert_eq!(pass.run(&mut func), 2);
930
931        let block = &func.blocks[func.entry_block];
932        assert_eq!(block.instructions.len(), 1);
933        let Some(Terminator::Return { values }) = &block.terminator else {
934            panic!("expected return terminator");
935        };
936        assert_eq!(values.as_slice(), &[zero]);
937    }
938
939    #[test]
940    fn gas_is_a_barrier() {
941        let mut func = test_func();
942        let mut builder = FunctionBuilder::new(&mut func);
943        let addr = builder.imm_u64(128);
944        let zero = builder.imm_u64(0);
945        let value = builder.imm_u64(42);
946        builder.mstore(addr, zero);
947        builder.gas();
948        builder.mstore(addr, value);
949        builder.stop();
950
951        let mut pass = MemoryStoreEliminator::new();
952        assert_eq!(pass.run(&mut func), 0);
953        assert_eq!(func.blocks[func.entry_block].instructions.len(), 3);
954    }
955
956    #[test]
957    fn handles_distinct_immediate_values_for_same_address() {
958        let mut func = test_func();
959        let mut builder = FunctionBuilder::new(&mut func);
960        let addr1 = builder.imm_u64(128);
961        let addr2 = builder.imm_u64(128);
962        let zero = builder.imm_u64(0);
963        let value = builder.imm_u64(42);
964        builder.mstore(addr1, zero);
965        builder.mstore(addr2, value);
966        builder.stop();
967
968        let mut pass = MemoryStoreEliminator::new();
969        assert_eq!(pass.run(&mut func), 1);
970        assert_eq!(func.blocks[func.entry_block].instructions.len(), 1);
971    }
972
973    #[test]
974    fn handles_equivalent_base_offset_addresses() {
975        let mut func = test_func();
976        let mut builder = FunctionBuilder::new(&mut func);
977        let base = builder.add_param(crate::mir::MirType::uint256());
978        let offset = builder.imm_u64(32);
979        let value = builder.imm_u64(42);
980        let addr1 = builder.add(base, offset);
981        builder.mstore(addr1, value);
982        let addr2 = builder.add(base, offset);
983        let loaded = builder.mload(addr2);
984        builder.ret(vec![loaded]);
985
986        let mut pass = MemoryStoreEliminator::new();
987        assert_eq!(pass.run(&mut func), 1);
988
989        let block = &func.blocks[func.entry_block];
990        assert_eq!(block.instructions.len(), 3);
991        let Some(Terminator::Return { values }) = &block.terminator else {
992            panic!("expected return terminator");
993        };
994        assert_eq!(values.as_slice(), &[value]);
995    }
996
997    #[test]
998    fn removes_overwritten_store_to_equivalent_base_offset_address() {
999        let mut func = test_func();
1000        let mut builder = FunctionBuilder::new(&mut func);
1001        let base = builder.add_param(crate::mir::MirType::uint256());
1002        let offset = builder.imm_u64(32);
1003        let zero = builder.imm_u64(0);
1004        let value = builder.imm_u64(42);
1005        let addr1 = builder.add(base, offset);
1006        builder.mstore(addr1, zero);
1007        let addr2 = builder.add(base, offset);
1008        builder.mstore(addr2, value);
1009        builder.stop();
1010
1011        let mut pass = MemoryStoreEliminator::new();
1012        assert_eq!(pass.run(&mut func), 1);
1013        assert_eq!(func.blocks[func.entry_block].instructions.len(), 3);
1014    }
1015
1016    #[test]
1017    fn removes_unused_internal_frame_store() {
1018        let mut func = test_func();
1019        let mut builder = FunctionBuilder::new(&mut func);
1020        let frame = builder.internal_frame_addr(192);
1021        let zero = builder.imm_u64(0);
1022        builder.mstore(frame, zero);
1023        builder.ret(vec![zero]);
1024
1025        let mut pass = MemoryStoreEliminator::new();
1026        assert_eq!(pass.run_to_fixpoint(&mut func), 1);
1027        assert_eq!(func.blocks[func.entry_block].instructions.len(), 1);
1028    }
1029
1030    #[test]
1031    fn reuses_dominated_immutable_copy_load() {
1032        let mut func = test_func();
1033        let mut builder = FunctionBuilder::new(&mut func);
1034        let dest = builder.imm_u64(0);
1035        let immutable_len = builder.imm_u64(64);
1036        let word_size = builder.imm_u64(32);
1037        let code_size = builder.codesize();
1038        let offset = builder.sub(code_size, immutable_len);
1039        builder.codecopy(dest, offset, word_size);
1040        let first = builder.mload(dest);
1041        let next = builder.create_block();
1042        builder.jump(next);
1043
1044        builder.switch_to_block(next);
1045        let code_size = builder.codesize();
1046        let offset = builder.sub(code_size, immutable_len);
1047        builder.codecopy(dest, offset, word_size);
1048        let second = builder.mload(dest);
1049        let sum = builder.add(first, second);
1050        builder.ret(vec![sum]);
1051
1052        let mut pass = MemoryStoreEliminator::new();
1053        assert_eq!(pass.run_to_fixpoint(&mut func), 1);
1054
1055        let active_insts = func.blocks.iter().flat_map(|block| block.instructions.iter().copied());
1056        let mut code_copies = 0;
1057        let mut loads = 0;
1058        let mut stores = 0;
1059        for inst_id in active_insts {
1060            match func.instructions[inst_id].kind {
1061                InstKind::CodeCopy(_, _, _) => code_copies += 1,
1062                InstKind::MLoad(_) => loads += 1,
1063                InstKind::MStore(_, _) => stores += 1,
1064                _ => {}
1065            }
1066        }
1067        assert_eq!(code_copies, 1);
1068        assert_eq!(loads, 1);
1069        assert_eq!(stores, 1);
1070
1071        let add = match &func.values[sum] {
1072            Value::Inst(inst_id) => &func.instructions[*inst_id].kind,
1073            _ => panic!("expected add instruction"),
1074        };
1075        let InstKind::Add(lhs, rhs) = *add else {
1076            panic!("expected add instruction");
1077        };
1078        assert_eq!(lhs, first);
1079        assert_eq!(rhs, first);
1080    }
1081
1082    #[test]
1083    fn overlapping_load_blocks_overwritten_store_elimination() {
1084        let mut func = test_func();
1085        let mut builder = FunctionBuilder::new(&mut func);
1086        let base = builder.imm_u64(128);
1087        let one = builder.imm_u64(1);
1088        let zero = builder.imm_u64(0);
1089        let value = builder.imm_u64(42);
1090        builder.mstore(base, zero);
1091        let overlap = builder.add(base, one);
1092        builder.mload(overlap);
1093        builder.mstore(base, value);
1094        builder.stop();
1095
1096        let mut pass = MemoryStoreEliminator::new();
1097        assert_eq!(pass.run(&mut func), 0);
1098        assert_eq!(func.blocks[func.entry_block].instructions.len(), 4);
1099    }
1100
1101    #[test]
1102    fn forwards_load_from_store() {
1103        let mut func = test_func();
1104        let mut builder = FunctionBuilder::new(&mut func);
1105        let addr = builder.imm_u64(128);
1106        let value = builder.imm_u64(42);
1107        builder.mstore(addr, value);
1108        let loaded = builder.mload(addr);
1109        builder.ret(vec![loaded]);
1110
1111        let mut pass = MemoryStoreEliminator::new();
1112        assert_eq!(pass.run(&mut func), 1);
1113
1114        let block = &func.blocks[func.entry_block];
1115        assert_eq!(block.instructions.len(), 1);
1116        let Some(Terminator::Return { values }) = &block.terminator else {
1117            panic!("expected return terminator");
1118        };
1119        assert_eq!(values.as_slice(), &[value]);
1120    }
1121
1122    #[test]
1123    fn forwards_load_through_non_memory_operation() {
1124        let mut func = test_func();
1125        let mut builder = FunctionBuilder::new(&mut func);
1126        let addr = builder.imm_u64(128);
1127        let value = builder.imm_u64(42);
1128        let one = builder.imm_u64(1);
1129        builder.mstore(addr, value);
1130        builder.add(one, one);
1131        let loaded = builder.mload(addr);
1132        builder.ret(vec![loaded]);
1133
1134        let mut pass = MemoryStoreEliminator::new();
1135        assert_eq!(pass.run(&mut func), 1);
1136
1137        let block = &func.blocks[func.entry_block];
1138        assert_eq!(block.instructions.len(), 2);
1139        let Some(Terminator::Return { values }) = &block.terminator else {
1140            panic!("expected return terminator");
1141        };
1142        assert_eq!(values.as_slice(), &[value]);
1143    }
1144
1145    #[test]
1146    fn does_not_forward_load_across_memory_write() {
1147        let mut func = test_func();
1148        let mut builder = FunctionBuilder::new(&mut func);
1149        let addr = builder.imm_u64(128);
1150        let other_addr = builder.imm_u64(160);
1151        let value = builder.imm_u64(42);
1152        let len = builder.imm_u64(32);
1153        builder.mstore(addr, value);
1154        builder.mcopy(other_addr, addr, len);
1155        let loaded = builder.mload(addr);
1156        builder.ret(vec![loaded]);
1157
1158        let mut pass = MemoryStoreEliminator::new();
1159        assert_eq!(pass.run(&mut func), 0);
1160        assert_eq!(func.blocks[func.entry_block].instructions.len(), 3);
1161    }
1162
1163    #[test]
1164    fn does_not_forward_load_across_internal_call() {
1165        let mut func = test_func();
1166        let mut builder = FunctionBuilder::new(&mut func);
1167        let addr = builder.imm_u64(128);
1168        let value = builder.imm_u64(42);
1169        builder.mstore(addr, value);
1170        builder.internal_call_void(FunctionId::from_usize(0), Vec::new(), 0);
1171        let loaded = builder.mload(addr);
1172        builder.ret(vec![loaded]);
1173
1174        let mut pass = MemoryStoreEliminator::new();
1175        assert_eq!(pass.run(&mut func), 0);
1176        assert_eq!(func.blocks[func.entry_block].instructions.len(), 3);
1177    }
1178
1179    #[test]
1180    fn resolves_chained_forwarded_loads() {
1181        let mut func = test_func();
1182        let mut builder = FunctionBuilder::new(&mut func);
1183        let addr1 = builder.imm_u64(128);
1184        let addr2 = builder.imm_u64(160);
1185        let value = builder.imm_u64(42);
1186        builder.mstore(addr1, value);
1187        let loaded1 = builder.mload(addr1);
1188        builder.mstore(addr2, loaded1);
1189        let loaded2 = builder.mload(addr2);
1190        builder.ret(vec![loaded2]);
1191
1192        let mut pass = MemoryStoreEliminator::new();
1193        assert_eq!(pass.run(&mut func), 2);
1194
1195        let block = &func.blocks[func.entry_block];
1196        assert_eq!(block.instructions.len(), 2);
1197        let Some(Terminator::Return { values }) = &block.terminator else {
1198            panic!("expected return terminator");
1199        };
1200        assert_eq!(values.as_slice(), &[value]);
1201    }
1202}