[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, ¶m.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, ¶m.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
+}