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 {
9
+ BlockId,
10
+ GotoVariant,
11
+ HIRFunction,
12
+ IdentifierId,
13
+ InstructionValue,
14
+ markInstructionIds,
15
+ markPredecessors,
16
+ mergeConsecutiveBlocks,
17
+ Place,
18
+ Primitive,
19
+ removeUnreachableFallthroughs,
20
+ reversePostorderBlocks,
21
+ shrink,
22
+} from "../HIR";
23
+import { eliminateRedundantPhi } from "../SSA";
24
+
25
+/**
26
+ * Applies constant propagation and constant folding to the given function.
27
+ * Note that because HIR operands are always a Place, constants cannot be directly
28
+ * propagated into the HIR itself (the closest option would be to copy constants to
29
+ * new temporaries just before each use, and update usage sites to reference those
30
+ * new temporaries).
31
+ *
32
+ * Instead this pass implements constant folding, in which constant values are
33
+ * propagated internally to the pass and subsequent operations are removed/folded where
34
+ * possible.
35
+ *
36
+ * Note that this pass may prune control flow blocks that are unreachable, for example
37
+ * a consequent or alternate branch if an `if` test is provably truthy or falsey.
38
+ * If (and only if) terminals change, the pass re-runs various stages to ensure the
39
+ * CFG is in minimal form. This means instruction ids *may* change as a result of this
40
+ * pass.
41
+ */
42
+export function constantPropagation(fn: HIRFunction): void {
43
+ const haveTerminalsChanged = applyConstantPropagation(fn);
44
+ if (haveTerminalsChanged) {
45
+ // If terminals have changed then blocks may have become newly unreachable.
46
+ // Re-run minification of the graph (incl reordering instruction ids)
47
+ shrink(fn.body);
48
+ reversePostorderBlocks(fn.body);
49
+ removeUnreachableFallthroughs(fn.body);
50
+ markInstructionIds(fn.body);
51
+ markPredecessors(fn.body);
52
+
53
+ // Now that predecessors are updated, prune phi operands that can never be reached
54
+ for (const [, block] of fn.body.blocks) {
55
+ for (const phi of block.phis) {
56
+ for (const [predecessor] of phi.operands) {
57
+ if (!block.preds.has(predecessor)) {
58
+ phi.operands.delete(predecessor);
59
+ }
60
+ }
61
+ }
62
+ }
63
+ // By removing some phi operands, there may be phis that were not previously
64
+ // redundant but now are
65
+ eliminateRedundantPhi(fn);
66
+ // Finally, merge together any blocks that are now guaranteed to execute
67
+ // consecutively
68
+ mergeConsecutiveBlocks(fn);
69
+ }
70
+}
71
+
72
+function applyConstantPropagation(fn: HIRFunction): boolean {
73
+ let hasChanges = false;
74
+
75
+ // A set of blocks whose terminals can't (yet) be safely rewritten
76
+ const valueBlocks = new Set<BlockId>();
77
+
78
+ const constants: Constants = new Map();
79
+ for (const [, block] of fn.body.blocks) {
80
+ // Initialize phi values if all operands have the same known constant value.
81
+ // Note that this analysis uses a single-pass only, so it will never fill in
82
+ // phi values for blocks that have a back-edge.
83
+ for (const phi of block.phis) {
84
+ let value: Primitive | null = null;
85
+ for (const [, operand] of phi.operands) {
86
+ const operandValue = constants.get(operand.id) ?? null;
87
+ if (operandValue === null) {
88
+ value = null;
89
+ break;
90
+ }
91
+ if (value === null) {
92
+ value = operandValue;
93
+ } else if (operandValue.value !== value.value) {
94
+ value = null;
95
+ break;
96
+ }
97
+ }
98
+ if (value !== null) {
99
+ constants.set(phi.id.id, value);
100
+ }
101
+ }
102
+
103
+ for (const instr of block.instructions) {
104
+ const value = evaluateInstruction(constants, instr.value);
105
+ if (value !== null) {
106
+ instr.value = value;
107
+ constants.set(instr.lvalue.place.identifier.id, value);
108
+ }
109
+ }
110
+
111
+ if (valueBlocks.has(block.id)) {
112
+ // can't rewrite terminals in value blocks yet
113
+ continue;
114
+ }
115
+ const terminal = block.terminal;
116
+ switch (terminal.kind) {
117
+ case "if": {
118
+ const testValue = read(constants, terminal.test);
119
+ if (testValue !== null && testValue.kind === "Primitive") {
120
+ hasChanges = true;
121
+ const targetBlockId = Boolean(testValue.value)
122
+ ? terminal.consequent
123
+ : terminal.alternate;
124
+ block.terminal = {
125
+ kind: "goto",
126
+ variant: GotoVariant.Break,
127
+ block: targetBlockId,
128
+ id: terminal.id,
129
+ };
130
+ }
131
+ break;
132
+ }
133
+ case "while": {
134
+ valueBlocks.add(terminal.test);
135
+ break;
136
+ }
137
+ case "for": {
138
+ valueBlocks.add(terminal.init);
139
+ valueBlocks.add(terminal.test);
140
+ valueBlocks.add(terminal.update);
141
+ break;
142
+ }
143
+ default: {
144
+ // no-op
145
+ }
146
+ }
147
+ }
148
+
149
+ return hasChanges;
150
+}
151
+
152
+function evaluateInstruction(
153
+ constants: Constants,
154
+ instr: InstructionValue
155
+): Constant | null {
156
+ switch (instr.kind) {
157
+ case "Primitive": {
158
+ return instr;
159
+ }
160
+ case "BinaryExpression": {
161
+ const lhsValue = read(constants, instr.left);
162
+ const rhsValue = read(constants, instr.right);
163
+ if (lhsValue !== null && rhsValue !== null) {
164
+ const lhs = lhsValue.value;
165
+ const rhs = rhsValue.value;
166
+ switch (instr.operator) {
167
+ case "+": {
168
+ if (typeof lhs === "number" && typeof rhs === "number") {
169
+ return { kind: "Primitive", value: lhs + rhs, loc: instr.loc };
170
+ }
171
+ return null;
172
+ }
173
+ case "-": {
174
+ if (typeof lhs === "number" && typeof rhs === "number") {
175
+ return { kind: "Primitive", value: lhs - rhs, loc: instr.loc };
176
+ }
177
+ return null;
178
+ }
179
+ case "*": {
180
+ if (typeof lhs === "number" && typeof rhs === "number") {
181
+ return { kind: "Primitive", value: lhs * rhs, loc: instr.loc };
182
+ }
183
+ return null;
184
+ }
185
+ case "/": {
186
+ if (typeof lhs === "number" && typeof rhs === "number") {
187
+ return { kind: "Primitive", value: lhs / rhs, loc: instr.loc };
188
+ }
189
+ return null;
190
+ }
191
+ case "<": {
192
+ if (typeof lhs === "number" && typeof rhs === "number") {
193
+ return { kind: "Primitive", value: lhs < rhs, loc: instr.loc };
194
+ }
195
+ return null;
196
+ }
197
+ case "<=": {
198
+ if (typeof lhs === "number" && typeof rhs === "number") {
199
+ return { kind: "Primitive", value: lhs <= rhs, loc: instr.loc };
200
+ }
201
+ return null;
202
+ }
203
+ case ">": {
204
+ if (typeof lhs === "number" && typeof rhs === "number") {
205
+ return { kind: "Primitive", value: lhs > rhs, loc: instr.loc };
206
+ }
207
+ return null;
208
+ }
209
+ case ">=": {
210
+ if (typeof lhs === "number" && typeof rhs === "number") {
211
+ return { kind: "Primitive", value: lhs >= rhs, loc: instr.loc };
212
+ }
213
+ return null;
214
+ }
215
+ case "==": {
216
+ return { kind: "Primitive", value: lhs == rhs, loc: instr.loc };
217
+ }
218
+ case "===": {
219
+ return { kind: "Primitive", value: lhs === rhs, loc: instr.loc };
220
+ }
221
+ case "!=": {
222
+ return { kind: "Primitive", value: lhs != rhs, loc: instr.loc };
223
+ }
224
+ case "!==": {
225
+ return { kind: "Primitive", value: lhs !== rhs, loc: instr.loc };
226
+ }
227
+ default: {
228
+ // TODO: handle more cases
229
+ return null;
230
+ }
231
+ }
232
+ }
233
+ return null;
234
+ }
235
+ case "PropertyLoad": {
236
+ const objectValue = read(constants, instr.object);
237
+ if (objectValue !== null) {
238
+ if (
239
+ typeof objectValue.value === "string" &&
240
+ instr.property === "length"
241
+ ) {
242
+ return {
243
+ kind: "Primitive",
244
+ value: objectValue.value.length,
245
+ loc: instr.loc,
246
+ };
247
+ }
248
+ }
249
+ return null;
250
+ }
251
+ case "Identifier": {
252
+ return read(constants, instr);
253
+ }
254
+ default: {
255
+ // TODO: handle more cases
256
+ return null;
257
+ }
258
+ }
259
+}
260
+
261
+/**
262
+ * Recursively read the value of a place: if it is a constant place, attempt to read
263
+ * from that place until reaching a primitive or finding a value that is unset.
264
+ */
265
+function read(constants: Constants, place: Place): Constant | null {
266
+ return constants.get(place.identifier.id) ?? null;
267
+}
268
+
269
+type Constant = Primitive;
270
+type Constants = Map<IdentifierId, Constant>;