use super::{CodeGen, CodeGenError};
use crate::ast::{Statement, WordDef};
use crate::call_graph::CallGraph;
use std::fmt::Write as _;
pub(super) struct LoopPattern<'a> {
prelude: &'a [Statement],
base_branch: &'a [Statement],
rec_branch: &'a [Statement],
base_is_then: bool,
}
pub(super) fn detect_loop_pattern<'a>(
word: &'a WordDef,
call_graph: &CallGraph,
) -> Option<LoopPattern<'a>> {
if word.name == "main" {
return None;
}
if !call_graph.is_self_recursive(&word.name) {
return None;
}
let last = word.body.last()?;
let (then_branch, else_branch) = match last {
Statement::If {
then_branch,
else_branch: Some(else_branch),
..
} => (then_branch.as_slice(), else_branch.as_slice()),
_ => return None,
};
let then_recs = ends_in_self_call(then_branch, &word.name);
let else_recs = ends_in_self_call(else_branch, &word.name);
let (rec_branch, base_branch, base_is_then) = match (then_recs, else_recs) {
(true, false) => (then_branch, else_branch, false), (false, true) => (else_branch, then_branch, true), _ => return None,
};
if calls_self(base_branch, &word.name) {
return None;
}
let prelude = &word.body[..word.body.len() - 1];
Some(LoopPattern {
prelude,
base_branch,
rec_branch,
base_is_then,
})
}
fn ends_in_self_call(branch: &[Statement], self_name: &str) -> bool {
matches!(
branch.last(),
Some(Statement::WordCall { name, .. }) if name == self_name
)
}
fn calls_self(branch: &[Statement], self_name: &str) -> bool {
branch.iter().any(|s| stmt_calls_self(s, self_name))
}
fn stmt_calls_self(stmt: &Statement, self_name: &str) -> bool {
match stmt {
Statement::WordCall { name, .. } => name == self_name,
Statement::If {
then_branch,
else_branch,
..
} => {
calls_self(then_branch, self_name)
|| else_branch
.as_ref()
.is_some_and(|eb| calls_self(eb, self_name))
}
Statement::Quotation { body, .. } => calls_self(body, self_name),
Statement::Match { arms, .. } => arms.iter().any(|arm| calls_self(&arm.body, self_name)),
_ => false,
}
}
impl CodeGen {
pub(super) fn codegen_loop_body(
&mut self,
entry_stack: &str,
pattern: &LoopPattern,
) -> Result<(), CodeGenError> {
let loop_lbl = self.fresh_block("loop");
let base_lbl = self.fresh_block("loop_base");
let cont_lbl = self.fresh_block("loop_cont");
let yield_lbl = self.fresh_block("loop_yield");
let loop_sp = self.fresh_named("loop_sp");
let cont_sp = self.fresh_named("loop_cont_sp");
let iter_phi = self.fresh_named("loop_iter");
let iter_next = self.fresh_named("loop_iter_next");
let mask = (self.loop_yield_cadence as u64).saturating_sub(1);
let aux_sp_entry = self.current_aux_sp;
writeln!(&mut self.output, " br label %{}", loop_lbl)?;
writeln!(&mut self.output, "{}:", loop_lbl)?;
writeln!(
&mut self.output,
" %{} = phi ptr [ %{}, %entry ], [ %{}, %{} ], [ %{}, %{} ]",
loop_sp, entry_stack, cont_sp, cont_lbl, cont_sp, yield_lbl
)?;
writeln!(
&mut self.output,
" %{} = phi i64 [ 0, %entry ], [ %{}, %{} ], [ %{}, %{} ]",
iter_phi, iter_next, cont_lbl, iter_next, yield_lbl
)?;
self.virtual_stack.clear();
self.current_aux_sp = aux_sp_entry;
let cond_stack = self.codegen_statements(pattern.prelude, &loop_sp, false)?;
let cond_stack = self.spill_virtual_stack(&cond_stack)?;
let top_ptr = self.emit_stack_gep(&cond_stack, -1)?;
let cond_val = self.emit_load_int_payload(&top_ptr)?;
let popped = top_ptr; let cmp = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = icmp ne i64 %{}, 0",
cmp, cond_val
)?;
let (true_lbl, false_lbl) = if pattern.base_is_then {
(base_lbl.as_str(), cont_lbl.as_str())
} else {
(cont_lbl.as_str(), base_lbl.as_str())
};
writeln!(
&mut self.output,
" br i1 %{}, label %{}, label %{}",
cmp, true_lbl, false_lbl
)?;
writeln!(&mut self.output, "{}:", base_lbl)?;
self.virtual_stack.clear();
self.current_aux_sp = aux_sp_entry;
let base_stack = self.codegen_statements(pattern.base_branch, &popped, false)?;
let base_stack = self.spill_virtual_stack(&base_stack)?;
writeln!(&mut self.output, " ret ptr %{}", base_stack)?;
writeln!(&mut self.output, "{}:", cont_lbl)?;
self.virtual_stack.clear();
self.current_aux_sp = aux_sp_entry;
let rec_prefix = &pattern.rec_branch[..pattern.rec_branch.len() - 1];
let cont_stack = self.codegen_statements(rec_prefix, &popped, false)?;
let cont_stack = self.spill_virtual_stack(&cont_stack)?;
writeln!(
&mut self.output,
" %{} = bitcast ptr %{} to ptr",
cont_sp, cont_stack
)?;
writeln!(
&mut self.output,
" %{} = add i64 %{}, 1",
iter_next, iter_phi
)?;
let masked = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = and i64 %{}, {}",
masked, iter_next, mask
)?;
let need_yield = self.fresh_temp();
writeln!(
&mut self.output,
" %{} = icmp eq i64 %{}, 0",
need_yield, masked
)?;
writeln!(
&mut self.output,
" br i1 %{}, label %{}, label %{}",
need_yield, yield_lbl, loop_lbl
)?;
writeln!(&mut self.output, "{}:", yield_lbl)?;
writeln!(&mut self.output, " call void @patch_seq_maybe_yield()")?;
writeln!(&mut self.output, " br label %{}", loop_lbl)?;
Ok(())
}
}