suicmez/src/lambda_lower.rs
2025-12-20 01:55:12 +05:30

396 lines
17 KiB
Rust

use crate::ast::*;
use std::cell::RefCell;
use std::rc::Rc;
/// LambdaLowerer converts lambda expressions into generated functions
/// that are hoisted to the top level of the program.
pub struct LambdaLowerer {
lambda_counter: Rc<RefCell<usize>>,
generated_functions: Rc<RefCell<Vec<ASTNode>>>,
}
impl LambdaLowerer {
pub fn new() -> Self {
LambdaLowerer {
lambda_counter: Rc::new(RefCell::new(0)),
generated_functions: Rc::new(RefCell::new(Vec::new())),
}
}
fn collect_free_vars(
&self,
expr: &Expr,
lambda_params: &[String],
) -> std::collections::HashSet<String> {
let mut free_vars = std::collections::HashSet::new();
let mut local_scope = std::collections::HashSet::new();
self.collect_free_vars_expr(expr, lambda_params, &mut free_vars, &mut local_scope);
free_vars
}
fn collect_free_vars_expr(
&self,
expr: &Expr,
lambda_params: &[String],
free_vars: &mut std::collections::HashSet<String>,
local_scope: &mut std::collections::HashSet<String>,
) {
match &expr.kind {
ExprKind::Variable(name) => {
if !lambda_params.contains(name) && !local_scope.contains(name) {
free_vars.insert(name.clone());
}
}
ExprKind::Lambda(args, body) => {
let param_names: Vec<String> = args.iter().map(|(name, _)| name.clone()).collect();
// For nested lambdas, we don't enter a new scope here since we're just collecting free vars
self.collect_free_vars_expr(body, &param_names, free_vars, local_scope);
}
ExprKind::Let(name, _, _, body) => {
// Collect from body before adding the binding
self.collect_free_vars_expr(body, lambda_params, free_vars, local_scope);
// Add to local scope
local_scope.insert(name.clone());
}
ExprKind::Call(func, args) => {
self.collect_free_vars_expr(func, lambda_params, free_vars, local_scope);
for arg in args {
self.collect_free_vars_expr(arg, lambda_params, free_vars, local_scope);
}
}
// Handle other expression types that contain subexpressions
ExprKind::If(cond, then_expr, else_expr) => {
self.collect_free_vars_expr(cond, lambda_params, free_vars, local_scope);
self.collect_free_vars_expr(then_expr, lambda_params, free_vars, local_scope);
if let Some(else_expr) = else_expr {
self.collect_free_vars_expr(else_expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::Match(scrutinee, arms) => {
self.collect_free_vars_expr(scrutinee, lambda_params, free_vars, local_scope);
for (_, expr) in arms {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::While(cond, body) => {
self.collect_free_vars_expr(cond, lambda_params, free_vars, local_scope);
self.collect_free_vars_expr(body, lambda_params, free_vars, local_scope);
}
ExprKind::For(_, iter, body) => {
self.collect_free_vars_expr(iter, lambda_params, free_vars, local_scope);
self.collect_free_vars_expr(body, lambda_params, free_vars, local_scope);
}
ExprKind::Do(exprs) => {
for expr in exprs {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::BinOp(left, _, right) => {
self.collect_free_vars_expr(left, lambda_params, free_vars, local_scope);
self.collect_free_vars_expr(right, lambda_params, free_vars, local_scope);
}
ExprKind::UnOp(_, operand) => {
self.collect_free_vars_expr(operand, lambda_params, free_vars, local_scope);
}
ExprKind::Assign(target, value) => {
self.collect_free_vars_expr(target, lambda_params, free_vars, local_scope);
self.collect_free_vars_expr(value, lambda_params, free_vars, local_scope);
}
ExprKind::Cast(operand, _) => {
self.collect_free_vars_expr(operand, lambda_params, free_vars, local_scope);
}
ExprKind::Index(obj, index) => {
self.collect_free_vars_expr(obj, lambda_params, free_vars, local_scope);
self.collect_free_vars_expr(index, lambda_params, free_vars, local_scope);
}
ExprKind::Dot(obj, _) => {
self.collect_free_vars_expr(obj, lambda_params, free_vars, local_scope);
}
ExprKind::EarlyReturn(expr) => {
if let Some(expr) = expr {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::OptionalChain(obj, _) => {
if let Some(obj) = obj {
self.collect_free_vars_expr(obj, lambda_params, free_vars, local_scope);
}
}
ExprKind::Return(expr) => {
if let Some(expr) = expr {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::Array(exprs) => {
for expr in exprs {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::FixedArray(expr, _) => {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
ExprKind::Tuple(exprs) => {
for expr in exprs {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::StructLit(_, fields) => {
for (_, expr) in fields {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
}
ExprKind::EnumLit(_, _, args) => {
for arg in args {
self.collect_free_vars_expr(arg, lambda_params, free_vars, local_scope);
}
}
ExprKind::Range(start, end) => {
self.collect_free_vars_expr(start, lambda_params, free_vars, local_scope);
self.collect_free_vars_expr(end, lambda_params, free_vars, local_scope);
}
ExprKind::Defer(expr) => {
self.collect_free_vars_expr(expr, lambda_params, free_vars, local_scope);
}
// Terminal expressions don't contain variables
ExprKind::Int(_)
| ExprKind::TypedInt(_, _)
| ExprKind::Float(_)
| ExprKind::Bool(_)
| ExprKind::String(_)
| ExprKind::Break
| ExprKind::Continue => {}
}
}
/// Lower all lambdas in a program by hoisting them to functions
pub fn lower_program(&self, nodes: &[ASTNode]) -> Result<Vec<ASTNode>, String> {
let mut lowered_nodes = Vec::new();
// Process each top-level node
for node in nodes {
let lowered = self.lower_node(node)?;
lowered_nodes.push(lowered);
}
// Add all generated lambda functions to the end
let generated = self.generated_functions.borrow();
lowered_nodes.extend(generated.iter().cloned());
Ok(lowered_nodes)
}
fn lower_node(&self, node: &ASTNode) -> Result<ASTNode, String> {
let new_kind = match &node.kind {
ASTNodeKind::Function(func) => {
let lowered_body = self.lower_expr(&func.body)?;
ASTNodeKind::Function(Function {
name: func.name.clone(),
parameters: func.parameters.clone(),
args: func.args.clone(),
return_type: func.return_type.clone(),
body: lowered_body,
})
}
other => other.clone(),
};
Ok(ASTNode {
kind: new_kind,
span: node.span.clone(),
attributes: node.attributes.clone(),
})
}
fn lower_expr(&self, expr: &Expr) -> Result<Expr, String> {
let new_kind = match &expr.kind {
ExprKind::Lambda(args, body) => {
// Collect free variables (captured variables)
let lambda_params: Vec<String> =
args.iter().map(|(name, _)| name.clone()).collect();
let free_vars = self.collect_free_vars(body, &lambda_params);
if free_vars.is_empty() {
// No captures - hoist to function like the original implementation
let lambda_id = {
let mut counter = self.lambda_counter.borrow_mut();
*counter += 1;
*counter
};
let lambda_name = format!("__suic_gen_lambda_{}", lambda_id);
// Lower the lambda body recursively
let lowered_body = self.lower_expr(body)?;
// Create a new function for this lambda
let lambda_func = ASTNode {
kind: ASTNodeKind::Function(Function {
name: lambda_name.clone(),
parameters: Vec::new(), // No type parameters for now
args: args.clone(),
return_type: None, // Let typechecker infer return type
body: lowered_body,
}),
span: expr.span.clone(),
attributes: Vec::new(),
};
// Store the generated function
self.generated_functions.borrow_mut().push(lambda_func);
// Replace the lambda with a reference to the generated function
ExprKind::Variable(lambda_name)
} else {
// Has captures - keep as lambda, but recursively lower the body
let lowered_body = self.lower_expr(body)?;
ExprKind::Lambda(args.clone(), Box::new(lowered_body))
}
}
ExprKind::Call(func, args) => {
let lowered_func = self.lower_expr(func)?;
let lowered_args = args
.iter()
.map(|arg| self.lower_expr(arg))
.collect::<Result<Vec<_>, _>>()?;
ExprKind::Call(Box::new(lowered_func), lowered_args)
}
ExprKind::Let(name, kind, type_annot, body) => {
let lowered_body = self.lower_expr(body)?;
ExprKind::Let(
name.clone(),
kind.clone(),
type_annot.clone(),
Box::new(lowered_body),
)
}
ExprKind::If(cond, then_expr, else_expr) => {
let lowered_cond = self.lower_expr(cond)?;
let lowered_then = self.lower_expr(then_expr)?;
let lowered_else = else_expr.as_ref().map(|e| self.lower_expr(e)).transpose()?;
ExprKind::If(
Box::new(lowered_cond),
Box::new(lowered_then),
lowered_else.map(Box::new),
)
}
ExprKind::Match(scrutinee, arms) => {
let lowered_scrutinee = self.lower_expr(scrutinee)?;
let mut lowered_arms = Vec::new();
for (pattern, expr) in arms {
let lowered_expr = self.lower_expr(expr)?;
lowered_arms.push((pattern.clone(), lowered_expr));
}
ExprKind::Match(Box::new(lowered_scrutinee), lowered_arms)
}
ExprKind::While(cond, body) => {
let lowered_cond = self.lower_expr(cond)?;
let lowered_body = self.lower_expr(body)?;
ExprKind::While(Box::new(lowered_cond), Box::new(lowered_body))
}
ExprKind::For(var, iter, body) => {
let lowered_iter = self.lower_expr(iter)?;
let lowered_body = self.lower_expr(body)?;
ExprKind::For(var.clone(), Box::new(lowered_iter), Box::new(lowered_body))
}
ExprKind::Do(exprs) => {
let lowered_exprs = exprs
.iter()
.map(|e| self.lower_expr(e))
.collect::<Result<Vec<_>, _>>()?;
ExprKind::Do(lowered_exprs)
}
ExprKind::BinOp(left, op, right) => {
let lowered_left = self.lower_expr(left)?;
let lowered_right = self.lower_expr(right)?;
ExprKind::BinOp(Box::new(lowered_left), op.clone(), Box::new(lowered_right))
}
ExprKind::UnOp(op, operand) => {
let lowered_operand = self.lower_expr(operand)?;
ExprKind::UnOp(op.clone(), Box::new(lowered_operand))
}
ExprKind::Assign(target, value) => {
let lowered_target = self.lower_expr(target)?;
let lowered_value = self.lower_expr(value)?;
ExprKind::Assign(Box::new(lowered_target), Box::new(lowered_value))
}
ExprKind::Cast(operand, type_annot) => {
let lowered_operand = self.lower_expr(operand)?;
ExprKind::Cast(Box::new(lowered_operand), type_annot.clone())
}
ExprKind::Index(obj, index) => {
let lowered_obj = self.lower_expr(obj)?;
let lowered_index = self.lower_expr(index)?;
ExprKind::Index(Box::new(lowered_obj), Box::new(lowered_index))
}
ExprKind::Dot(obj, field) => {
let lowered_obj = self.lower_expr(obj)?;
ExprKind::Dot(Box::new(lowered_obj), field.clone())
}
ExprKind::EarlyReturn(expr) => {
let lowered_expr = expr.as_ref().map(|e| self.lower_expr(e)).transpose()?;
ExprKind::EarlyReturn(lowered_expr.map(Box::new))
}
ExprKind::OptionalChain(obj, field) => {
let lowered_obj = obj.as_ref().map(|e| self.lower_expr(e)).transpose()?;
ExprKind::OptionalChain(lowered_obj.map(Box::new), field.clone())
}
ExprKind::Return(expr) => {
let lowered_expr = expr.as_ref().map(|e| self.lower_expr(e)).transpose()?;
ExprKind::Return(lowered_expr.map(Box::new))
}
ExprKind::Array(exprs) => {
let lowered_exprs = exprs
.iter()
.map(|e| self.lower_expr(e))
.collect::<Result<Vec<_>, _>>()?;
ExprKind::Array(lowered_exprs)
}
ExprKind::FixedArray(expr, size) => {
let lowered_expr = self.lower_expr(expr)?;
ExprKind::FixedArray(Box::new(lowered_expr), *size)
}
ExprKind::Tuple(exprs) => {
let lowered_exprs = exprs
.iter()
.map(|e| self.lower_expr(e))
.collect::<Result<Vec<_>, _>>()?;
ExprKind::Tuple(lowered_exprs)
}
ExprKind::StructLit(name, fields) => {
let mut lowered_fields = Vec::new();
for (field_name, field_expr) in fields {
let lowered_expr = self.lower_expr(field_expr)?;
lowered_fields.push((field_name.clone(), lowered_expr));
}
ExprKind::StructLit(name.clone(), lowered_fields)
}
ExprKind::EnumLit(enum_name, variant, args) => {
let lowered_args = args
.iter()
.map(|arg| self.lower_expr(arg))
.collect::<Result<Vec<_>, _>>()?;
ExprKind::EnumLit(enum_name.clone(), variant.clone(), lowered_args)
}
ExprKind::Range(start, end) => {
let lowered_start = self.lower_expr(start)?;
let lowered_end = self.lower_expr(end)?;
ExprKind::Range(Box::new(lowered_start), Box::new(lowered_end))
}
// Terminal expressions that don't contain other expressions
ExprKind::Int(_)
| ExprKind::TypedInt(_, _)
| ExprKind::Float(_)
| ExprKind::Bool(_)
| ExprKind::String(_)
| ExprKind::Variable(_)
| ExprKind::Break
| ExprKind::Continue => expr.kind.clone(),
ExprKind::Defer(inner) => ExprKind::Defer(Box::new(self.lower_expr(inner)?)),
};
Ok(Expr {
kind: new_kind,
span: expr.span.clone(),
attributes: expr.attributes.clone(),
})
}
}