Lambda lowering

This commit is contained in:
Masashi 2025-12-15 16:01:29 +05:30
commit 6a1362390c
4 changed files with 264 additions and 2 deletions

245
src/lambda_lower.rs Normal file
View file

@ -0,0 +1,245 @@
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())),
}
}
/// 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) => {
// Generate a unique name for this lambda function
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)
}
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::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::Float(_)
| ExprKind::Bool(_)
| ExprKind::String(_)
| ExprKind::Variable(_)
| ExprKind::Break
| ExprKind::Continue => expr.kind.clone(),
};
Ok(Expr {
kind: new_kind,
span: expr.span.clone(),
attributes: expr.attributes.clone(),
})
}
}

View file

@ -3,5 +3,6 @@ pub const EXTENSION: &str = ".sui";
pub mod ast;
pub mod lexer;
pub mod parser;
pub mod lambda_lower;
pub mod typechecker;
pub mod monomorphize;

View file

@ -1,7 +1,7 @@
use logos::Logos;
use std::fs;
use suicmez::{
lexer::Token, parser::Parser, typechecker::TypeChecker,
lexer::Token, parser::Parser, lambda_lower::LambdaLowerer, typechecker::TypeChecker,
monomorphize::{Monomorphizer, check_no_typevars},
};
@ -73,9 +73,16 @@ fn run_file(filename: &str) -> Result<(), String> {
println!("Parsed {} AST nodes successfully", ast_nodes.len());
// Lower lambdas to generated functions
let lowerer = LambdaLowerer::new();
let lowered_nodes = lowerer.lower_program(&ast_nodes)
.map_err(|e| format!("Lambda lowering error: {}", e))?;
println!("Lambda lowering passed! {} nodes after lowering.", lowered_nodes.len());
// Typecheck the AST
let mut typechecker = TypeChecker::new();
let typed_nodes = typechecker.typecheck_program(&ast_nodes).map_err(|e| {
let typed_nodes = typechecker.typecheck_program(&lowered_nodes).map_err(|e| {
format!(
"Type error at {}:{}: {:?}",
e.span.file, e.span.start, e.kind

9
tests/lamba.sui Normal file
View file

@ -0,0 +1,9 @@
# OCaml my Caml, our fearful trip is `done`
fn main(_) do
let mylamba = lambda (x) x+1;
mylamba(5)
let id = lambda (y) y;
id("hi")
id(1)
end