396 lines
17 KiB
Rust
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, ¶m_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(),
|
|
})
|
|
}
|
|
}
|