@samitouri / QOS-React-2 / commits / dbe1af601b

[rust] Start of function expression support

Start of function expression support: * Basic structure for representing function expressions in the HIR * Printer support * swc -> estree -> hir conversion for function expression _bodies_. Dependencies and context are not handled yet.

Joe Savona committed Jul 13, 2023 at 11:14 UTC dbe1af601bc9282382a6e1ec17cd3f5c2c94228d
14 files changed +321 -193
compiler/forget/Cargo.lock
+2
@@ -638,6 +638,7 @@ dependencies = [
638 "estree",
639 "indexmap 2.0.0",
640 "serde",
641 + "utils",
642 ]
643
644 [[package]]
@@ -2806,6 +2807,7 @@ name = "utils"
2807 version = "0.1.0"
2808 dependencies = [
2809 "bumpalo",
2810 + "stacker",
2811 ]
2812
2813 [[package]]
compiler/forget/crates/build-hir/src/build.rs
+32 -18
@@ -1,8 +1,10 @@
1 -use bumpalo::collections::{String, Vec};
1 +use bumpalo::{
2 + boxed::Box,
3 + collections::{String, Vec},
4 +};
5 use estree::{
6 AssignmentTarget, BinaryExpression, BlockStatement, Expression, ForInit, ForStatement,
4 - FunctionDeclaration, IfStatement, JsValue, Literal, Pattern, Statement,
5 - VariableDeclarationKind,
7 + FunctionExpression, IfStatement, JsValue, Literal, Pattern, Statement, VariableDeclarationKind,
8 };
9 use hir::{
10 ArrayElement, BlockKind, BranchTerminal, Environment, ForTerminal, Function, GotoKind,
@@ -24,11 +26,11 @@ use crate::{
26 /// that is not yet supported.
27 pub fn build<'a>(
28 env: &'a Environment<'a>,
27 - fun: FunctionDeclaration,
28 -) -> Result<&'a mut Function<'a>, BuildDiagnostic> {
29 + fun: estree::Function,
30 +) -> Result<Box<'a, Function<'a>>, BuildDiagnostic> {
31 let mut builder = Builder::new(env);
32
31 - match fun.function.body {
33 + match fun.body {
34 Some(estree::FunctionBody::BlockStatement(body)) => {
35 lower_block_statement(env, &mut builder, *body)?
36 }
@@ -44,8 +46,8 @@ pub fn build<'a>(
46 }
47 }
48
47 - let mut params = Vec::with_capacity_in(fun.function.params.len(), &env.allocator);
48 - for param in fun.function.params {
49 + let mut params = Vec::with_capacity_in(fun.params.len(), &env.allocator);
50 + for param in fun.params {
51 match param {
52 Pattern::Identifier(param) => {
53 let identifier = lower_identifier_for_assignment(
@@ -76,16 +78,18 @@ pub fn build<'a>(
78 );
79
80 let body = builder.build()?;
79 - Ok(env.alloc(Function {
80 - id: fun
81 - .function
82 - .id
83 - .map(|id| String::from_str_in(&id.name, &env.allocator)),
84 - body,
85 - params,
86 - is_async: fun.function.is_async,
87 - is_generator: fun.function.is_generator,
88 - }))
81 + Ok(Box::new_in(
82 + Function {
83 + id: fun
84 + .id
85 + .map(|id| String::from_str_in(&id.name, &env.allocator)),
86 + body,
87 + params,
88 + is_async: fun.is_async,
89 + is_generator: fun.is_generator,
90 + },
91 + &env.allocator,
92 + ))
93 }
94
95 fn lower_block_statement<'a>(
@@ -421,6 +425,16 @@ fn lower_expression<'a>(
425 })
426 }
427
428 + Expression::FunctionExpression(expr) => {
429 + let FunctionExpression { function, .. } = *expr;
430 + let fun = build(env, function)?;
431 + InstructionValue::Function(hir::FunctionExpression {
432 + // TODO: collect dependencies!
433 + dependencies: Vec::new_in(&env.allocator),
434 + lowered_function: fun,
435 + })
436 + }
437 +
438 _ => todo!("Lower expr {expr:#?}"),
439 };
440 Ok(builder.push(value))
compiler/forget/crates/estree-codegen/src/ecmascript.json
+8
@@ -43,6 +43,14 @@
43 "type": "bool",
44 "optional": true,
45 "rename": "async"
46 + },
47 + "loc": {
48 + "type": "Option<SourceLocation>",
49 + "optional": true
50 + },
51 + "range": {
52 + "type": "Option<SourceRange>",
53 + "optional": true
54 }
55 }
56 },
compiler/forget/crates/estree-swc/src/lib.rs
+34 -24
@@ -6,8 +6,9 @@ use swc_core::common::errors::Handler;
6 use swc_core::common::source_map::Pos;
7 use swc_core::common::{FileName, FilePathMapping, Mark, SourceMap, Span, SyntaxContext, GLOBALS};
8 use swc_core::ecma::ast::{
9 - AssignOp, BinaryOp, BlockStmt, Decl, EsVersion, Expr, Ident, Lit, MemberExpr, MemberProp,
10 - ModuleItem, Pat, PatOrExpr, Program, Stmt, UnaryOp, VarDecl, VarDeclKind, VarDeclOrExpr,
9 + AssignOp, BinaryOp, BlockStmt, Decl, EsVersion, Expr, Function, Ident, Lit, MemberExpr,
10 + MemberProp, ModuleItem, Pat, PatOrExpr, Program, Stmt, UnaryOp, VarDecl, VarDeclKind,
11 + VarDeclOrExpr,
12 };
13 use swc_core::ecma::parser::Syntax;
14 use swc_core::ecma::transforms::base::resolver;
@@ -122,32 +123,34 @@ fn convert_block_statement(cx: &Context, stmt: &BlockStmt) -> estree::BlockState
123 }
124 }
125
126 +fn convert_function(cx: &Context, id: Option<&Ident>, fun: &Function) -> estree::Function {
127 + estree::Function {
128 + id: id.map(|id| estree::Identifier {
129 + name: id.sym.to_string(),
130 + binding: convert_binding(cx, id.span.ctxt),
131 + loc: None,
132 + range: convert_span(&id.span),
133 + }),
134 + params: fun
135 + .params
136 + .iter()
137 + .map(|param| convert_pattern(cx, &param.pat))
138 + .collect(),
139 + body: fun.body.as_ref().map(|body| {
140 + estree::FunctionBody::BlockStatement(Box::new(convert_block_statement(cx, body)))
141 + }),
142 + is_async: fun.is_async,
143 + is_generator: fun.is_generator,
144 + loc: None,
145 + range: convert_span(&fun.span),
146 + }
147 +}
148 +
149 fn convert_statement(cx: &Context, stmt: &Stmt) -> estree::Statement {
150 match stmt {
151 Stmt::Decl(Decl::Fn(item)) => {
128 - let name = item.ident.sym.to_string();
152 estree::Statement::FunctionDeclaration(Box::new(estree::FunctionDeclaration {
130 - function: estree::Function {
131 - id: Some(estree::Identifier {
132 - name,
133 - binding: convert_binding(cx, item.ident.span.ctxt),
134 - loc: None,
135 - range: convert_span(&item.ident.span),
136 - }),
137 - params: item
138 - .function
139 - .params
140 - .iter()
141 - .map(|param| convert_pattern(cx, &param.pat))
142 - .collect(),
143 - body: item.function.body.as_ref().map(|body| {
144 - estree::FunctionBody::BlockStatement(Box::new(convert_block_statement(
145 - cx, body,
146 - )))
147 - }),
148 - is_async: item.function.is_async,
149 - is_generator: item.function.is_generator,
150 - },
153 + function: convert_function(cx, Some(&item.ident), &item.function),
154 loc: None,
155 range: convert_span(&item.function.span),
156 }))
@@ -368,6 +371,13 @@ fn convert_expression(cx: &Context, expr: &Expr) -> estree::Expression {
371 Expr::Member(expr) => {
372 estree::Expression::MemberExpression(Box::new(convert_member_expression(cx, expr)))
373 }
374 + Expr::Fn(expr) => {
375 + estree::Expression::FunctionExpression(Box::new(estree::FunctionExpression {
376 + function: convert_function(cx, expr.ident.as_ref(), &expr.function),
377 + loc: None,
378 + range: convert_span(&expr.function.span),
379 + }))
380 + }
381 _ => todo!("translate expression {:#?}", expr),
382 }
383 }
compiler/forget/crates/estree/src/generated.rs
+4
@@ -23,6 +23,10 @@ pub struct Function {
23 #[serde(rename = "async")]
24 #[serde(default)]
25 pub is_async: bool,
26 + #[serde(default)]
27 + pub loc: Option<SourceLocation>,
28 + #[serde(default)]
29 + pub range: Option<SourceRange>,
30 }
31 #[derive(Serialize, Deserialize, Clone, Debug)]
32 pub struct RegExpValue {
compiler/forget/crates/fixtures/tests/fixtures/function-expressions.js new
+5
@@ -0,0 +1,5 @@
1 +function Component(props) {
2 + const foo = function foo(x) {
3 + return x + 1;
4 + };
5 +}
compiler/forget/crates/fixtures/tests/fixtures_test.rs
+2 -1
@@ -33,7 +33,7 @@ fn fixtures() {
33 if ix != 0 {
34 output.push_str("\n\n");
35 }
36 - match build(&environment, *fun) {
36 + match build(&environment, fun.function) {
37 Ok(mut fun) => {
38 enter_ssa(&environment, &mut fun).unwrap();
39 eliminate_redundant_phis(&environment, &mut fun);
@@ -56,6 +56,7 @@ fn fixtures() {
56 }
57 }
58
59 + let output = output.trim();
60 assert_snapshot!(format!("Input:\n{input}\n\nOutput:\n{output}"));
61 });
62 }
compiler/forget/crates/fixtures/tests/snapshots/fixtures_test__fixtures@function-expressions.js.snap new
+34
@@ -0,0 +1,34 @@
1 +---
2 +source: crates/fixtures/tests/fixtures_test.rs
3 +expression: "format!(\"Input:\\n{input}\\n\\nOutput:\\n{output}\")"
4 +input_file: crates/fixtures/tests/fixtures/function-expressions.js
5 +---
6 +Input:
7 +function Component(props) {
8 + const foo = function foo(x) {
9 + return x + 1;
10 + };
11 +}
12 +
13 +
14 +Output:
15 +function Component(
16 + unknown props$3,
17 +)
18 +entry bb0
19 +bb0 (block)
20 + [0] #0 = Function @deps[] @context[]:
21 + function foo(
22 + unknown x$0,
23 + )
24 + entry bb1
25 + bb1 (block)
26 + [0] #0 = LoadLocal unknown x$0
27 + [1] #1 = 1
28 + [2] #2 = Binary unknown #0 + unknown #1
29 + [3] Return unknown #2
30 + [1] #1 = StoreLocal Const unknown foo$4 = unknown #0
31 + [2] #2 = <undefined>
32 + [3] Return unknown #2
33 +
34 +
compiler/forget/crates/hir/Cargo.toml
+1
@@ -10,3 +10,4 @@ bumpalo = { version = "3.13.0", features = ["boxed", "collections"] }
10 estree = { path = "../estree" }
11 indexmap = "2.0.0"
12 serde = "1.0.164"
13 +utils = { path = "../utils" }
compiler/forget/crates/hir/src/instruction.rs
+30 -19
@@ -1,9 +1,12 @@
1 use std::{cell::RefCell, fmt::Display, rc::Rc};
2
3 -use bumpalo::collections::{String, Vec};
3 +use bumpalo::{
4 + boxed::Box,
5 + collections::{String, Vec},
6 +};
7 use estree::BinaryOperator;
8
6 -use crate::{IdentifierId, InstrIx, InstructionId, ScopeId, Type};
9 +use crate::{Function, IdentifierId, InstrIx, InstructionId, ScopeId, Type};
10
11 #[derive(Debug)]
12 pub struct Instruction<'a> {
@@ -17,22 +20,23 @@ impl<'a> Instruction<'a> {
20 F: FnMut(&mut LValue<'a>) -> (),
21 {
22 match &mut self.value {
20 - InstructionValue::Array(_) => {}
21 - InstructionValue::Binary(_) => {}
23 InstructionValue::DeclareContext(instr) => {
24 f(&mut instr.lvalue);
25 }
26 InstructionValue::DeclareLocal(instr) => {
27 f(&mut instr.lvalue);
28 }
28 - InstructionValue::LoadContext(_) => {}
29 - InstructionValue::LoadGlobal(_) => {}
30 - InstructionValue::LoadLocal(_) => {}
31 - InstructionValue::Primitive(_) => {}
29 InstructionValue::StoreLocal(instr) => {
30 f(&mut instr.lvalue);
31 }
35 - InstructionValue::Tombstone => {}
32 + InstructionValue::Array(_)
33 + | InstructionValue::Binary(_)
34 + | InstructionValue::LoadContext(_)
35 + | InstructionValue::LoadGlobal(_)
36 + | InstructionValue::LoadLocal(_)
37 + | InstructionValue::Primitive(_)
38 + | InstructionValue::Function(_)
39 + | InstructionValue::Tombstone => {}
40 }
41 }
42
@@ -41,16 +45,17 @@ impl<'a> Instruction<'a> {
45 F: FnMut(&mut IdentifierOperand<'a>) -> (),
46 {
47 match &mut self.value {
44 - InstructionValue::Array(_) => {}
45 - InstructionValue::Binary(_) => {}
46 - InstructionValue::DeclareContext(_) => {}
47 - InstructionValue::DeclareLocal(_) => {}
48 - InstructionValue::LoadContext(_) => {}
49 - InstructionValue::LoadGlobal(_) => {}
48 InstructionValue::LoadLocal(instr) => f(&mut instr.place),
51 - InstructionValue::Primitive(_) => {}
52 - InstructionValue::StoreLocal(_) => {}
53 - InstructionValue::Tombstone => {}
49 + InstructionValue::Array(_)
50 + | InstructionValue::Binary(_)
51 + | InstructionValue::DeclareContext(_)
52 + | InstructionValue::DeclareLocal(_)
53 + | InstructionValue::LoadContext(_)
54 + | InstructionValue::LoadGlobal(_)
55 + | InstructionValue::Primitive(_)
56 + | InstructionValue::StoreLocal(_)
57 + | InstructionValue::Function(_)
58 + | InstructionValue::Tombstone => {}
59 }
60 }
61 }
@@ -68,7 +73,7 @@ pub enum InstructionValue<'a> {
73 DeclareContext(DeclareContext<'a>),
74 DeclareLocal(DeclareLocal<'a>),
75 // Destructure(Destructure<'a>),
71 - // Function(Function<'a>),
76 + Function(FunctionExpression<'a>),
77 // JsxFragment(JsxFragment<'a>),
78 // JsxText(JsxText<'a>),
79 LoadContext(LoadContext),
@@ -110,6 +115,12 @@ pub struct Binary {
115 pub right: Operand,
116 }
117
118 +#[derive(Debug)]
119 +pub struct FunctionExpression<'a> {
120 + pub dependencies: Vec<'a, Operand>,
121 + pub lowered_function: Box<'a, Function<'a>>,
122 +}
123 +
124 #[derive(Debug, Clone, PartialEq, Eq)]
125 pub struct Primitive<'a> {
126 pub value: PrimitiveValue<'a>,
compiler/forget/crates/hir/src/print.rs
+50 -18
@@ -1,5 +1,7 @@
1 use std::fmt::{Result, Write};
2
3 +use utils::ensure_sufficient_stack;
4 +
5 use crate::{
6 ArrayElement, BasicBlock, Function, Identifier, IdentifierOperand, Instruction,
7 InstructionValue, LValue, Operand, Phi, PrimitiveValue, Terminal, TerminalValue, HIR,
@@ -16,25 +18,28 @@ pub trait Print<'a> {
18
19 impl<'a> Print<'a> for Function<'a> {
20 fn print(&self, hir: &HIR<'a>, out: &mut impl Write) -> Result {
19 - writeln!(
20 - out,
21 - "function {}(",
22 - match &self.id {
23 - Some(id) => id,
24 - None => "<anonymous>",
21 + ensure_sufficient_stack(|| {
22 + writeln!(
23 + out,
24 + "function {}(",
25 + match &self.id {
26 + Some(id) => id,
27 + None => "<anonymous>",
28 + }
29 + )?;
30 + for param in &self.params {
31 + write!(out, " ")?;
32 + param.print(hir, out)?;
33 + writeln!(out, ",")?;
34 }
26 - )?;
27 - for param in &self.params {
28 - write!(out, " ")?;
29 - param.print(hir, out)?;
30 - writeln!(out, ",")?;
31 - }
32 - writeln!(out, ")")?;
33 - writeln!(out, "entry {}", self.body.entry)?;
34 - for (_, block) in self.body.blocks.iter() {
35 - block.print(hir, out)?;
36 - }
37 - Ok(())
35 + writeln!(out, ")")?;
36 + writeln!(out, "entry {}", self.body.entry)?;
37 + for (_, block) in self.body.blocks.iter() {
38 + block.print(hir, out)?;
39 + }
40 + writeln!(out)?;
41 + Ok(())
42 + })
43 }
44 }
45
@@ -146,6 +151,33 @@ impl<'a> Print<'a> for InstructionValue<'a> {
151 write!(out, " {} ", value.operator)?;
152 value.right.print(hir, out)?;
153 }
154 + InstructionValue::Function(value) => {
155 + write!(out, "Function @deps[")?;
156 + for (ix, dep) in value.dependencies.iter().enumerate() {
157 + if ix != 0 {
158 + write!(out, ", ")?;
159 + }
160 + dep.print(hir, out)?;
161 + }
162 + write!(out, "] @context[")?;
163 + // for (ix, dep) in value.lowered_function.context.iter().enumerate() {
164 + // if ix != 0 {
165 + // write!(out, ", ")?;
166 + // }
167 + // dep.print(hir, out)?;
168 + // }
169 + writeln!(out, "]:")?;
170 + let mut inner_output = String::new();
171 + value
172 + .lowered_function
173 + .print(&value.lowered_function.body, &mut inner_output)?;
174 + let lines: Vec<_> = inner_output
175 + .split("\n")
176 + .map(|line| format!(" {}", line))
177 + .filter(|line| line.trim().len() != 0)
178 + .collect();
179 + write!(out, "{}", lines.join("\n"))?;
180 + }
181 InstructionValue::Tombstone => {
182 write!(out, "Tombstone!")?;
183 }
compiler/forget/crates/utils/Cargo.toml
+2 -1
@@ -6,4 +6,5 @@ edition = "2021"
6 # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
7
8 [dependencies]
9 -bumpalo = { version = "3.13.0", features = ["boxed", "collections"] }
\ No newline at end of file
9 +bumpalo = { version = "3.13.0", features = ["boxed", "collections"] }
10 +stacker = "0.1.15"
compiler/forget/crates/utils/src/lib.rs
+4 -112
@@ -1,113 +1,5 @@
1 -use bumpalo::collections::Vec;
1 +mod ensure_sufficient_stack;
2 +mod retain_mut;
3
3 -pub trait RetainMut<T> {
4 - fn retain_mut<F>(&mut self, f: F) -> ()
5 - where
6 - F: FnMut(&mut T) -> bool;
7 -}
8 -
9 -impl<'a, T> RetainMut<T> for Vec<'a, T> {
10 - fn retain_mut<F>(&mut self, mut f: F) -> ()
11 - where
12 - F: FnMut(&mut T) -> bool,
13 - {
14 - // NOTE: implementation adapted from retain_mut crate
15 - // which is in turn adapted from Rust stdlib
16 - // https://docs.rs/retain_mut/latest/src/retain_mut/lib.rs.html#68-69
17 -
18 - let original_len = self.len();
19 - // Avoid double drop if the drop guard is not executed,
20 - // since we may make some holes during the process.
21 - unsafe { self.set_len(0) };
22 -
23 - // Vec: [Kept, Kept, Hole, Hole, Hole, Hole, Unchecked, Unchecked]
24 - // |<- processed len ->| ^- next to check
25 - // |<- deleted cnt ->|
26 - // |<- original_len ->|
27 - // Kept: Elements which predicate returns true on.
28 - // Hole: Moved or dropped element slot.
29 - // Unchecked: Unchecked valid elements.
30 - //
31 - // This drop guard will be invoked when predicate or `drop` of element panicked.
32 - // It shifts unchecked elements to cover holes and `set_len` to the correct length.
33 - // In cases when predicate and `drop` never panick, it will be optimized out.
34 - struct BackshiftOnDrop<'a, 'b, T> {
35 - v: &'b mut Vec<'a, T>,
36 - processed_len: usize,
37 - deleted_cnt: usize,
38 - original_len: usize,
39 - }
40 -
41 - impl<T> Drop for BackshiftOnDrop<'_, '_, T> {
42 - fn drop(&mut self) {
43 - if self.deleted_cnt > 0 {
44 - // SAFETY: Trailing unchecked items must be valid since we never touch them.
45 - unsafe {
46 - std::ptr::copy(
47 - self.v.as_ptr().add(self.processed_len),
48 - self.v
49 - .as_mut_ptr()
50 - .add(self.processed_len - self.deleted_cnt),
51 - self.original_len - self.processed_len,
52 - );
53 - }
54 - }
55 - // SAFETY: After filling holes, all items are in contiguous memory.
56 - unsafe {
57 - self.v.set_len(self.original_len - self.deleted_cnt);
58 - }
59 - }
60 - }
61 -
62 - let mut g = BackshiftOnDrop {
63 - v: self,
64 - processed_len: 0,
65 - deleted_cnt: 0,
66 - original_len,
67 - };
68 -
69 - fn process_loop<F, T, const DELETED: bool>(
70 - original_len: usize,
71 - f: &mut F,
72 - g: &mut BackshiftOnDrop<'_, '_, T>,
73 - ) where
74 - F: FnMut(&mut T) -> bool,
75 - {
76 - while g.processed_len != original_len {
77 - // SAFETY: Unchecked element must be valid.
78 - let cur = unsafe { &mut *g.v.as_mut_ptr().add(g.processed_len) };
79 - if !f(cur) {
80 - // Advance early to avoid double drop if `drop_in_place` panicked.
81 - g.processed_len += 1;
82 - g.deleted_cnt += 1;
83 - // SAFETY: We never touch this element again after dropped.
84 - unsafe { std::ptr::drop_in_place(cur) };
85 - // We already advanced the counter.
86 - if DELETED {
87 - continue;
88 - } else {
89 - break;
90 - }
91 - }
92 - if DELETED {
93 - // SAFETY: `deleted_cnt` > 0, so the hole slot must not overlap with current element.
94 - // We use copy for move, and never touch this element again.
95 - unsafe {
96 - let hole_slot = g.v.as_mut_ptr().add(g.processed_len - g.deleted_cnt);
97 - std::ptr::copy_nonoverlapping(cur, hole_slot, 1);
98 - }
99 - }
100 - g.processed_len += 1;
101 - }
102 - }
103 -
104 - // Stage 1: Nothing was deleted.
105 - process_loop::<F, T, false>(original_len, &mut f, &mut g);
106 -
107 - // Stage 2: Some elements were deleted.
108 - process_loop::<F, T, true>(original_len, &mut f, &mut g);
109 -
110 - // All item are processed. This can be optimized to `set_len` by LLVM.
111 - drop(g);
112 - }
113 -}
4 +pub use ensure_sufficient_stack::*;
5 +pub use retain_mut::*;
compiler/forget/crates/utils/src/retain_mut.rs new
+113
@@ -0,0 +1,113 @@
1 +use bumpalo::collections::Vec;
2 +
3 +pub trait RetainMut<T> {
4 + fn retain_mut<F>(&mut self, f: F) -> ()
5 + where
6 + F: FnMut(&mut T) -> bool;
7 +}
8 +
9 +impl<'a, T> RetainMut<T> for Vec<'a, T> {
10 + fn retain_mut<F>(&mut self, mut f: F) -> ()
11 + where
12 + F: FnMut(&mut T) -> bool,
13 + {
14 + // NOTE: implementation adapted from retain_mut crate
15 + // which is in turn adapted from Rust stdlib
16 + // https://docs.rs/retain_mut/latest/src/retain_mut/lib.rs.html#68-69
17 +
18 + let original_len = self.len();
19 + // Avoid double drop if the drop guard is not executed,
20 + // since we may make some holes during the process.
21 + unsafe { self.set_len(0) };
22 +
23 + // Vec: [Kept, Kept, Hole, Hole, Hole, Hole, Unchecked, Unchecked]
24 + // |<- processed len ->| ^- next to check
25 + // |<- deleted cnt ->|
26 + // |<- original_len ->|
27 + // Kept: Elements which predicate returns true on.
28 + // Hole: Moved or dropped element slot.
29 + // Unchecked: Unchecked valid elements.
30 + //
31 + // This drop guard will be invoked when predicate or `drop` of element panicked.
32 + // It shifts unchecked elements to cover holes and `set_len` to the correct length.
33 + // In cases when predicate and `drop` never panick, it will be optimized out.
34 + struct BackshiftOnDrop<'a, 'b, T> {
35 + v: &'b mut Vec<'a, T>,
36 + processed_len: usize,
37 + deleted_cnt: usize,
38 + original_len: usize,
39 + }
40 +
41 + impl<T> Drop for BackshiftOnDrop<'_, '_, T> {
42 + fn drop(&mut self) {
43 + if self.deleted_cnt > 0 {
44 + // SAFETY: Trailing unchecked items must be valid since we never touch them.
45 + unsafe {
46 + std::ptr::copy(
47 + self.v.as_ptr().add(self.processed_len),
48 + self.v
49 + .as_mut_ptr()
50 + .add(self.processed_len - self.deleted_cnt),
51 + self.original_len - self.processed_len,
52 + );
53 + }
54 + }
55 + // SAFETY: After filling holes, all items are in contiguous memory.
56 + unsafe {
57 + self.v.set_len(self.original_len - self.deleted_cnt);
58 + }
59 + }
60 + }
61 +
62 + let mut g = BackshiftOnDrop {
63 + v: self,
64 + processed_len: 0,
65 + deleted_cnt: 0,
66 + original_len,
67 + };
68 +
69 + fn process_loop<F, T, const DELETED: bool>(
70 + original_len: usize,
71 + f: &mut F,
72 + g: &mut BackshiftOnDrop<'_, '_, T>,
73 + ) where
74 + F: FnMut(&mut T) -> bool,
75 + {
76 + while g.processed_len != original_len {
77 + // SAFETY: Unchecked element must be valid.
78 + let cur = unsafe { &mut *g.v.as_mut_ptr().add(g.processed_len) };
79 + if !f(cur) {
80 + // Advance early to avoid double drop if `drop_in_place` panicked.
81 + g.processed_len += 1;
82 + g.deleted_cnt += 1;
83 + // SAFETY: We never touch this element again after dropped.
84 + unsafe { std::ptr::drop_in_place(cur) };
85 + // We already advanced the counter.
86 + if DELETED {
87 + continue;
88 + } else {
89 + break;
90 + }
91 + }
92 + if DELETED {
93 + // SAFETY: `deleted_cnt` > 0, so the hole slot must not overlap with current element.
94 + // We use copy for move, and never touch this element again.
95 + unsafe {
96 + let hole_slot = g.v.as_mut_ptr().add(g.processed_len - g.deleted_cnt);
97 + std::ptr::copy_nonoverlapping(cur, hole_slot, 1);
98 + }
99 + }
100 + g.processed_len += 1;
101 + }
102 + }
103 +
104 + // Stage 1: Nothing was deleted.
105 + process_loop::<F, T, false>(original_len, &mut f, &mut g);
106 +
107 + // Stage 2: Some elements were deleted.
108 + process_loop::<F, T, true>(original_len, &mut f, &mut g);
109 +
110 + // All item are processed. This can be optimized to `set_len` by LLVM.
111 + drop(g);
112 + }
113 +}