use std::collections::{BTreeMap, BTreeSet};
use super::super::common::{AstBindingRef, AstBlock, AstModule, AstStmt};
use super::ReadabilityContext;
use super::binding_flow::{BindingUseIndex, binding_mentions_in_stmt};
use super::expr_analysis::is_discard_safe_expr;
use super::walk::{self, AstRewritePass, BlockKind};
pub(super) fn apply(module: &mut AstModule, _context: ReadabilityContext) -> bool {
walk::rewrite_module(module, &mut CleanupPass)
}
struct CleanupPass;
impl AstRewritePass for CleanupPass {
fn rewrite_block(&mut self, block: &mut AstBlock, kind: BlockKind) -> bool {
cleanup_block(
block,
matches!(kind, BlockKind::ModuleBody | BlockKind::FunctionBody),
)
}
}
fn cleanup_block(block: &mut AstBlock, allow_trailing_empty_return_elision: bool) -> bool {
let mut changed = false;
let old_stmts = std::mem::take(&mut block.stmts);
let mut flattened_stmts = Vec::with_capacity(old_stmts.len());
for stmt in old_stmts {
match stmt {
AstStmt::DoBlock(nested)
if nested.stmts.len() == 1 && can_elide_single_stmt_do_block(&nested.stmts[0]) =>
{
flattened_stmts.push(nested.stmts[0].clone());
changed = true;
}
other => flattened_stmts.push(other),
}
}
block.stmts = flattened_stmts;
while let Some(AstStmt::DoBlock(nested)) = block.stmts.last()
&& !nested
.stmts
.iter()
.any(|s| matches!(s, AstStmt::GlobalDecl(_)))
{
let Some(AstStmt::DoBlock(nested)) = block.stmts.pop() else {
unreachable!();
};
block.stmts.extend(nested.stmts);
changed = true;
}
let binding_flow = BlockBindingFlow::new(block);
let discardable_unused_locals = collect_discardable_unused_locals(block, &binding_flow);
let original_len = block.stmts.len();
block.stmts.retain(|stmt| {
!matches!(
stmt,
AstStmt::LocalDecl(local_decl)
if local_decl.bindings.len() == 1
&& local_decl.values.len() == 1
&& discardable_unused_locals.contains(&local_decl.bindings[0].id)
)
});
changed |= block.stmts.len() != original_len;
let binding_flow = BlockBindingFlow::new(block);
let live_mechanical_bindings = collect_live_mechanical_bindings(block, &binding_flow);
for stmt in &mut block.stmts {
let AstStmt::LocalDecl(local_decl) = stmt else {
continue;
};
if !local_decl.values.is_empty() {
continue;
}
let original_len = local_decl.bindings.len();
local_decl.bindings.retain(|binding| match binding.id {
AstBindingRef::Temp(_) | AstBindingRef::SyntheticLocal(_) => {
live_mechanical_bindings.contains(&binding.id)
}
AstBindingRef::Local(_) => true,
});
changed |= local_decl.bindings.len() != original_len;
}
let original_len = block.stmts.len();
block.stmts.retain(|stmt| match stmt {
AstStmt::LocalDecl(local_decl) => {
!(local_decl.bindings.is_empty() && local_decl.values.is_empty())
}
_ => true,
});
changed |= block.stmts.len() != original_len;
if allow_trailing_empty_return_elision
&& matches!(
block.stmts.last(),
Some(AstStmt::Return(ret)) if ret.values.is_empty()
)
{
block.stmts.pop();
changed = true;
}
changed
}
fn can_elide_single_stmt_do_block(stmt: &AstStmt) -> bool {
matches!(
stmt,
AstStmt::Assign(_)
| AstStmt::CallStmt(_)
| AstStmt::Return(_)
| AstStmt::If(_)
| AstStmt::While(_)
| AstStmt::Repeat(_)
| AstStmt::NumericFor(_)
| AstStmt::GenericFor(_)
| AstStmt::Break
| AstStmt::Continue
| AstStmt::Goto(_)
| AstStmt::FunctionDecl(_)
)
}
struct BlockBindingFlow {
mention_counts: BTreeMap<AstBindingRef, usize>,
use_index: BindingUseIndex,
}
impl BlockBindingFlow {
fn new(block: &AstBlock) -> Self {
let mut mention_counts = BTreeMap::<AstBindingRef, usize>::new();
for stmt in &block.stmts {
for binding in binding_mentions_in_stmt(stmt) {
*mention_counts.entry(binding).or_default() += 1;
}
}
Self {
mention_counts,
use_index: BindingUseIndex::for_stmts(&block.stmts),
}
}
fn mentioned_outside_own_decl(&self, binding: AstBindingRef) -> bool {
self.mention_counts.get(&binding).copied().unwrap_or(0) > 1
}
fn used_or_captured(&self, binding: AstBindingRef) -> bool {
self.use_index.count_uses_in_suffix(0, binding) != 0
}
fn keeps_decl_alive(&self, binding: AstBindingRef) -> bool {
self.mentioned_outside_own_decl(binding) || self.used_or_captured(binding)
}
}
fn collect_live_mechanical_bindings(
block: &AstBlock,
binding_flow: &BlockBindingFlow,
) -> BTreeSet<AstBindingRef> {
let mut live_bindings = BTreeSet::new();
for stmt in &block.stmts {
let AstStmt::LocalDecl(local_decl) = stmt else {
continue;
};
for binding in &local_decl.bindings {
if matches!(
binding.id,
AstBindingRef::Temp(_) | AstBindingRef::SyntheticLocal(_)
) && binding_flow.keeps_decl_alive(binding.id)
{
live_bindings.insert(binding.id);
}
}
}
live_bindings
}
fn collect_discardable_unused_locals(
block: &AstBlock,
binding_flow: &BlockBindingFlow,
) -> std::collections::BTreeSet<AstBindingRef> {
let mut bindings = std::collections::BTreeSet::new();
for stmt in &block.stmts {
let AstStmt::LocalDecl(local_decl) = stmt else {
continue;
};
let [binding] = local_decl.bindings.as_slice() else {
continue;
};
let [value] = local_decl.values.as_slice() else {
continue;
};
if !matches!(binding.origin, crate::ast::AstLocalOrigin::Recovered) {
continue;
}
if binding_flow.keeps_decl_alive(binding.id) {
continue;
}
if is_discard_safe_expr(value) {
bindings.insert(binding.id);
}
}
bindings
}