use crate::{SsaFormingInput, common::SymbolAccessCollector, static_single_assignment::visitor::SsaFormingVisitor};
use super::transform::TransformVisitor;
use indexmap::IndexSet;
use leo_ast::*;
use leo_span::Symbol;
pub enum InlineBinder {
Definition(DefinitionPlace),
Discard,
}
pub struct PendingIterativeInline {
statements: Vec<Statement>,
binder: InlineBinder,
}
impl PendingIterativeInline {
pub fn try_from_body(statements: Vec<Statement>) -> Result<Self, Vec<Statement>> {
let under_live_control_flow = match statements.last() {
Some(Statement::Return(_)) => &statements[..statements.len() - 1],
_ => &statements[..],
};
if statements_contain_return(under_live_control_flow) {
Ok(Self { statements, binder: InlineBinder::Discard })
} else {
Err(statements)
}
}
pub fn bind_to(&mut self, place: DefinitionPlace) {
self.binder = InlineBinder::Definition(place);
}
}
impl TransformVisitor<'_> {
pub(super) fn lower_iterative_inline(
&mut self,
inline: PendingIterativeInline,
continuation: Vec<Statement>,
) -> Vec<Statement> {
let (prefix, suffix) = self.split_continuation(&inline.binder, continuation);
let mut out = self.weave(inline.statements, &inline.binder, &prefix);
out.extend(suffix);
out
}
fn weave(
&mut self,
statements: Vec<Statement>,
binder: &InlineBinder,
continuation: &[Statement],
) -> Vec<Statement> {
let mut out = Vec::new();
let mut iter = statements.into_iter();
while let Some(statement) = iter.next() {
match statement {
Statement::Return(ret) => {
out.extend(self.bind_and_continue(ret.expression, binder, continuation));
return out;
}
Statement::Block(block) if statements_contain_return(&block.statements) => {
let mut merged = block.statements;
merged.extend(iter);
out.extend(self.weave(merged, binder, continuation));
return out;
}
Statement::Conditional(conditional) if conditional_contains_return(&conditional) => {
let rest: Vec<Statement> = iter.collect();
let mut then_statements = conditional.then.statements;
then_statements.extend(self.fresh_clone(rest.clone()));
let mut else_statements = match conditional.otherwise.map(|s| *s) {
Some(Statement::Block(block)) => block.statements,
Some(other) => vec![other],
None => Vec::new(),
};
else_statements.extend(rest);
let then = Block {
statements: self.weave(then_statements, binder, continuation),
span: conditional.then.span,
id: conditional.then.id,
};
let otherwise = Block {
statements: self.weave(else_statements, binder, continuation),
span: Default::default(),
id: self.state.node_builder.next_id(),
};
out.push(
ConditionalStatement {
condition: conditional.condition,
then,
otherwise: Some(Box::new(otherwise.into())),
span: conditional.span,
id: conditional.id,
}
.into(),
);
return out;
}
other => out.push(other),
}
}
let id = self.state.node_builder.next_id();
self.state.type_table.insert(id, Type::UNIT);
out.extend(self.bind_and_continue(
UnitExpression { span: Default::default(), id }.into(),
binder,
continuation,
));
out
}
fn bind_and_continue(
&mut self,
value: Expression,
binder: &InlineBinder,
continuation: &[Statement],
) -> Vec<Statement> {
let mut fragment = Vec::with_capacity(continuation.len() + 1);
match binder {
InlineBinder::Discard => {}
InlineBinder::Definition(place) => match (place, value) {
(DefinitionPlace::Multiple(left), Expression::Tuple(right)) => {
assert_eq!(left.len(), right.elements.len());
for (identifier, rhs_value) in left.iter().zip(right.elements) {
fragment.push(
DefinitionStatement {
place: DefinitionPlace::Single(*identifier),
type_: None,
value: rhs_value,
span: Default::default(),
id: self.state.node_builder.next_id(),
}
.into(),
);
}
}
(place, value) => fragment.push(
DefinitionStatement {
place: place.clone(),
type_: None,
value,
span: Default::default(),
id: self.state.node_builder.next_id(),
}
.into(),
),
},
}
fragment.extend(continuation.to_vec());
self.fresh_clone(fragment)
}
fn split_continuation(
&mut self,
binder: &InlineBinder,
continuation: Vec<Statement>,
) -> (Vec<Statement>, Vec<Statement>) {
let mut dup_names: IndexSet<Symbol> = match binder {
InlineBinder::Discard => IndexSet::new(),
InlineBinder::Definition(DefinitionPlace::Single(identifier)) => IndexSet::from([identifier.name]),
InlineBinder::Definition(DefinitionPlace::Multiple(identifiers)) => {
identifiers.iter().map(|i| i.name).collect()
}
};
let uses: Vec<IndexSet<Symbol>> = continuation.iter().map(|s| self.local_uses(s)).collect();
let defs: Vec<Vec<Symbol>> = continuation.iter().map(defs_of).collect();
let mut cut: Option<usize> = None;
for (i, statement_uses) in uses.iter().enumerate() {
if !statement_uses.is_disjoint(&dup_names) {
cut = Some(i);
dup_names.extend(defs[i].iter().copied());
}
}
while let Some(c) = cut {
let prefix_defs: IndexSet<Symbol> = defs[..=c].iter().flatten().copied().collect();
match uses.iter().enumerate().skip(c + 1).find(|(_, u)| !u.is_disjoint(&prefix_defs)) {
Some((i, _)) => cut = Some(i),
None => break,
}
}
match cut {
None => (Vec::new(), continuation),
Some(c) => {
let mut prefix = continuation;
let suffix = prefix.split_off(c + 1);
(prefix, suffix)
}
}
}
fn local_uses(&mut self, statement: &Statement) -> IndexSet<Symbol> {
let mut collector = SymbolAccessCollector::new(self.state);
collector.visit_statement(statement);
collector
.symbol_accesses
.iter()
.filter(|(path, _)| path.try_global_location().is_none())
.map(|(path, _)| path.identifier().name)
.collect()
}
fn fresh_clone(&mut self, statements: Vec<Statement>) -> Vec<Statement> {
if statements.is_empty() {
return Vec::new();
}
let block = Block { statements, span: Default::default(), id: self.state.node_builder.next_id() };
SsaFormingVisitor::new(self.state, SsaFormingInput { rename_defs: true }, self.program).consume_block(block)
}
}
fn defs_of(statement: &Statement) -> Vec<Symbol> {
match statement {
Statement::Definition(def) => match &def.place {
DefinitionPlace::Single(identifier) => vec![identifier.name],
DefinitionPlace::Multiple(identifiers) => identifiers.iter().map(|i| i.name).collect(),
},
Statement::Block(block) => block.statements.iter().flat_map(defs_of).collect(),
Statement::Conditional(conditional) => {
let mut defs: Vec<Symbol> = conditional.then.statements.iter().flat_map(defs_of).collect();
if let Some(otherwise) = &conditional.otherwise {
defs.extend(defs_of(otherwise));
}
defs
}
_ => Vec::new(),
}
}
fn statements_contain_return(statements: &[Statement]) -> bool {
statements.iter().any(statement_contains_return)
}
fn statement_contains_return(statement: &Statement) -> bool {
match statement {
Statement::Return(_) => true,
Statement::Block(block) => statements_contain_return(&block.statements),
Statement::Conditional(conditional) => conditional_contains_return(conditional),
_ => false,
}
}
fn conditional_contains_return(conditional: &ConditionalStatement) -> bool {
statements_contain_return(&conditional.then.statements)
|| conditional.otherwise.as_ref().is_some_and(|s| statement_contains_return(s))
}