use brink_format::CountingFlags;
use rowan::TextRange;
use crate::hir;
use crate::symbols::SymbolKind;
use crate::{Diagnostic, DiagnosticCode};
use super::context::{LowerCtx, TempMap};
use super::lir;
pub(super) fn lower_lambda(l: &hir::LambdaExpr, ctx: &mut LowerCtx<'_>) -> lir::Expr {
let captures = captured_locals(l, ctx);
let mut param_names: Vec<String> = captures.iter().map(|(n, _)| n.clone()).collect();
for p in &l.params {
param_names.push(p.name.text.clone());
}
let mut temps = TempMap::new();
for (i, name) in param_names.iter().enumerate() {
#[expect(
clippy::cast_possible_truncation,
reason = "a lambda's captures + params never approach u16::MAX"
)]
temps.insert(name.clone(), i as u16);
}
#[expect(
clippy::cast_possible_truncation,
reason = "a lambda's captures + params never approach u16::MAX"
)]
let mut block_slot = param_names.len() as u16;
let params: Vec<lir::Param> = param_names
.iter()
.enumerate()
.map(|(i, name)| lir::Param {
name: ctx.names.intern(name),
#[expect(
clippy::cast_possible_truncation,
reason = "a lambda's captures + params never approach u16::MAX"
)]
slot: i as u16,
is_ref: false,
is_divert: false,
})
.collect();
let relative = lambda_scope_path(l, ctx);
let path = ctx.ids.qualify_lambda_path(&relative);
let id = l
.container_id
.unwrap_or_else(|| ctx.ids.alloc_address(&relative));
let borrowed_names: Vec<&str> = param_names.iter().map(String::as_str).collect();
let (body, children) = {
let mut lctx = LowerCtx {
file: ctx.file,
native: ctx.native,
resolutions: ctx.resolutions,
index: ctx.index,
temps: &temps,
names: ctx.names,
ids: ctx.ids,
scope_path: relative.clone(),
is_root_content_scope: false,
pending_children: Vec::new(),
visible_temps: borrowed_names.iter().map(|s| (*s).to_string()).collect(),
file_paths: ctx.file_paths,
root_id: ctx.root_id,
choice_gather_target: None,
next_block_slot: &mut block_slot,
block_scopes: Vec::new(),
as_binding_slots: crate::determinism::LookupSet::new(),
block_scoped_temp_names: crate::determinism::LookupSet::new(),
diagnostics: ctx.diagnostics,
loop_depth: 0,
structs: ctx.structs,
temp_shapes: crate::determinism::LookupMap::new(),
tables: ctx.tables,
lifted: ctx.lifted,
current_stmt_provenance: l.ptr,
};
let body = lower_body(&l.body, &mut lctx);
(body, core::mem::take(&mut lctx.pending_children))
};
ctx.lifted.push(lir::Container {
id,
provenance: l.ptr,
name: Some(path),
kind: lir::ContainerKind::Knot,
params,
body,
children,
counting_flags: CountingFlags::empty(),
temp_slot_count: block_slot,
labeled: false,
inline: false,
is_function: true,
local: false,
});
lir::ExprKind::MakeFnValue {
target: id,
bound: captures
.into_iter()
.map(|(name, slot)| lir::ExprKind::GetTemp(slot, ctx.names.intern(&name)).at(l.ptr))
.map(lir::CallArg::Value)
.collect(),
}
.at(l.ptr)
}
fn lambda_scope_path(l: &hir::LambdaExpr, ctx: &LowerCtx<'_>) -> String {
let offset = u32::from(l.ptr.text_range().start());
if ctx.scope_path.is_empty() {
format!("#lambda-{offset}")
} else {
format!("{}.#lambda-{offset}", ctx.scope_path)
}
}
fn lower_body(body: &hir::LambdaBody, ctx: &mut LowerCtx<'_>) -> Vec<lir::Stmt> {
match body {
hir::LambdaBody::Expr(e) => {
let value = super::expr::lower_expr(e, ctx);
vec![value_return(value, tail_provenance(e, ctx))]
}
hir::LambdaBody::Block { stmts, tail } => {
ctx.push_block_scope();
let mut out = super::blocks::lower_block_stmt_list(stmts, ctx);
if let Some(t) = tail {
let value = super::expr::lower_expr(t, ctx);
out.push(value_return(value, tail_provenance(t, ctx)));
}
ctx.pop_block_scope();
out
}
}
}
fn tail_provenance(tail: &hir::Expr, ctx: &LowerCtx<'_>) -> crate::Provenance {
crate::hir::expr_span(tail).map_or(ctx.current_stmt_provenance, |r| {
ctx.provenance_at(r, crate::NodeClass::Expr)
})
}
fn value_return(value: lir::Expr, provenance: crate::Provenance) -> lir::Stmt {
lir::Stmt::new(
lir::StmtKind::Return {
value: Some(value),
is_tunnel: false,
args: Vec::new(),
},
provenance,
)
}
fn captured_locals(l: &hir::LambdaExpr, ctx: &mut LowerCtx<'_>) -> Vec<(String, u16)> {
let mut scan = FreeScan {
bound: vec![l.params.iter().map(|p| p.name.text.clone()).collect()],
free: Vec::new(),
};
scan.body(&l.body);
scan.free
.into_iter()
.filter_map(|(name, range)| {
if let Some(slot) = ctx.temp_slot(&name) {
return Some((name, slot));
}
reject_unliftable_capture(&name, range, ctx);
None
})
.collect()
}
fn reject_unliftable_capture(name: &str, range: TextRange, ctx: &mut LowerCtx<'_>) {
let Some(info) = ctx.resolve_path(range) else {
return;
};
if !matches!(info.kind, SymbolKind::Temp | SymbolKind::Param) {
return;
}
ctx.diagnostics.push(Diagnostic {
file: ctx.file,
range,
message: format!(
"{}: `{name}` is a local the lambda cannot capture here — most likely its own \
`let` name, read before the `let` finishes binding (recursion is not supported)",
DiagnosticCode::E158.title(),
),
code: DiagnosticCode::E158,
});
}
struct FreeScan {
bound: Vec<Vec<String>>,
free: Vec<(String, TextRange)>,
}
impl FreeScan {
fn is_bound(&self, name: &str) -> bool {
self.bound.iter().any(|f| f.iter().any(|n| n == name))
}
fn read(&mut self, name: &str, range: TextRange) {
if !self.is_bound(name) && !self.free.iter().any(|(n, _)| n == name) {
self.free.push((name.to_string(), range));
}
}
fn bind(&mut self, name: &str) {
if let Some(frame) = self.bound.last_mut() {
frame.push(name.to_string());
}
}
fn body(&mut self, body: &hir::LambdaBody) {
match body {
hir::LambdaBody::Expr(e) => self.expr(e),
hir::LambdaBody::Block { stmts, tail } => {
self.bound.push(Vec::new());
self.stmts(stmts);
if let Some(t) = tail {
self.expr(t);
}
self.bound.pop();
}
}
}
fn stmts(&mut self, stmts: &[hir::BlockStmt]) {
for s in stmts {
self.stmt(s);
}
}
fn nested(&mut self, stmts: &[hir::BlockStmt]) {
self.bound.push(Vec::new());
self.stmts(stmts);
self.bound.pop();
}
fn stmt(&mut self, stmt: &hir::BlockStmt) {
match stmt {
hir::BlockStmt::TempDecl(t) => {
if let Some(e) = &t.value {
self.expr(e);
}
self.bind(&t.name.text);
}
hir::BlockStmt::Assignment(a) => {
self.expr(&a.target);
self.expr(&a.value);
}
hir::BlockStmt::Return(r) => {
if let Some(e) = &r.value {
self.expr(e);
}
for a in &r.onwards_args {
self.expr(a);
}
}
hir::BlockStmt::If(i) => self.if_stmt(i),
hir::BlockStmt::While(w) => {
self.expr(&w.condition);
self.bound.push(Vec::new());
if let Some(b) = &w.binding {
self.bind(&b.text);
}
self.stmts(&w.body);
self.bound.pop();
}
hir::BlockStmt::For(f) => {
self.expr(&f.iterable);
self.bound.push(Vec::new());
self.bind(&f.var_name.text);
if let Some(v) = &f.val_name {
self.bind(&v.text);
}
self.stmts(&f.body);
self.bound.pop();
}
hir::BlockStmt::ExprStmt(e) => self.expr(e),
hir::BlockStmt::Await(a) => {
if let Some(c) = &a.condition {
self.expr(c);
}
}
hir::BlockStmt::Break(_) | hir::BlockStmt::Continue(_) => {}
}
}
fn if_stmt(&mut self, i: &hir::IfStmt) {
self.expr(&i.condition);
self.bound.push(Vec::new());
if let Some(b) = &i.binding {
self.bind(&b.text);
}
self.stmts(&i.body);
self.bound.pop();
match &i.else_branch {
Some(hir::ElseBranch::ElseIf(nested)) => self.if_stmt(nested),
Some(hir::ElseBranch::Else(stmts)) => self.nested(stmts),
None => {}
}
}
fn expr(&mut self, e: &hir::Expr) {
match e {
hir::Expr::Path(p) => {
if let Some(head) = p.segments.first() {
self.read(&head.text, p.range);
}
}
hir::Expr::Prefix(_, inner) | hir::Expr::Postfix(inner, _) => self.expr(inner),
hir::Expr::Infix(ie) => {
self.expr(&ie.lhs);
self.expr(&ie.rhs);
}
hir::Expr::Call(p, args) => {
if let Some(head) = p.segments.first() {
self.read(&head.text, p.range);
}
for a in args {
self.expr(a);
}
}
hir::Expr::String(s) => {
for part in &s.parts {
if let hir::StringPart::Interpolation(inner) = part {
self.expr(inner);
}
}
}
hir::Expr::ArrayLiteral(a) => {
for el in &a.elements {
self.expr(el);
}
}
hir::Expr::MapLiteral(m) => {
for (k, v) in &m.entries {
self.expr(k);
self.expr(v);
}
}
hir::Expr::Index(idx) => {
self.expr(&idx.base);
self.expr(&idx.index);
}
hir::Expr::StructLiteral(sl) => {
for (_, v) in &sl.fields {
self.expr(v);
}
}
hir::Expr::FieldAccess(fa) => self.expr(&fa.base),
hir::Expr::FnLiteral(fl) => {
for a in &fl.args {
self.expr(a);
}
}
hir::Expr::RefArg(ra) => self.expr(&ra.operand),
hir::Expr::Lambda(inner) => {
self.bound
.push(inner.params.iter().map(|p| p.name.text.clone()).collect());
self.body(&inner.body);
self.bound.pop();
}
hir::Expr::Range(r) => {
self.expr(&r.start);
self.expr(&r.end);
}
hir::Expr::Int(_)
| hir::Expr::Float(_)
| hir::Expr::Bool(_)
| hir::Expr::Null
| hir::Expr::DivertTarget(_)
| hir::Expr::ListLiteral(_)
| hir::Expr::Fragment(_) => {}
}
}
}