1
+use std::cell::RefCell;
2
+use std::collections::HashSet;
3
+use std::rc::Rc;
4
+
5
+use bumpalo::collections::String;
6
+use forget_diagnostics::Diagnostic;
7
+
8
+use crate::{
9
+ initialize_hir, BasicBlock, BlockRewriter, BlockRewriterAction, DeclareLocal, Environment,
10
+ Function, GotoKind, GotoTerminal, Identifier, IdentifierData, IdentifierOperand, InstrIx,
11
+ Instruction, InstructionKind, InstructionValue, LValue, LabelTerminal, LoadLocal, MutableRange,
12
+ Operand, PlaceOrSpread, ReturnTerminal, StoreLocal, Terminal, TerminalValue, Type,
13
+};
14
+
15
+/// Inlines `useMemo()` calls, rewriting so that the lambda body becomes part of the
16
+/// outer block's instructions. To account for complex control flow, the inlining works
17
+/// as follows:
18
+/// * First, block ids are guaranteed to be unique for all blocks within a function and
19
+/// its recursive function expressions. Thus, the function expression's blocks can be
20
+/// directly moved into the outer function's `blocks` map.
21
+/// * To account for complex control flow, we create a "label" terminal just prior to
22
+/// the useMemo call, with the useMemo function's entry block as the body of the
23
+/// label terminal. The code following the useMemo call becomes the fallthrough.
24
+/// All returns within the useMemo are translated to instead:
25
+/// * Assign to a temporary identifier representing the useMemo result
26
+/// * Break to the label's fallthrough.
27
+///
28
+/// ## Example
29
+///
30
+/// Input:
31
+/// ```javascript
32
+/// foo();
33
+/// const x = useMemo(() => {
34
+/// if (a) {
35
+/// return b;
36
+/// }
37
+/// return c;
38
+/// })
39
+/// x;
40
+/// ```
41
+///
42
+/// HIR after translation:
43
+/// ```hir
44
+/// bb0:
45
+/// [ 1] #1 = LoadLocal 'foo'
46
+/// [ 2] #2 = Call #1()
47
+/// // label to allow substituting return -> goto in the lambda body
48
+/// [ 3] Label body=bb1 fallthrough=bb4
49
+/// bb1:
50
+/// [ 4] #3 = LoadLocal 'a'
51
+/// [ 5] If test=#3 consequent=bb2 alternate=bb3
52
+/// bb2:
53
+/// [ 6] #4 = LoadLocal 'b'
54
+/// [ 7] StoreLocal '<tmp>', #4
55
+/// [ 8] Goto bb4
56
+/// bb3:
57
+/// [ 9] #5 = LoadLocal 'c'
58
+/// [10] StoreLocal '<tmp>' #5
59
+/// [11] Goto bb4
60
+/// bb4:
61
+/// // code after the useMemo. save the temporary
62
+/// [12] #6 = LoadLocal '<tmp>'
63
+/// [13] StoreLocal 'x', #6
64
+/// ```
65
+///
66
+pub fn inline_use_memo<'a>(
67
+ env: &Environment<'a>,
68
+ fun: &mut Function<'a>,
69
+) -> Result<(), Diagnostic> {
70
+ let mut use_memo_globals: HashSet<InstrIx> = Default::default();
71
+ let mut functions: HashSet<InstrIx> = Default::default();
72
+
73
+ let blocks = &mut fun.body.blocks;
74
+ let instructions = &mut fun.body.instructions;
75
+ let mut rewriter = BlockRewriter::new(blocks, fun.body.entry);
76
+
77
+ let mut inlined = Vec::new();
78
+
79
+ rewriter.try_each_block(|mut block, rewriter| {
80
+ for (i, instr_ix) in block.instructions.iter().cloned().enumerate() {
81
+ let instr = &mut instructions[usize::from(instr_ix)];
82
+ match &mut instr.value {
83
+ InstructionValue::LoadGlobal(value) => {
84
+ if value.name.as_str() == "useMemo" {
85
+ use_memo_globals.insert(instr_ix);
86
+ }
87
+ }
88
+ InstructionValue::Function(_) => {
89
+ functions.insert(instr_ix);
90
+ }
91
+ InstructionValue::Call(value) => {
92
+ if !use_memo_globals.contains(&value.callee.ix) {
93
+ continue;
94
+ }
95
+ // Skip useMemo calls where the argument is a spread element
96
+ let lambda_ix = match &value.arguments.get(0) {
97
+ Some(PlaceOrSpread::Place(place)) => place.ix,
98
+ _ => continue,
99
+ };
100
+ // Skip useMemo where the argument is not a function expression
101
+ if !functions.contains(&lambda_ix) {
102
+ continue;
103
+ }
104
+ let instr_id = instr.id;
105
+
106
+ // Create a temporary variable to store the useMemo result into
107
+ let temporary_id = env.next_identifier_id();
108
+ let temporary = Identifier {
109
+ id: temporary_id,
110
+ // NOTE: for memoization to work correctly this variable has to be named
111
+ name: Some(String::from_str_in("t", &env.allocator)),
112
+ data: Rc::new(RefCell::new(IdentifierData {
113
+ mutable_range: MutableRange::new(),
114
+ scope: None,
115
+ type_: Type::Var(env.next_type_var_id()),
116
+ })),
117
+ };
118
+ // Replace the call with a load of the temporary
119
+ // this is convenient since consumers of the useMemo call
120
+ // already point to this instruction id, so by reusing the
121
+ // instruction we don't have to update the consumer(s) to
122
+ // look at a different instruction
123
+ instr.value = InstructionValue::LoadLocal(LoadLocal {
124
+ place: IdentifierOperand {
125
+ identifier: temporary.clone(),
126
+ effect: None,
127
+ },
128
+ });
129
+
130
+ // Move the function expression out of its instruction so that we own
131
+ // the value and can modify and inline its contents into the outer
132
+ // function. We replace with a tombstone value that we can filter out later
133
+ let lambda = std::mem::replace(
134
+ &mut instructions[usize::from(lambda_ix)].value,
135
+ InstructionValue::Tombstone,
136
+ );
137
+ let mut lambda = if let InstructionValue::Function(lambda) = lambda {
138
+ lambda
139
+ } else {
140
+ unreachable!("Must be a function, checked above")
141
+ };
142
+
143
+ // Additional validation
144
+ // TODO: this should be part of a separate validation pass
145
+ if !lambda.lowered_function.params.is_empty() {
146
+ return Err(Diagnostic::invalid_react(
147
+ "useMemo callbacks may not accept any arguments",
148
+ None,
149
+ ));
150
+ }
151
+ if lambda.lowered_function.is_async || lambda.lowered_function.is_generator {
152
+ return Err(Diagnostic::invalid_react(
153
+ "useMemo callbacks may not be async or generator functions",
154
+ None,
155
+ ));
156
+ }
157
+
158
+ // Set aside a BlockId for the code that follows the useMemo call
159
+ let continuation_block_id = env.next_block_id();
160
+
161
+ // Rewrite the body of the lambda to replace any return terminals
162
+ // with an assignment to the useMemo temporary followed by a break
163
+ // to the continuation block
164
+ for block in lambda.lowered_function.body.blocks.iter_mut() {
165
+ if let TerminalValue::Return(ReturnTerminal { value }) =
166
+ &mut block.terminal.value
167
+ {
168
+ let store_ix = InstrIx::new(
169
+ lambda.lowered_function.body.instructions.len() as u32,
170
+ );
171
+ lambda.lowered_function.body.instructions.push(Instruction {
172
+ id: instr_id,
173
+ value: InstructionValue::StoreLocal(StoreLocal {
174
+ lvalue: LValue {
175
+ identifier: IdentifierOperand {
176
+ identifier: temporary.clone(),
177
+ effect: None,
178
+ },
179
+ kind: InstructionKind::Reassign,
180
+ },
181
+ value: Operand {
182
+ ix: value.ix,
183
+ effect: None,
184
+ },
185
+ }),
186
+ });
187
+ block.instructions.push(store_ix);
188
+ block.terminal.value = TerminalValue::Goto(GotoTerminal {
189
+ block: continuation_block_id,
190
+ kind: GotoKind::Break,
191
+ });
192
+ }
193
+ }
194
+
195
+ // Extract the block's original terminal, which we will move to the
196
+ // continuation block. Replace it with a label terminal, necessary to
197
+ // allow the goto statements to have a target.
198
+ let terminal_id = block.terminal.id;
199
+ let terminal = std::mem::replace(
200
+ &mut block.terminal,
201
+ Terminal {
202
+ id: terminal_id,
203
+ value: TerminalValue::Label(LabelTerminal {
204
+ block: lambda.lowered_function.body.entry,
205
+ fallthrough: Some(continuation_block_id),
206
+ }),
207
+ },
208
+ );
209
+
210
+ // Extract the instructions for the continuation block
211
+ let continuation_instructions = block.instructions.split_off(i);
212
+
213
+ // Declare the temporary variable at the end of the block preceding
214
+ // the useMemo invocation
215
+ let declare_ix = InstrIx::new(instructions.len() as u32);
216
+ instructions.push(Instruction {
217
+ id: instr_id,
218
+ value: InstructionValue::DeclareLocal(DeclareLocal {
219
+ lvalue: LValue {
220
+ identifier: IdentifierOperand {
221
+ identifier: temporary.clone(),
222
+ effect: None,
223
+ },
224
+ kind: InstructionKind::Let,
225
+ },
226
+ }),
227
+ });
228
+ block.instructions.push(declare_ix);
229
+
230
+ // Add the continuation block
231
+ let continuation_block = Box::new(BasicBlock {
232
+ id: continuation_block_id,
233
+ instructions: continuation_instructions,
234
+ kind: block.kind,
235
+ phis: env.vec_new(),
236
+ predecessors: Default::default(),
237
+ terminal,
238
+ });
239
+ rewriter.add_block(continuation_block);
240
+
241
+ inlined.push(lambda);
242
+ break;
243
+ }
244
+ _ => {}
245
+ }
246
+ }
247
+ Ok(BlockRewriterAction::Keep(block))
248
+ })?;
249
+
250
+ if !inlined.is_empty() {
251
+ for lambda in inlined {
252
+ fun.body.inline(lambda);
253
+ }
254
+ initialize_hir(&mut fun.body)?;
255
+ }
256
+
257
+ Ok(())
258
+}