1
+/**
2
+ * Copyright (c) Facebook, Inc. and its affiliates.
3
+ *
4
+ * This source code is licensed under the MIT license found in the
5
+ * LICENSE file in the root directory of this source tree.
6
+ */
7
+
8
+import { CompilerError } from "../CompilerError";
9
+import {
10
+ BasicBlock,
11
+ BlockId,
12
+ Effect,
13
+ Environment,
14
+ FunctionExpression,
15
+ GotoTerminal,
16
+ GotoVariant,
17
+ HIR,
18
+ HIRFunction,
19
+ IdentifierId,
20
+ InstructionKind,
21
+ makeInstructionId,
22
+ makeType,
23
+ Place,
24
+ reversePostorderBlocks,
25
+ shrink,
26
+} from "../HIR";
27
+import { markInstructionIds, markPredecessors } from "../HIR/HIRBuilder";
28
+import { assertExhaustive, retainWhere } from "../Utils/utils";
29
+
30
+/**
31
+ * Rewrites `useMemo()` calls, rewriting so that the lambda body becomes part of the
32
+ * outer block's instructions.
33
+ *
34
+ * Example:
35
+ *
36
+ * ```javascript
37
+ * // Before
38
+ * const x = useMemo(() => foo(y, z), [y, z])
39
+ *
40
+ * // After
41
+ * const x = foo(y, z);
42
+ * ```
43
+ *
44
+ * The main challenge is dealing with the possibility of complex control flow within
45
+ * the lambda body. The approach is roughly:
46
+ * - split the block with the useMemo call in two:
47
+ * - the first block is everything up to the memo call plus the lambda body
48
+ * - the second block is everything after the memo call
49
+ * - use the temporary from the useMemo call result value as the place to store
50
+ * the useMemo result
51
+ * - for every return terminal in the lambda body:
52
+ * - add a StoreLocal to the temporary, assigning the return value
53
+ * - replace the terminal w a goto to the second block
54
+ *
55
+ * NOTE: *this pass must be run prior to EnterSSA*. Prior to entering SSA form identifiers
56
+ * in the top-level function and any function expressions will have consistent
57
+ * correlation between `Identifier` instances and IdentifierIds. After entering SSA
58
+ * form we drop this correspondence. It's much easier to write this inlining pass
59
+ * without having to worry about SSA form.
60
+ */
61
+export function inlineUseMemo(fn: HIRFunction): void {
62
+ // Track all function expressions in case they appear as the argument to a useMemo
63
+ const functions = new Map<IdentifierId, FunctionExpression>();
64
+ // Track all references to `useMemo`
65
+ const useMemoGlobals = new Set<IdentifierId>();
66
+ // Identifiers (lvalues) for known useMemo functions, so that we can prune them
67
+ // at the end of the pass
68
+ const useMemoFunctions = new Set<IdentifierId>();
69
+
70
+ // Iterate the *existing* blocks from the outer component to find useMemo calls
71
+ // and inline them. During iteration we will modify `fn` (by inlining the CFG
72
+ // of useMemo callbacks) so we explicitly copy references to just the original
73
+ // function's blocks first. As blocks are split to make room for useMemo calls,
74
+ // the split portions of the blocks will be added to this queue.
75
+ const queue = Array.from(fn.body.blocks.values());
76
+ queue: for (const block of queue) {
77
+ for (let ii = 0; ii < block.instructions.length; ii++) {
78
+ const instr = block.instructions[ii]!;
79
+ switch (instr.value.kind) {
80
+ case "LoadGlobal": {
81
+ if (instr.value.name === "useMemo") {
82
+ useMemoGlobals.add(instr.lvalue.identifier.id);
83
+ }
84
+ break;
85
+ }
86
+ case "FunctionExpression": {
87
+ functions.set(instr.lvalue.identifier.id, instr.value);
88
+ break;
89
+ }
90
+ case "CallExpression": {
91
+ if (useMemoGlobals.has(instr.value.callee.identifier.id)) {
92
+ const [lambda] = instr.value.args;
93
+ if (lambda.kind === "Spread") {
94
+ continue;
95
+ }
96
+ const body = functions.get(lambda.identifier.id);
97
+ if (body === undefined) {
98
+ CompilerError.invariant(
99
+ "Expected first argument to useMemo() to be a function expression",
100
+ fn.loc
101
+ );
102
+ }
103
+ // We know this function is used for useMemo and can prune it later
104
+ useMemoFunctions.add(lambda.identifier.id);
105
+
106
+ // Create a new block which will contain code following the useMemo call
107
+ const continuationBlockId = fn.env.nextBlockId;
108
+ const continuationBlock: BasicBlock = {
109
+ id: continuationBlockId,
110
+ instructions: block.instructions.slice(ii + 1),
111
+ kind: block.kind,
112
+ phis: new Set(),
113
+ preds: new Set(),
114
+ terminal: block.terminal,
115
+ };
116
+ fn.body.blocks.set(continuationBlockId, continuationBlock);
117
+
118
+ // Trim the original block to contain instructions up to (but not including)
119
+ // the useMemo
120
+ block.instructions.length = ii;
121
+
122
+ // The block leading up to the useMemo needs to jump to the entry block of
123
+ // the useMemo control flow graph. These will be merged into a single block
124
+ // via MergeConsectuveBlocks
125
+ const newTerminal: GotoTerminal = {
126
+ block: body.loweredFunc.body.entry,
127
+ id: makeInstructionId(0),
128
+ kind: "goto",
129
+ variant: GotoVariant.Break,
130
+ loc: block.terminal.loc,
131
+ };
132
+ block.terminal = newTerminal;
133
+
134
+ // If the final terminal type has a fallthrough, update it to point to the
135
+ // continuation block
136
+ const terminalBlock = getTerminalBlock(
137
+ body.loweredFunc.body,
138
+ body.loweredFunc.body.entry
139
+ );
140
+ switch (terminalBlock.terminal.kind) {
141
+ case "if":
142
+ case "switch":
143
+ case "label": {
144
+ // These terminals can all appear as the final top-level terminal
145
+ // *and* have fallthroughs. If they are final, their fallthrough
146
+ // must be updated to point to the continuation block to main
147
+ // proper CFG structure (a block that succeeds all branches of a conditional
148
+ // must be marked as that conditional's fallthrough)
149
+ terminalBlock.terminal.fallthrough = continuationBlockId;
150
+ break;
151
+ }
152
+ case "return":
153
+ case "throw": {
154
+ // These can appear as the final top-level terminal
155
+ break;
156
+ }
157
+ // These all have non-nullable fallthroughs: there is always some code in the
158
+ // CFG that succeeds them which we should find instead
159
+ case "optional-call":
160
+ case "ternary":
161
+ case "logical":
162
+ case "while":
163
+ case "for":
164
+ case "for-of":
165
+ case "do-while":
166
+ // These are invalid terminals for a top-level block
167
+ case "branch":
168
+ case "goto":
169
+ case "unsupported": {
170
+ CompilerError.invariant(
171
+ `Unexpected final top-level terminal`,
172
+ terminalBlock.terminal.loc,
173
+ `Found ${terminalBlock.terminal.kind}, expected one of if, switch, label, return, or throw`
174
+ );
175
+ }
176
+ default: {
177
+ assertExhaustive(
178
+ terminalBlock.terminal,
179
+ `Unexpected terminal kind '${
180
+ (terminalBlock.terminal as any).kind
181
+ }'`
182
+ );
183
+ }
184
+ }
185
+
186
+ // Rewrite blocks from the lambda to replace any `return` with a
187
+ // store the useMemo temporary and `goto` the continuation block
188
+ for (const [id, block] of body.loweredFunc.body.blocks) {
189
+ block.preds.clear();
190
+ rewriteBlock(fn.env, block, continuationBlockId, instr.lvalue);
191
+ fn.body.blocks.set(id, block);
192
+ }
193
+
194
+ // Ensure we visit the continuation block, since there may have been
195
+ // sequential useMemos that need to be visited.
196
+ queue.push(continuationBlock);
197
+ continue queue;
198
+ }
199
+ }
200
+ }
201
+ }
202
+ }
203
+
204
+ if (useMemoFunctions.size !== 0) {
205
+ // Remove instructions that define lambdas which we inlined
206
+ for (const [, block] of fn.body.blocks) {
207
+ retainWhere(
208
+ block.instructions,
209
+ (instr) => !useMemoFunctions.has(instr.lvalue.identifier.id)
210
+ );
211
+ }
212
+
213
+ // If terminals have changed then blocks may have become newly unreachable.
214
+ // Re-run minification of the graph (incl reordering instruction ids)
215
+ shrink(fn.body);
216
+ reversePostorderBlocks(fn.body);
217
+ markInstructionIds(fn.body);
218
+ markPredecessors(fn.body);
219
+ }
220
+}
221
+
222
+// Finds the final top-level terminal node for a CFG, by following any
223
+// fallthrough nodes.
224
+function getTerminalBlock(cfg: HIR, start: BlockId): BasicBlock {
225
+ let current = cfg.blocks.get(start)!;
226
+ while (true) {
227
+ const { terminal } = current;
228
+ switch (terminal.kind) {
229
+ case "if": {
230
+ if (
231
+ terminal.fallthrough !== null &&
232
+ terminal.fallthrough === terminal.alternate
233
+ ) {
234
+ // Here we don't know if the fallthrough and alternate are the same because there was
235
+ // no alternate or because both the alternate exists and the fallthrough is just unreachable
236
+ // So we check if the fallthrough returns/throws (the if is the final top-level terminal)
237
+ // or whether execution actually may continue.
238
+ const fallthrough = getTerminalBlock(cfg, terminal.fallthrough);
239
+ if (
240
+ fallthrough.terminal.kind === "return" ||
241
+ fallthrough.terminal.kind === "throw"
242
+ ) {
243
+ return current;
244
+ } else {
245
+ current = fallthrough;
246
+ continue;
247
+ }
248
+ } else {
249
+ return current;
250
+ }
251
+ }
252
+ case "switch":
253
+ case "label": {
254
+ if (terminal.fallthrough !== null) {
255
+ current = cfg.blocks.get(terminal.fallthrough)!;
256
+ continue;
257
+ } else {
258
+ return current;
259
+ }
260
+ }
261
+ case "optional-call":
262
+ case "ternary":
263
+ case "logical":
264
+ case "while":
265
+ case "for":
266
+ case "for-of":
267
+ case "do-while": {
268
+ current = cfg.blocks.get(terminal.fallthrough)!;
269
+ continue;
270
+ }
271
+ case "return":
272
+ case "throw": {
273
+ return current;
274
+ }
275
+ case "unsupported":
276
+ case "branch":
277
+ case "goto": {
278
+ CompilerError.invariant(
279
+ `Unexpected block terminal`,
280
+ terminal.loc,
281
+ `Top-level blocks may not end in a ${terminal.kind} terminal`
282
+ );
283
+ }
284
+ default: {
285
+ assertExhaustive(
286
+ terminal,
287
+ `Unexpected terminal kind '${(terminal as any).kind}'`
288
+ );
289
+ }
290
+ }
291
+ }
292
+}
293
+
294
+/**
295
+ * Rewrites the block so that all `return` terminals are replaced:
296
+ * * Add a StoreLocal <returnValue> = <terminal.value>
297
+ * * Replace the terminal with a Goto to <returnTarget>
298
+ */
299
+function rewriteBlock(
300
+ env: Environment,
301
+ block: BasicBlock,
302
+ returnTarget: BlockId,
303
+ returnValue: Place
304
+): void {
305
+ const { terminal } = block;
306
+ if (terminal.kind !== "return") {
307
+ return;
308
+ }
309
+ if (terminal.value !== null) {
310
+ block.instructions.push({
311
+ id: makeInstructionId(0),
312
+ loc: terminal.loc,
313
+ lvalue: {
314
+ effect: Effect.Unknown,
315
+ identifier: {
316
+ id: env.nextIdentifierId,
317
+ mutableRange: {
318
+ start: makeInstructionId(0),
319
+ end: makeInstructionId(0),
320
+ },
321
+ name: null,
322
+ scope: null,
323
+ type: makeType(),
324
+ },
325
+ kind: "Identifier",
326
+ loc: terminal.loc,
327
+ },
328
+ value: {
329
+ kind: "StoreLocal",
330
+ lvalue: { kind: InstructionKind.Const, place: { ...returnValue } },
331
+ value: terminal.value,
332
+ loc: terminal.loc,
333
+ },
334
+ });
335
+ }
336
+ block.terminal = {
337
+ kind: "goto",
338
+ block: returnTarget,
339
+ id: makeInstructionId(0),
340
+ variant: GotoVariant.Break,
341
+ loc: block.terminal.loc,
342
+ };
343
+}