[rust] Use new block helper for more passes
Joe Savona committed
Jul 14, 2023 at 22:47 UTC
bac24ffd78129e6f7dce2b22d7ba44dd1978d0b0
3 files changed
+46
-22
compiler/forget/crates/forget_hir/src/function.rs
+27
-3
@@ -62,7 +62,7 @@ impl<'a> Blocks<'a> {
62
self.data.keys().cloned().collect()
63
}
64
65
- pub fn take(&mut self, id: BlockId) -> Box<BasicBlock<'a>> {
65
+ pub fn remove(&mut self, id: BlockId) -> Box<BasicBlock<'a>> {
66
self.data.remove(&id).unwrap().unwrap()
67
}
68
@@ -146,7 +146,26 @@ impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
146
}
147
}
148
149
- pub fn each_block<F>(&mut self, mut f: F) -> Result<(), Diagnostic>
149
+ pub fn each_block<F>(&mut self, mut f: F) -> ()
150
+ where
151
+ F: FnMut(Box<BasicBlock<'a>>, &mut Self) -> BlockRewriterAction<'a>,
152
+ {
153
+ let keys = self.blocks.block_ids();
154
+ for block_id in keys {
155
+ self.current = block_id;
156
+ let block = self.blocks.data.get_mut(&block_id).unwrap().take().unwrap();
157
+ match f(block, self) {
158
+ BlockRewriterAction::Keep(block) => {
159
+ self.blocks.data.insert(block_id, Some(block));
160
+ }
161
+ BlockRewriterAction::Remove => {
162
+ self.blocks.data.remove(&block_id);
163
+ }
164
+ }
165
+ }
166
+ }
167
+
168
+ pub fn try_each_block<F>(&mut self, mut f: F) -> Result<(), Diagnostic>
169
where
170
F: FnMut(Box<BasicBlock<'a>>, &mut Self) -> Result<BlockRewriterAction<'a>, Diagnostic>,
171
{
@@ -159,13 +178,18 @@ impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
178
self.blocks.data.insert(block_id, Some(block));
179
}
180
BlockRewriterAction::Remove => {
162
- // nothing to do, already removed from the blocks
181
+ self.blocks.data.remove(&block_id);
182
}
183
}
184
}
185
Ok(())
186
}
187
188
+ pub fn contains(&self, block_id: BlockId) -> bool {
189
+ assert_ne!(block_id, self.current);
190
+ self.blocks.data.contains_key(&block_id)
191
+ }
192
+
193
pub fn block(&self, block_id: BlockId) -> &BasicBlock<'a> {
194
assert_ne!(block_id, self.current);
195
self.blocks.block(block_id)
compiler/forget/crates/forget_hir/src/initialize.rs
+18
-18
@@ -3,7 +3,10 @@ use std::collections::HashSet;
3
use forget_diagnostics::{invariant, Diagnostic};
4
use thiserror::Error;
5
6
-use crate::{BlockId, Blocks, GotoKind, GotoTerminal, InstructionIdGenerator, TerminalValue, HIR};
6
+use crate::{
7
+ BlockId, BlockRewriter, BlockRewriterAction, Blocks, GotoKind, GotoTerminal,
8
+ InstructionIdGenerator, TerminalValue, HIR,
9
+};
10
11
/// Runs a variety of passes to put the HIR in canonical form. This should be called
12
/// after initial HIR construction and after any transformations that change the
@@ -65,7 +68,7 @@ pub fn reverse_postorder_blocks<'a>(hir: &mut HIR<'a>) {
68
// NOTE: could consider sorting the blocks in-place by key
69
let mut blocks = Blocks::with_capacity(hir.blocks.len());
70
for id in postorder.iter().rev().cloned() {
68
- blocks.insert(hir.blocks.take(id));
71
+ blocks.insert(hir.blocks.remove(id));
72
}
73
74
hir.blocks = blocks;
@@ -73,52 +76,49 @@ pub fn reverse_postorder_blocks<'a>(hir: &mut HIR<'a>) {
76
77
/// Prunes ForTerminal.update values (sets to None) if they are unreachable
78
pub fn remove_unreachable_for_updates<'a>(hir: &mut HIR<'a>) {
76
- let block_ids = hir.blocks.block_ids();
77
-
78
- for block in hir.blocks.iter_mut() {
79
+ BlockRewriter::new(&mut hir.blocks, hir.entry).each_block(|mut block, rewriter| {
80
if let TerminalValue::For(terminal) = &mut block.terminal.value {
81
if let Some(update) = terminal.update {
81
- if !block_ids.contains(&update) {
82
+ if !rewriter.contains(update) {
83
terminal.update = None;
84
}
85
}
86
}
86
- }
87
+ BlockRewriterAction::Keep(block)
88
+ });
89
}
90
91
/// Prunes unreachable fallthrough values, setting them to None if the referenced
92
/// block was not otherwise reachable.
93
pub fn remove_unreachable_fallthroughs<'a>(hir: &mut HIR<'a>) {
92
- let block_ids = hir.blocks.block_ids();
93
-
94
- for block in hir.blocks.iter_mut() {
94
+ BlockRewriter::new(&mut hir.blocks, hir.entry).each_block(|mut block, rewriter| {
95
block
96
.terminal
97
.value
98
.map_optional_fallthroughs(|fallthrough| {
99
- if block_ids.contains(&fallthrough) {
99
+ if rewriter.contains(fallthrough) {
100
Some(fallthrough)
101
} else {
102
None
103
}
104
- })
105
- }
104
+ });
105
+ BlockRewriterAction::Keep(block)
106
+ });
107
}
108
109
/// Rewrites DoWhile statements into Gotos if the test block is not reachable
110
pub fn remove_unreachable_do_while_statements<'a>(hir: &mut HIR<'a>) {
110
- let block_ids = hir.blocks.block_ids();
111
-
112
- for block in hir.blocks.iter_mut() {
111
+ BlockRewriter::new(&mut hir.blocks, hir.entry).each_block(|mut block, rewriter| {
112
if let TerminalValue::DoWhile(terminal) = &mut block.terminal.value {
114
- if !block_ids.contains(&terminal.test) {
113
+ if !rewriter.contains(terminal.test) {
114
block.terminal.value = TerminalValue::Goto(GotoTerminal {
115
block: terminal.body,
116
kind: GotoKind::Break,
117
});
118
}
119
}
121
- }
120
+ BlockRewriterAction::Keep(block)
121
+ });
122
}
123
124
/// Updates the instruction ids for all instructions and blocks
compiler/forget/crates/forget_hir/src/merge_consecutive_blocks.rs
+1
-1
@@ -27,7 +27,7 @@ pub fn merge_consecutive_blocks<'a>(
27
let mut rewriter = BlockRewriter::new(blocks, fun.body.entry);
28
let mut has_changes = false;
29
30
- rewriter.each_block(|mut block, rewriter| {
30
+ rewriter.try_each_block(|mut block, rewriter| {
31
let block_id = block.id;
32
// Visit instructions to merge blocks within function expressions
33
for instr_ix in &block.instructions {