diff --git a/src/lambda_lower.rs b/src/lambda_lower.rs new file mode 100644 index 0000000..6c4a132 --- /dev/null +++ b/src/lambda_lower.rs @@ -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>, + generated_functions: Rc>>, +} + +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, 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 { + 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 { + 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::, _>>()?; + 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::, _>>()?; + 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::, _>>()?; + ExprKind::Array(lowered_exprs) + } + ExprKind::Tuple(exprs) => { + let lowered_exprs = exprs + .iter() + .map(|e| self.lower_expr(e)) + .collect::, _>>()?; + 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::, _>>()?; + 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(), + }) + } +} diff --git a/src/lib.rs b/src/lib.rs index 7c8a819..352e25a 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -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; diff --git a/src/main.rs b/src/main.rs index faf0876..67af040 100644 --- a/src/main.rs +++ b/src/main.rs @@ -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 diff --git a/tests/lamba.sui b/tests/lamba.sui new file mode 100644 index 0000000..d31fc47 --- /dev/null +++ b/tests/lamba.sui @@ -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