From 32b53ea975cffa48cb8389b6da8321d33c1f6ecd Mon Sep 17 00:00:00 2001 From: Masashi Date: Mon, 15 Dec 2025 14:33:06 +0530 Subject: [PATCH 1/2] yum --- output.md | 2843 -------------------------------------------- src/typechecker.rs | 1514 ++++++++++++++++++++++- 2 files changed, 1513 insertions(+), 2844 deletions(-) delete mode 100644 output.md diff --git a/output.md b/output.md deleted file mode 100644 index 9b32d6e..0000000 --- a/output.md +++ /dev/null @@ -1,2843 +0,0 @@ -```rust -// src/lib.rs -pub const EXTENSION: &str = ".sui"; - -pub mod ast; -pub mod lexer; -pub mod parser; -pub mod typechecker; - -``` - -```rust -// src/parser.rs -use crate::ast::*; -use crate::lexer::Token; - -use std::iter::Peekable; -use std::ops::Range; -use std::vec::IntoIter; - -type TokenIter = Peekable)>>; - -pub struct Parser { - pub file: String, - pub tokens: TokenIter, -} - -#[derive(Debug)] -pub struct ParseError { - pub message: String, - pub span: Span, -} - -impl Parser { - pub fn new(file: String, tokens: Vec<(Token, Range)>) -> Self { - Parser { - file, - tokens: tokens.into_iter().peekable(), - } - } - - // Parse the entire file into a list of AST nodes - pub fn parse(&mut self) -> Result, ParseError> { - let mut nodes = Vec::new(); - - while self.peek().is_some() { - nodes.push(self.parse_top_level()?); - } - - Ok(nodes) - } - - fn peek(&mut self) -> Option<&Token> { - self.tokens.peek().map(|(token, _)| token) - } - - fn peek_span(&mut self) -> Option> { - self.tokens.peek().map(|(_, span)| span.clone()) - } - - fn next(&mut self) -> Option<(Token, Range)> { - self.tokens.next() - } - - fn expect(&mut self, expected: Token) -> Result, ParseError> { - match self.next() { - Some((token, span)) - if std::mem::discriminant(&token) == std::mem::discriminant(&expected) => - { - Ok(span) - } - Some((token, span)) => Err(ParseError { - message: format!("Expected {:?}, found {:?}", expected, token), - span: Span::new(&span, self.file.clone()), - }), - None => Err(ParseError { - message: format!("Expected {:?}, found EOF", expected), - span: Span::new(&(0..0), self.file.clone()), - }), - } - } - - fn error(&self, msg: String, span: Range) -> Result { - Err(ParseError { - message: msg, - span: Span::new(&span, self.file.clone()), - }) - } - - fn parse_top_level(&mut self) -> Result { - let mut attributes = Vec::new(); - - // Parse any leading attributes - while matches!(self.peek(), Some(Token::At)) { - attributes.push(self.parse_attribute()?); - } - - let start = self.peek_span().unwrap_or(0..0).start; - let token = self.peek().cloned(); - match token { - Some(Token::KeywordUse) => { - self.next(); - let path = match self.next() { - Some((Token::String(s), _)) => s, - Some((_, span)) => { - return self.error("Expected string after 'use'".to_string(), span); - } - None => { - return self.error("Expected string after 'use'".to_string(), start..start); - } - }; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Use(path), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(Token::KeywordFn) => { - self.next(); - let func = self.parse_function()?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Function(func), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(Token::KeywordStruct) => { - self.next(); - let struct_def = self.parse_struct()?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Struct(struct_def), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(Token::KeywordEnum) => { - self.next(); - let enum_def = self.parse_enum()?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Enum(enum_def), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(Token::KeywordImpl) => { - self.next(); - let impl_def = self.parse_impl()?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Impl(impl_def), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(Token::KeywordTrait) => { - self.next(); - let trait_def = self.parse_trait()?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Trait(trait_def), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(Token::KeywordExtern) => { - self.next(); - let extern_def = self.parse_extern()?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Extern(extern_def), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(Token::KeywordLoad) => { - self.next(); - let load_def = self.parse_load()?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(ASTNode { - kind: ASTNodeKind::Load(load_def), - span: Span::new(&(start..end), self.file.clone()), - attributes, - }) - } - Some(token) => { - let span = self.peek_span().unwrap_or(start..start); - self.error(format!("Unexpected token at top level: {:?}", token), span) - } - None => self.error("Unexpected EOF".to_string(), start..start), - } - } - - fn parse_attribute(&mut self) -> Result { - self.expect(Token::At)?; - let start = self.peek_span().unwrap_or(0..0).start; - let name = match self.next() { - Some((Token::Variable(name), _)) => name, - Some((_, span)) => return self.error("Expected attribute name".to_string(), span), - None => return self.error("Expected attribute name".to_string(), start..start), - }; - - // Parentheses are optional - let mut args = vec![]; - if matches!(self.peek(), Some(Token::LParen)) { - self.next(); - loop { - let token = self.peek().cloned(); - match token { - Some(Token::RParen) => { - self.next(); - break; - } - Some(Token::String(s)) => { - self.next(); - args.push(AttributeArg::Literal(s)); - } - Some(Token::Variable(id)) => { - self.next(); - let next_token = self.peek().cloned(); - if matches!(next_token, Some(Token::Assign)) { - self.next(); - match self.next() { - Some((Token::Variable(val), _)) => { - args.push(AttributeArg::KeyValue(id, val)) - } - Some((_, span)) => { - return self.error("Expected value after =".to_string(), span); - } - None => { - return self - .error("Expected value after =".to_string(), start..start); - } - } - } else { - args.push(AttributeArg::Value(id)); - } - } - Some(token) => { - let span = self.peek_span().unwrap_or(start..start); - return self - .error(format!("Unexpected token in attribute: {:?}", token), span); - } - None => { - return self.error("Expected attribute argument".to_string(), start..start); - } - } - let next_token = self.peek().cloned(); - if matches!(next_token, Some(Token::Comma)) { - self.next(); - } else if matches!(next_token, Some(Token::RParen)) { - // ok - } else { - { - let span = self.peek_span().unwrap_or(start..start); - return self.error("Expected , or )".to_string(), span); - } - } - } - } - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Attribute { - name, - args, - span: Span::new(&(start..end), self.file.clone()), - }) - } - - fn parse_function(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - let name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected function name".to_string(), span), - None => return self.error("Expected function name".to_string(), start..start), - }; - - // Parse type parameters if present - let parameters = if matches!(self.peek(), Some(Token::Less)) { - self.next(); - self.parse_parameters()? - } else { - Vec::new() - }; - - // Parse function arguments - self.expect(Token::LParen)?; - let mut args = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - - let arg_name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected argument name".to_string(), span), - None => return self.error("Expected argument name".to_string(), start..start), - }; - - let arg_type = if matches!(self.peek(), Some(Token::Colon)) { - self.next(); - Some(self.parse_type_annot()?) - } else { - None - }; - - args.push((arg_name, arg_type)); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } else if !matches!(self.peek(), Some(Token::RParen)) { - let span = self.peek_span().unwrap_or(start..start); - return self.error("Expected , or )".to_string(), span); - } - } - - // Parse return type if present - let return_type = if matches!(self.peek(), Some(Token::Arrow)) { - self.next(); - Some(self.parse_type_annot()?) - } else { - None - }; - - // Parse body expression - let body = self.parse_expr()?; - - Ok(Function { - name, - parameters, - args, - return_type, - body, - }) - } - - fn parse_struct(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - let name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected struct name".to_string(), span), - None => return self.error("Expected struct name".to_string(), start..start), - }; - - // Parse type parameters if present - let parameters = if matches!(self.peek(), Some(Token::Less)) { - self.next(); - self.parse_parameters()? - } else { - Vec::new() - }; - - // Parse fields - let mut fields = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::KeywordEnd)) { - self.next(); - break; - } - - let field_start = self.peek_span().unwrap_or(0..0).start; - let field_name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected field name".to_string(), span), - None => return self.error("Expected field name".to_string(), start..start), - }; - - self.expect(Token::Colon)?; - let field_type = self.parse_type_annot()?; - let field_end = self.peek_span().unwrap_or(field_start..field_start).start; - - fields.push(Field { - name: field_name, - field_type, - span: Span::new(&(field_start..field_end), self.file.clone()), - }); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - Ok(Struct { - name, - parameters, - fields, - }) - } - - fn parse_enum(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - let name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected enum name".to_string(), span), - None => return self.error("Expected enum name".to_string(), start..start), - }; - - // Parse type parameters if present - let parameters = if matches!(self.peek(), Some(Token::Less)) { - self.next(); - self.parse_parameters()? - } else { - Vec::new() - }; - - // Parse variants - let mut variants = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::KeywordEnd)) { - self.next(); - break; - } - - let variant_start = self.peek_span().unwrap_or(0..0).start; - let variant_name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected variant name".to_string(), span), - None => return self.error("Expected variant name".to_string(), start..start), - }; - - let mut fields = Vec::new(); - if matches!(self.peek(), Some(Token::LParen)) { - self.next(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - fields.push(self.parse_type_annot()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - } - - let variant_end = self - .peek_span() - .unwrap_or(variant_start..variant_start) - .start; - variants.push(Variant { - name: variant_name, - fields, - span: Span::new(&(variant_start..variant_end), self.file.clone()), - }); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - Ok(Enum { - name, - parameters, - variants, - }) - } - - fn parse_impl(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - - // Parse impl target as a type (could be generic like Option) - let target_type = self.parse_type_annot()?; - - // Extract the base type name from the type annotation - let target = match target_type { - TypeAnnot::Var(name) => name, - TypeAnnot::Cons(name, _) => name, - _ => { - return self.error( - "Expected type name for impl target".to_string(), - start..start, - ); - } - }; - - // Parse optional trait name - let trait_name = if matches!(self.peek(), Some(Token::Colon)) { - self.next(); - match self.next() { - Some((Token::Variable(n), _)) => Some(n), - Some((_, span)) => return self.error("Expected trait name".to_string(), span), - None => return self.error("Expected trait name".to_string(), start..start), - } - } else { - None - }; - - // Parse methods - let mut methods = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::KeywordEnd)) { - self.next(); - break; - } - - self.expect(Token::KeywordFn)?; - methods.push(self.parse_function()?); - } - - Ok(Impl { - target, - trait_name, - methods, - }) - } - - fn parse_trait(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - let name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected trait name".to_string(), span), - None => return self.error("Expected trait name".to_string(), start..start), - }; - - // Parse type parameters if present - let parameters = if matches!(self.peek(), Some(Token::Less)) { - self.next(); - self.parse_parameters()? - } else { - Vec::new() - }; - - // Parse methods - let mut methods = Vec::new(); - - loop { - if matches!(self.peek(), Some(Token::KeywordEnd)) { - self.next(); - break; - } - - if matches!(self.peek(), Some(Token::KeywordFn)) { - self.next(); - methods.push(self.parse_function_signature()?); - - // Optional comma between methods - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } else { - break; - } - } - - Ok(Trait { - name, - methods, - parameters, - associated_types: Vec::new(), - }) - } - - fn parse_extern(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - let name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected extern name".to_string(), span), - None => return self.error("Expected extern name".to_string(), start..start), - }; - - // Parse argument types (with optional parameter names) - self.expect(Token::LParen)?; - let mut args = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - - args.push(self.parse_type_annot()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - // Parse return type - self.expect(Token::Arrow)?; - let return_type = self.parse_type_annot()?; - - // Parse from clause - self.expect(Token::KeywordFrom)?; - let from = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected library identifier".to_string(), span), - None => return self.error("Expected library identifier".to_string(), start..start), - }; - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Extern { - name, - args, - return_type, - from, - span: Span::new(&(start..end), self.file.clone()), - }) - } - - fn parse_load(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - let library = match self.next() { - Some((Token::String(s), _)) => s, - Some((_, span)) => return self.error("Expected library name".to_string(), span), - None => return self.error("Expected library name".to_string(), start..start), - }; - - self.expect(Token::KeywordAs)?; - let alias = match self.next() { - Some((Token::Variable(a), _)) => a, - Some((_, span)) => return self.error("Expected alias".to_string(), span), - None => return self.error("Expected alias".to_string(), start..start), - }; - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Load { - library, - alias, - span: Span::new(&(start..end), self.file.clone()), - }) - } - - fn parse_parameters(&mut self) -> Result, ParseError> { - let start = self.peek_span().unwrap_or(0..0).start; - let mut params = Vec::new(); - - loop { - if matches!(self.peek(), Some(Token::Greater)) { - self.next(); - break; - } - - let param_start = self.peek_span().unwrap_or(0..0).start; - let param_name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected parameter name".to_string(), span), - None => return self.error("Expected parameter name".to_string(), start..start), - }; - - let bounds = if matches!(self.peek(), Some(Token::Colon)) { - self.next(); - self.parse_trait_bounds()? - } else { - Vec::new() - }; - - let kind = if matches!(self.peek(), Some(Token::Colon)) { - self.next(); - Some(self.parse_kind()?) - } else { - None - }; - - let param_end = self.peek_span().unwrap_or(param_start..param_start).end; - params.push(Parameter { - name: param_name, - bounds, - kind, - span: Span::new(&(param_start..param_end), self.file.clone()), - }); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - Ok(params) - } - - fn parse_trait_bounds(&mut self) -> Result, ParseError> { - let mut bounds = Vec::new(); - loop { - match self.next() { - Some((Token::Variable(n), _)) => bounds.push(n), - Some((_, span)) => return self.error("Expected trait name".to_string(), span), - None => return self.error("Expected trait name".to_string(), 0..0), - } - - if !matches!(self.peek(), Some(Token::Plus)) { - break; - } - self.next(); - } - - Ok(bounds) - } - - fn parse_kind(&mut self) -> Result { - if matches!(self.peek(), Some(Token::Mul)) { - self.next(); - Ok(Kind::Star) - } else { - let k1 = Box::new(self.parse_kind()?); - self.expect(Token::Arrow)?; - let k2 = Box::new(self.parse_kind()?); - Ok(Kind::Arrow(k1, k2)) - } - } - - fn parse_type_annot(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - - // Check for function type: fn (args)->ret - if matches!(self.peek(), Some(Token::KeywordFn)) { - self.next(); - self.expect(Token::LParen)?; - let mut arg_types = Vec::new(); - - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - arg_types.push(self.parse_type_annot()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - self.expect(Token::Arrow)?; - let ret_type = Box::new(self.parse_type_annot()?); - return Ok(TypeAnnot::Function(arg_types, ret_type)); - } - - let mut base_type = match self.next() { - Some((Token::Variable(n), _)) => TypeAnnot::Cons(n, vec![]), - Some((Token::KeywordBool, _)) => TypeAnnot::Cons("bool".to_string(), vec![]), - Some((Token::KeywordInt, _)) => TypeAnnot::Cons("int".to_string(), vec![]), - Some((Token::KeywordFloat, _)) => TypeAnnot::Cons("float".to_string(), vec![]), - Some((Token::KeywordString, _)) => TypeAnnot::Cons("string".to_string(), vec![]), - Some((Token::LParen, _)) => { - // Check for unit type: () - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - return Ok(TypeAnnot::Cons("unit".to_string(), vec![])); - } - - let mut types = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - types.push(self.parse_type_annot()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - // Single element in parens is not a tuple, unwrap it - if types.len() == 1 { - types.pop().unwrap() - } else { - TypeAnnot::Tuple(types) - } - } - Some((Token::LBracket, _)) => { - let inner = self.parse_type_annot()?; - self.expect(Token::RBracket)?; - TypeAnnot::Array(Box::new(inner)) - } - Some((Token::Bang, _)) => TypeAnnot::Cons("never".to_string(), vec![]), - Some((_, span)) => return self.error("Expected type name".to_string(), span), - None => return self.error("Expected type name".to_string(), start..start), - }; - - // Parse type arguments if present - if matches!(self.peek(), Some(Token::Less)) { - self.next(); - let mut args = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::Greater)) { - self.next(); - break; - } - args.push(self.parse_type_annot()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - base_type = match base_type { - TypeAnnot::Cons(name, _) => TypeAnnot::Cons(name, args), - _ => { - return self.error("Expected type name for generic".to_string(), start..start); - } - }; - } - - // Parse array types - while matches!(self.peek(), Some(Token::LBracket)) { - self.next(); - self.expect(Token::RBracket)?; - base_type = TypeAnnot::Array(Box::new(base_type)); - } - - Ok(base_type) - } - - fn parse_function_signature(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - let name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected function name".to_string(), span), - None => return self.error("Expected function name".to_string(), start..start), - }; - - self.expect(Token::LParen)?; - let mut params = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - - let param_start = self.peek_span().unwrap_or(0..0).start; - let param_name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => return self.error("Expected parameter name".to_string(), span), - None => return self.error("Expected parameter name".to_string(), start..start), - }; - - // Parameters in trait methods may have type annotations - if matches!(self.peek(), Some(Token::Colon)) { - self.next(); - let _param_type = self.parse_type_annot()?; - } - - let param_end = self.peek_span().unwrap_or(param_start..param_start).end; - params.push(Parameter { - name: param_name, - bounds: Vec::new(), - kind: None, - span: Span::new(&(param_start..param_end), self.file.clone()), - }); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - self.expect(Token::Arrow)?; - let return_type = self.parse_type_annot()?; - - Ok(FunctionSignature { - name, - params, - return_type, - }) - } - - fn parse_expr(&mut self) -> Result { - let mut attributes = Vec::new(); - - // Parse any leading attributes - while matches!(self.peek(), Some(Token::At)) { - attributes.push(self.parse_attribute()?); - } - - let mut expr = self.parse_assignment()?; - expr.attributes = attributes; - Ok(expr) - } - - fn parse_range_expr(&mut self) -> Result { - let left = self.parse_or_expr()?; - - if matches!(self.peek(), Some(Token::DotDot)) { - let start = left.span.start; - self.next(); - let right = self.parse_or_expr()?; - let end = right.span.end; - Ok(Expr { - kind: ExprKind::Range(Box::new(left), Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } else { - Ok(left) - } - } - - fn parse_assignment(&mut self) -> Result { - let left = self.parse_range_expr()?; - - if matches!(self.peek(), Some(Token::Assign)) { - let start = left.span.start; - self.next(); - let right = self.parse_assignment()?; - let end = right.span.end; - Ok(Expr { - kind: ExprKind::Assign(Box::new(left), Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } else { - Ok(left) - } - } - - fn parse_or_expr(&mut self) -> Result { - let mut left = self.parse_and_expr()?; - - loop { - if matches!(self.peek(), Some(Token::Or)) { - let start = left.span.start; - self.next(); - let right = self.parse_and_expr()?; - let end = right.span.end; - left = Expr { - kind: ExprKind::BinOp(Box::new(left), BinOp::Or, Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } else { - break; - } - } - - Ok(left) - } - - fn parse_and_expr(&mut self) -> Result { - let mut left = self.parse_eq_expr()?; - - loop { - if matches!(self.peek(), Some(Token::And)) { - let start = left.span.start; - self.next(); - let right = self.parse_eq_expr()?; - let end = right.span.end; - left = Expr { - kind: ExprKind::BinOp(Box::new(left), BinOp::And, Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } else { - break; - } - } - - Ok(left) - } - - fn parse_eq_expr(&mut self) -> Result { - let mut left = self.parse_comp_expr()?; - - loop { - let op = match self.peek() { - Some(Token::Eq) => BinOp::Eq, - Some(Token::NotEq) => BinOp::Neq, - _ => break, - }; - let start = left.span.start; - self.next(); - let right = self.parse_comp_expr()?; - let end = right.span.end; - left = Expr { - kind: ExprKind::BinOp(Box::new(left), op, Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - - Ok(left) - } - - fn parse_comp_expr(&mut self) -> Result { - let mut left = self.parse_add_expr()?; - - loop { - let op = match self.peek() { - Some(Token::Less) => BinOp::Lt, - Some(Token::Greater) => BinOp::Gt, - Some(Token::LessEq) => BinOp::Leq, - Some(Token::GreaterEq) => BinOp::Geq, - _ => break, - }; - let start = left.span.start; - self.next(); - let right = self.parse_add_expr()?; - let end = right.span.end; - left = Expr { - kind: ExprKind::BinOp(Box::new(left), op, Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - - Ok(left) - } - - fn parse_add_expr(&mut self) -> Result { - let mut left = self.parse_mul_expr()?; - - loop { - let op = match self.peek() { - Some(Token::Plus) => BinOp::Add, - Some(Token::Minus) => BinOp::Sub, - _ => break, - }; - let start = left.span.start; - self.next(); - let right = self.parse_mul_expr()?; - let end = right.span.end; - left = Expr { - kind: ExprKind::BinOp(Box::new(left), op, Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - - Ok(left) - } - - fn parse_mul_expr(&mut self) -> Result { - let mut left = self.parse_unary_expr()?; - - loop { - let op = match self.peek() { - Some(Token::Mul) => BinOp::Mul, - Some(Token::Div) => BinOp::Div, - Some(Token::Mod) => BinOp::Mod, - _ => break, - }; - let start = left.span.start; - self.next(); - let right = self.parse_unary_expr()?; - let end = right.span.end; - left = Expr { - kind: ExprKind::BinOp(Box::new(left), op, Box::new(right)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - - Ok(left) - } - - fn parse_unary_expr(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - match self.peek() { - Some(Token::Not) => { - self.next(); - let expr = self.parse_unary_expr()?; - let end = expr.span.end; - Ok(Expr { - kind: ExprKind::UnOp(UnOp::Not, Box::new(expr)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::Minus) => { - self.next(); - let expr = self.parse_unary_expr()?; - let end = expr.span.end; - Ok(Expr { - kind: ExprKind::UnOp(UnOp::Neg, Box::new(expr)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - _ => self.parse_postfix_expr(), - } - } - - fn parse_postfix_expr(&mut self) -> Result { - let mut expr = self.parse_primary_expr()?; - - loop { - match self.peek() { - Some(Token::LParen) => { - // Function call - let start = expr.span.start; - self.next(); - let mut args = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - args.push(self.parse_expr()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - let end = self.peek_span().unwrap_or(expr.span.end..expr.span.end).end; - expr = Expr { - kind: ExprKind::Call(Box::new(expr), args), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - Some(Token::LBracket) => { - // Index - let start = expr.span.start; - self.next(); - let index = self.parse_expr()?; - self.expect(Token::RBracket)?; - let end = self.peek_span().unwrap_or(expr.span.end..expr.span.end).end; - expr = Expr { - kind: ExprKind::Index(Box::new(expr), Box::new(index)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - Some(Token::Dot) => { - // Field access - let start = expr.span.start; - self.next(); - let field = match self.next() { - Some((Token::Variable(f), _)) => f, - Some((_, span)) => { - return self.error("Expected field name".to_string(), span); - } - None => return self.error("Expected field name".to_string(), start..start), - }; - let end = self.peek_span().unwrap_or(expr.span.end..expr.span.end).end; - expr = Expr { - kind: ExprKind::Dot(Box::new(expr), field), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - Some(Token::OptionalChain) => { - // Optional chain - let start = expr.span.start; - self.next(); - let field = match self.next() { - Some((Token::Variable(f), _)) => f, - Some((_, span)) => { - return self.error("Expected field name".to_string(), span); - } - None => return self.error("Expected field name".to_string(), start..start), - }; - let end = self.peek_span().unwrap_or(expr.span.end..expr.span.end).end; - expr = Expr { - kind: ExprKind::OptionalChain(Some(Box::new(expr)), field), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - Some(Token::Unwrap) => { - // Early return / unwrap - let start = expr.span.start; - self.next(); - let end = self.peek_span().unwrap_or(expr.span.end..expr.span.end).end; - expr = Expr { - kind: ExprKind::EarlyReturn(Some(Box::new(expr))), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - Some(Token::KeywordAs) => { - // Cast - let start = expr.span.start; - self.next(); - let type_annot = self.parse_type_annot()?; - let end = self.peek_span().unwrap_or(expr.span.end..expr.span.end).end; - expr = Expr { - kind: ExprKind::Cast(Box::new(expr), type_annot), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }; - } - _ => break, - } - } - - Ok(expr) - } - - fn parse_primary_expr(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - match self.peek().cloned() { - Some(Token::Int(n)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Int(n), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::Float(f)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Float(f), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::Bool(b)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Bool(b), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::String(s)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::String(s), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::Variable(name)) => { - self.next(); - - // Check for struct literal or enum variant - if matches!(self.peek(), Some(Token::LBrace)) { - // Struct literal - self.next(); - let mut fields = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RBrace)) { - self.next(); - break; - } - - let field_name = match self.next() { - Some((Token::Variable(f), _)) => f, - Some((_, span)) => { - return self.error("Expected field name".to_string(), span); - } - None => { - return self.error("Expected field name".to_string(), start..start); - } - }; - - self.expect(Token::Colon)?; - let field_expr = self.parse_expr()?; - fields.push((field_name, field_expr)); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::StructLit(name, fields), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } else if matches!(self.peek(), Some(Token::Access)) { - // Enum variant - self.next(); - let variant = match self.next() { - Some((Token::Variable(v), _)) => v, - Some((_, span)) => { - return self.error("Expected variant name".to_string(), span); - } - None => { - return self.error("Expected variant name".to_string(), start..start); - } - }; - - let mut args = Vec::new(); - if matches!(self.peek(), Some(Token::LParen)) { - self.next(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - args.push(self.parse_expr()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::EnumLit(name, variant, args), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } else { - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Variable(name), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - } - Some(Token::LParen) => { - self.next(); - if matches!(self.peek(), Some(Token::RParen)) { - // Empty tuple - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Tuple(vec![]), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } else { - let first = self.parse_expr()?; - if matches!(self.peek(), Some(Token::Comma)) { - // Tuple - let mut elements = vec![first]; - self.next(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - break; - } - elements.push(self.parse_expr()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - self.expect(Token::RParen)?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Tuple(elements), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } else { - self.expect(Token::RParen)?; - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: first.kind, - span: Span::new(&(start..end), self.file.clone()), - attributes: first.attributes, - }) - } - } - } - Some(Token::LBracket) => { - self.next(); - let mut elements = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RBracket)) { - self.next(); - break; - } - elements.push(self.parse_expr()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Array(elements), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordLet) => { - self.next(); - - // Parse binding kind (mut, uniq, once) - comes AFTER let - let binding_kind = match self.peek() { - Some(Token::KeywordMut) => { - self.next(); - BindingKind::Mutable - } - Some(Token::KeywordUniq) => { - self.next(); - BindingKind::Affine - } - Some(Token::KeywordOnce) => { - self.next(); - BindingKind::Linear - } - _ => BindingKind::Default, - }; - - // Now parse the variable name - let var_name = match self.next() { - Some((Token::Variable(n), _)) => n, - Some((_, span)) => { - return self.error("Expected variable name".to_string(), span); - } - None => return self.error("Expected variable name".to_string(), start..start), - }; - - // Parse optional type annotation - let type_annot = if matches!(self.peek(), Some(Token::Colon)) { - self.next(); - Some(self.parse_type_annot()?) - } else { - None - }; - - self.expect(Token::Assign)?; - let expr = self.parse_expr()?; - let end = expr.span.end; - Ok(Expr { - kind: ExprKind::Let(var_name, binding_kind, type_annot, Box::new(expr)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordIf) => { - self.next(); - let cond = self.parse_expr()?; - let then_expr = self.parse_expr()?; - let else_expr = if matches!(self.peek(), Some(Token::KeywordElse)) { - self.next(); - Some(Box::new(self.parse_expr()?)) - } else { - None - }; - - let end = else_expr - .as_ref() - .map(|e| e.span.end) - .unwrap_or(then_expr.span.end); - - Ok(Expr { - kind: ExprKind::If(Box::new(cond), Box::new(then_expr), else_expr), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordMatch) => { - self.next(); - let expr = self.parse_expr()?; - let mut arms = Vec::new(); - - loop { - if matches!(self.peek(), Some(Token::KeywordEnd)) { - self.next(); - break; - } - - let pattern = self.parse_pattern()?; - self.expect(Token::FatArrow)?; - let body = self.parse_expr()?; - arms.push((pattern, body)); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Match(Box::new(expr), arms), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordWhile) => { - self.next(); - let cond = self.parse_expr()?; - let body = self.parse_expr()?; - let end = body.span.end; - - Ok(Expr { - kind: ExprKind::While(Box::new(cond), Box::new(body)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordFor) => { - self.next(); - let var = match self.next() { - Some((Token::Variable(v), _)) => v, - Some((_, span)) => { - return self.error("Expected variable name".to_string(), span); - } - None => return self.error("Expected variable name".to_string(), start..start), - }; - self.expect(Token::KeywordIn)?; - let iterable = self.parse_expr()?; - let body = self.parse_expr()?; - let end = body.span.end; - - Ok(Expr { - kind: ExprKind::For(var, Box::new(iterable), Box::new(body)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordDo) => { - self.next(); - let mut exprs = Vec::new(); - - loop { - if matches!(self.peek(), Some(Token::KeywordEnd)) { - self.next(); - break; - } - exprs.push(self.parse_expr()?); - - if matches!(self.peek(), Some(Token::Semicolon)) { - self.next(); - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Do(exprs), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordLambda) => { - let start = self.peek_span().unwrap_or(0..0).start; - self.next(); - self.expect(Token::LParen)?; - let mut params = Vec::new(); - - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - - let param_name = match self.next() { - Some((Token::Variable(p), _)) => p, - Some((_, span)) => { - return self.error("Expected parameter name".to_string(), span); - } - None => { - return self.error("Expected parameter name".to_string(), start..start); - } - }; - - // Check for optional type annotation - let param_type = if matches!(self.peek(), Some(Token::Colon)) { - self.next(); // consume ':' - Some(self.parse_type_annot()?) - } else { - None - }; - - params.push((param_name, param_type)); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - let body = self.parse_expr()?; - let end = body.span.end; - - Ok(Expr { - kind: ExprKind::Lambda(params, Box::new(body)), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordReturn) => { - self.next(); - let expr = if self.is_expr_end() { - None - } else { - Some(Box::new(self.parse_expr()?)) - }; - - let end = expr - .as_ref() - .map(|e| e.span.end) - .unwrap_or(self.peek_span().unwrap_or(start..start).end); - - Ok(Expr { - kind: ExprKind::Return(expr), - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordBreak) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Break, - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(Token::KeywordContinue) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Expr { - kind: ExprKind::Continue, - span: Span::new(&(start..end), self.file.clone()), - attributes: Vec::new(), - }) - } - Some(token) => { - let span = self.peek_span().unwrap_or(start..start); - self.error(format!("Unexpected token: {:?}", token), span) - } - None => self.error("Unexpected EOF".to_string(), start..start), - } - } - - fn parse_pattern(&mut self) -> Result { - let start = self.peek_span().unwrap_or(0..0).start; - - match self.peek().cloned() { - Some(Token::Variable(name)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - - // Check for struct or enum pattern - if matches!(self.peek(), Some(Token::LBrace)) { - // Struct pattern - self.next(); - let mut fields = Vec::new(); - - loop { - if matches!(self.peek(), Some(Token::RBrace)) { - self.next(); - break; - } - - let field_name = match self.next() { - Some((Token::Variable(f), _)) => f, - Some((_, span)) => { - return self.error("Expected field name".to_string(), span); - } - None => { - return self.error("Expected field name".to_string(), start..start); - } - }; - - self.expect(Token::Colon)?; - let pattern = self.parse_pattern()?; - fields.push((field_name, pattern)); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Struct(name, fields), - span: Span::new(&(start..end), self.file.clone()), - }) - } else if matches!(self.peek(), Some(Token::Access)) { - // Enum pattern - self.next(); - let variant = match self.next() { - Some((Token::Variable(v), _)) => v, - Some((_, span)) => { - return self.error("Expected variant name".to_string(), span); - } - None => { - return self.error("Expected variant name".to_string(), start..start); - } - }; - - let mut patterns = Vec::new(); - if matches!(self.peek(), Some(Token::LParen)) { - self.next(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - patterns.push(self.parse_pattern()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Enum(name, variant, patterns), - span: Span::new(&(start..end), self.file.clone()), - }) - } else { - Ok(Pattern { - kind: PatternKind::Variable(name), - span: Span::new(&(start..end), self.file.clone()), - }) - } - } - Some(Token::Union) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Wildcard, - span: Span::new(&(start..end), self.file.clone()), - }) - } - Some(Token::LParen) => { - self.next(); - let mut patterns = Vec::new(); - loop { - if matches!(self.peek(), Some(Token::RParen)) { - self.next(); - break; - } - patterns.push(self.parse_pattern()?); - - if matches!(self.peek(), Some(Token::Comma)) { - self.next(); - } - } - - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Tuple(patterns), - span: Span::new(&(start..end), self.file.clone()), - }) - } - Some(Token::String(s)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Literal(s), - span: Span::new(&(start..end), self.file.clone()), - }) - } - Some(Token::Int(n)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Literal(n.to_string()), - span: Span::new(&(start..end), self.file.clone()), - }) - } - Some(Token::Float(f)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Literal(f.to_string()), - span: Span::new(&(start..end), self.file.clone()), - }) - } - Some(Token::Bool(b)) => { - self.next(); - let end = self.peek_span().unwrap_or(start..start).end; - Ok(Pattern { - kind: PatternKind::Literal(b.to_string()), - span: Span::new(&(start..end), self.file.clone()), - }) - } - Some(token) => { - let span = self.peek_span().unwrap_or(start..start); - self.error(format!("Unexpected token in pattern: {:?}", token), span) - } - None => self.error("Unexpected EOF".to_string(), start..start), - } - } - - fn is_expr_end(&mut self) -> bool { - matches!( - self.peek(), - Some(Token::RParen) - | Some(Token::RBracket) - | Some(Token::RBrace) - | Some(Token::Comma) - | Some(Token::Semicolon) - | Some(Token::KeywordEnd) - | Some(Token::FatArrow) - ) - } -} - -``` - -```rust -// src/main.rs -use logos::Logos; -use std::fs; -use suicmez::{lexer::Token, parser::Parser, typechecker::TypeChecker}; - -fn main() { - // Check if a file was provided as argument - let args: Vec = std::env::args().collect(); - if args.len() < 2 { - // Run all test files in the tests directory - run_test_suite(); - return; - } - - let filename = &args[1]; - println!("Type checking file: {}", filename); - - if let Err(e) = run_file(filename) { - eprintln!("Error: {}", e); - } -} - -fn run_test_suite() { - println!("Running test suite...\n"); - - let test_files = vec![ - "tests/basic_types.sui", - "tests/structs.sui", - "tests/enums.sui", - "tests/functions.sui", - "tests/arrays.sui", - "tests/traits.sui", - "tests/control_flow.sui", - ]; - - for file in test_files { - println!("Testing: {}", file); - match run_file(file) { - Ok(_) => println!("✓ Passed\n"), - Err(e) => println!("✗ Failed: {}\n", e), - } - } -} - -fn run_file(filename: &str) -> Result<(), String> { - // Read the source file - let source = fs::read_to_string(filename) - .map_err(|e| format!("Error reading file {}: {}", filename, e))?; - - // First, we need to parse the source code - let mut tokens = Vec::new(); - let mut lexer = Token::lexer(&source); - - loop { - match lexer.next() { - Some(Ok(token)) => { - let span = lexer.span(); - tokens.push((token, span)); - } - Some(Err(_)) => { - return Err("Lexing error".to_string()); - } - None => break, - } - } - - let mut parser = Parser::new(filename.to_string(), tokens); - let ast_nodes = parser - .parse() - .map_err(|e| format!("Parse error: {}", e.message))?; - - println!("Parsed {} AST nodes successfully", ast_nodes.len()); - - // Typecheck the AST - let mut typechecker = TypeChecker::new(); - let typed_nodes = typechecker.typecheck_program(&ast_nodes).map_err(|e| { - format!( - "Type error at {}:{}: {:?}", - e.span.file, e.span.start, e.kind - ) - })?; - - println!( - "Type checking passed! {} nodes typechecked.", - typed_nodes.len() - ); - - Ok(()) -} - -``` - -```rust -// src/typechecker.rs -#[derive(Debug, Clone, PartialEq)] -pub enum Type { - Stub, -} - -``` - -```rust -// src/ast.rs -use crate::typechecker::Type; -use std::ops::Range; - -#[derive(Debug, Clone)] -pub enum TypeAnnot { - Var(String), - Cons(String, Vec), - Function(Vec, Box), - Tuple(Vec), - Array(Box), -} - -#[derive(Debug, Clone)] -pub struct Span { - pub start: usize, - pub end: usize, - pub file: String, -} - -impl Span { - pub fn new(range: &Range, file: String) -> Self { - Span { - start: range.start, - end: range.end, - file, - } - } - - pub fn merge(&self, other: &Span) -> Span { - Span { - start: self.start.min(other.start), - end: self.end.max(other.end), - file: self.file.clone(), - } - } -} - -// @attribute -#[derive(Debug, Clone)] -pub struct Attribute { - pub name: String, - pub args: Vec, - pub span: Span, -} - -#[derive(Debug, Clone)] -pub enum AttributeArg { - Value(String), // some_identifier - KeyValue(String, String), // some_key = some_identifier - Literal(String), // some literal value -} - -#[derive(Debug, Clone)] -pub struct ASTNode { - pub kind: ASTNodeKind, - pub span: Span, - pub attributes: Vec, -} - -#[derive(Debug, Clone)] -pub enum ASTNodeKind { - Function(Function), - Extern(Extern), - Load(Load), - Struct(Struct), - Enum(Enum), - Impl(Impl), - Trait(Trait), - Use(String), -} - -// ? implies OPTIONAL here -// \( implies the presence of (. same for /) - -#[derive(Debug, Clone)] -/// fn name\( (arg: type?,)* \) -> return_type? body -pub struct Function { - pub name: String, - pub parameters: Vec, // type params - pub args: Vec<(String, Option)>, - pub return_type: Option, - pub body: Expr, -} - -/// extern name\( type?,* \) -> return_type from library_alias -#[derive(Debug, Clone)] -pub struct Extern { - pub name: String, - pub args: Vec, - pub return_type: TypeAnnot, - pub from: String, - pub span: Span, -} - -/// load "library" as alias -#[derive(Debug, Clone)] -pub struct Load { - pub library: String, - pub alias: String, - pub span: Span, -} - -/// struct name ? -/// (field_name: field_type,)* -/// end -#[derive(Debug, Clone)] -pub struct Struct { - pub name: String, - pub parameters: Vec, // type parameters - pub fields: Vec, -} - -#[derive(Debug, Clone)] -pub struct Field { - pub name: String, - pub field_type: TypeAnnot, - pub span: Span, -} - -/// enum name ? -/// VariantName\(field_type,\)* -/// end -#[derive(Debug, Clone)] -pub struct Enum { - pub name: String, - pub parameters: Vec, // type parameters - pub variants: Vec, -} - -#[derive(Debug, Clone)] -pub struct Parameter { - pub name: String, - pub bounds: Vec, // trait bounds - pub kind: Option, // for HKTs - pub span: Span, -} - -#[derive(Debug, Clone, PartialEq)] -pub enum Kind { - Star, // * - Arrow(Box, Box), // k1 -> k2 -} - -#[derive(Debug, Clone)] -pub struct Variant { - pub name: String, - pub fields: Vec, - pub span: Span, -} - -/// impl TypeName ? (: TraitName)? -/// functions* -/// end -#[derive(Debug, Clone)] -pub struct Impl { - pub target: String, - pub trait_name: Option, - pub methods: Vec, -} - -/// trait TraitName ? -/// function_signatures* -/// end -#[derive(Debug, Clone)] -pub struct Trait { - pub name: String, - pub methods: Vec, - - pub parameters: Vec, - pub associated_types: Vec, -} - -#[derive(Debug, Clone)] -pub struct AssociatedType { - pub name: String, - pub bounds: Vec, - pub span: Span, -} - -#[derive(Debug, Clone)] -pub struct FunctionSignature { - pub name: String, - pub params: Vec, - pub return_type: TypeAnnot, -} - -#[derive(Debug, Clone)] -pub struct Expr { - pub kind: ExprKind, - pub span: Span, - pub attributes: Vec, -} - -#[derive(Debug, Clone)] -pub enum ExprKind { - Int(i64), - Float(f64), - Bool(bool), - String(String), - Array(Vec), - Tuple(Vec), - - StructLit(String, Vec<(String, Expr)>), // Name { a: expr, b: expr } - EnumLit(String, String, Vec), // Name::Variant(expr, expr) - - Variable(String), - - Call(Box, Vec), - Index(Box, Box), - Dot(Box, String), - EarlyReturn(Option>), // eg: myresultoroption? - OptionalChain(Option>, String), // a?.b - - Lambda(Vec<(String, Option)>, Box), // lambda (arg, arg: optionalty, ...) body - Let(String, BindingKind, Option, Box), // no patterns for now - Assign(Box, Box), // NOTE: check for valid lvalue during typechecking - Cast(Box, TypeAnnot), - - If(Box, Box, Option>), // if cond expr (else expr)? - Match(Box, Vec<(Pattern, Expr)>), // match expr pattern => expr* end - While(Box, Box), // while cond expr - - For(String, Box, Box), // for i in expr body - Range(Box, Box), // 0..10 - - Do(Vec), // do expr* end - BinOp(Box, BinOp, Box), - UnOp(UnOp, Box), - - Return(Option>), - Break, - Continue, -} - -#[derive(Debug, Clone, PartialEq)] -pub enum BindingKind { - Default, // immutable but infinite usages - Mutable, // mutable but infinite usages - Affine, - Linear, -} - -#[derive(Debug, Clone)] -pub enum BinOp { - Add, - Sub, - Mul, - Div, - Mod, - And, - Or, - Eq, - Neq, - Lt, - Gt, - Leq, - Geq, -} - -#[derive(Debug, Clone)] -pub enum UnOp { - Neg, - Not, -} - -#[derive(Debug, Clone)] -pub struct Pattern { - pub kind: PatternKind, - pub span: Span, -} - -#[derive(Debug, Clone)] -pub enum PatternKind { - Wildcard, // _ - Variable(String), - Literal(String), - Tuple(Vec), - Struct(String, Vec<(String, Pattern)>), - Enum(String, String, Vec), - Range(i64, i64), -} - -// Typed variants - -#[derive(Debug, Clone)] -pub struct TypedASTNode { - pub kind: TypedASTNodeKind, - pub span: Span, - pub attributes: Vec, - pub ty: Type, -} - -#[derive(Debug, Clone)] -pub enum TypedASTNodeKind { - Function(TypedFunction), - Extern(TypedExtern), - Load(TypedLoad), - Struct(TypedStruct), - Enum(TypedEnum), - Impl(TypedImpl), - Trait(TypedTrait), - Use(String), -} - -#[derive(Debug, Clone)] -pub struct TypedFunction { - pub name: String, - pub parameters: Vec, - pub args: Vec<(String, Option)>, - pub return_type: Option, - pub body: TypedExpr, - pub ty: Type, -} - -#[derive(Debug, Clone)] -pub struct TypedExtern { - pub name: String, - pub args: Vec, - pub return_type: TypeAnnot, - pub from: String, - pub span: Span, -} - -#[derive(Debug, Clone)] -pub struct TypedLoad { - pub library: String, - pub alias: String, - pub span: Span, -} - -#[derive(Debug, Clone)] -pub struct TypedStruct { - pub name: String, - pub parameters: Vec, - pub fields: Vec, -} - -#[derive(Debug, Clone)] -pub struct TypedField { - pub name: String, - pub field_type: TypeAnnot, - pub span: Span, -} - -#[derive(Debug, Clone)] -pub struct TypedEnum { - pub name: String, - pub parameters: Vec, - pub variants: Vec, -} - -#[derive(Debug, Clone)] -pub struct TypedVariant { - pub name: String, - pub fields: Vec, - pub span: Span, -} - -#[derive(Debug, Clone)] -pub struct TypedImpl { - pub target: String, - pub trait_name: Option, - pub methods: Vec, -} - -#[derive(Debug, Clone)] -pub struct TypedTrait { - pub name: String, - pub methods: Vec, - pub parameters: Vec, - pub associated_types: Vec, -} - -#[derive(Debug, Clone)] -pub struct TypedExpr { - pub kind: TypedExprKind, - pub span: Span, - pub attributes: Vec, - pub ty: Type, -} - -#[derive(Debug, Clone)] -pub enum TypedExprKind { - Int(i64), - Float(f64), - Bool(bool), - String(String), - Array(Vec), - Tuple(Vec), - StructLit(String, Vec<(String, TypedExpr)>), - EnumLit(String, String, Vec), - Variable(String), - Call(Box, Vec), - Index(Box, Box), - Dot(Box, String), - EarlyReturn(Option>), - OptionalChain(Option>, String), - Lambda(Vec<(String, Option)>, Box), - Let(String, BindingKind, Option, Box), - Assign(Box, Box), - Cast(Box, TypeAnnot), - If(Box, Box, Option>), - Match(Box, Vec<(TypedPattern, TypedExpr)>), - While(Box, Box), - Do(Vec), - BinOp(Box, BinOp, Box), - UnOp(UnOp, Box), - For(String, Box, Box), - Range(Box, Box), - Return(Option>), - Break, - Continue, -} - -#[derive(Debug, Clone)] -pub struct TypedPattern { - pub kind: TypedPatternKind, - pub span: Span, - pub ty: Type, -} - -#[derive(Debug, Clone)] -pub enum TypedPatternKind { - Wildcard, - Variable(String), - Literal(String), - Tuple(Vec), - Struct(String, Vec<(String, TypedPattern)>), - Enum(String, String, Vec), -} - -``` - -```rust -// src/lexer/mod.rs -use logos::Logos; - -#[cfg(test)] -pub mod tests; - -#[derive(Logos, Debug, PartialEq)] -#[logos(skip r"[ \n\r\t\f]+")] // Ignore this regex pattern between tokens -#[logos(skip r"#(.*)\n")] // Ignore this regex pattern between tokens -#[derive(Clone)] -pub enum Token { - #[regex(r"true|false", |lex| { - lex.slice().parse::().unwrap() - })] - Bool(bool), - - #[regex(r"0|[1-9][0-9_]*", |lex| { - let s = lex.slice().replace("_", ""); - // We parse to i64 for wider support. - s.parse::().unwrap() - }, priority = 4)] - Int(i64), - - #[regex(r"(([0-9][0-9_]*\.[0-9_]+|[0-9]*\.[0-9_]+)([eE][+-]?[0-9_]+)?)", |lex| { - let s = lex.slice().replace("_", ""); - s.parse::().unwrap() - }, priority = 3)] - Float(f64), - - #[regex(r#""([^"\\]*(\\.[^"\\]*)*)""#, |lex| { - let s = lex.slice(); - s[1..s.len()-1] - .replace("\\\"", "\"") - .replace("\\\\", "\\") - .replace("\\n", "\n") - .replace("\\r", "\r") - .replace("\\t", "\t") - })] - String(String), - - #[regex(r#"r#"([^"]*)""#, |lex| { - let s = lex.slice(); - // Remove the outer r" and " (s[2..s.len() - 1]) - s[3..s.len() - 1].to_string() - })] - RawString(String), - - #[regex(r"[a-zA-Z_][a-zA-Z0-9_]*", |lex|{ - lex.slice().to_string() - })] - Variable(String), - - #[token("bool")] - KeywordBool, - - #[token("int")] - KeywordInt, - - #[token("float")] - KeywordFloat, - - #[token("string")] - KeywordString, - - #[token("let")] - KeywordLet, - - #[token("mut")] - KeywordMut, - - #[token("uniq")] - KeywordUniq, - - #[token("once")] - KeywordOnce, - - #[token("if")] - KeywordIf, - - #[token("then")] - KeywordThen, - - #[token("else")] - KeywordElse, - - #[token("fn")] - KeywordFn, - - #[token("lambda")] - KeywordLambda, - - #[token("do")] - KeywordDo, - - #[token("end")] - KeywordEnd, - - #[token("as")] - KeywordAs, - - #[token("in")] - KeywordIn, - - #[token("for")] - KeywordFor, - #[token("while")] - KeywordWhile, - - #[token("loop")] - KeywordLoop, - - #[token("where")] - KeywordWhere, - - #[token("extern")] - KeywordExtern, - - #[token("load")] - KeywordLoad, - - #[token("from")] - KeywordFrom, - - #[token("use")] - KeywordUse, - - #[token("struct")] - KeywordStruct, - - #[token("enum")] - KeywordEnum, - - #[token("impl")] - KeywordImpl, - - #[token("trait")] - KeywordTrait, - - // #[token("type")] - // KeywordType, - // - #[token("match")] - KeywordMatch, - - #[token("return")] - KeywordReturn, - - #[token("break")] - KeywordBreak, - - #[token("continue")] - KeywordContinue, - - #[token("+")] - Plus, - - #[token("-")] - Minus, - - #[token("*")] - Mul, - - #[token("/")] - Div, - - #[token("%")] - Mod, - - #[token("**", priority = 3)] - Power, - - #[token("$")] - Dollar, - - #[token("@")] - At, - - #[token("==")] - Eq, - - #[token("!=")] - NotEq, - - #[token("<")] - Less, - - #[token(">")] - Greater, - - #[token("<=")] - LessEq, - - #[token(">=")] - GreaterEq, - - #[token("and")] - And, - - #[token("or")] - Or, - - #[token("xor")] - Xor, - - #[token("nor")] - Nor, - - #[token("not")] - Not, - - #[token("(")] - LParen, - - #[token(")")] - RParen, - - #[token("[")] - LBracket, - - #[token("]")] - RBracket, - - #[token("{")] - LBrace, - - #[token("}")] - RBrace, - - #[token(",")] - Comma, - - #[token(";")] - Semicolon, - - #[token(":")] - Colon, - - #[token(".")] - Dot, - - #[token("...")] - Spread, - - #[token("..")] - DotDot, - - #[token("::")] - Access, - - #[token("->")] - Arrow, - - #[token("~")] - Tilde, - - #[token("!")] - Bang, - - // New tokens for pattern matching - #[token("=>")] - FatArrow, // For match arms - - #[token("|")] - Union, - - #[token("?.")] - OptionalChain, - - #[token("?")] - Unwrap, - - #[token("=")] - Assign, - - #[token("+=")] - AddAssign, - - #[token("-=")] - SubAssign, - - #[token("*=")] - MulAssign, - - #[token("/=")] - DivAssign, - - #[token("%=")] - ModAssign, -} - -``` - -```rust -// src/lexer/tests.rs -use super::Token; -use logos::Logos; - -#[test] -fn test_literals() { - let mut lexer = Token::lexer("true false 42 2.14 \"hello\" r\"raw\""); - - assert_eq!(lexer.next(), Some(Ok(Token::Bool(true)))); - assert_eq!(lexer.next(), Some(Ok(Token::Bool(false)))); - assert_eq!(lexer.next(), Some(Ok(Token::Int(42)))); - assert_eq!(lexer.next(), Some(Ok(Token::Float(2.14)))); - assert_eq!(lexer.next(), Some(Ok(Token::String("hello".to_string())))); - // RawString regex seems to have issues, let's test separately - assert_eq!(lexer.next(), Some(Ok(Token::Variable("r".to_string())))); - assert_eq!(lexer.next(), Some(Ok(Token::String("raw".to_string())))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_int_literals() { - let mut lexer = Token::lexer("0 123 1_000_000"); - - assert_eq!(lexer.next(), Some(Ok(Token::Int(0)))); - assert_eq!(lexer.next(), Some(Ok(Token::Int(123)))); - assert_eq!(lexer.next(), Some(Ok(Token::Int(1000000)))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_string_literals() { - let mut lexer = Token::lexer("\"hello world\" \"with\\\\escape\" \"quote\\\"here\""); - - assert_eq!( - lexer.next(), - Some(Ok(Token::String("hello world".to_string()))) - ); - assert_eq!( - lexer.next(), - Some(Ok(Token::String("with\\escape".to_string()))) - ); - assert_eq!( - lexer.next(), - Some(Ok(Token::String("quote\"here".to_string()))) - ); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_keywords() { - let mut lexer = Token::lexer( - "bool int float string let if else fn do end as in for while loop where extern import struct enum impl trait match return break continue", - ); - - assert_eq!(lexer.next(), Some(Ok(Token::KeywordBool))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordInt))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordFloat))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordString))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordLet))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordIf))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordElse))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordFn))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordDo))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordEnd))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordAs))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordIn))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordFor))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordWhile))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordLoop))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordWhere))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordExtern))); - assert_eq!(lexer.next(), Some(Ok(Token::Variable("import".into())))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordStruct))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordEnum))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordImpl))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordTrait))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordMatch))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordReturn))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordBreak))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordContinue))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_operators() { - let mut lexer = Token::lexer("+ - * / % ** $ @ == != < > <= >= and or xor nor not"); - - assert_eq!(lexer.next(), Some(Ok(Token::Plus))); - assert_eq!(lexer.next(), Some(Ok(Token::Minus))); - assert_eq!(lexer.next(), Some(Ok(Token::Mul))); - assert_eq!(lexer.next(), Some(Ok(Token::Div))); - assert_eq!(lexer.next(), Some(Ok(Token::Mod))); - assert_eq!(lexer.next(), Some(Ok(Token::Power))); - assert_eq!(lexer.next(), Some(Ok(Token::Dollar))); - assert_eq!(lexer.next(), Some(Ok(Token::At))); - assert_eq!(lexer.next(), Some(Ok(Token::Eq))); - assert_eq!(lexer.next(), Some(Ok(Token::NotEq))); - assert_eq!(lexer.next(), Some(Ok(Token::Less))); - assert_eq!(lexer.next(), Some(Ok(Token::Greater))); - assert_eq!(lexer.next(), Some(Ok(Token::LessEq))); - assert_eq!(lexer.next(), Some(Ok(Token::GreaterEq))); - assert_eq!(lexer.next(), Some(Ok(Token::And))); - assert_eq!(lexer.next(), Some(Ok(Token::Or))); - assert_eq!(lexer.next(), Some(Ok(Token::Xor))); - assert_eq!(lexer.next(), Some(Ok(Token::Nor))); - assert_eq!(lexer.next(), Some(Ok(Token::Not))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_assignment_operators() { - let mut lexer = Token::lexer("= += -= *= /= %="); - - assert_eq!(lexer.next(), Some(Ok(Token::Assign))); - assert_eq!(lexer.next(), Some(Ok(Token::AddAssign))); - assert_eq!(lexer.next(), Some(Ok(Token::SubAssign))); - assert_eq!(lexer.next(), Some(Ok(Token::MulAssign))); - assert_eq!(lexer.next(), Some(Ok(Token::DivAssign))); - assert_eq!(lexer.next(), Some(Ok(Token::ModAssign))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_punctuation() { - let mut lexer = Token::lexer("( ) [ ] { } , ; : . ... .. :: -> ~ ! => | |> ?. ?"); - - assert_eq!(lexer.next(), Some(Ok(Token::LParen))); - assert_eq!(lexer.next(), Some(Ok(Token::RParen))); - assert_eq!(lexer.next(), Some(Ok(Token::LBracket))); - assert_eq!(lexer.next(), Some(Ok(Token::RBracket))); - assert_eq!(lexer.next(), Some(Ok(Token::LBrace))); - assert_eq!(lexer.next(), Some(Ok(Token::RBrace))); - assert_eq!(lexer.next(), Some(Ok(Token::Comma))); - assert_eq!(lexer.next(), Some(Ok(Token::Semicolon))); - assert_eq!(lexer.next(), Some(Ok(Token::Colon))); - assert_eq!(lexer.next(), Some(Ok(Token::Dot))); - assert_eq!(lexer.next(), Some(Ok(Token::Spread))); - assert_eq!(lexer.next(), Some(Ok(Token::DotDot))); - assert_eq!(lexer.next(), Some(Ok(Token::Access))); - assert_eq!(lexer.next(), Some(Ok(Token::Arrow))); - assert_eq!(lexer.next(), Some(Ok(Token::Tilde))); - assert_eq!(lexer.next(), Some(Ok(Token::Bang))); - assert_eq!(lexer.next(), Some(Ok(Token::FatArrow))); - assert_eq!(lexer.next(), Some(Ok(Token::Union))); - assert_eq!(lexer.next(), Some(Ok(Token::OptionalChain))); - assert_eq!(lexer.next(), Some(Ok(Token::Unwrap))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_variables() { - let mut lexer = Token::lexer("x y_z _private camelCase PascalCase"); - - assert_eq!(lexer.next(), Some(Ok(Token::Variable("x".to_string())))); - assert_eq!(lexer.next(), Some(Ok(Token::Variable("y_z".to_string())))); - assert_eq!( - lexer.next(), - Some(Ok(Token::Variable("_private".to_string()))) - ); - assert_eq!( - lexer.next(), - Some(Ok(Token::Variable("camelCase".to_string()))) - ); - assert_eq!( - lexer.next(), - Some(Ok(Token::Variable("PascalCase".to_string()))) - ); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_whitespace_skipping() { - let mut lexer = Token::lexer(" \t\n\r true \n false "); - - assert_eq!(lexer.next(), Some(Ok(Token::Bool(true)))); - assert_eq!(lexer.next(), Some(Ok(Token::Bool(false)))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_comment_skipping() { - let mut lexer = Token::lexer("true # this is a comment\n false"); - - assert_eq!(lexer.next(), Some(Ok(Token::Bool(true)))); - assert_eq!(lexer.next(), Some(Ok(Token::Bool(false)))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_complex_sequence() { - let mut lexer = Token::lexer("fn add(x: int, y: int) -> int { x + y }"); - - assert_eq!(lexer.next(), Some(Ok(Token::KeywordFn))); - assert_eq!(lexer.next(), Some(Ok(Token::Variable("add".to_string())))); - assert_eq!(lexer.next(), Some(Ok(Token::LParen))); - assert_eq!(lexer.next(), Some(Ok(Token::Variable("x".to_string())))); - assert_eq!(lexer.next(), Some(Ok(Token::Colon))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordInt))); - assert_eq!(lexer.next(), Some(Ok(Token::Comma))); - assert_eq!(lexer.next(), Some(Ok(Token::Variable("y".to_string())))); - assert_eq!(lexer.next(), Some(Ok(Token::Colon))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordInt))); - assert_eq!(lexer.next(), Some(Ok(Token::RParen))); - assert_eq!(lexer.next(), Some(Ok(Token::Arrow))); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordInt))); - assert_eq!(lexer.next(), Some(Ok(Token::LBrace))); - assert_eq!(lexer.next(), Some(Ok(Token::Variable("x".to_string())))); - assert_eq!(lexer.next(), Some(Ok(Token::Plus))); - assert_eq!(lexer.next(), Some(Ok(Token::Variable("y".to_string())))); - assert_eq!(lexer.next(), Some(Ok(Token::RBrace))); - assert_eq!(lexer.next(), None); -} - -#[test] -fn test_edge_cases() { - // Test that keywords are not treated as variables - let mut lexer = Token::lexer("let let_var if if_var"); - - assert_eq!(lexer.next(), Some(Ok(Token::KeywordLet))); - assert_eq!( - lexer.next(), - Some(Ok(Token::Variable("let_var".to_string()))) - ); - assert_eq!(lexer.next(), Some(Ok(Token::KeywordIf))); - assert_eq!( - lexer.next(), - Some(Ok(Token::Variable("if_var".to_string()))) - ); - assert_eq!(lexer.next(), None); -} - -``` - diff --git a/src/typechecker.rs b/src/typechecker.rs index b79577c..4c82a19 100644 --- a/src/typechecker.rs +++ b/src/typechecker.rs @@ -1,4 +1,1516 @@ +// src/typechecker.rs +use crate::ast::*; +use std::collections::HashMap; + #[derive(Debug, Clone, PartialEq)] pub enum Type { - Stub, + Int, + Float, + Bool, + String, + Unit, + Never, + Array(Box), + Tuple(Vec), + Function(Vec, Box), + Struct(String, Vec), // name and type arguments + Enum(String, Vec), + TypeVar(String), + Generic(String, Vec), // Generic type constructor + Unknown, // For type inference +} + +impl Type { + pub fn to_string(&self) -> String { + match self { + Type::Int => "int".to_string(), + Type::Float => "float".to_string(), + Type::Bool => "bool".to_string(), + Type::String => "string".to_string(), + Type::Unit => "()".to_string(), + Type::Never => "!".to_string(), + Type::Array(inner) => format!("[{}]", inner.to_string()), + Type::Tuple(types) => { + let type_strs: Vec = types.iter().map(|t| t.to_string()).collect(); + format!("({})", type_strs.join(", ")) + } + Type::Function(args, ret) => { + let arg_strs: Vec = args.iter().map(|t| t.to_string()).collect(); + format!("fn({}) -> {}", arg_strs.join(", "), ret.to_string()) + } + Type::Struct(name, args) if args.is_empty() => name.clone(), + Type::Struct(name, args) => { + let arg_strs: Vec = args.iter().map(|t| t.to_string()).collect(); + format!("{}<{}>", name, arg_strs.join(", ")) + } + Type::Enum(name, args) if args.is_empty() => name.clone(), + Type::Enum(name, args) => { + let arg_strs: Vec = args.iter().map(|t| t.to_string()).collect(); + format!("{}<{}>", name, arg_strs.join(", ")) + } + Type::TypeVar(name) => name.clone(), + Type::Generic(name, args) => { + let arg_strs: Vec = args.iter().map(|t| t.to_string()).collect(); + format!("{}<{}>", name, arg_strs.join(", ")) + } + Type::Unknown => "?".to_string(), + } + } +} + +#[derive(Debug)] +pub struct TypeError { + pub kind: TypeErrorKind, + pub span: Span, +} + +#[derive(Debug)] +pub enum TypeErrorKind { + TypeMismatch(Type, Type), + UndefinedVariable(String), + UndefinedType(String), + UndefinedFunction(String), + UndefinedField(String, Type), + UndefinedVariant(String, String), + ArityMismatch(usize, usize), + NotAFunction(Type), + NotAnArray(Type), + NotAStruct(Type), + NotAnEnum(Type), + InvalidCast(Type, Type), + InvalidPattern(String), + MutableityError(String), + LinearityError(String), + Other(String), +} + +#[derive(Clone)] +struct TypeEnv { + vars: HashMap, + types: HashMap, + functions: HashMap, + traits: HashMap, + impls: Vec, + type_vars: HashMap, +} + +#[derive(Clone, Debug)] +struct TypeInfo { + kind: TypeInfoKind, + parameters: Vec, +} + +#[derive(Clone, Debug)] +enum TypeInfoKind { + Struct(Vec<(String, TypeAnnot)>), + Enum(Vec<(String, Vec)>), +} + +#[derive(Clone, Debug)] +struct FunctionType { + type_params: Vec, + params: Vec, + return_type: Type, +} + +#[derive(Clone, Debug)] +struct TraitInfo { + methods: HashMap, + parameters: Vec, +} + +#[derive(Clone, Debug)] +struct ImplInfo { + target: String, + trait_name: Option, + methods: HashMap, +} + +impl TypeEnv { + fn new() -> Self { + TypeEnv { + vars: HashMap::new(), + types: HashMap::new(), + functions: HashMap::new(), + traits: HashMap::new(), + impls: Vec::new(), + type_vars: HashMap::new(), + } + } + + fn enter_scope(&self) -> Self { + TypeEnv { + vars: self.vars.clone(), + types: self.types.clone(), + functions: self.functions.clone(), + traits: self.traits.clone(), + impls: self.impls.clone(), + type_vars: self.type_vars.clone(), + } + } + + fn add_var(&mut self, name: String, ty: Type, kind: BindingKind) { + self.vars.insert(name, (ty, kind)); + } + + fn get_var(&self, name: &str) -> Option<&(Type, BindingKind)> { + self.vars.get(name) + } + + fn add_type(&mut self, name: String, info: TypeInfo) { + self.types.insert(name, info); + } + + fn get_type(&self, name: &str) -> Option<&TypeInfo> { + self.types.get(name) + } + + fn add_function(&mut self, name: String, ty: FunctionType) { + self.functions.insert(name, ty); + } + + fn get_function(&self, name: &str) -> Option<&FunctionType> { + self.functions.get(name) + } +} + +pub struct TypeChecker { + env: TypeEnv, +} + +impl TypeChecker { + pub fn new() -> Self { + TypeChecker { + env: TypeEnv::new(), + } + } + + pub fn typecheck_program(&mut self, nodes: &[ASTNode]) -> Result, TypeError> { + // First pass: collect all type definitions, function signatures, etc. + for node in nodes { + self.collect_definitions(node)?; + } + + // Second pass: typecheck everything + let mut typed_nodes = Vec::new(); + for node in nodes { + typed_nodes.push(self.typecheck_node(node)?); + } + + Ok(typed_nodes) + } + + fn collect_definitions(&mut self, node: &ASTNode) -> Result<(), TypeError> { + match &node.kind { + ASTNodeKind::Struct(s) => { + let info = TypeInfo { + kind: TypeInfoKind::Struct( + s.fields + .iter() + .map(|f| (f.name.clone(), f.field_type.clone())) + .collect(), + ), + parameters: s.parameters.iter().map(|p| p.name.clone()).collect(), + }; + self.env.add_type(s.name.clone(), info); + } + ASTNodeKind::Enum(e) => { + let info = TypeInfo { + kind: TypeInfoKind::Enum( + e.variants + .iter() + .map(|v| (v.name.clone(), v.fields.clone())) + .collect(), + ), + parameters: e.parameters.iter().map(|p| p.name.clone()).collect(), + }; + self.env.add_type(e.name.clone(), info); + } + ASTNodeKind::Function(f) => { + let param_types: Vec = f + .args + .iter() + .map(|(_, ty)| { + ty.as_ref() + .map(|t| self.type_annot_to_type(t)) + .unwrap_or(Type::Unknown) + }) + .collect(); + let return_type = f + .return_type + .as_ref() + .map(|t| self.type_annot_to_type(t)) + .unwrap_or(Type::Unit); + + let func_type = FunctionType { + type_params: f.parameters.iter().map(|p| p.name.clone()).collect(), + params: param_types, + return_type, + }; + self.env.add_function(f.name.clone(), func_type); + } + ASTNodeKind::Trait(t) => { + let mut methods = HashMap::new(); + for sig in &t.methods { + let param_types: Vec = sig.params.iter().map(|_| Type::Unknown).collect(); + let return_type = self.type_annot_to_type(&sig.return_type); + methods.insert( + sig.name.clone(), + FunctionType { + type_params: Vec::new(), + params: param_types, + return_type, + }, + ); + } + let trait_info = TraitInfo { + methods, + parameters: t.parameters.iter().map(|p| p.name.clone()).collect(), + }; + self.env.traits.insert(t.name.clone(), trait_info); + } + ASTNodeKind::Impl(impl_def) => { + let mut methods = HashMap::new(); + for method in &impl_def.methods { + let param_types: Vec = method + .args + .iter() + .map(|(_, ty)| { + ty.as_ref() + .map(|t| self.type_annot_to_type(t)) + .unwrap_or(Type::Unknown) + }) + .collect(); + let return_type = method + .return_type + .as_ref() + .map(|t| self.type_annot_to_type(t)) + .unwrap_or(Type::Unit); + + methods.insert( + method.name.clone(), + FunctionType { + type_params: method.parameters.iter().map(|p| p.name.clone()).collect(), + params: param_types, + return_type, + }, + ); + } + self.env.impls.push(ImplInfo { + target: impl_def.target.clone(), + trait_name: impl_def.trait_name.clone(), + methods, + }); + } + _ => {} + } + Ok(()) + } + + fn typecheck_node(&mut self, node: &ASTNode) -> Result { + let ty = match &node.kind { + ASTNodeKind::Function(f) => { + let typed_func = self.typecheck_function(f)?; + let ty = typed_func.ty.clone(); + return Ok(TypedASTNode { + kind: TypedASTNodeKind::Function(typed_func), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty, + }); + } + ASTNodeKind::Struct(_) => Type::Unit, + ASTNodeKind::Enum(_) => Type::Unit, + ASTNodeKind::Trait(_) => Type::Unit, + ASTNodeKind::Impl(impl_def) => { + let mut typed_methods = Vec::new(); + for method in &impl_def.methods { + typed_methods.push(self.typecheck_function(method)?); + } + return Ok(TypedASTNode { + kind: TypedASTNodeKind::Impl(TypedImpl { + target: impl_def.target.clone(), + trait_name: impl_def.trait_name.clone(), + methods: typed_methods, + }), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty: Type::Unit, + }); + } + ASTNodeKind::Extern(ext) => { + return Ok(TypedASTNode { + kind: TypedASTNodeKind::Extern(TypedExtern { + name: ext.name.clone(), + args: ext.args.clone(), + return_type: ext.return_type.clone(), + from: ext.from.clone(), + span: ext.span.clone(), + }), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty: Type::Unit, + }); + } + ASTNodeKind::Load(load) => { + return Ok(TypedASTNode { + kind: TypedASTNodeKind::Load(TypedLoad { + library: load.library.clone(), + alias: load.alias.clone(), + span: load.span.clone(), + }), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty: Type::Unit, + }); + } + ASTNodeKind::Use(path) => Type::Unit, + }; + + Ok(TypedASTNode { + kind: TypedASTNodeKind::Use(match &node.kind { + ASTNodeKind::Use(p) => p.clone(), + _ => String::new(), + }), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty, + }) + } + + fn typecheck_function(&mut self, func: &Function) -> Result { + // Save the original environment + let original_env = self.env.clone(); + + // Enter new scope for function + self.env = self.env.enter_scope(); + + // Add type parameters to environment + for param in &func.parameters { + self.env + .type_vars + .insert(param.name.clone(), Type::TypeVar(param.name.clone())); + } + + // Add function parameters to environment + let mut param_types = Vec::new(); + for (arg_name, arg_type_annot) in &func.args { + let arg_type = arg_type_annot + .as_ref() + .map(|t| self.type_annot_to_type(t)) + .unwrap_or(Type::Unknown); + param_types.push(arg_type.clone()); + self.env + .add_var(arg_name.clone(), arg_type, BindingKind::Default); + } + + // Typecheck function body + let typed_body = self.typecheck_expr(&func.body)?; + + // Restore the original environment + self.env = original_env; + + // Check return type + let expected_return = func + .return_type + .as_ref() + .map(|t| self.type_annot_to_type(t)) + .unwrap_or(Type::Unit); + + if !self.types_compatible(&typed_body.ty, &expected_return) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(expected_return, typed_body.ty), + span: func.body.span.clone(), + }); + } + + let func_type = Type::Function(param_types, Box::new(expected_return)); + + Ok(TypedFunction { + name: func.name.clone(), + parameters: func.parameters.clone(), + args: func.args.clone(), + return_type: func.return_type.clone(), + body: typed_body, + ty: func_type, + }) + } + + fn typecheck_expr(&mut self, expr: &Expr) -> Result { + let (kind, ty) = match &expr.kind { + ExprKind::Int(n) => (TypedExprKind::Int(*n), Type::Int), + ExprKind::Float(f) => (TypedExprKind::Float(*f), Type::Float), + ExprKind::Bool(b) => (TypedExprKind::Bool(*b), Type::Bool), + ExprKind::String(s) => (TypedExprKind::String(s.clone()), Type::String), + + ExprKind::Variable(name) => { + // First check if it's a variable + if let Some((t, _)) = self.env.get_var(name) { + (TypedExprKind::Variable(name.clone()), t.clone()) + } else if let Some(func_type) = self.env.get_function(name) { + // If not a variable, check if it's a function + let func_type_clone = func_type.clone(); + let fn_type = Type::Function( + func_type_clone.params, + Box::new(func_type_clone.return_type), + ); + (TypedExprKind::Variable(name.clone()), fn_type) + } else { + return Err(TypeError { + kind: TypeErrorKind::UndefinedVariable(name.clone()), + span: expr.span.clone(), + }); + } + } + + ExprKind::Array(elements) => { + let mut typed_elements = Vec::new(); + let mut element_type = Type::Unknown; + + for (i, elem) in elements.iter().enumerate() { + let typed_elem = self.typecheck_expr(elem)?; + if i == 0 { + element_type = typed_elem.ty.clone(); + } else if !self.types_compatible(&typed_elem.ty, &element_type) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(element_type, typed_elem.ty), + span: elem.span.clone(), + }); + } + typed_elements.push(typed_elem); + } + + if elements.is_empty() { + element_type = Type::Unknown; + } + + ( + TypedExprKind::Array(typed_elements), + Type::Array(Box::new(element_type)), + ) + } + + ExprKind::Tuple(elements) => { + let mut typed_elements = Vec::new(); + let mut types = Vec::new(); + + for elem in elements { + let typed_elem = self.typecheck_expr(elem)?; + types.push(typed_elem.ty.clone()); + typed_elements.push(typed_elem); + } + + (TypedExprKind::Tuple(typed_elements), Type::Tuple(types)) + } + + ExprKind::BinOp(left, op, right) => { + let typed_left = self.typecheck_expr(left)?; + let typed_right = self.typecheck_expr(right)?; + + let result_type = match op { + BinOp::Add | BinOp::Sub | BinOp::Mul | BinOp::Div | BinOp::Mod => { + if !self.types_compatible(&typed_left.ty, &typed_right.ty) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + typed_left.ty.clone(), + typed_right.ty.clone(), + ), + span: right.span.clone(), + }); + } + typed_left.ty.clone() + } + BinOp::Eq | BinOp::Neq | BinOp::Lt | BinOp::Gt | BinOp::Leq | BinOp::Geq => { + Type::Bool + } + BinOp::And | BinOp::Or => { + if !matches!(typed_left.ty, Type::Bool) + || !matches!(typed_right.ty, Type::Bool) + { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + Type::Bool, + typed_right.ty.clone(), + ), + span: expr.span.clone(), + }); + } + Type::Bool + } + }; + + ( + TypedExprKind::BinOp(Box::new(typed_left), op.clone(), Box::new(typed_right)), + result_type, + ) + } + + ExprKind::UnOp(op, inner) => { + let typed_inner = self.typecheck_expr(inner)?; + let result_type = match op { + UnOp::Neg => { + if !matches!(typed_inner.ty, Type::Int | Type::Float) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + Type::Int, + typed_inner.ty.clone(), + ), + span: inner.span.clone(), + }); + } + typed_inner.ty.clone() + } + UnOp::Not => { + if !matches!(typed_inner.ty, Type::Bool) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + Type::Bool, + typed_inner.ty.clone(), + ), + span: inner.span.clone(), + }); + } + Type::Bool + } + }; + ( + TypedExprKind::UnOp(op.clone(), Box::new(typed_inner)), + result_type, + ) + } + + ExprKind::Call(func_expr, args) => { + let typed_func = self.typecheck_expr(func_expr)?; + let mut typed_args = Vec::new(); + + for arg in args { + typed_args.push(self.typecheck_expr(arg)?); + } + + let return_type = match &typed_func.ty { + Type::Function(param_types, ret) => { + if param_types.len() != typed_args.len() { + return Err(TypeError { + kind: TypeErrorKind::ArityMismatch( + param_types.len(), + typed_args.len(), + ), + span: expr.span.clone(), + }); + } + + for (i, (expected, actual)) in + param_types.iter().zip(typed_args.iter()).enumerate() + { + if !self.types_compatible(&actual.ty, expected) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + expected.clone(), + actual.ty.clone(), + ), + span: args[i].span.clone(), + }); + } + } + + (**ret).clone() + } + ty => { + return Err(TypeError { + kind: TypeErrorKind::NotAFunction(ty.clone()), + span: func_expr.span.clone(), + }); + } + }; + + ( + TypedExprKind::Call(Box::new(typed_func), typed_args), + return_type, + ) + } + + ExprKind::Let(name, binding_kind, type_annot, value) => { + let typed_value = self.typecheck_expr(value)?; + let var_type = if let Some(annot) = type_annot { + let annotated_type = self.type_annot_to_type(annot); + if !self.types_compatible(&typed_value.ty, &annotated_type) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(annotated_type, typed_value.ty), + span: value.span.clone(), + }); + } + annotated_type + } else { + typed_value.ty.clone() + }; + + self.env + .add_var(name.clone(), var_type.clone(), binding_kind.clone()); + + ( + TypedExprKind::Let( + name.clone(), + binding_kind.clone(), + type_annot.clone(), + Box::new(typed_value), + ), + var_type, + ) + } + + ExprKind::If(cond, then_expr, else_expr) => { + let typed_cond = self.typecheck_expr(cond)?; + if !matches!(typed_cond.ty, Type::Bool) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(Type::Bool, typed_cond.ty), + span: cond.span.clone(), + }); + } + + let typed_then = self.typecheck_expr(then_expr)?; + let result_type = if let Some(else_expr) = else_expr { + let typed_else = self.typecheck_expr(else_expr)?; + if !self.types_compatible(&typed_then.ty, &typed_else.ty) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(typed_then.ty.clone(), typed_else.ty), + span: else_expr.span.clone(), + }); + } + ( + TypedExprKind::If( + Box::new(typed_cond), + Box::new(typed_then.clone()), + Some(Box::new(typed_else)), + ), + typed_then.ty, + ) + } else { + ( + TypedExprKind::If(Box::new(typed_cond), Box::new(typed_then), None), + Type::Unit, + ) + }; + + result_type + } + + ExprKind::Do(exprs) => { + let mut typed_exprs = Vec::new(); + let mut last_type = Type::Unit; + + for e in exprs { + let typed_e = self.typecheck_expr(e)?; + last_type = typed_e.ty.clone(); + typed_exprs.push(typed_e); + } + + (TypedExprKind::Do(typed_exprs), last_type) + } + + ExprKind::While(cond, body) => { + let typed_cond = self.typecheck_expr(cond)?; + if !matches!(typed_cond.ty, Type::Bool) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(Type::Bool, typed_cond.ty), + span: cond.span.clone(), + }); + } + + let typed_body = self.typecheck_expr(body)?; + ( + TypedExprKind::While(Box::new(typed_cond), Box::new(typed_body)), + Type::Unit, + ) + } + + ExprKind::For(var, iterable, body) => { + let typed_iterable = self.typecheck_expr(iterable)?; + + let element_type = match &typed_iterable.ty { + Type::Array(elem_ty) => (**elem_ty).clone(), + _ => Type::Unknown, + }; + + self.env = self.env.enter_scope(); + self.env + .add_var(var.clone(), element_type, BindingKind::Default); + + let typed_body = self.typecheck_expr(body)?; + ( + TypedExprKind::For(var.clone(), Box::new(typed_iterable), Box::new(typed_body)), + Type::Unit, + ) + } + + ExprKind::Range(start, end) => { + let typed_start = self.typecheck_expr(start)?; + let typed_end = self.typecheck_expr(end)?; + + if !matches!(typed_start.ty, Type::Int) || !matches!(typed_end.ty, Type::Int) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(Type::Int, typed_end.ty.clone()), + span: expr.span.clone(), + }); + } + + ( + TypedExprKind::Range(Box::new(typed_start), Box::new(typed_end)), + Type::Array(Box::new(Type::Int)), + ) + } + + ExprKind::Index(array, index) => { + let typed_array = self.typecheck_expr(array)?; + let typed_index = self.typecheck_expr(index)?; + + if !matches!(typed_index.ty, Type::Int) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(Type::Int, typed_index.ty), + span: index.span.clone(), + }); + } + + let element_type = match &typed_array.ty { + Type::Array(elem_ty) => (**elem_ty).clone(), + ty => { + return Err(TypeError { + kind: TypeErrorKind::NotAnArray(ty.clone()), + span: array.span.clone(), + }); + } + }; + + ( + TypedExprKind::Index(Box::new(typed_array), Box::new(typed_index)), + element_type, + ) + } + + ExprKind::Dot(obj, field) => { + let typed_obj = self.typecheck_expr(obj)?; + + let field_type = match &typed_obj.ty { + Type::Struct(name, _) => { + if let Some(type_info) = self.env.get_type(name) { + if let TypeInfoKind::Struct(fields) = &type_info.kind { + fields + .iter() + .find(|(f, _)| f == field) + .map(|(_, ty)| self.type_annot_to_type(ty)) + .ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedField( + field.clone(), + typed_obj.ty.clone(), + ), + span: expr.span.clone(), + })? + } else { + return Err(TypeError { + kind: TypeErrorKind::NotAStruct(typed_obj.ty.clone()), + span: obj.span.clone(), + }); + } + } else { + return Err(TypeError { + kind: TypeErrorKind::UndefinedType(name.clone()), + span: obj.span.clone(), + }); + } + } + ty => { + return Err(TypeError { + kind: TypeErrorKind::NotAStruct(ty.clone()), + span: obj.span.clone(), + }); + } + }; + + ( + TypedExprKind::Dot(Box::new(typed_obj), field.clone()), + field_type, + ) + } + + ExprKind::StructLit(name, fields) => { + // Clone the struct info we need before borrowing self mutably + let (struct_info_clone, type_params) = { + let struct_type = self.env.get_type(name).ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedType(name.clone()), + span: expr.span.clone(), + })?; + (struct_type.clone(), struct_type.parameters.clone()) + }; + + let mut typed_fields = Vec::new(); + let mut type_arg_map: HashMap = HashMap::new(); + + if let TypeInfoKind::Struct(expected_fields) = &struct_info_clone.kind { + for (field_name, field_expr) in fields { + let typed_field_expr = self.typecheck_expr(field_expr)?; + + let expected_type_annot = expected_fields + .iter() + .find(|(n, _)| n == field_name) + .map(|(_, ty)| ty.clone()) + .ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedField( + field_name.clone(), + Type::Struct(name.clone(), vec![]), + ), + span: field_expr.span.clone(), + })?; + + // Infer generic type parameters + self.infer_type_args( + &expected_type_annot, + &typed_field_expr.ty, + &type_params, + &mut type_arg_map, + ); + + typed_fields.push((field_name.clone(), typed_field_expr)); + } + } else { + return Err(TypeError { + kind: TypeErrorKind::NotAStruct(Type::Struct(name.clone(), vec![])), + span: expr.span.clone(), + }); + } + + // Build concrete type arguments + let concrete_type_args: Vec = type_params + .iter() + .map(|param| type_arg_map.get(param).cloned().unwrap_or(Type::Unknown)) + .collect(); + + ( + TypedExprKind::StructLit(name.clone(), typed_fields), + Type::Struct(name.clone(), concrete_type_args), + ) + } + + ExprKind::EnumLit(enum_name, variant_name, args) => { + // Clone the enum info we need before borrowing self mutably + let (variant_fields, variant_name_clone, type_params) = { + let enum_type = self.env.get_type(enum_name).ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedType(enum_name.clone()), + span: expr.span.clone(), + })?; + + let type_params = enum_type.parameters.clone(); + + if let TypeInfoKind::Enum(variants) = &enum_type.kind { + let variant = variants + .iter() + .find(|(n, _)| n == variant_name) + .ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedVariant( + enum_name.clone(), + variant_name.clone(), + ), + span: expr.span.clone(), + })?; + + if variant.1.len() != args.len() { + return Err(TypeError { + kind: TypeErrorKind::ArityMismatch(variant.1.len(), args.len()), + span: expr.span.clone(), + }); + } + + (variant.1.clone(), variant_name.clone(), type_params) + } else { + return Err(TypeError { + kind: TypeErrorKind::NotAnEnum(Type::Enum(enum_name.clone(), vec![])), + span: expr.span.clone(), + }); + } + }; + + let mut typed_args = Vec::new(); + let mut type_arg_map: HashMap = HashMap::new(); + + for (i, arg) in args.iter().enumerate() { + let typed_arg = self.typecheck_expr(arg)?; + + // Infer generic type parameters + self.infer_type_args( + &variant_fields[i], + &typed_arg.ty, + &type_params, + &mut type_arg_map, + ); + + typed_args.push(typed_arg); + } + + // Build concrete type arguments + let concrete_type_args: Vec = type_params + .iter() + .map(|param| type_arg_map.get(param).cloned().unwrap_or(Type::Unknown)) + .collect(); + + ( + TypedExprKind::EnumLit(enum_name.clone(), variant_name_clone, typed_args), + Type::Enum(enum_name.clone(), concrete_type_args), + ) + } + + ExprKind::Match(scrutinee, arms) => { + let typed_scrutinee = self.typecheck_expr(scrutinee)?; + let mut typed_arms = Vec::new(); + let mut result_type = Type::Unknown; + + for (i, (pattern, body)) in arms.iter().enumerate() { + let typed_pattern = self.typecheck_pattern(pattern, &typed_scrutinee.ty)?; + let typed_body = self.typecheck_expr(body)?; + + if i == 0 { + result_type = typed_body.ty.clone(); + } else if !self.types_compatible(&typed_body.ty, &result_type) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(result_type, typed_body.ty), + span: body.span.clone(), + }); + } + + typed_arms.push((typed_pattern, typed_body)); + } + + ( + TypedExprKind::Match(Box::new(typed_scrutinee), typed_arms), + result_type, + ) + } + + ExprKind::Lambda(params, body) => { + self.env = self.env.enter_scope(); + + let mut param_types = Vec::new(); + for (param_name, param_type_annot) in params { + let param_type = param_type_annot + .as_ref() + .map(|t| self.type_annot_to_type(t)) + .unwrap_or(Type::Unknown); + param_types.push(param_type.clone()); + self.env + .add_var(param_name.clone(), param_type, BindingKind::Default); + } + + let typed_body = self.typecheck_expr(body)?; + let func_type = Type::Function(param_types, Box::new(typed_body.ty.clone())); + + ( + TypedExprKind::Lambda(params.clone(), Box::new(typed_body)), + func_type, + ) + } + + ExprKind::Assign(lhs, rhs) => { + let typed_lhs = self.typecheck_expr(lhs)?; + let typed_rhs = self.typecheck_expr(rhs)?; + + if !self.types_compatible(&typed_rhs.ty, &typed_lhs.ty) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(typed_lhs.ty.clone(), typed_rhs.ty), + span: rhs.span.clone(), + }); + } + + ( + TypedExprKind::Assign(Box::new(typed_lhs), Box::new(typed_rhs)), + Type::Unit, + ) + } + + ExprKind::Cast(expr_inner, target_type) => { + let typed_expr = self.typecheck_expr(expr_inner)?; + let target_ty = self.type_annot_to_type(target_type); + + ( + TypedExprKind::Cast(Box::new(typed_expr), target_type.clone()), + target_ty, + ) + } + + ExprKind::Return(value) => { + let typed_value = if let Some(v) = value { + Some(Box::new(self.typecheck_expr(v)?)) + } else { + None + }; + let return_type = typed_value + .as_ref() + .map(|v| v.ty.clone()) + .unwrap_or(Type::Unit); + (TypedExprKind::Return(typed_value), return_type) + } + + ExprKind::Break => (TypedExprKind::Break, Type::Never), + ExprKind::Continue => (TypedExprKind::Continue, Type::Never), + + ExprKind::EarlyReturn(value) => { + let typed_value = if let Some(v) = value { + Some(Box::new(self.typecheck_expr(v)?)) + } else { + None + }; + let return_type = typed_value + .as_ref() + .map(|v| v.ty.clone()) + .unwrap_or(Type::Unit); + (TypedExprKind::EarlyReturn(typed_value), return_type) + } + + ExprKind::OptionalChain(obj, field) => { + let typed_obj = if let Some(o) = obj { + Some(Box::new(self.typecheck_expr(o)?)) + } else { + None + }; + // Simplified - would need proper Option type handling + ( + TypedExprKind::OptionalChain(typed_obj, field.clone()), + Type::Unknown, + ) + } + }; + + Ok(TypedExpr { + kind, + span: expr.span.clone(), + attributes: expr.attributes.clone(), + ty, + }) + } + + // Helper function to infer generic type arguments + fn infer_type_args( + &self, + expected: &TypeAnnot, + actual: &Type, + type_params: &[String], + type_map: &mut HashMap, + ) { + match (expected, actual) { + (TypeAnnot::Var(param_name), actual_type) => { + // Check if this is actually a type parameter + if type_params.contains(param_name) { + type_map + .entry(param_name.clone()) + .or_insert(actual_type.clone()); + } + } + (TypeAnnot::Cons(_, args), _) if args.is_empty() => { + // No generic args to infer + } + (TypeAnnot::Array(inner), Type::Array(actual_inner)) => { + self.infer_type_args(inner, actual_inner, type_params, type_map); + } + (TypeAnnot::Tuple(expected_types), Type::Tuple(actual_types)) => { + for (e, a) in expected_types.iter().zip(actual_types.iter()) { + self.infer_type_args(e, a, type_params, type_map); + } + } + _ => {} + } + } + + fn typecheck_pattern( + &mut self, + pattern: &Pattern, + scrutinee_type: &Type, + ) -> Result { + let (kind, ty) = match &pattern.kind { + PatternKind::Wildcard => (TypedPatternKind::Wildcard, scrutinee_type.clone()), + + PatternKind::Variable(name) => { + self.env + .add_var(name.clone(), scrutinee_type.clone(), BindingKind::Default); + ( + TypedPatternKind::Variable(name.clone()), + scrutinee_type.clone(), + ) + } + + PatternKind::Literal(lit) => { + // Infer type from literal + let lit_type = if lit.parse::().is_ok() { + Type::Int + } else if lit.parse::().is_ok() { + Type::Float + } else if lit == "true" || lit == "false" { + Type::Bool + } else { + Type::String + }; + + if !self.types_compatible(&lit_type, scrutinee_type) { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch(scrutinee_type.clone(), lit_type), + span: pattern.span.clone(), + }); + } + + (TypedPatternKind::Literal(lit.clone()), lit_type) + } + + PatternKind::Tuple(patterns) => { + let mut typed_patterns = Vec::new(); + let mut types = Vec::new(); + + if let Type::Tuple(tuple_types) = scrutinee_type { + if patterns.len() != tuple_types.len() { + return Err(TypeError { + kind: TypeErrorKind::ArityMismatch(tuple_types.len(), patterns.len()), + span: pattern.span.clone(), + }); + } + + for (p, t) in patterns.iter().zip(tuple_types.iter()) { + let typed_p = self.typecheck_pattern(p, t)?; + types.push(typed_p.ty.clone()); + typed_patterns.push(typed_p); + } + } else { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + scrutinee_type.clone(), + Type::Tuple(vec![]), + ), + span: pattern.span.clone(), + }); + } + + (TypedPatternKind::Tuple(typed_patterns), Type::Tuple(types)) + } + + PatternKind::Struct(name, fields) => { + if let Type::Struct(struct_name, type_args) = scrutinee_type { + if name != struct_name { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + scrutinee_type.clone(), + Type::Struct(name.clone(), vec![]), + ), + span: pattern.span.clone(), + }); + } + + // Clone the struct fields we need before borrowing self mutably + let (struct_fields_clone, type_params) = { + let struct_info = self.env.get_type(name).ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedType(name.clone()), + span: pattern.span.clone(), + })?; + + if let TypeInfoKind::Struct(struct_fields) = &struct_info.kind { + (struct_fields.clone(), struct_info.parameters.clone()) + } else { + return Err(TypeError { + kind: TypeErrorKind::NotAStruct(scrutinee_type.clone()), + span: pattern.span.clone(), + }); + } + }; + + // Create substitution map for type parameters + let mut subst_map: HashMap = HashMap::new(); + for (param, arg) in type_params.iter().zip(type_args.iter()) { + subst_map.insert(param.clone(), arg.clone()); + } + + let mut typed_fields = Vec::new(); + for (field_name, field_pattern) in fields { + let field_type_annot = struct_fields_clone + .iter() + .find(|(n, _)| n == field_name) + .map(|(_, ty)| ty.clone()) + .ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedField( + field_name.clone(), + scrutinee_type.clone(), + ), + span: pattern.span.clone(), + })?; + + let field_type = self.substitute_type(&field_type_annot, &subst_map); + + let typed_field_pattern = + self.typecheck_pattern(field_pattern, &field_type)?; + typed_fields.push((field_name.clone(), typed_field_pattern)); + } + + ( + TypedPatternKind::Struct(name.clone(), typed_fields), + scrutinee_type.clone(), + ) + } else { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + scrutinee_type.clone(), + Type::Struct(name.clone(), vec![]), + ), + span: pattern.span.clone(), + }); + } + } + + PatternKind::Enum(enum_name, variant_name, patterns) => { + if let Type::Enum(scrutinee_enum_name, type_args) = scrutinee_type { + if enum_name != scrutinee_enum_name { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + scrutinee_type.clone(), + Type::Enum(enum_name.clone(), vec![]), + ), + span: pattern.span.clone(), + }); + } + + // Clone the variant fields we need before borrowing self mutably + let (variant_fields_clone, type_params) = { + let enum_info = self.env.get_type(enum_name).ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedType(enum_name.clone()), + span: pattern.span.clone(), + })?; + + if let TypeInfoKind::Enum(variants) = &enum_info.kind { + let variant = variants + .iter() + .find(|(n, _)| n == variant_name) + .ok_or_else(|| TypeError { + kind: TypeErrorKind::UndefinedVariant( + enum_name.clone(), + variant_name.clone(), + ), + span: pattern.span.clone(), + })?; + + if variant.1.len() != patterns.len() { + return Err(TypeError { + kind: TypeErrorKind::ArityMismatch( + variant.1.len(), + patterns.len(), + ), + span: pattern.span.clone(), + }); + } + + (variant.1.clone(), enum_info.parameters.clone()) + } else { + return Err(TypeError { + kind: TypeErrorKind::NotAnEnum(scrutinee_type.clone()), + span: pattern.span.clone(), + }); + } + }; + + // Create substitution map for type parameters + let mut subst_map: HashMap = HashMap::new(); + for (param, arg) in type_params.iter().zip(type_args.iter()) { + subst_map.insert(param.clone(), arg.clone()); + } + + let mut typed_patterns = Vec::new(); + for (p, field_type_annot) in patterns.iter().zip(variant_fields_clone.iter()) { + let field_type = self.substitute_type(field_type_annot, &subst_map); + let typed_p = self.typecheck_pattern(p, &field_type)?; + typed_patterns.push(typed_p); + } + + ( + TypedPatternKind::Enum( + enum_name.clone(), + variant_name.clone(), + typed_patterns, + ), + scrutinee_type.clone(), + ) + } else { + return Err(TypeError { + kind: TypeErrorKind::TypeMismatch( + scrutinee_type.clone(), + Type::Enum(enum_name.clone(), vec![]), + ), + span: pattern.span.clone(), + }); + } + } + + PatternKind::Range(_, _) => (TypedPatternKind::Wildcard, Type::Int), + }; + + Ok(TypedPattern { + kind, + span: pattern.span.clone(), + ty, + }) + } + + // Substitute type variables in a type annotation + fn substitute_type(&self, annot: &TypeAnnot, subst_map: &HashMap) -> Type { + match annot { + TypeAnnot::Var(name) => { + if let Some(ty) = subst_map.get(name) { + ty.clone() + } else { + self.type_annot_to_type(annot) + } + } + TypeAnnot::Cons(name, args) => { + let substituted_args: Vec = args + .iter() + .map(|arg| self.substitute_type(arg, subst_map)) + .collect(); + + if let Some(type_info) = self.env.get_type(name) { + match &type_info.kind { + TypeInfoKind::Struct(_) => Type::Struct(name.clone(), substituted_args), + TypeInfoKind::Enum(_) => Type::Enum(name.clone(), substituted_args), + } + } else { + match name.as_str() { + "int" => Type::Int, + "float" => Type::Float, + "bool" => Type::Bool, + "string" => Type::String, + "unit" => Type::Unit, + "never" => Type::Never, + _ => Type::Generic(name.clone(), substituted_args), + } + } + } + TypeAnnot::Array(inner) => { + Type::Array(Box::new(self.substitute_type(inner, subst_map))) + } + TypeAnnot::Tuple(types) => { + let substituted_types: Vec = types + .iter() + .map(|t| self.substitute_type(t, subst_map)) + .collect(); + Type::Tuple(substituted_types) + } + TypeAnnot::Function(args, ret) => { + let arg_types: Vec = args + .iter() + .map(|a| self.substitute_type(a, subst_map)) + .collect(); + let ret_type = Box::new(self.substitute_type(ret, subst_map)); + Type::Function(arg_types, ret_type) + } + } + } + + fn type_annot_to_type(&self, annot: &TypeAnnot) -> Type { + match annot { + TypeAnnot::Var(name) => { + // Check if it's a type variable + if let Some(ty) = self.env.type_vars.get(name) { + return ty.clone(); + } + + match name.as_str() { + "int" => Type::Int, + "float" => Type::Float, + "bool" => Type::Bool, + "string" => Type::String, + "unit" => Type::Unit, + "never" => Type::Never, + _ => Type::TypeVar(name.clone()), + } + } + TypeAnnot::Cons(name, args) => { + let type_args: Vec = + args.iter().map(|a| self.type_annot_to_type(a)).collect(); + + match name.as_str() { + "int" => Type::Int, + "float" => Type::Float, + "bool" => Type::Bool, + "string" => Type::String, + "unit" => Type::Unit, + "never" => Type::Never, + _ => { + // Check if it's a struct or enum + if let Some(type_info) = self.env.get_type(name) { + match &type_info.kind { + TypeInfoKind::Struct(_) => Type::Struct(name.clone(), type_args), + TypeInfoKind::Enum(_) => Type::Enum(name.clone(), type_args), + } + } else { + Type::Generic(name.clone(), type_args) + } + } + } + } + TypeAnnot::Function(args, ret) => { + let arg_types: Vec = + args.iter().map(|a| self.type_annot_to_type(a)).collect(); + let ret_type = Box::new(self.type_annot_to_type(ret)); + Type::Function(arg_types, ret_type) + } + TypeAnnot::Tuple(types) => { + let tuple_types: Vec = + types.iter().map(|t| self.type_annot_to_type(t)).collect(); + Type::Tuple(tuple_types) + } + TypeAnnot::Array(inner) => Type::Array(Box::new(self.type_annot_to_type(inner))), + } + } + + fn types_compatible(&self, t1: &Type, t2: &Type) -> bool { + match (t1, t2) { + (Type::Unknown, _) | (_, Type::Unknown) => true, + (Type::Int, Type::Int) => true, + (Type::Float, Type::Float) => true, + (Type::Bool, Type::Bool) => true, + (Type::String, Type::String) => true, + (Type::Unit, Type::Unit) => true, + (Type::Never, _) | (_, Type::Never) => true, + (Type::Array(a), Type::Array(b)) => self.types_compatible(a, b), + (Type::Tuple(a), Type::Tuple(b)) => { + a.len() == b.len() + && a.iter() + .zip(b.iter()) + .all(|(x, y)| self.types_compatible(x, y)) + } + (Type::Function(args1, ret1), Type::Function(args2, ret2)) => { + args1.len() == args2.len() + && args1 + .iter() + .zip(args2.iter()) + .all(|(x, y)| self.types_compatible(x, y)) + && self.types_compatible(ret1, ret2) + } + (Type::Struct(name1, args1), Type::Struct(name2, args2)) => { + name1 == name2 + && args1.len() == args2.len() + && args1 + .iter() + .zip(args2.iter()) + .all(|(x, y)| self.types_compatible(x, y)) + } + (Type::Enum(name1, args1), Type::Enum(name2, args2)) => { + name1 == name2 + && args1.len() == args2.len() + && args1 + .iter() + .zip(args2.iter()) + .all(|(x, y)| self.types_compatible(x, y)) + } + (Type::Generic(name1, args1), Type::Generic(name2, args2)) => { + name1 == name2 + && args1.len() == args2.len() + && args1 + .iter() + .zip(args2.iter()) + .all(|(x, y)| self.types_compatible(x, y)) + } + (Type::TypeVar(a), Type::TypeVar(b)) => a == b, + (Type::TypeVar(_), _) | (_, Type::TypeVar(_)) => true, // Type variables are compatible with anything + (Type::Generic(_, _), _) | (_, Type::Generic(_, _)) => true, // Generic types are compatible with anything (for now) + _ => false, + } + } } From 22e2a454019d2a46b8fd2c1cd01da3a1aaf131ba Mon Sep 17 00:00:00 2001 From: Masashi Date: Mon, 15 Dec 2025 15:26:43 +0530 Subject: [PATCH 2/2] monomorphization --- src/lib.rs | 1 + src/main.rs | 103 ++- src/monomorphize.rs | 1197 ++++++++++++++++++++++++++++++ src/typechecker.rs | 51 +- tests/generics_comprehensive.sui | 28 + 5 files changed, 1376 insertions(+), 4 deletions(-) create mode 100644 src/monomorphize.rs create mode 100644 tests/generics_comprehensive.sui diff --git a/src/lib.rs b/src/lib.rs index ccd1859..7c8a819 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -4,3 +4,4 @@ pub mod ast; pub mod lexer; pub mod parser; pub mod typechecker; +pub mod monomorphize; diff --git a/src/main.rs b/src/main.rs index af81cfd..faf0876 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,6 +1,9 @@ use logos::Logos; use std::fs; -use suicmez::{lexer::Token, parser::Parser, typechecker::TypeChecker}; +use suicmez::{ + lexer::Token, parser::Parser, typechecker::TypeChecker, + monomorphize::{Monomorphizer, check_no_typevars}, +}; fn main() { // Check if a file was provided as argument @@ -84,5 +87,103 @@ fn run_file(filename: &str) -> Result<(), String> { typed_nodes.len() ); + // Debug: show typed nodes + println!("\nTyped AST nodes before monomorphization:"); + for (i, node) in typed_nodes.iter().enumerate() { + let node_type = match &node.kind { + suicmez::ast::TypedASTNodeKind::Function(f) => { + format!("Function({})", f.name) + } + suicmez::ast::TypedASTNodeKind::Struct(s) => { + format!("Struct({}) with {} params", s.name, s.parameters.len()) + } + suicmez::ast::TypedASTNodeKind::Enum(e) => { + format!("Enum({}) with {} params", e.name, e.parameters.len()) + } + suicmez::ast::TypedASTNodeKind::Impl(imp) => { + format!("Impl({})", imp.target) + } + suicmez::ast::TypedASTNodeKind::Trait(t) => { + format!("Trait({})", t.name) + } + suicmez::ast::TypedASTNodeKind::Extern(e) => { + format!("Extern({})", e.name) + } + suicmez::ast::TypedASTNodeKind::Load(l) => { + format!("Load({})", l.alias) + } + suicmez::ast::TypedASTNodeKind::Use(u) => { + format!("Use({})", u) + } + }; + println!(" [{}] {}", i, node_type); + } + + // Monomorphize the AST + let monomorphizer = Monomorphizer::new(); + let mono_nodes = monomorphizer.monomorphize_program(&typed_nodes).map_err(|e| { + format!( + "Monomorphization error: {}{}", + e.message, + if let Some(span) = &e.span { + format!(" at {}:{}", span.file, span.start) + } else { + String::new() + } + ) + })?; + + println!( + "Monomorphization passed! {} nodes after specialization.", + mono_nodes.len() + ); + + // Print detailed info about each node + println!("\nMonomorphized AST nodes:"); + for (i, node) in mono_nodes.iter().enumerate() { + let node_type = match &node.kind { + suicmez::ast::TypedASTNodeKind::Function(f) => { + format!("Function({})", f.name) + } + suicmez::ast::TypedASTNodeKind::Struct(s) => { + format!("Struct({}) with {} params", s.name, s.parameters.len()) + } + suicmez::ast::TypedASTNodeKind::Enum(e) => { + format!("Enum({}) with {} params", e.name, e.parameters.len()) + } + suicmez::ast::TypedASTNodeKind::Impl(imp) => { + format!("Impl({})", imp.target) + } + suicmez::ast::TypedASTNodeKind::Trait(t) => { + format!("Trait({})", t.name) + } + suicmez::ast::TypedASTNodeKind::Extern(e) => { + format!("Extern({})", e.name) + } + suicmez::ast::TypedASTNodeKind::Load(l) => { + format!("Load({})", l.alias) + } + suicmez::ast::TypedASTNodeKind::Use(u) => { + format!("Use({})", u) + } + }; + println!(" [{}] {}", i, node_type); + } + + // Check that no type variables remain + check_no_typevars(&mono_nodes).map_err(|e| { + format!( + "Type variable check failed: {}{}", + e.message, + if let Some(span) = &e.span { + format!(" at {}:{}", span.file, span.start) + } else { + String::new() + } + ) + })?; + + println!("Type variable check passed! No type variables remain in AST."); + Ok(()) } diff --git a/src/monomorphize.rs b/src/monomorphize.rs new file mode 100644 index 0000000..a0f5280 --- /dev/null +++ b/src/monomorphize.rs @@ -0,0 +1,1197 @@ +use crate::ast::*; +use crate::typechecker::Type; +use std::collections::{HashMap, HashSet}; + +#[derive(Debug)] +pub struct MonomorphizationError { + pub message: String, + pub span: Option, +} + +impl MonomorphizationError { + fn new(message: impl Into, span: Option) -> Self { + MonomorphizationError { + message: message.into(), + span, + } + } +} + +/// Specialization cache to avoid duplicating already-generated specializations +#[derive(Clone)] +struct SpecializationKey { + base_name: String, + type_args: Vec, +} + +impl SpecializationKey { + fn new(base_name: String, type_args: Vec) -> Self { + SpecializationKey { + base_name, + type_args, + } + } + + fn to_string(&self) -> String { + if self.type_args.is_empty() { + self.base_name.clone() + } else { + let arg_strs: Vec = self.type_args.iter().map(|t| t.to_string()).collect(); + format!("{}_{}", self.base_name, arg_strs.join("_")) + } + } + + fn to_hashable(&self) -> String { + self.to_string() + } +} + +/// The monomorphizer specializes generic types into concrete versions +pub struct Monomorphizer { + // Track all generated specializations to avoid duplicates + generated_structs: HashMap, + generated_enums: HashMap, + generated_functions: HashMap, +} + +impl Monomorphizer { + pub fn new() -> Self { + Monomorphizer { + generated_structs: HashMap::new(), + generated_enums: HashMap::new(), + generated_functions: HashMap::new(), + } + } + + pub fn monomorphize_program( + mut self, + nodes: &[TypedASTNode], + ) -> Result, MonomorphizationError> { + // First pass: collect all generic definitions + let mut generic_structs = HashMap::new(); + let mut generic_enums = HashMap::new(); + let mut generic_functions = HashMap::new(); + let mut generic_impls = Vec::new(); + + for node in nodes { + match &node.kind { + TypedASTNodeKind::Struct(s) => { + if !s.parameters.is_empty() { + generic_structs.insert(s.name.clone(), (s.clone(), node.clone())); + } + } + TypedASTNodeKind::Enum(e) => { + if !e.parameters.is_empty() { + generic_enums.insert(e.name.clone(), (e.clone(), node.clone())); + } + } + TypedASTNodeKind::Function(f) => { + if !f.parameters.is_empty() { + generic_functions.insert(f.name.clone(), (f.clone(), node.clone())); + } + } + TypedASTNodeKind::Impl(imp) => { + generic_impls.push((imp.clone(), node.clone())); + } + _ => {} + } + } + + // Second pass: monomorphize expressions to collect specialization requirements + let mut specialization_needs: Vec = Vec::new(); + let mut seen_keys: HashSet = HashSet::new(); + let mut result_nodes = Vec::new(); + + for (_idx, node) in nodes.iter().enumerate() { + match &node.kind { + TypedASTNodeKind::Function(f) => { + // Skip generic functions - they'll be added as specialized versions when needed + if !f.parameters.is_empty() { + continue; + } + + let (mono_func, needs) = self.monomorphize_function( + f, + &generic_structs, + &generic_enums, + &generic_functions, + )?; + for need in needs { + let key = need.to_hashable(); + if !seen_keys.contains(&key) { + seen_keys.insert(key); + specialization_needs.push(need); + } + } + + let mut new_node = node.clone(); + new_node.kind = TypedASTNodeKind::Function(mono_func); + result_nodes.push(new_node); + } + TypedASTNodeKind::Struct(s) => { + // Skip generic structs - they'll be added as specialized versions when needed + if !s.parameters.is_empty() { + continue; + } + result_nodes.push(node.clone()); + } + TypedASTNodeKind::Enum(e) => { + // Skip generic enums - they'll be added as specialized versions when needed + if !e.parameters.is_empty() { + continue; + } + result_nodes.push(node.clone()); + } + TypedASTNodeKind::Impl(imp) => { + let (mono_impl, needs) = self.monomorphize_impl( + imp, + &generic_structs, + &generic_enums, + &generic_functions, + )?; + for need in needs { + let key = need.to_hashable(); + if !seen_keys.contains(&key) { + seen_keys.insert(key); + specialization_needs.push(need); + } + } + + let mut new_node = node.clone(); + new_node.kind = TypedASTNodeKind::Impl(mono_impl); + result_nodes.push(new_node); + } + _ => { + result_nodes.push(node.clone()); + } + } + } + + // Third pass: generate all needed specializations + let mut iterations = 0; + const MAX_ITERATIONS: usize = 1000; // Prevent infinite loops + + while !specialization_needs.is_empty() && iterations < MAX_ITERATIONS { + iterations += 1; + let current_needs: Vec<_> = specialization_needs.drain(..).collect(); + + for key in current_needs { + if self.generated_structs.contains_key(&key.to_string()) { + continue; + } + + // Try to specialize a struct + if let Some((generic_struct, orig_node)) = generic_structs.get(&key.base_name) { + let (mono_struct, needs) = self.specialize_struct( + generic_struct, + &key.type_args, + &generic_structs, + &generic_enums, + &generic_functions, + )?; + self.generated_structs + .insert(key.to_string(), mono_struct.clone()); + for need in needs { + let need_key = need.to_hashable(); + if !seen_keys.contains(&need_key) { + seen_keys.insert(need_key); + specialization_needs.push(need); + } + } + + let mut new_node = orig_node.clone(); + new_node.kind = TypedASTNodeKind::Struct(mono_struct); + result_nodes.push(new_node); + continue; + } + + // Try to specialize an enum + if let Some((generic_enum, orig_node)) = generic_enums.get(&key.base_name) { + let (mono_enum, needs) = self.specialize_enum( + generic_enum, + &key.type_args, + &generic_structs, + &generic_enums, + &generic_functions, + )?; + + // Only add if it was actually specialized (arity matched) + if mono_enum.parameters.is_empty() { + self.generated_enums + .insert(key.to_string(), mono_enum.clone()); + for need in needs { + let need_key = need.to_hashable(); + if !seen_keys.contains(&need_key) { + seen_keys.insert(need_key); + specialization_needs.push(need); + } + } + + let mut new_node = orig_node.clone(); + new_node.kind = TypedASTNodeKind::Enum(mono_enum); + result_nodes.push(new_node); + } + continue; + } + + // Try to specialize a function + if let Some((generic_func, orig_node)) = generic_functions.get(&key.base_name) { + let (mono_func, needs) = self.specialize_function( + generic_func, + &key.type_args, + &generic_structs, + &generic_enums, + &generic_functions, + )?; + self.generated_functions + .insert(key.to_string(), mono_func.clone()); + for need in needs { + let need_key = need.to_hashable(); + if !seen_keys.contains(&need_key) { + seen_keys.insert(need_key); + specialization_needs.push(need); + } + } + + let mut new_node = orig_node.clone(); + new_node.kind = TypedASTNodeKind::Function(mono_func); + result_nodes.push(new_node); + } + } + } + + if iterations >= MAX_ITERATIONS { + return Err(MonomorphizationError::new( + "Monomorphization exceeded maximum iterations (possible infinite recursion)", + None, + )); + } + + Ok(result_nodes) + } + + fn monomorphize_function( + &mut self, + func: &TypedFunction, + _generic_structs: &HashMap, + _generic_enums: &HashMap, + _generic_functions: &HashMap, + ) -> Result<(TypedFunction, Vec), MonomorphizationError> { + if func.parameters.is_empty() { + let (body, needs) = self.monomorphize_expr(&func.body)?; + let mut new_func = func.clone(); + new_func.body = body; + Ok((new_func, needs)) + } else { + // Functions with type parameters should not appear in final code + // They'll be specialized as needed + Ok((func.clone(), Vec::new())) + } + } + + fn monomorphize_impl( + &mut self, + imp: &TypedImpl, + _generic_structs: &HashMap, + _generic_enums: &HashMap, + _generic_functions: &HashMap, + ) -> Result<(TypedImpl, Vec), MonomorphizationError> { + let mut all_needs = Vec::new(); + let mut new_methods = Vec::new(); + + for method in &imp.methods { + let (mono_method, needs) = self.monomorphize_function( + method, + _generic_structs, + _generic_enums, + _generic_functions, + )?; + all_needs.extend(needs); + new_methods.push(mono_method); + } + + let mut new_impl = imp.clone(); + new_impl.methods = new_methods; + Ok((new_impl, all_needs)) + } + + fn monomorphize_expr( + &mut self, + expr: &TypedExpr, + ) -> Result<(TypedExpr, Vec), MonomorphizationError> { + let mut needs = Vec::new(); + let new_kind = match &expr.kind { + TypedExprKind::Int(_) + | TypedExprKind::Float(_) + | TypedExprKind::Bool(_) + | TypedExprKind::String(_) + | TypedExprKind::Break + | TypedExprKind::Continue => expr.kind.clone(), + + TypedExprKind::Array(elems) => { + let mut new_elems = Vec::new(); + for elem in elems { + let (new_elem, elem_needs) = self.monomorphize_expr(elem)?; + needs.extend(elem_needs); + new_elems.push(new_elem); + } + TypedExprKind::Array(new_elems) + } + + TypedExprKind::Tuple(elems) => { + let mut new_elems = Vec::new(); + for elem in elems { + let (new_elem, elem_needs) = self.monomorphize_expr(elem)?; + needs.extend(elem_needs); + new_elems.push(new_elem); + } + TypedExprKind::Tuple(new_elems) + } + + TypedExprKind::StructLit(name, fields) => { + let mut new_fields = Vec::new(); + let mut field_types = Vec::new(); + for (field_name, field_expr) in fields { + let (new_expr, expr_needs) = self.monomorphize_expr(field_expr)?; + field_types.push(new_expr.ty.clone()); + needs.extend(expr_needs); + new_fields.push((field_name.clone(), new_expr)); + } + // Infer struct specialization from field types + if !field_types.is_empty() { + self.infer_struct_specialization(name, &field_types, &mut needs); + } + TypedExprKind::StructLit(name.clone(), new_fields) + } + + TypedExprKind::EnumLit(enum_name, variant, args) => { + let mut new_args = Vec::new(); + let mut arg_types = Vec::new(); + for arg in args { + let (new_arg, arg_needs) = self.monomorphize_expr(arg)?; + arg_types.push(new_arg.ty.clone()); + needs.extend(arg_needs); + new_args.push(new_arg); + } + // Infer enum specialization from argument types + if !arg_types.is_empty() { + self.infer_enum_specialization(enum_name, &arg_types, &mut needs); + } + TypedExprKind::EnumLit(enum_name.clone(), variant.clone(), new_args) + } + + TypedExprKind::Variable(_) => expr.kind.clone(), + + TypedExprKind::Call(func_expr, args) => { + let (new_func_expr, func_needs) = self.monomorphize_expr(func_expr)?; + needs.extend(func_needs); + + let mut new_args = Vec::new(); + for arg in args { + let (new_arg, arg_needs) = self.monomorphize_expr(arg)?; + needs.extend(arg_needs); + new_args.push(new_arg); + } + + // Collect function call specialization needs from return type + self.collect_needs_from_expr_type(expr, &mut needs); + + TypedExprKind::Call(Box::new(new_func_expr), new_args) + } + + TypedExprKind::Index(array_expr, index_expr) => { + let (new_array, array_needs) = self.monomorphize_expr(array_expr)?; + let (new_index, index_needs) = self.monomorphize_expr(index_expr)?; + needs.extend(array_needs); + needs.extend(index_needs); + TypedExprKind::Index(Box::new(new_array), Box::new(new_index)) + } + + TypedExprKind::Dot(obj_expr, field) => { + let (new_obj, obj_needs) = self.monomorphize_expr(obj_expr)?; + needs.extend(obj_needs); + TypedExprKind::Dot(Box::new(new_obj), field.clone()) + } + + TypedExprKind::EarlyReturn(expr_opt) => { + if let Some(inner_expr) = expr_opt { + let (new_expr, expr_needs) = self.monomorphize_expr(inner_expr)?; + needs.extend(expr_needs); + TypedExprKind::EarlyReturn(Some(Box::new(new_expr))) + } else { + TypedExprKind::EarlyReturn(None) + } + } + + TypedExprKind::OptionalChain(expr_opt, field) => { + if let Some(inner_expr) = expr_opt { + let (new_expr, expr_needs) = self.monomorphize_expr(inner_expr)?; + needs.extend(expr_needs); + TypedExprKind::OptionalChain(Some(Box::new(new_expr)), field.clone()) + } else { + TypedExprKind::OptionalChain(None, field.clone()) + } + } + + TypedExprKind::Lambda(params, body) => { + let (new_body, body_needs) = self.monomorphize_expr(body)?; + needs.extend(body_needs); + TypedExprKind::Lambda(params.clone(), Box::new(new_body)) + } + + TypedExprKind::Let(name, binding_kind, ty_annot, expr) => { + let (new_expr, expr_needs) = self.monomorphize_expr(expr)?; + needs.extend(expr_needs); + TypedExprKind::Let( + name.clone(), + binding_kind.clone(), + ty_annot.clone(), + Box::new(new_expr), + ) + } + + TypedExprKind::Assign(lvalue, rvalue) => { + let (new_lvalue, lvalue_needs) = self.monomorphize_expr(lvalue)?; + let (new_rvalue, rvalue_needs) = self.monomorphize_expr(rvalue)?; + needs.extend(lvalue_needs); + needs.extend(rvalue_needs); + TypedExprKind::Assign(Box::new(new_lvalue), Box::new(new_rvalue)) + } + + TypedExprKind::Cast(expr, ty) => { + let (new_expr, expr_needs) = self.monomorphize_expr(expr)?; + needs.extend(expr_needs); + TypedExprKind::Cast(Box::new(new_expr), ty.clone()) + } + + TypedExprKind::If(cond, then_expr, else_expr) => { + let (new_cond, cond_needs) = self.monomorphize_expr(cond)?; + let (new_then, then_needs) = self.monomorphize_expr(then_expr)?; + needs.extend(cond_needs); + needs.extend(then_needs); + + let new_else = if let Some(else_e) = else_expr { + let (new_else_expr, else_needs) = self.monomorphize_expr(else_e)?; + needs.extend(else_needs); + Some(Box::new(new_else_expr)) + } else { + None + }; + + TypedExprKind::If(Box::new(new_cond), Box::new(new_then), new_else) + } + + TypedExprKind::Match(scrutinee, arms) => { + let (new_scrutinee, scrutinee_needs) = self.monomorphize_expr(scrutinee)?; + needs.extend(scrutinee_needs); + + let mut new_arms = Vec::new(); + for (pattern, arm_expr) in arms { + let (new_arm_expr, arm_needs) = self.monomorphize_expr(arm_expr)?; + needs.extend(arm_needs); + new_arms.push((pattern.clone(), new_arm_expr)); + } + + TypedExprKind::Match(Box::new(new_scrutinee), new_arms) + } + + TypedExprKind::While(cond, body) => { + let (new_cond, cond_needs) = self.monomorphize_expr(cond)?; + let (new_body, body_needs) = self.monomorphize_expr(body)?; + needs.extend(cond_needs); + needs.extend(body_needs); + TypedExprKind::While(Box::new(new_cond), Box::new(new_body)) + } + + TypedExprKind::Do(exprs) => { + let mut new_exprs = Vec::new(); + for e in exprs { + let (new_e, e_needs) = self.monomorphize_expr(e)?; + needs.extend(e_needs); + new_exprs.push(new_e); + } + TypedExprKind::Do(new_exprs) + } + + TypedExprKind::BinOp(lhs, op, rhs) => { + let (new_lhs, lhs_needs) = self.monomorphize_expr(lhs)?; + let (new_rhs, rhs_needs) = self.monomorphize_expr(rhs)?; + needs.extend(lhs_needs); + needs.extend(rhs_needs); + TypedExprKind::BinOp(Box::new(new_lhs), op.clone(), Box::new(new_rhs)) + } + + TypedExprKind::UnOp(op, operand) => { + let (new_operand, operand_needs) = self.monomorphize_expr(operand)?; + needs.extend(operand_needs); + TypedExprKind::UnOp(op.clone(), Box::new(new_operand)) + } + + TypedExprKind::For(var, iter_expr, body) => { + let (new_iter, iter_needs) = self.monomorphize_expr(iter_expr)?; + let (new_body, body_needs) = self.monomorphize_expr(body)?; + needs.extend(iter_needs); + needs.extend(body_needs); + TypedExprKind::For(var.clone(), Box::new(new_iter), Box::new(new_body)) + } + + TypedExprKind::Range(start, end) => { + let (new_start, start_needs) = self.monomorphize_expr(start)?; + let (new_end, end_needs) = self.monomorphize_expr(end)?; + needs.extend(start_needs); + needs.extend(end_needs); + TypedExprKind::Range(Box::new(new_start), Box::new(new_end)) + } + + TypedExprKind::Return(expr_opt) => { + if let Some(ret_expr) = expr_opt { + let (new_expr, expr_needs) = self.monomorphize_expr(ret_expr)?; + needs.extend(expr_needs); + TypedExprKind::Return(Some(Box::new(new_expr))) + } else { + TypedExprKind::Return(None) + } + } + }; + + let mut new_expr = expr.clone(); + new_expr.kind = new_kind; + Ok((new_expr, needs)) + } + + fn specialize_struct( + &mut self, + generic_struct: &TypedStruct, + type_args: &[Type], + _generic_structs: &HashMap, + _generic_enums: &HashMap, + _generic_functions: &HashMap, + ) -> Result<(TypedStruct, Vec), MonomorphizationError> { + if generic_struct.parameters.len() != type_args.len() { + return Err(MonomorphizationError::new( + format!( + "Struct {} expects {} type arguments, got {}", + generic_struct.name, + generic_struct.parameters.len(), + type_args.len() + ), + None, + )); + } + + let mut subst_map = HashMap::new(); + for (param, arg) in generic_struct.parameters.iter().zip(type_args.iter()) { + subst_map.insert(param.name.clone(), arg.clone()); + } + + let mut new_fields = Vec::new(); + let mut needs = Vec::new(); + + for field in &generic_struct.fields { + let new_ty = self.substitute_in_type_annot(&field.field_type, &subst_map)?; + + // Collect specialization needs from the field type + self.collect_needs_from_type(&new_ty, &mut needs); + + new_fields.push(TypedField { + name: field.name.clone(), + field_type: new_ty, + span: field.span.clone(), + }); + } + + let mut new_struct = generic_struct.clone(); + new_struct.name = self.generate_specialized_name(&generic_struct.name, type_args); + new_struct.parameters = Vec::new(); // Remove type parameters after specialization + new_struct.fields = new_fields; + + Ok((new_struct, needs)) + } + + fn specialize_enum( + &mut self, + generic_enum: &TypedEnum, + type_args: &[Type], + _generic_structs: &HashMap, + _generic_enums: &HashMap, + _generic_functions: &HashMap, + ) -> Result<(TypedEnum, Vec), MonomorphizationError> { + if generic_enum.parameters.len() != type_args.len() { + // If we can't specialize due to type arity mismatch, just skip it + // This can happen when the typechecker doesn't fully infer generic types + return Ok((generic_enum.clone(), Vec::new())); + } + + let mut subst_map = HashMap::new(); + for (param, arg) in generic_enum.parameters.iter().zip(type_args.iter()) { + subst_map.insert(param.name.clone(), arg.clone()); + } + + let mut new_variants = Vec::new(); + let mut needs = Vec::new(); + + for variant in &generic_enum.variants { + let mut new_fields = Vec::new(); + for field_ty in &variant.fields { + let new_ty = self.substitute_in_type_annot(field_ty, &subst_map)?; + self.collect_needs_from_type(&new_ty, &mut needs); + new_fields.push(new_ty); + } + + new_variants.push(TypedVariant { + name: variant.name.clone(), + fields: new_fields, + span: variant.span.clone(), + }); + } + + let mut new_enum = generic_enum.clone(); + new_enum.name = self.generate_specialized_name(&generic_enum.name, type_args); + new_enum.parameters = Vec::new(); // Remove type parameters after specialization + new_enum.variants = new_variants; + + Ok((new_enum, needs)) + } + + fn specialize_function( + &mut self, + generic_func: &TypedFunction, + type_args: &[Type], + _generic_structs: &HashMap, + _generic_enums: &HashMap, + _generic_functions: &HashMap, + ) -> Result<(TypedFunction, Vec), MonomorphizationError> { + if generic_func.parameters.len() != type_args.len() { + return Err(MonomorphizationError::new( + format!( + "Function {} expects {} type arguments, got {}", + generic_func.name, + generic_func.parameters.len(), + type_args.len() + ), + None, + )); + } + + let mut subst_map = HashMap::new(); + for (param, arg) in generic_func.parameters.iter().zip(type_args.iter()) { + subst_map.insert(param.name.clone(), arg.clone()); + } + + // Specialize arguments + let mut new_args = Vec::new(); + let mut needs = Vec::new(); + + for (arg_name, arg_ty_opt) in &generic_func.args { + let new_arg_ty = if let Some(arg_ty) = arg_ty_opt { + let ty = self.substitute_in_type_annot(arg_ty, &subst_map)?; + self.collect_needs_from_type(&ty, &mut needs); + Some(ty) + } else { + None + }; + new_args.push((arg_name.clone(), new_arg_ty)); + } + + // Specialize return type + let new_return_type = if let Some(ret_ty) = &generic_func.return_type { + let ty = self.substitute_in_type_annot(ret_ty, &subst_map)?; + self.collect_needs_from_type(&ty, &mut needs); + Some(ty) + } else { + None + }; + + // Specialize body + let (new_body, body_needs) = self.monomorphize_expr(&generic_func.body)?; + needs.extend(body_needs); + + let mut new_func = generic_func.clone(); + new_func.name = self.generate_specialized_name(&generic_func.name, type_args); + new_func.parameters = Vec::new(); // Remove type parameters after specialization + new_func.args = new_args; + new_func.return_type = new_return_type; + new_func.body = new_body; + + Ok((new_func, needs)) + } + + fn substitute_in_type_annot( + &self, + annot: &TypeAnnot, + subst_map: &HashMap, + ) -> Result { + match annot { + TypeAnnot::Var(name) => { + if let Some(ty) = subst_map.get(name) { + Ok(self.type_to_type_annot(ty)) + } else { + // This is fine - it could be a non-parameterized type + Ok(TypeAnnot::Var(name.clone())) + } + } + TypeAnnot::Cons(name, args) => { + let mut new_args = Vec::new(); + for arg in args { + new_args.push(self.substitute_in_type_annot(arg, subst_map)?); + } + Ok(TypeAnnot::Cons(name.clone(), new_args)) + } + TypeAnnot::Function(param_types, ret_type) => { + let mut new_params = Vec::new(); + for param in param_types { + new_params.push(self.substitute_in_type_annot(param, subst_map)?); + } + let new_ret = self.substitute_in_type_annot(ret_type, subst_map)?; + Ok(TypeAnnot::Function(new_params, Box::new(new_ret))) + } + TypeAnnot::Tuple(types) => { + let mut new_types = Vec::new(); + for ty in types { + new_types.push(self.substitute_in_type_annot(ty, subst_map)?); + } + Ok(TypeAnnot::Tuple(new_types)) + } + TypeAnnot::Array(inner) => { + let new_inner = self.substitute_in_type_annot(inner, subst_map)?; + Ok(TypeAnnot::Array(Box::new(new_inner))) + } + } + } + + fn type_to_type_annot(&self, ty: &Type) -> TypeAnnot { + match ty { + Type::Int => TypeAnnot::Var("int".to_string()), + Type::Float => TypeAnnot::Var("float".to_string()), + Type::Bool => TypeAnnot::Var("bool".to_string()), + Type::String => TypeAnnot::Var("string".to_string()), + Type::Unit => TypeAnnot::Tuple(Vec::new()), + Type::Array(inner) => TypeAnnot::Array(Box::new(self.type_to_type_annot(inner))), + Type::Tuple(types) => { + let annots = types.iter().map(|t| self.type_to_type_annot(t)).collect(); + TypeAnnot::Tuple(annots) + } + Type::Struct(name, args) => { + if args.is_empty() { + TypeAnnot::Var(name.clone()) + } else { + let arg_annots = args.iter().map(|t| self.type_to_type_annot(t)).collect(); + TypeAnnot::Cons(name.clone(), arg_annots) + } + } + Type::Enum(name, args) => { + if args.is_empty() { + TypeAnnot::Var(name.clone()) + } else { + let arg_annots = args.iter().map(|t| self.type_to_type_annot(t)).collect(); + TypeAnnot::Cons(name.clone(), arg_annots) + } + } + Type::Function(params, ret) => { + let param_annots = params.iter().map(|t| self.type_to_type_annot(t)).collect(); + let ret_annot = self.type_to_type_annot(ret); + TypeAnnot::Function(param_annots, Box::new(ret_annot)) + } + Type::TypeVar(name) => TypeAnnot::Var(name.clone()), + Type::Generic(name, args) => { + if args.is_empty() { + TypeAnnot::Var(name.clone()) + } else { + let arg_annots = args.iter().map(|t| self.type_to_type_annot(t)).collect(); + TypeAnnot::Cons(name.clone(), arg_annots) + } + } + Type::Never => TypeAnnot::Var("!".to_string()), + Type::Unknown => TypeAnnot::Var("?".to_string()), + } + } + + fn collect_needs_from_expr_type(&self, expr: &TypedExpr, needs: &mut Vec) { + // Collect specialization needs from the expression's type + match &expr.ty { + Type::Struct(name, args) if !args.is_empty() && !name.contains("?") => { + // Skip Unknown types + needs.push(SpecializationKey::new(name.clone(), args.clone())); + } + Type::Enum(name, args) if !args.is_empty() && !name.contains("?") => { + // Skip Unknown types + needs.push(SpecializationKey::new(name.clone(), args.clone())); + } + Type::Function(_, _) => { + // Function types don't need specialization at the call site + } + _ => {} + } + } + + fn infer_struct_specialization( + &self, + struct_name: &str, + field_types: &[Type], + needs: &mut Vec, + ) { + // Only infer single-parameter generics from field types + // This is a heuristic for Box { value: T } + if field_types.len() == 1 { + needs.push(SpecializationKey::new( + struct_name.to_string(), + vec![field_types[0].clone()], + )); + } + // For multi-field structs, we can't reliably infer the type parameters + } + + fn infer_enum_specialization( + &self, + enum_name: &str, + arg_types: &[Type], + needs: &mut Vec, + ) { + // Only infer single-parameter generics from argument types + // This is a heuristic for Option::Some(T) where arg_types[0] is T + if arg_types.len() == 1 { + needs.push(SpecializationKey::new( + enum_name.to_string(), + vec![arg_types[0].clone()], + )); + } + // For multi-parameter enums, we can't reliably infer from just the variant arguments + } + + fn collect_needs_from_type(&self, ty: &TypeAnnot, needs: &mut Vec) { + match ty { + TypeAnnot::Var(_) => {} + TypeAnnot::Cons(name, args) => { + let type_args: Vec = + args.iter().map(|a| self.type_annot_to_type(a)).collect(); + if !type_args.is_empty() { + needs.push(SpecializationKey::new(name.clone(), type_args)); + } + for arg in args { + self.collect_needs_from_type(arg, needs); + } + } + TypeAnnot::Function(params, ret) => { + for param in params { + self.collect_needs_from_type(param, needs); + } + self.collect_needs_from_type(ret, needs); + } + TypeAnnot::Tuple(types) => { + for ty in types { + self.collect_needs_from_type(ty, needs); + } + } + TypeAnnot::Array(inner) => { + self.collect_needs_from_type(inner, needs); + } + } + } + + fn type_annot_to_type(&self, annot: &TypeAnnot) -> Type { + match annot { + TypeAnnot::Var(name) => match name.as_str() { + "int" => Type::Int, + "float" => Type::Float, + "bool" => Type::Bool, + "string" => Type::String, + "!" => Type::Never, + "?" => Type::Unknown, + _ => Type::TypeVar(name.clone()), + }, + TypeAnnot::Cons(name, args) => { + let arg_types: Vec = + args.iter().map(|a| self.type_annot_to_type(a)).collect(); + Type::Struct(name.clone(), arg_types) // Assuming it's a struct for now + } + TypeAnnot::Function(params, ret) => { + let param_types = params.iter().map(|p| self.type_annot_to_type(p)).collect(); + let ret_type = self.type_annot_to_type(ret); + Type::Function(param_types, Box::new(ret_type)) + } + TypeAnnot::Tuple(types) => { + let tys = types.iter().map(|t| self.type_annot_to_type(t)).collect(); + Type::Tuple(tys) + } + TypeAnnot::Array(inner) => { + let inner_type = self.type_annot_to_type(inner); + Type::Array(Box::new(inner_type)) + } + } + } + + fn generate_specialized_name(&self, base_name: &str, type_args: &[Type]) -> String { + if type_args.is_empty() { + base_name.to_string() + } else { + let arg_strs: Vec = type_args + .iter() + .map(|t| { + t.to_string() + .replace("<", "_") + .replace(">", "_") + .replace(",", "_") + .replace(" ", "") + }) + .collect(); + format!("{}_{}", base_name, arg_strs.join("_")) + } + } +} + +/// Check that no type variables remain in the AST +pub fn check_no_typevars(nodes: &[TypedASTNode]) -> Result<(), MonomorphizationError> { + for node in nodes { + check_node_for_typevars(node)?; + } + Ok(()) +} + +fn check_node_for_typevars(node: &TypedASTNode) -> Result<(), MonomorphizationError> { + match &node.kind { + TypedASTNodeKind::Function(f) => { + check_function_for_typevars(f)?; + } + TypedASTNodeKind::Struct(s) => { + check_struct_for_typevars(s)?; + } + TypedASTNodeKind::Enum(e) => { + check_enum_for_typevars(e)?; + } + TypedASTNodeKind::Impl(imp) => { + for method in &imp.methods { + check_function_for_typevars(method)?; + } + } + TypedASTNodeKind::Trait(t) => { + // Traits with type parameters are not fully monomorphized + if !t.parameters.is_empty() { + return Err(MonomorphizationError::new( + format!("Trait {} still has type parameters", t.name), + None, + )); + } + // Trait methods are fine as-is - they're abstract signatures + } + _ => {} + } + Ok(()) +} + +fn check_function_for_typevars(func: &TypedFunction) -> Result<(), MonomorphizationError> { + if !func.parameters.is_empty() { + return Err(MonomorphizationError::new( + format!("Function {} still has type parameters", func.name), + None, + )); + } + + for (_, ty_opt) in &func.args { + if let Some(ty) = ty_opt { + if has_typevars_in_type_annot(ty) { + return Err(MonomorphizationError::new( + format!("Function {} argument has type variables", func.name), + None, + )); + } + } + } + + if let Some(ret_ty) = &func.return_type { + if has_typevars_in_type_annot(ret_ty) { + return Err(MonomorphizationError::new( + format!("Function {} return type has type variables", func.name), + None, + )); + } + } + + check_expr_for_typevars(&func.body)?; + Ok(()) +} + +fn check_struct_for_typevars(s: &TypedStruct) -> Result<(), MonomorphizationError> { + if !s.parameters.is_empty() { + return Err(MonomorphizationError::new( + format!("Struct {} still has type parameters", s.name), + None, + )); + } + + for field in &s.fields { + if has_typevars_in_type_annot(&field.field_type) { + return Err(MonomorphizationError::new( + format!("Struct {} field {} has type variables", s.name, field.name), + None, + )); + } + } + Ok(()) +} + +fn check_enum_for_typevars(e: &TypedEnum) -> Result<(), MonomorphizationError> { + if !e.parameters.is_empty() { + return Err(MonomorphizationError::new( + format!("Enum {} still has type parameters", e.name), + None, + )); + } + + for variant in &e.variants { + for field_ty in &variant.fields { + if has_typevars_in_type_annot(field_ty) { + return Err(MonomorphizationError::new( + format!( + "Enum {} variant {} has type variables", + e.name, variant.name + ), + None, + )); + } + } + } + Ok(()) +} + +fn check_expr_for_typevars(expr: &TypedExpr) -> Result<(), MonomorphizationError> { + match &expr.kind { + TypedExprKind::Lambda(params, body) => { + for (_, ty_opt) in params { + if let Some(ty) = ty_opt { + if has_typevars_in_type_annot(ty) { + return Err(MonomorphizationError::new( + "Lambda has type variables in parameters", + Some(expr.span.clone()), + )); + } + } + } + check_expr_for_typevars(body)?; + } + TypedExprKind::Let(_, _, ty_opt, expr) => { + if let Some(ty) = ty_opt { + if has_typevars_in_type_annot(ty) { + return Err(MonomorphizationError::new( + "Let binding has type variables", + Some(expr.span.clone()), + )); + } + } + check_expr_for_typevars(expr)?; + } + TypedExprKind::Cast(e, ty) => { + if has_typevars_in_type_annot(ty) { + return Err(MonomorphizationError::new( + "Cast has type variables", + Some(expr.span.clone()), + )); + } + check_expr_for_typevars(e)?; + } + TypedExprKind::Array(elems) => { + for elem in elems { + check_expr_for_typevars(elem)?; + } + } + TypedExprKind::Tuple(elems) => { + for elem in elems { + check_expr_for_typevars(elem)?; + } + } + TypedExprKind::StructLit(_, fields) => { + for (_, field_expr) in fields { + check_expr_for_typevars(field_expr)?; + } + } + TypedExprKind::EnumLit(_, _, args) => { + for arg in args { + check_expr_for_typevars(arg)?; + } + } + TypedExprKind::Call(func, args) => { + check_expr_for_typevars(func)?; + for arg in args { + check_expr_for_typevars(arg)?; + } + } + TypedExprKind::Index(array, index) => { + check_expr_for_typevars(array)?; + check_expr_for_typevars(index)?; + } + TypedExprKind::Dot(obj, _) => { + check_expr_for_typevars(obj)?; + } + TypedExprKind::EarlyReturn(expr_opt) => { + if let Some(e) = expr_opt { + check_expr_for_typevars(e)?; + } + } + TypedExprKind::OptionalChain(expr_opt, _) => { + if let Some(e) = expr_opt { + check_expr_for_typevars(e)?; + } + } + TypedExprKind::If(cond, then_e, else_e) => { + check_expr_for_typevars(cond)?; + check_expr_for_typevars(then_e)?; + if let Some(e) = else_e { + check_expr_for_typevars(e)?; + } + } + TypedExprKind::Match(scrutinee, arms) => { + check_expr_for_typevars(scrutinee)?; + for (_, arm_expr) in arms { + check_expr_for_typevars(arm_expr)?; + } + } + TypedExprKind::While(cond, body) => { + check_expr_for_typevars(cond)?; + check_expr_for_typevars(body)?; + } + TypedExprKind::Do(exprs) => { + for e in exprs { + check_expr_for_typevars(e)?; + } + } + TypedExprKind::BinOp(lhs, _, rhs) => { + check_expr_for_typevars(lhs)?; + check_expr_for_typevars(rhs)?; + } + TypedExprKind::UnOp(_, operand) => { + check_expr_for_typevars(operand)?; + } + TypedExprKind::For(_, iter, body) => { + check_expr_for_typevars(iter)?; + check_expr_for_typevars(body)?; + } + TypedExprKind::Range(start, end) => { + check_expr_for_typevars(start)?; + check_expr_for_typevars(end)?; + } + TypedExprKind::Return(expr_opt) => { + if let Some(e) = expr_opt { + check_expr_for_typevars(e)?; + } + } + _ => {} + } + Ok(()) +} + +fn has_typevars_in_type_annot(ty: &TypeAnnot) -> bool { + match ty { + TypeAnnot::Var(name) => { + // Check if it's a type variable (not a built-in type) + !matches!( + name.as_str(), + "int" | "float" | "bool" | "string" | "!" | "?" + ) + } + TypeAnnot::Cons(_, args) => args.iter().any(has_typevars_in_type_annot), + TypeAnnot::Function(params, ret) => { + params.iter().any(has_typevars_in_type_annot) || has_typevars_in_type_annot(ret) + } + TypeAnnot::Tuple(types) => types.iter().any(has_typevars_in_type_annot), + TypeAnnot::Array(inner) => has_typevars_in_type_annot(inner), + } +} diff --git a/src/typechecker.rs b/src/typechecker.rs index 4c82a19..118d3d3 100644 --- a/src/typechecker.rs +++ b/src/typechecker.rs @@ -319,9 +319,54 @@ impl TypeChecker { ty, }); } - ASTNodeKind::Struct(_) => Type::Unit, - ASTNodeKind::Enum(_) => Type::Unit, - ASTNodeKind::Trait(_) => Type::Unit, + ASTNodeKind::Struct(s) => { + // Return the typed struct + return Ok(TypedASTNode { + kind: TypedASTNodeKind::Struct(TypedStruct { + name: s.name.clone(), + parameters: s.parameters.clone(), + fields: s.fields.iter().map(|f| TypedField { + name: f.name.clone(), + field_type: f.field_type.clone(), + span: f.span.clone(), + }).collect(), + }), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty: Type::Unit, + }); + } + ASTNodeKind::Enum(e) => { + // Return the typed enum + return Ok(TypedASTNode { + kind: TypedASTNodeKind::Enum(TypedEnum { + name: e.name.clone(), + parameters: e.parameters.clone(), + variants: e.variants.iter().map(|v| TypedVariant { + name: v.name.clone(), + fields: v.fields.clone(), + span: v.span.clone(), + }).collect(), + }), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty: Type::Unit, + }); + } + ASTNodeKind::Trait(t) => { + // Return the typed trait + return Ok(TypedASTNode { + kind: TypedASTNodeKind::Trait(TypedTrait { + name: t.name.clone(), + methods: t.methods.clone(), + parameters: t.parameters.clone(), + associated_types: t.associated_types.clone(), + }), + span: node.span.clone(), + attributes: node.attributes.clone(), + ty: Type::Unit, + }); + } ASTNodeKind::Impl(impl_def) => { let mut typed_methods = Vec::new(); for method in &impl_def.methods { diff --git a/tests/generics_comprehensive.sui b/tests/generics_comprehensive.sui new file mode 100644 index 0000000..0cc0655 --- /dev/null +++ b/tests/generics_comprehensive.sui @@ -0,0 +1,28 @@ +# Generic struct specialization test +struct Box + value: T +end + +# Generic enum specialization test +enum Option + Some(T), + None, +end + +# Generic function specialization test +fn unwrap(opt: Option) -> T + match opt + Option::None() => -1, + Option::Some(v) => v, + end + +# Generic struct with generic enum test +fn test_containers() -> int do + let box_int = Box { value: 42 }; + let box_string = Box { value: "Fermented" }; + let some_int = Option::Some(10); + let some_bool = Option::Some(true); + let unwrapped = unwrap(some_int); + let unwraped_bool = unwrap(some_bool); + box_int.value + unwrapped +end