mirror of
https://github.com/facebook/react.git
synced 2025-11-01 09:12:30 +00:00
[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.
This commit is contained in:
Generated
+2
@@ -638,6 +638,7 @@ dependencies = [
|
||||
"estree",
|
||||
"indexmap 2.0.0",
|
||||
"serde",
|
||||
"utils",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -2806,6 +2807,7 @@ name = "utils"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"bumpalo",
|
||||
"stacker",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
use bumpalo::collections::{String, Vec};
|
||||
use bumpalo::{
|
||||
boxed::Box,
|
||||
collections::{String, Vec},
|
||||
};
|
||||
use estree::{
|
||||
AssignmentTarget, BinaryExpression, BlockStatement, Expression, ForInit, ForStatement,
|
||||
FunctionDeclaration, IfStatement, JsValue, Literal, Pattern, Statement,
|
||||
VariableDeclarationKind,
|
||||
FunctionExpression, IfStatement, JsValue, Literal, Pattern, Statement, VariableDeclarationKind,
|
||||
};
|
||||
use hir::{
|
||||
ArrayElement, BlockKind, BranchTerminal, Environment, ForTerminal, Function, GotoKind,
|
||||
@@ -24,11 +26,11 @@ use crate::{
|
||||
/// that is not yet supported.
|
||||
pub fn build<'a>(
|
||||
env: &'a Environment<'a>,
|
||||
fun: FunctionDeclaration,
|
||||
) -> Result<&'a mut Function<'a>, BuildDiagnostic> {
|
||||
fun: estree::Function,
|
||||
) -> Result<Box<'a, Function<'a>>, BuildDiagnostic> {
|
||||
let mut builder = Builder::new(env);
|
||||
|
||||
match fun.function.body {
|
||||
match fun.body {
|
||||
Some(estree::FunctionBody::BlockStatement(body)) => {
|
||||
lower_block_statement(env, &mut builder, *body)?
|
||||
}
|
||||
@@ -44,8 +46,8 @@ pub fn build<'a>(
|
||||
}
|
||||
}
|
||||
|
||||
let mut params = Vec::with_capacity_in(fun.function.params.len(), &env.allocator);
|
||||
for param in fun.function.params {
|
||||
let mut params = Vec::with_capacity_in(fun.params.len(), &env.allocator);
|
||||
for param in fun.params {
|
||||
match param {
|
||||
Pattern::Identifier(param) => {
|
||||
let identifier = lower_identifier_for_assignment(
|
||||
@@ -76,16 +78,18 @@ pub fn build<'a>(
|
||||
);
|
||||
|
||||
let body = builder.build()?;
|
||||
Ok(env.alloc(Function {
|
||||
id: fun
|
||||
.function
|
||||
.id
|
||||
.map(|id| String::from_str_in(&id.name, &env.allocator)),
|
||||
body,
|
||||
params,
|
||||
is_async: fun.function.is_async,
|
||||
is_generator: fun.function.is_generator,
|
||||
}))
|
||||
Ok(Box::new_in(
|
||||
Function {
|
||||
id: fun
|
||||
.id
|
||||
.map(|id| String::from_str_in(&id.name, &env.allocator)),
|
||||
body,
|
||||
params,
|
||||
is_async: fun.is_async,
|
||||
is_generator: fun.is_generator,
|
||||
},
|
||||
&env.allocator,
|
||||
))
|
||||
}
|
||||
|
||||
fn lower_block_statement<'a>(
|
||||
@@ -421,6 +425,16 @@ fn lower_expression<'a>(
|
||||
})
|
||||
}
|
||||
|
||||
Expression::FunctionExpression(expr) => {
|
||||
let FunctionExpression { function, .. } = *expr;
|
||||
let fun = build(env, function)?;
|
||||
InstructionValue::Function(hir::FunctionExpression {
|
||||
// TODO: collect dependencies!
|
||||
dependencies: Vec::new_in(&env.allocator),
|
||||
lowered_function: fun,
|
||||
})
|
||||
}
|
||||
|
||||
_ => todo!("Lower expr {expr:#?}"),
|
||||
};
|
||||
Ok(builder.push(value))
|
||||
|
||||
@@ -43,6 +43,14 @@
|
||||
"type": "bool",
|
||||
"optional": true,
|
||||
"rename": "async"
|
||||
},
|
||||
"loc": {
|
||||
"type": "Option<SourceLocation>",
|
||||
"optional": true
|
||||
},
|
||||
"range": {
|
||||
"type": "Option<SourceRange>",
|
||||
"optional": true
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -6,8 +6,9 @@ use swc_core::common::errors::Handler;
|
||||
use swc_core::common::source_map::Pos;
|
||||
use swc_core::common::{FileName, FilePathMapping, Mark, SourceMap, Span, SyntaxContext, GLOBALS};
|
||||
use swc_core::ecma::ast::{
|
||||
AssignOp, BinaryOp, BlockStmt, Decl, EsVersion, Expr, Ident, Lit, MemberExpr, MemberProp,
|
||||
ModuleItem, Pat, PatOrExpr, Program, Stmt, UnaryOp, VarDecl, VarDeclKind, VarDeclOrExpr,
|
||||
AssignOp, BinaryOp, BlockStmt, Decl, EsVersion, Expr, Function, Ident, Lit, MemberExpr,
|
||||
MemberProp, ModuleItem, Pat, PatOrExpr, Program, Stmt, UnaryOp, VarDecl, VarDeclKind,
|
||||
VarDeclOrExpr,
|
||||
};
|
||||
use swc_core::ecma::parser::Syntax;
|
||||
use swc_core::ecma::transforms::base::resolver;
|
||||
@@ -122,32 +123,34 @@ fn convert_block_statement(cx: &Context, stmt: &BlockStmt) -> estree::BlockState
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_function(cx: &Context, id: Option<&Ident>, fun: &Function) -> estree::Function {
|
||||
estree::Function {
|
||||
id: id.map(|id| estree::Identifier {
|
||||
name: id.sym.to_string(),
|
||||
binding: convert_binding(cx, id.span.ctxt),
|
||||
loc: None,
|
||||
range: convert_span(&id.span),
|
||||
}),
|
||||
params: fun
|
||||
.params
|
||||
.iter()
|
||||
.map(|param| convert_pattern(cx, ¶m.pat))
|
||||
.collect(),
|
||||
body: fun.body.as_ref().map(|body| {
|
||||
estree::FunctionBody::BlockStatement(Box::new(convert_block_statement(cx, body)))
|
||||
}),
|
||||
is_async: fun.is_async,
|
||||
is_generator: fun.is_generator,
|
||||
loc: None,
|
||||
range: convert_span(&fun.span),
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_statement(cx: &Context, stmt: &Stmt) -> estree::Statement {
|
||||
match stmt {
|
||||
Stmt::Decl(Decl::Fn(item)) => {
|
||||
let name = item.ident.sym.to_string();
|
||||
estree::Statement::FunctionDeclaration(Box::new(estree::FunctionDeclaration {
|
||||
function: estree::Function {
|
||||
id: Some(estree::Identifier {
|
||||
name,
|
||||
binding: convert_binding(cx, item.ident.span.ctxt),
|
||||
loc: None,
|
||||
range: convert_span(&item.ident.span),
|
||||
}),
|
||||
params: item
|
||||
.function
|
||||
.params
|
||||
.iter()
|
||||
.map(|param| convert_pattern(cx, ¶m.pat))
|
||||
.collect(),
|
||||
body: item.function.body.as_ref().map(|body| {
|
||||
estree::FunctionBody::BlockStatement(Box::new(convert_block_statement(
|
||||
cx, body,
|
||||
)))
|
||||
}),
|
||||
is_async: item.function.is_async,
|
||||
is_generator: item.function.is_generator,
|
||||
},
|
||||
function: convert_function(cx, Some(&item.ident), &item.function),
|
||||
loc: None,
|
||||
range: convert_span(&item.function.span),
|
||||
}))
|
||||
@@ -368,6 +371,13 @@ fn convert_expression(cx: &Context, expr: &Expr) -> estree::Expression {
|
||||
Expr::Member(expr) => {
|
||||
estree::Expression::MemberExpression(Box::new(convert_member_expression(cx, expr)))
|
||||
}
|
||||
Expr::Fn(expr) => {
|
||||
estree::Expression::FunctionExpression(Box::new(estree::FunctionExpression {
|
||||
function: convert_function(cx, expr.ident.as_ref(), &expr.function),
|
||||
loc: None,
|
||||
range: convert_span(&expr.function.span),
|
||||
}))
|
||||
}
|
||||
_ => todo!("translate expression {:#?}", expr),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -23,6 +23,10 @@ pub struct Function {
|
||||
#[serde(rename = "async")]
|
||||
#[serde(default)]
|
||||
pub is_async: bool,
|
||||
#[serde(default)]
|
||||
pub loc: Option<SourceLocation>,
|
||||
#[serde(default)]
|
||||
pub range: Option<SourceRange>,
|
||||
}
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
pub struct RegExpValue {
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
function Component(props) {
|
||||
const foo = function foo(x) {
|
||||
return x + 1;
|
||||
};
|
||||
}
|
||||
@@ -33,7 +33,7 @@ fn fixtures() {
|
||||
if ix != 0 {
|
||||
output.push_str("\n\n");
|
||||
}
|
||||
match build(&environment, *fun) {
|
||||
match build(&environment, fun.function) {
|
||||
Ok(mut fun) => {
|
||||
enter_ssa(&environment, &mut fun).unwrap();
|
||||
eliminate_redundant_phis(&environment, &mut fun);
|
||||
@@ -56,6 +56,7 @@ fn fixtures() {
|
||||
}
|
||||
}
|
||||
|
||||
let output = output.trim();
|
||||
assert_snapshot!(format!("Input:\n{input}\n\nOutput:\n{output}"));
|
||||
});
|
||||
}
|
||||
|
||||
+34
@@ -0,0 +1,34 @@
|
||||
---
|
||||
source: crates/fixtures/tests/fixtures_test.rs
|
||||
expression: "format!(\"Input:\\n{input}\\n\\nOutput:\\n{output}\")"
|
||||
input_file: crates/fixtures/tests/fixtures/function-expressions.js
|
||||
---
|
||||
Input:
|
||||
function Component(props) {
|
||||
const foo = function foo(x) {
|
||||
return x + 1;
|
||||
};
|
||||
}
|
||||
|
||||
|
||||
Output:
|
||||
function Component(
|
||||
unknown props$3,
|
||||
)
|
||||
entry bb0
|
||||
bb0 (block)
|
||||
[0] #0 = Function @deps[] @context[]:
|
||||
function foo(
|
||||
unknown x$0,
|
||||
)
|
||||
entry bb1
|
||||
bb1 (block)
|
||||
[0] #0 = LoadLocal unknown x$0
|
||||
[1] #1 = 1
|
||||
[2] #2 = Binary unknown #0 + unknown #1
|
||||
[3] Return unknown #2
|
||||
[1] #1 = StoreLocal Const unknown foo$4 = unknown #0
|
||||
[2] #2 = <undefined>
|
||||
[3] Return unknown #2
|
||||
|
||||
|
||||
@@ -10,3 +10,4 @@ bumpalo = { version = "3.13.0", features = ["boxed", "collections"] }
|
||||
estree = { path = "../estree" }
|
||||
indexmap = "2.0.0"
|
||||
serde = "1.0.164"
|
||||
utils = { path = "../utils" }
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
use std::{cell::RefCell, fmt::Display, rc::Rc};
|
||||
|
||||
use bumpalo::collections::{String, Vec};
|
||||
use bumpalo::{
|
||||
boxed::Box,
|
||||
collections::{String, Vec},
|
||||
};
|
||||
use estree::BinaryOperator;
|
||||
|
||||
use crate::{IdentifierId, InstrIx, InstructionId, ScopeId, Type};
|
||||
use crate::{Function, IdentifierId, InstrIx, InstructionId, ScopeId, Type};
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct Instruction<'a> {
|
||||
@@ -17,22 +20,23 @@ impl<'a> Instruction<'a> {
|
||||
F: FnMut(&mut LValue<'a>) -> (),
|
||||
{
|
||||
match &mut self.value {
|
||||
InstructionValue::Array(_) => {}
|
||||
InstructionValue::Binary(_) => {}
|
||||
InstructionValue::DeclareContext(instr) => {
|
||||
f(&mut instr.lvalue);
|
||||
}
|
||||
InstructionValue::DeclareLocal(instr) => {
|
||||
f(&mut instr.lvalue);
|
||||
}
|
||||
InstructionValue::LoadContext(_) => {}
|
||||
InstructionValue::LoadGlobal(_) => {}
|
||||
InstructionValue::LoadLocal(_) => {}
|
||||
InstructionValue::Primitive(_) => {}
|
||||
InstructionValue::StoreLocal(instr) => {
|
||||
f(&mut instr.lvalue);
|
||||
}
|
||||
InstructionValue::Tombstone => {}
|
||||
InstructionValue::Array(_)
|
||||
| InstructionValue::Binary(_)
|
||||
| InstructionValue::LoadContext(_)
|
||||
| InstructionValue::LoadGlobal(_)
|
||||
| InstructionValue::LoadLocal(_)
|
||||
| InstructionValue::Primitive(_)
|
||||
| InstructionValue::Function(_)
|
||||
| InstructionValue::Tombstone => {}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -41,16 +45,17 @@ impl<'a> Instruction<'a> {
|
||||
F: FnMut(&mut IdentifierOperand<'a>) -> (),
|
||||
{
|
||||
match &mut self.value {
|
||||
InstructionValue::Array(_) => {}
|
||||
InstructionValue::Binary(_) => {}
|
||||
InstructionValue::DeclareContext(_) => {}
|
||||
InstructionValue::DeclareLocal(_) => {}
|
||||
InstructionValue::LoadContext(_) => {}
|
||||
InstructionValue::LoadGlobal(_) => {}
|
||||
InstructionValue::LoadLocal(instr) => f(&mut instr.place),
|
||||
InstructionValue::Primitive(_) => {}
|
||||
InstructionValue::StoreLocal(_) => {}
|
||||
InstructionValue::Tombstone => {}
|
||||
InstructionValue::Array(_)
|
||||
| InstructionValue::Binary(_)
|
||||
| InstructionValue::DeclareContext(_)
|
||||
| InstructionValue::DeclareLocal(_)
|
||||
| InstructionValue::LoadContext(_)
|
||||
| InstructionValue::LoadGlobal(_)
|
||||
| InstructionValue::Primitive(_)
|
||||
| InstructionValue::StoreLocal(_)
|
||||
| InstructionValue::Function(_)
|
||||
| InstructionValue::Tombstone => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -68,7 +73,7 @@ pub enum InstructionValue<'a> {
|
||||
DeclareContext(DeclareContext<'a>),
|
||||
DeclareLocal(DeclareLocal<'a>),
|
||||
// Destructure(Destructure<'a>),
|
||||
// Function(Function<'a>),
|
||||
Function(FunctionExpression<'a>),
|
||||
// JsxFragment(JsxFragment<'a>),
|
||||
// JsxText(JsxText<'a>),
|
||||
LoadContext(LoadContext),
|
||||
@@ -110,6 +115,12 @@ pub struct Binary {
|
||||
pub right: Operand,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct FunctionExpression<'a> {
|
||||
pub dependencies: Vec<'a, Operand>,
|
||||
pub lowered_function: Box<'a, Function<'a>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct Primitive<'a> {
|
||||
pub value: PrimitiveValue<'a>,
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use std::fmt::{Result, Write};
|
||||
|
||||
use utils::ensure_sufficient_stack;
|
||||
|
||||
use crate::{
|
||||
ArrayElement, BasicBlock, Function, Identifier, IdentifierOperand, Instruction,
|
||||
InstructionValue, LValue, Operand, Phi, PrimitiveValue, Terminal, TerminalValue, HIR,
|
||||
@@ -16,25 +18,28 @@ pub trait Print<'a> {
|
||||
|
||||
impl<'a> Print<'a> for Function<'a> {
|
||||
fn print(&self, hir: &HIR<'a>, out: &mut impl Write) -> Result {
|
||||
writeln!(
|
||||
out,
|
||||
"function {}(",
|
||||
match &self.id {
|
||||
Some(id) => id,
|
||||
None => "<anonymous>",
|
||||
ensure_sufficient_stack(|| {
|
||||
writeln!(
|
||||
out,
|
||||
"function {}(",
|
||||
match &self.id {
|
||||
Some(id) => id,
|
||||
None => "<anonymous>",
|
||||
}
|
||||
)?;
|
||||
for param in &self.params {
|
||||
write!(out, " ")?;
|
||||
param.print(hir, out)?;
|
||||
writeln!(out, ",")?;
|
||||
}
|
||||
)?;
|
||||
for param in &self.params {
|
||||
write!(out, " ")?;
|
||||
param.print(hir, out)?;
|
||||
writeln!(out, ",")?;
|
||||
}
|
||||
writeln!(out, ")")?;
|
||||
writeln!(out, "entry {}", self.body.entry)?;
|
||||
for (_, block) in self.body.blocks.iter() {
|
||||
block.print(hir, out)?;
|
||||
}
|
||||
Ok(())
|
||||
writeln!(out, ")")?;
|
||||
writeln!(out, "entry {}", self.body.entry)?;
|
||||
for (_, block) in self.body.blocks.iter() {
|
||||
block.print(hir, out)?;
|
||||
}
|
||||
writeln!(out)?;
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -146,6 +151,33 @@ impl<'a> Print<'a> for InstructionValue<'a> {
|
||||
write!(out, " {} ", value.operator)?;
|
||||
value.right.print(hir, out)?;
|
||||
}
|
||||
InstructionValue::Function(value) => {
|
||||
write!(out, "Function @deps[")?;
|
||||
for (ix, dep) in value.dependencies.iter().enumerate() {
|
||||
if ix != 0 {
|
||||
write!(out, ", ")?;
|
||||
}
|
||||
dep.print(hir, out)?;
|
||||
}
|
||||
write!(out, "] @context[")?;
|
||||
// for (ix, dep) in value.lowered_function.context.iter().enumerate() {
|
||||
// if ix != 0 {
|
||||
// write!(out, ", ")?;
|
||||
// }
|
||||
// dep.print(hir, out)?;
|
||||
// }
|
||||
writeln!(out, "]:")?;
|
||||
let mut inner_output = String::new();
|
||||
value
|
||||
.lowered_function
|
||||
.print(&value.lowered_function.body, &mut inner_output)?;
|
||||
let lines: Vec<_> = inner_output
|
||||
.split("\n")
|
||||
.map(|line| format!(" {}", line))
|
||||
.filter(|line| line.trim().len() != 0)
|
||||
.collect();
|
||||
write!(out, "{}", lines.join("\n"))?;
|
||||
}
|
||||
InstructionValue::Tombstone => {
|
||||
write!(out, "Tombstone!")?;
|
||||
}
|
||||
|
||||
@@ -6,4 +6,5 @@ edition = "2021"
|
||||
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
|
||||
|
||||
[dependencies]
|
||||
bumpalo = { version = "3.13.0", features = ["boxed", "collections"] }
|
||||
bumpalo = { version = "3.13.0", features = ["boxed", "collections"] }
|
||||
stacker = "0.1.15"
|
||||
|
||||
@@ -1,113 +1,5 @@
|
||||
use bumpalo::collections::Vec;
|
||||
mod ensure_sufficient_stack;
|
||||
mod retain_mut;
|
||||
|
||||
pub trait RetainMut<T> {
|
||||
fn retain_mut<F>(&mut self, f: F) -> ()
|
||||
where
|
||||
F: FnMut(&mut T) -> bool;
|
||||
}
|
||||
|
||||
impl<'a, T> RetainMut<T> for Vec<'a, T> {
|
||||
fn retain_mut<F>(&mut self, mut f: F) -> ()
|
||||
where
|
||||
F: FnMut(&mut T) -> bool,
|
||||
{
|
||||
// NOTE: implementation adapted from retain_mut crate
|
||||
// which is in turn adapted from Rust stdlib
|
||||
// https://docs.rs/retain_mut/latest/src/retain_mut/lib.rs.html#68-69
|
||||
|
||||
let original_len = self.len();
|
||||
// Avoid double drop if the drop guard is not executed,
|
||||
// since we may make some holes during the process.
|
||||
unsafe { self.set_len(0) };
|
||||
|
||||
// Vec: [Kept, Kept, Hole, Hole, Hole, Hole, Unchecked, Unchecked]
|
||||
// |<- processed len ->| ^- next to check
|
||||
// |<- deleted cnt ->|
|
||||
// |<- original_len ->|
|
||||
// Kept: Elements which predicate returns true on.
|
||||
// Hole: Moved or dropped element slot.
|
||||
// Unchecked: Unchecked valid elements.
|
||||
//
|
||||
// This drop guard will be invoked when predicate or `drop` of element panicked.
|
||||
// It shifts unchecked elements to cover holes and `set_len` to the correct length.
|
||||
// In cases when predicate and `drop` never panick, it will be optimized out.
|
||||
struct BackshiftOnDrop<'a, 'b, T> {
|
||||
v: &'b mut Vec<'a, T>,
|
||||
processed_len: usize,
|
||||
deleted_cnt: usize,
|
||||
original_len: usize,
|
||||
}
|
||||
|
||||
impl<T> Drop for BackshiftOnDrop<'_, '_, T> {
|
||||
fn drop(&mut self) {
|
||||
if self.deleted_cnt > 0 {
|
||||
// SAFETY: Trailing unchecked items must be valid since we never touch them.
|
||||
unsafe {
|
||||
std::ptr::copy(
|
||||
self.v.as_ptr().add(self.processed_len),
|
||||
self.v
|
||||
.as_mut_ptr()
|
||||
.add(self.processed_len - self.deleted_cnt),
|
||||
self.original_len - self.processed_len,
|
||||
);
|
||||
}
|
||||
}
|
||||
// SAFETY: After filling holes, all items are in contiguous memory.
|
||||
unsafe {
|
||||
self.v.set_len(self.original_len - self.deleted_cnt);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut g = BackshiftOnDrop {
|
||||
v: self,
|
||||
processed_len: 0,
|
||||
deleted_cnt: 0,
|
||||
original_len,
|
||||
};
|
||||
|
||||
fn process_loop<F, T, const DELETED: bool>(
|
||||
original_len: usize,
|
||||
f: &mut F,
|
||||
g: &mut BackshiftOnDrop<'_, '_, T>,
|
||||
) where
|
||||
F: FnMut(&mut T) -> bool,
|
||||
{
|
||||
while g.processed_len != original_len {
|
||||
// SAFETY: Unchecked element must be valid.
|
||||
let cur = unsafe { &mut *g.v.as_mut_ptr().add(g.processed_len) };
|
||||
if !f(cur) {
|
||||
// Advance early to avoid double drop if `drop_in_place` panicked.
|
||||
g.processed_len += 1;
|
||||
g.deleted_cnt += 1;
|
||||
// SAFETY: We never touch this element again after dropped.
|
||||
unsafe { std::ptr::drop_in_place(cur) };
|
||||
// We already advanced the counter.
|
||||
if DELETED {
|
||||
continue;
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
if DELETED {
|
||||
// SAFETY: `deleted_cnt` > 0, so the hole slot must not overlap with current element.
|
||||
// We use copy for move, and never touch this element again.
|
||||
unsafe {
|
||||
let hole_slot = g.v.as_mut_ptr().add(g.processed_len - g.deleted_cnt);
|
||||
std::ptr::copy_nonoverlapping(cur, hole_slot, 1);
|
||||
}
|
||||
}
|
||||
g.processed_len += 1;
|
||||
}
|
||||
}
|
||||
|
||||
// Stage 1: Nothing was deleted.
|
||||
process_loop::<F, T, false>(original_len, &mut f, &mut g);
|
||||
|
||||
// Stage 2: Some elements were deleted.
|
||||
process_loop::<F, T, true>(original_len, &mut f, &mut g);
|
||||
|
||||
// All item are processed. This can be optimized to `set_len` by LLVM.
|
||||
drop(g);
|
||||
}
|
||||
}
|
||||
pub use ensure_sufficient_stack::*;
|
||||
pub use retain_mut::*;
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
use bumpalo::collections::Vec;
|
||||
|
||||
pub trait RetainMut<T> {
|
||||
fn retain_mut<F>(&mut self, f: F) -> ()
|
||||
where
|
||||
F: FnMut(&mut T) -> bool;
|
||||
}
|
||||
|
||||
impl<'a, T> RetainMut<T> for Vec<'a, T> {
|
||||
fn retain_mut<F>(&mut self, mut f: F) -> ()
|
||||
where
|
||||
F: FnMut(&mut T) -> bool,
|
||||
{
|
||||
// NOTE: implementation adapted from retain_mut crate
|
||||
// which is in turn adapted from Rust stdlib
|
||||
// https://docs.rs/retain_mut/latest/src/retain_mut/lib.rs.html#68-69
|
||||
|
||||
let original_len = self.len();
|
||||
// Avoid double drop if the drop guard is not executed,
|
||||
// since we may make some holes during the process.
|
||||
unsafe { self.set_len(0) };
|
||||
|
||||
// Vec: [Kept, Kept, Hole, Hole, Hole, Hole, Unchecked, Unchecked]
|
||||
// |<- processed len ->| ^- next to check
|
||||
// |<- deleted cnt ->|
|
||||
// |<- original_len ->|
|
||||
// Kept: Elements which predicate returns true on.
|
||||
// Hole: Moved or dropped element slot.
|
||||
// Unchecked: Unchecked valid elements.
|
||||
//
|
||||
// This drop guard will be invoked when predicate or `drop` of element panicked.
|
||||
// It shifts unchecked elements to cover holes and `set_len` to the correct length.
|
||||
// In cases when predicate and `drop` never panick, it will be optimized out.
|
||||
struct BackshiftOnDrop<'a, 'b, T> {
|
||||
v: &'b mut Vec<'a, T>,
|
||||
processed_len: usize,
|
||||
deleted_cnt: usize,
|
||||
original_len: usize,
|
||||
}
|
||||
|
||||
impl<T> Drop for BackshiftOnDrop<'_, '_, T> {
|
||||
fn drop(&mut self) {
|
||||
if self.deleted_cnt > 0 {
|
||||
// SAFETY: Trailing unchecked items must be valid since we never touch them.
|
||||
unsafe {
|
||||
std::ptr::copy(
|
||||
self.v.as_ptr().add(self.processed_len),
|
||||
self.v
|
||||
.as_mut_ptr()
|
||||
.add(self.processed_len - self.deleted_cnt),
|
||||
self.original_len - self.processed_len,
|
||||
);
|
||||
}
|
||||
}
|
||||
// SAFETY: After filling holes, all items are in contiguous memory.
|
||||
unsafe {
|
||||
self.v.set_len(self.original_len - self.deleted_cnt);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut g = BackshiftOnDrop {
|
||||
v: self,
|
||||
processed_len: 0,
|
||||
deleted_cnt: 0,
|
||||
original_len,
|
||||
};
|
||||
|
||||
fn process_loop<F, T, const DELETED: bool>(
|
||||
original_len: usize,
|
||||
f: &mut F,
|
||||
g: &mut BackshiftOnDrop<'_, '_, T>,
|
||||
) where
|
||||
F: FnMut(&mut T) -> bool,
|
||||
{
|
||||
while g.processed_len != original_len {
|
||||
// SAFETY: Unchecked element must be valid.
|
||||
let cur = unsafe { &mut *g.v.as_mut_ptr().add(g.processed_len) };
|
||||
if !f(cur) {
|
||||
// Advance early to avoid double drop if `drop_in_place` panicked.
|
||||
g.processed_len += 1;
|
||||
g.deleted_cnt += 1;
|
||||
// SAFETY: We never touch this element again after dropped.
|
||||
unsafe { std::ptr::drop_in_place(cur) };
|
||||
// We already advanced the counter.
|
||||
if DELETED {
|
||||
continue;
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
if DELETED {
|
||||
// SAFETY: `deleted_cnt` > 0, so the hole slot must not overlap with current element.
|
||||
// We use copy for move, and never touch this element again.
|
||||
unsafe {
|
||||
let hole_slot = g.v.as_mut_ptr().add(g.processed_len - g.deleted_cnt);
|
||||
std::ptr::copy_nonoverlapping(cur, hole_slot, 1);
|
||||
}
|
||||
}
|
||||
g.processed_len += 1;
|
||||
}
|
||||
}
|
||||
|
||||
// Stage 1: Nothing was deleted.
|
||||
process_loop::<F, T, false>(original_len, &mut f, &mut g);
|
||||
|
||||
// Stage 2: Some elements were deleted.
|
||||
process_loop::<F, T, true>(original_len, &mut f, &mut g);
|
||||
|
||||
// All item are processed. This can be optimized to `set_len` by LLVM.
|
||||
drop(g);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user