@samitouri / QOS-React-1 / commits / 59c1f22057

[rust] port inline_use_memo

Ports InlineUseMemo to Rust, this is the only missing transform pass on the pipeline up through constant propagation. I wanted to finish this so that we could do benchmarking of the early phase of compilation. UseMemo does a bunch of rewriting so it seemed worth comparing performance. Included: * Add swc -> estree and estree -> HIR conversions for arrow functions and call expressions * Add HIR definition for labeled terminals and call expressions * Utility for inlining one function into another — note that this requires remapping InstrIx operands since instruction indices will change. An alternative would be to store all the instructions for both outer/inner functions in the same array, in which case we wouldn't need to remap. * Helpers for mutably iterating the operands of an instruction or terminal. * The actual inline_use_memo() function. The overall logic is similar, i just had to slightly shuffle the order of operations to satisfy the borrow checker.

Joe Savona committed Jul 25, 2023 at 09:24 UTC 59c1f22057fe1be0181b79a096084a4d7b846510
13 files changed +704 -63
compiler/forget/crates/forget_build_hir/src/build.rs
+65 -13
@@ -1,16 +1,15 @@
1 use std::collections::HashSet;
2
3 -use bumpalo::boxed::Box;
3 use bumpalo::collections::String;
4 use forget_diagnostics::Diagnostic;
5 use forget_estree::{
7 - AssignmentTarget, BinaryExpression, BlockStatement, Expression, ForInit, ForStatement,
8 - Function, FunctionExpression, IfStatement, JsValue, Literal, Pattern, Statement,
9 - VariableDeclarationKind,
6 + AssignmentTarget, BinaryExpression, BlockStatement, Expression, ExpressionOrSpread,
7 + ExpressionOrSuper, ForInit, ForStatement, Function, IfStatement, JsValue, Literal, Pattern,
8 + Statement, VariableDeclarationKind,
9 };
10 use forget_hir::{
12 - ArrayElement, BlockKind, BranchTerminal, Environment, ForTerminal, GotoKind, IdentifierOperand,
13 - InstrIx, InstructionKind, InstructionValue, LValue, LoadGlobal, LoadLocal, Operand,
11 + BlockKind, BranchTerminal, Environment, ForTerminal, GotoKind, IdentifierOperand, InstrIx,
12 + InstructionKind, InstructionValue, LValue, LoadGlobal, LoadLocal, Operand, PlaceOrSpread,
13 PrimitiveValue, TerminalValue,
14 };
15
@@ -27,7 +26,7 @@ use crate::error::BuildHIRError;
26 pub fn build<'a>(
27 env: &'a Environment<'a>,
28 fun: Function,
30 -) -> Result<Box<'a, forget_hir::Function<'a>>, Diagnostic> {
29 +) -> Result<Box<forget_hir::Function<'a>>, Diagnostic> {
30 let mut builder = Builder::new(env);
31
32 match fun.body {
@@ -77,7 +76,7 @@ pub fn build<'a>(
76 );
77
78 let body = builder.build()?;
80 - Ok(env.box_new(forget_hir::Function {
79 + Ok(Box::new(forget_hir::Function {
80 id: fun
81 .id
82 .map(|id| String::from_str_in(&id.name, &env.allocator)),
@@ -365,13 +364,13 @@ fn lower_expression<'a>(
364 for expr in expr.elements {
365 let element = match expr {
366 Some(forget_estree::ExpressionOrSpread::SpreadElement(expr)) => {
368 - Some(ArrayElement::Spread(Operand {
367 + Some(PlaceOrSpread::Spread(Operand {
368 ix: lower_expression(env, builder, expr.argument)?,
369 effect: None,
370 }))
371 }
372 Some(forget_estree::ExpressionOrSpread::Expression(expr)) => {
374 - Some(ArrayElement::Place(Operand {
373 + Some(PlaceOrSpread::Place(Operand {
374 ix: lower_expression(env, builder, expr)?,
375 effect: None,
376 }))
@@ -420,7 +419,37 @@ fn lower_expression<'a>(
419 }
420
421 Expression::FunctionExpression(expr) => {
423 - InstructionValue::Function(lower_function(env, builder, *expr)?)
422 + InstructionValue::Function(lower_function(env, builder, expr.function)?)
423 + }
424 +
425 + Expression::ArrowFunctionExpression(expr) => {
426 + InstructionValue::Function(lower_function(env, builder, expr.function)?)
427 + }
428 +
429 + Expression::CallExpression(expr) => {
430 + let callee_expr = match expr.callee {
431 + ExpressionOrSuper::Super(callee) => {
432 + return Err(Diagnostic::unsupported(
433 + BuildHIRError::UnsupportedSuperExpression,
434 + callee.range,
435 + ));
436 + }
437 + ExpressionOrSuper::Expression(callee) => callee,
438 + };
439 +
440 + if matches!(&callee_expr, Expression::MemberExpression(_)) {
441 + return Err(Diagnostic::todo("Support method calls", expr.range));
442 + }
443 +
444 + let callee = lower_expression(env, builder, callee_expr)?;
445 + let arguments = lower_arguments(env, builder, expr.arguments)?;
446 + InstructionValue::Call(forget_hir::Call {
447 + callee: Operand {
448 + ix: callee,
449 + effect: None,
450 + },
451 + arguments,
452 + })
453 }
454
455 _ => todo!("Lower expr {expr:#?}"),
@@ -428,12 +457,35 @@ fn lower_expression<'a>(
457 Ok(builder.push(value))
458 }
459
460 +fn lower_arguments<'a>(
461 + env: &'a Environment<'a>,
462 + builder: &mut Builder<'a>,
463 + args: Vec<ExpressionOrSpread>,
464 +) -> Result<bumpalo::collections::Vec<'a, PlaceOrSpread>, Diagnostic> {
465 + let mut arguments = env.vec_with_capacity(args.len());
466 + for arg in args {
467 + let element = match arg {
468 + forget_estree::ExpressionOrSpread::SpreadElement(arg) => {
469 + PlaceOrSpread::Spread(Operand {
470 + ix: lower_expression(env, builder, arg.argument)?,
471 + effect: None,
472 + })
473 + }
474 + forget_estree::ExpressionOrSpread::Expression(arg) => PlaceOrSpread::Place(Operand {
475 + ix: lower_expression(env, builder, arg)?,
476 + effect: None,
477 + }),
478 + };
479 + arguments.push(element);
480 + }
481 + Ok(arguments)
482 +}
483 +
484 fn lower_function<'a>(
485 env: &'a Environment<'a>,
486 builder: &mut Builder<'a>,
434 - expr: FunctionExpression,
487 + function: forget_estree::Function,
488 ) -> Result<forget_hir::FunctionExpression<'a>, Diagnostic> {
436 - let FunctionExpression { function, .. } = expr;
489 println!("get_context_identifiers() ...");
490 let context_identifiers = get_context_identifiers(env, &function);
491 println!("ok");
compiler/forget/crates/forget_build_hir/src/error.rs
+3
@@ -50,4 +50,7 @@ pub enum BuildHIRError {
50 /// ErrorSeverity::InvalidSyntax
51 #[error("Expected function to have a body")]
52 EmptyFunction,
53 +
54 + #[error("`super` is not suppported")]
55 + UnsupportedSuperExpression,
56 }
compiler/forget/crates/forget_estree_swc/src/lib.rs
+81 -17
@@ -8,9 +8,9 @@ use swc_core::common::errors::Handler;
8 use swc_core::common::source_map::Pos;
9 use swc_core::common::{FileName, FilePathMapping, Mark, SourceMap, Span, SyntaxContext, GLOBALS};
10 use swc_core::ecma::ast::{
11 - AssignOp, BinaryOp, BlockStmt, Decl, EsVersion, Expr, Function, Ident, Lit, MemberExpr,
12 - MemberProp, ModuleItem, Pat, PatOrExpr, Program, Stmt, UnaryOp, VarDecl, VarDeclKind,
13 - VarDeclOrExpr,
11 + AssignOp, BinaryOp, BlockStmt, BlockStmtOrExpr, Callee, Decl, EsVersion, Expr, ExprOrSpread,
12 + Function, Ident, Lit, MemberExpr, MemberProp, ModuleItem, Pat, PatOrExpr, Program, Stmt,
13 + UnaryOp, VarDecl, VarDeclKind, VarDeclOrExpr,
14 };
15 use swc_core::ecma::parser::Syntax;
16 use swc_core::ecma::transforms::base::resolver;
@@ -303,20 +303,55 @@ fn convert_expression(cx: &Context, expr: &Expr) -> forget_estree::Expression {
303 elements: expr
304 .elems
305 .iter()
306 - .map(|item| {
307 - // TODO: represent holes in array expressions
308 - let value = item.as_ref()?;
309 - match value.spread {
310 - Some(spread) => Some(forget_estree::ExpressionOrSpread::SpreadElement(
311 - Box::new(forget_estree::SpreadElement {
312 - argument: convert_expression(cx, &value.expr),
313 - loc: None,
314 - range: convert_span(&spread),
315 - }),
316 - )),
317 - None => Some(forget_estree::ExpressionOrSpread::Expression(
318 - convert_expression(cx, &value.expr),
319 - )),
306 + .map(|item| match item {
307 + Some(ExprOrSpread {
308 + spread: Some(spread),
309 + expr,
310 + }) => Some(forget_estree::ExpressionOrSpread::SpreadElement(Box::new(
311 + forget_estree::SpreadElement {
312 + argument: convert_expression(cx, expr),
313 + loc: None,
314 + range: convert_span(&spread),
315 + },
316 + ))),
317 + Some(ExprOrSpread { spread: None, expr }) => {
318 + Some(forget_estree::ExpressionOrSpread::Expression(
319 + convert_expression(cx, expr),
320 + ))
321 + }
322 + None => None,
323 + })
324 + .collect(),
325 + loc: None,
326 + range: convert_span(&expr.span),
327 + }))
328 + }
329 + Expr::Call(expr) => {
330 + forget_estree::Expression::CallExpression(Box::new(forget_estree::CallExpression {
331 + callee: match &expr.callee {
332 + Callee::Expr(callee) => {
333 + forget_estree::ExpressionOrSuper::Expression(convert_expression(cx, callee))
334 + }
335 + _ => todo!(),
336 + },
337 + arguments: expr
338 + .args
339 + .iter()
340 + .map(|arg| match arg {
341 + ExprOrSpread {
342 + spread: Some(spread),
343 + expr,
344 + } => forget_estree::ExpressionOrSpread::SpreadElement(Box::new(
345 + forget_estree::SpreadElement {
346 + argument: convert_expression(cx, expr),
347 + loc: None,
348 + range: convert_span(&spread),
349 + },
350 + )),
351 + ExprOrSpread { spread: None, expr } => {
352 + forget_estree::ExpressionOrSpread::Expression(convert_expression(
353 + cx, expr,
354 + ))
355 }
356 })
357 .collect(),
@@ -405,6 +440,35 @@ fn convert_expression(cx: &Context, expr: &Expr) -> forget_estree::Expression {
440 range: convert_span(&expr.function.span),
441 },
442 )),
443 + Expr::Arrow(expr) => forget_estree::Expression::ArrowFunctionExpression(Box::new(
444 + forget_estree::ArrowFunctionExpression {
445 + function: forget_estree::Function {
446 + id: None,
447 + body: match expr.body.as_ref() {
448 + BlockStmtOrExpr::Expr(body) => Some(
449 + forget_estree::FunctionBody::Expression(convert_expression(cx, body)),
450 + ),
451 + BlockStmtOrExpr::BlockStmt(body) => {
452 + Some(forget_estree::FunctionBody::BlockStatement(Box::new(
453 + convert_block_statement(cx, body),
454 + )))
455 + }
456 + },
457 + params: expr
458 + .params
459 + .iter()
460 + .map(|param| convert_pattern(cx, param))
461 + .collect(),
462 + is_generator: expr.is_generator,
463 + is_async: expr.is_async,
464 + loc: None,
465 + range: convert_span(&expr.span),
466 + },
467 + is_expression: true, // TODO
468 + loc: None,
469 + range: convert_span(&expr.span),
470 + },
471 + )),
472 _ => todo!("translate expression {:#?}", expr),
473 }
474 }
compiler/forget/crates/forget_fixtures/tests/fixtures/use-memo.js new
+6
@@ -0,0 +1,6 @@
1 +function Component(x) {
2 + const x = useMemo(() => {
3 + return y;
4 + });
5 + return x;
6 +}
compiler/forget/crates/forget_fixtures/tests/fixtures_test.rs
+3 -1
@@ -5,7 +5,7 @@ use bumpalo::Bump;
5 use forget_build_hir::build;
6 use forget_estree::{ModuleItem, Statement};
7 use forget_estree_swc::parse;
8 -use forget_hir::{Environment, Features, Print, Registry};
8 +use forget_hir::{inline_use_memo, Environment, Features, Print, Registry};
9 use forget_optimization::constant_propagation;
10 use forget_ssa::{eliminate_redundant_phis, enter_ssa};
11 use insta::{assert_snapshot, glob};
@@ -48,6 +48,8 @@ fn fixtures() {
48 println!("ok eliminate_redundant_phis");
49 constant_propagation(&environment, &mut fun).unwrap();
50 println!("ok constant_propagation");
51 + inline_use_memo(&environment, &mut fun).unwrap();
52 + println!("ok inline_use_memo");
53 fun.print(&fun.body, &mut output).unwrap();
54 println!("ok print");
55 }
compiler/forget/crates/forget_fixtures/tests/snapshots/fixtures_test__fixtures@use-memo.js.snap new
+34
@@ -0,0 +1,34 @@
1 +---
2 +source: crates/forget_fixtures/tests/fixtures_test.rs
3 +expression: "format!(\"Input:\\n{input}\\n\\nOutput:\\n{output}\")"
4 +input_file: crates/forget_fixtures/tests/fixtures/use-memo.js
5 +---
6 +Input:
7 +function Component(x) {
8 + const x = useMemo(() => {
9 + return y;
10 + });
11 + return x;
12 +}
13 +
14 +
15 +Output:
16 +function Component(
17 + unknown x$1,
18 +)
19 +entry bb0
20 +bb0 (block)
21 + [0] #0 = LoadGlobal useMemo
22 + [1] #6 = DeclareLocal Let unknown t$3
23 + [2] Label block=bb1 fallthrough=bb6
24 +bb1 (block)
25 + predecessors: bb0
26 + [3] #7 = LoadGlobal y
27 + [4] #9 = StoreLocal Reassign unknown t$3 = unknown #7
28 + [5] Goto bb6
29 +bb6 (block)
30 + predecessors: bb1
31 + [6] #2 = LoadLocal unknown t$3
32 + [7] #3 = StoreLocal Const unknown x$2 = unknown #2
33 + [8] #4 = LoadLocal unknown x$2
34 + [9] Return unknown #4
compiler/forget/crates/forget_hir/src/function.rs
+103 -19
@@ -2,7 +2,7 @@ use bumpalo::collections::{String, Vec};
2 use forget_diagnostics::Diagnostic;
3 use indexmap::IndexMap;
4
5 -use crate::{BasicBlock, BlockId, IdentifierOperand, Instruction};
5 +use crate::{BasicBlock, BlockId, FunctionExpression, IdentifierOperand, InstrIx, Instruction};
6
7 /// Represents either a React function or a function expression
8 #[derive(Debug)]
@@ -32,6 +32,27 @@ pub struct HIR<'a> {
32 pub instructions: Vec<'a, Instruction<'a>>,
33 }
34
35 +impl<'a> HIR<'a> {
36 + pub fn inline(&mut self, other: FunctionExpression<'a>) -> () {
37 + let offset = self.instructions.len();
38 + for mut instr in other.lowered_function.body.instructions.into_iter() {
39 + instr.each_operand(|operand| {
40 + operand.ix = InstrIx::new((offset + usize::from(operand.ix)) as u32);
41 + });
42 + self.instructions.push(instr);
43 + }
44 + for mut block in other.lowered_function.body.blocks.into_iter() {
45 + for ix in block.instructions.iter_mut() {
46 + *ix = InstrIx::new((offset + usize::from(*ix)) as u32);
47 + }
48 + block.terminal.value.each_operand(|operand| {
49 + operand.ix = InstrIx::new((offset + usize::from(operand.ix)) as u32);
50 + });
51 + self.blocks.insert(block);
52 + }
53 + }
54 +}
55 +
56 #[derive(Default, Debug)]
57 pub struct Blocks<'a> {
58 data: IndexMap<BlockId, Option<Box<BasicBlock<'a>>>>,
@@ -66,6 +87,14 @@ impl<'a> Blocks<'a> {
87 self.data.remove(&id).unwrap().unwrap()
88 }
89
90 + pub fn extend(&mut self, other: Self) {
91 + self.data.extend(other.data);
92 + }
93 +
94 + pub fn into_iter(self) -> BlocksIntoIter<'a> {
95 + BlocksIntoIter::new(self.data.into_iter())
96 + }
97 +
98 pub fn block(&self, id: BlockId) -> &BasicBlock<'a> {
99 self.data.get(&id).unwrap().as_ref().unwrap()
100 }
@@ -82,6 +111,33 @@ impl<'a> Blocks<'a> {
111 BlocksIterMut::new(self.data.iter_mut())
112 }
113 }
114 +pub struct BlocksIntoIter<'a> {
115 + iter: indexmap::map::IntoIter<BlockId, Option<Box<BasicBlock<'a>>>>,
116 +}
117 +
118 +impl<'a> BlocksIntoIter<'a> {
119 + fn new(iter: indexmap::map::IntoIter<BlockId, Option<Box<BasicBlock<'a>>>>) -> Self {
120 + Self { iter }
121 + }
122 +}
123 +
124 +impl<'a> Iterator for BlocksIntoIter<'a> {
125 + type Item = Box<BasicBlock<'a>>;
126 +
127 + fn size_hint(&self) -> (usize, Option<usize>) {
128 + self.iter.size_hint()
129 + }
130 +
131 + fn next(&mut self) -> Option<Self::Item> {
132 + loop {
133 + match self.iter.next() {
134 + Some((_, Some(next))) => return Some(next),
135 + Some((_, None)) => continue,
136 + None => return None,
137 + }
138 + }
139 + }
140 +}
141
142 pub struct BlocksIter<'b, 'a> {
143 iter: indexmap::map::Iter<'b, BlockId, Option<Box<BasicBlock<'a>>>>,
@@ -136,6 +192,7 @@ impl<'b, 'a> Iterator for BlocksIterMut<'b, 'a> {
192 pub struct BlockRewriter<'blocks, 'a> {
193 blocks: &'blocks mut Blocks<'a>,
194 current: BlockId,
195 + new_blocks: std::vec::Vec<Box<BasicBlock<'a>>>,
196 }
197
198 impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
@@ -143,6 +200,7 @@ impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
200 Self {
201 blocks,
202 current: entry,
203 + new_blocks: Default::default(),
204 }
205 }
206
@@ -150,17 +208,28 @@ impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
208 where
209 F: FnMut(Box<BasicBlock<'a>>, &mut Self) -> BlockRewriterAction<'a>,
210 {
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));
211 + let mut keys = self.blocks.block_ids();
212 + loop {
213 + for block_id in keys {
214 + self.current = block_id;
215 + let block = self.blocks.data.get_mut(&block_id).unwrap().take().unwrap();
216 + match f(block, self) {
217 + BlockRewriterAction::Keep(block) => {
218 + self.blocks.data.insert(block_id, Some(block));
219 + }
220 + BlockRewriterAction::Remove => {
221 + self.blocks.data.remove(&block_id);
222 + }
223 }
161 - BlockRewriterAction::Remove => {
162 - self.blocks.data.remove(&block_id);
224 + }
225 + if !self.new_blocks.is_empty() {
226 + keys = self.new_blocks.iter().map(|block| block.id).collect();
227 + for block in self.new_blocks.drain(..) {
228 + self.blocks.insert(block);
229 }
230 + continue;
231 + } else {
232 + break;
233 }
234 }
235 }
@@ -169,17 +238,28 @@ impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
238 where
239 F: FnMut(Box<BasicBlock<'a>>, &mut Self) -> Result<BlockRewriterAction<'a>, Diagnostic>,
240 {
172 - let keys = self.blocks.block_ids();
173 - for block_id in keys {
174 - self.current = block_id;
175 - let block = self.blocks.data.get_mut(&block_id).unwrap().take().unwrap();
176 - match f(block, self)? {
177 - BlockRewriterAction::Keep(block) => {
178 - self.blocks.data.insert(block_id, Some(block));
241 + let mut keys = self.blocks.block_ids();
242 + loop {
243 + for block_id in keys {
244 + self.current = block_id;
245 + let block = self.blocks.data.get_mut(&block_id).unwrap().take().unwrap();
246 + match f(block, self)? {
247 + BlockRewriterAction::Keep(block) => {
248 + self.blocks.data.insert(block_id, Some(block));
249 + }
250 + BlockRewriterAction::Remove => {
251 + self.blocks.data.remove(&block_id);
252 + }
253 }
180 - BlockRewriterAction::Remove => {
181 - self.blocks.data.remove(&block_id);
254 + }
255 + if !self.new_blocks.is_empty() {
256 + keys = self.new_blocks.iter().map(|block| block.id).collect();
257 + for block in self.new_blocks.drain(..) {
258 + self.blocks.insert(block);
259 }
260 + continue;
261 + } else {
262 + break;
263 }
264 }
265 Ok(())
@@ -199,6 +279,10 @@ impl<'blocks, 'a> BlockRewriter<'blocks, 'a> {
279 assert_ne!(block_id, self.current);
280 self.blocks.block_mut(block_id)
281 }
282 +
283 + pub fn add_block(&mut self, block: Box<BasicBlock<'a>>) {
284 + self.new_blocks.push(block);
285 + }
286 }
287
288 pub enum BlockRewriterAction<'a> {
compiler/forget/crates/forget_hir/src/initialize.rs
+15 -1
@@ -5,13 +5,14 @@ use thiserror::Error;
5
6 use crate::{
7 BlockId, BlockRewriter, BlockRewriterAction, Blocks, GotoKind, GotoTerminal,
8 - InstructionIdGenerator, TerminalValue, HIR,
8 + InstructionIdGenerator, InstructionValue, 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
13 /// shape of the control-flow graph.
14 pub fn initialize_hir<'a>(hir: &mut HIR<'a>) -> Result<(), Diagnostic> {
15 + prune_tombstones(hir);
16 reverse_postorder_blocks(hir);
17 remove_unreachable_for_updates(hir);
18 remove_unreachable_fallthroughs(hir);
@@ -21,6 +22,16 @@ pub fn initialize_hir<'a>(hir: &mut HIR<'a>) -> Result<(), Diagnostic> {
22 Ok(())
23 }
24
25 +pub fn prune_tombstones<'a>(hir: &mut HIR<'a>) {
26 + for block in hir.blocks.iter_mut() {
27 + block.instructions.retain(|ix| {
28 + let instr = &hir.instructions[usize::from(*ix)];
29 + // Retain all values that are not the tombstone
30 + !matches!(instr.value, InstructionValue::Tombstone)
31 + });
32 + }
33 +}
34 +
35 /// Modifies the HIR to put the blocks in reverse postorder, with predecessors before
36 /// successors (except for the case of loops)
37 pub fn reverse_postorder_blocks<'a>(hir: &mut HIR<'a>) {
@@ -56,6 +67,9 @@ pub fn reverse_postorder_blocks<'a>(hir: &mut HIR<'a>) {
67 TerminalValue::Goto(terminal) => {
68 visit(terminal.block, hir, visited, postorder);
69 }
70 + TerminalValue::Label(terminal) => {
71 + visit(terminal.block, hir, visited, postorder);
72 + }
73 TerminalValue::Return(..) => { /* no-op */ }
74 TerminalValue::Unsupported(..) => {
75 panic!("Unexpected unsupported terminal")
compiler/forget/crates/forget_hir/src/inline_use_memo.rs new
+258
@@ -0,0 +1,258 @@
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 +}
compiler/forget/crates/forget_hir/src/instruction.rs
+58 -5
@@ -2,7 +2,6 @@ use std::cell::RefCell;
2 use std::fmt::Display;
3 use std::rc::Rc;
4
5 -use bumpalo::boxed::Box;
5 use bumpalo::collections::{String, Vec};
6 use forget_estree::BinaryOperator;
7
@@ -31,6 +30,7 @@ impl<'a> Instruction<'a> {
30 }
31 InstructionValue::Array(_)
32 | InstructionValue::Binary(_)
33 + | InstructionValue::Call(_)
34 | InstructionValue::LoadContext(_)
35 | InstructionValue::LoadGlobal(_)
36 | InstructionValue::LoadLocal(_)
@@ -56,6 +56,7 @@ impl<'a> Instruction<'a> {
56 }
57 InstructionValue::Array(_)
58 | InstructionValue::Binary(_)
59 + | InstructionValue::Call(_)
60 | InstructionValue::LoadContext(_)
61 | InstructionValue::LoadGlobal(_)
62 | InstructionValue::LoadLocal(_)
@@ -74,6 +75,7 @@ impl<'a> Instruction<'a> {
75 InstructionValue::LoadLocal(instr) => f(&mut instr.place),
76 InstructionValue::Array(_)
77 | InstructionValue::Binary(_)
78 + | InstructionValue::Call(_)
79 | InstructionValue::DeclareContext(_)
80 | InstructionValue::DeclareLocal(_)
81 | InstructionValue::LoadContext(_)
@@ -84,6 +86,51 @@ impl<'a> Instruction<'a> {
86 | InstructionValue::Tombstone => {}
87 }
88 }
89 +
90 + pub fn each_operand<F>(&mut self, mut f: F) -> ()
91 + where
92 + F: FnMut(&mut Operand) -> (),
93 + {
94 + match &mut self.value {
95 + InstructionValue::Array(value) => {
96 + for item in &mut value.elements {
97 + match item {
98 + Some(PlaceOrSpread::Place(item)) => f(item),
99 + Some(PlaceOrSpread::Spread(item)) => f(item),
100 + None => {}
101 + }
102 + }
103 + }
104 + InstructionValue::Binary(value) => {
105 + f(&mut value.left);
106 + f(&mut value.right);
107 + }
108 + InstructionValue::Call(value) => {
109 + f(&mut value.callee);
110 + for arg in &mut value.arguments {
111 + match arg {
112 + PlaceOrSpread::Place(item) => f(item),
113 + PlaceOrSpread::Spread(item) => f(item),
114 + }
115 + }
116 + }
117 + InstructionValue::StoreLocal(value) => {
118 + f(&mut value.value);
119 + }
120 + InstructionValue::Function(value) => {
121 + for dep in &mut value.dependencies {
122 + f(dep)
123 + }
124 + }
125 + InstructionValue::DeclareContext(_)
126 + | InstructionValue::LoadContext(_)
127 + | InstructionValue::LoadGlobal(_)
128 + | InstructionValue::DeclareLocal(_)
129 + | InstructionValue::LoadLocal(_)
130 + | InstructionValue::Primitive(_)
131 + | InstructionValue::Tombstone => {}
132 + }
133 + }
134 }
135
136 #[derive(Debug)]
@@ -91,7 +138,7 @@ pub enum InstructionValue<'a> {
138 Array(Array<'a>),
139 // Await(Await<'a>),
140 Binary(Binary),
94 - // Call(Call<'a>),
141 + Call(Call<'a>),
142 // ComputedDelete(ComputedDelete<'a>),
143 // ComputedLoad(ComputedLoad<'a>),
144 // ComputedStore(ComputedStore<'a>),
@@ -125,11 +172,11 @@ pub enum InstructionValue<'a> {
172
173 #[derive(Debug)]
174 pub struct Array<'a> {
128 - pub elements: Vec<'a, Option<ArrayElement>>,
175 + pub elements: Vec<'a, Option<PlaceOrSpread>>,
176 }
177
178 #[derive(Debug)]
132 -pub enum ArrayElement {
179 +pub enum PlaceOrSpread {
180 Place(Operand),
181 Spread(Operand),
182 }
@@ -141,10 +188,16 @@ pub struct Binary {
188 pub right: Operand,
189 }
190
191 +#[derive(Debug)]
192 +pub struct Call<'a> {
193 + pub callee: Operand,
194 + pub arguments: Vec<'a, PlaceOrSpread>,
195 +}
196 +
197 #[derive(Debug)]
198 pub struct FunctionExpression<'a> {
199 pub dependencies: Vec<'a, Operand>,
147 - pub lowered_function: Box<'a, Function<'a>>,
200 + pub lowered_function: Box<Function<'a>>,
201 }
202
203 #[derive(Debug, Clone, PartialEq, Eq)]
compiler/forget/crates/forget_hir/src/lib.rs
+2
@@ -4,6 +4,7 @@ mod features;
4 mod function;
5 mod id_types;
6 mod initialize;
7 +mod inline_use_memo;
8 mod instruction;
9 mod merge_consecutive_blocks;
10 mod print;
@@ -21,6 +22,7 @@ pub use initialize::{
22 remove_unreachable_do_while_statements, remove_unreachable_fallthroughs,
23 remove_unreachable_for_updates, reverse_postorder_blocks,
24 };
25 +pub use inline_use_memo::inline_use_memo;
26 pub use instruction::*;
27 pub use merge_consecutive_blocks::merge_consecutive_blocks;
28 pub use print::Print;
compiler/forget/crates/forget_hir/src/print.rs
+44 -6
@@ -3,8 +3,8 @@ use std::fmt::{Result, Write};
3 use forget_utils::ensure_sufficient_stack;
4
5 use crate::{
6 - ArrayElement, BasicBlock, Function, Identifier, IdentifierOperand, Instruction,
7 - InstructionValue, LValue, Operand, Phi, PrimitiveValue, Terminal, TerminalValue, HIR,
6 + BasicBlock, Function, Identifier, IdentifierOperand, Instruction, InstructionValue, LValue,
7 + Operand, Phi, PlaceOrSpread, PrimitiveValue, Terminal, TerminalValue, HIR,
8 };
9
10 /// Trait for HIR types to describe how they print themselves.
@@ -16,6 +16,14 @@ pub trait Print<'a> {
16 fn print(&self, hir: &HIR<'a>, out: &mut impl Write) -> Result;
17 }
18
19 +impl<'a> Function<'a> {
20 + pub fn debug(&self) {
21 + let mut out = String::new();
22 + self.print(&self.body, &mut out).unwrap();
23 + println!("{out}");
24 + }
25 +}
26 +
27 impl<'a> Print<'a> for Function<'a> {
28 fn print(&self, hir: &HIR<'a>, out: &mut impl Write) -> Result {
29 ensure_sufficient_stack(|| {
@@ -61,6 +69,10 @@ impl<'a> Print<'a> for BasicBlock<'a> {
69 writeln!(out)?;
70 }
71 for ix in &self.instructions {
72 + if usize::from(*ix) >= hir.instructions.len() {
73 + writeln!(out, " <out of bounds {}>", ix)?;
74 + continue;
75 + }
76 let instr = &hir.instructions[usize::from(*ix)];
77 write!(out, " {} {} = ", instr.id, ix)?;
78 instr.value.print(hir, out)?;
@@ -114,6 +126,18 @@ impl<'a> Print<'a> for InstructionValue<'a> {
126 }
127 write!(out, "]")?;
128 }
129 + InstructionValue::Call(value) => {
130 + write!(out, "Call ")?;
131 + value.callee.print(hir, out)?;
132 + write!(out, "(")?;
133 + for (ix, arg) in value.arguments.iter().enumerate() {
134 + if ix != 0 {
135 + write!(out, ", ")?;
136 + }
137 + arg.print(hir, out)?;
138 + }
139 + write!(out, ")")?;
140 + }
141 InstructionValue::LoadGlobal(value) => {
142 write!(out, "LoadGlobal {}", &value.name)?;
143 }
@@ -187,11 +211,11 @@ impl<'a> Print<'a> for InstructionValue<'a> {
211 }
212 }
213
190 -impl<'a> Print<'a> for ArrayElement {
214 +impl<'a> Print<'a> for PlaceOrSpread {
215 fn print(&self, hir: &HIR<'a>, out: &mut impl Write) -> Result {
216 match self {
193 - ArrayElement::Place(place) => place.print(hir, out),
194 - ArrayElement::Spread(place) => {
217 + PlaceOrSpread::Place(place) => place.print(hir, out),
218 + PlaceOrSpread::Spread(place) => {
219 write!(out, "...")?;
220 place.print(hir, out)?;
221 Ok(())
@@ -298,13 +322,27 @@ impl<'a> Print<'a> for TerminalValue<'a> {
322 terminal.init,
323 terminal.test,
324 match terminal.update {
301 - Some(fallthrough) => format!("{fallthrough}"),
325 + Some(update) => format!("{update}"),
326 None => "<none>".to_string(),
327 },
328 terminal.body,
329 terminal.fallthrough,
330 )?;
331 }
332 + TerminalValue::Label(terminal) => {
333 + write!(
334 + out,
335 + "Label block={} fallthrough={}",
336 + terminal.block,
337 + match terminal.fallthrough {
338 + Some(fallthrough) => format!("{fallthrough}"),
339 + None => "<none>".to_string(),
340 + },
341 + )?;
342 + }
343 + TerminalValue::Unsupported(_) => {
344 + write!(out, "Unsupported")?;
345 + }
346 _ => write!(out, "{:?}", self)?,
347 }
348 Ok(())
compiler/forget/crates/forget_hir/src/terminal.rs
+32 -1
@@ -17,7 +17,7 @@ pub enum TerminalValue<'a> {
17 For(ForTerminal),
18 Goto(GotoTerminal),
19 If(IfTerminal),
20 - // Label(LabelTerminal),
20 + Label(LabelTerminal),
21 // Logical(LogicalTerminal),
22 // Optional(OptionalTerminal),
23 Return(ReturnTerminal),
@@ -41,6 +41,12 @@ impl<'a> TerminalValue<'a> {
41 _ => None,
42 }
43 }
44 + Self::Label(terminal) => {
45 + terminal.fallthrough = match terminal.fallthrough {
46 + Some(fallthrough) => f(fallthrough),
47 + _ => None,
48 + }
49 + }
50 Self::DoWhile(DoWhileTerminal { fallthrough, .. })
51 | Self::For(ForTerminal { fallthrough, .. }) => {
52 // statically detect if fallthrough is changed to Option so
@@ -69,12 +75,31 @@ impl<'a> TerminalValue<'a> {
75 Self::Goto(terminal) => {
76 vec![terminal.block]
77 }
78 + Self::Label(terminal) => {
79 + vec![terminal.block]
80 + }
81 Self::Return(_) => {
82 vec![]
83 }
84 Self::Unsupported(_) => panic!("Unexpected unsupported terminal"),
85 }
86 }
87 +
88 + pub fn each_operand<F>(&mut self, mut f: F) -> ()
89 + where
90 + F: FnMut(&mut Operand) -> (),
91 + {
92 + match self {
93 + TerminalValue::Branch(terminal) => f(&mut terminal.test),
94 + TerminalValue::If(terminal) => f(&mut terminal.test),
95 + TerminalValue::Return(terminal) => f(&mut terminal.value),
96 + TerminalValue::DoWhile(_)
97 + | TerminalValue::For(_)
98 + | TerminalValue::Label(_)
99 + | TerminalValue::Goto(_)
100 + | TerminalValue::Unsupported(_) => {}
101 + }
102 + }
103 }
104
105 #[derive(Debug)]
@@ -129,3 +154,9 @@ pub struct ForTerminal {
154 pub body: BlockId,
155 pub fallthrough: BlockId,
156 }
157 +
158 +#[derive(Debug)]
159 +pub struct LabelTerminal {
160 + pub block: BlockId,
161 + pub fallthrough: Option<BlockId>,
162 +}