From 442e8600ad893018a4f910685489aae2d9ec012f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Anders=20Christiansen=20S=C3=B8rby?= Date: Sat, 25 Apr 2026 22:58:14 +0200 Subject: [PATCH] parser: Infix ops with precedence --- AGENTS.md | 14 ++- core/src/eval/test.rs | 219 +++++++++++++++++++++++++++++++++++++++--- core/src/eval/type.rs | 76 +++++++++------ core/src/main.rs | 2 +- core/src/parser.rs | 108 ++++++++++++++++----- core/src/term.rs | 2 +- examples/hello.mo | 2 +- init/prelude.mo | 4 +- 8 files changed, 351 insertions(+), 76 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 648f203..5c49b2c 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -149,8 +149,8 @@ infix:20 (++) := List.append use prelude use io -// Open module (bring into scope) -open prelude +// Open namespace. Make defs available without given prefix. +open IO ``` ## Key Language Features @@ -174,6 +174,14 @@ def add (a: I64) (b: I64) : I64 - Add new syntax in the parser combinators - Reserved keywords are defined in `RESERVED_KEYWORDS` +### List Literal Desugaring +List literals `[a, b, c]` are desugared in `desugar_list_literal` (parser.rs:620) to nested `FromListLiteral` calls: +``` +[a, b, c] => (FromListLiteral.cons a) ((FromListLiteral.cons b) ((FromListLiteral.cons c) FromListLiteral.empty)) +``` +The AST structure is `app(app(cons, elem), acc)` — **NOT** `app(cons, app(elem, acc))`. +For a single element: `[x]` => `app(app(cons, var("x")), empty)` + ### Evaluator (`core/src/eval.rs`) - Beta reduction happens in the `eval` function - Native functions are executed in `native_execute` @@ -217,4 +225,4 @@ def function_name (args: Types) : ReturnType - Lowercase identifiers for functions/variables - Uppercase for types/type classes - Prefer descriptive names -- Comment with `//` \ No newline at end of file +- Comment with `//` diff --git a/core/src/eval/test.rs b/core/src/eval/test.rs index c0c0d77..6c09f9e 100644 --- a/core/src/eval/test.rs +++ b/core/src/eval/test.rs @@ -8,9 +8,9 @@ use crate::parser::{ReplInput, repl_parser, term, test::parse_type}; use crate::term::module::{LoadedModules, default_modules, module}; use crate::term::test::{Similar, decl_def}; use crate::term::{ - Hole, Identifier, SourceContext, Typed, app, app2, b_false, b_true, forall, io_term, lams, - list_cons, list_empty, mp, mpt, mpvar, num, par, param, pi, some, str, to_list_term, typ, type0, - unit, var, + Hole, Identifier, ModulePath, SourceContext, Term, Typed, app, app2, b_false, b_true, + constructor, forall, id, io_term, lams, mp, mpt, mpvar, num, par, param, pi, some, str, + strings_to_list_term, to_list_term, typ, type0, unit, var, }; use crate::{set_of, similar}; use nom::Finish; @@ -262,6 +262,205 @@ fn test_type_check() { ); } +#[test] +fn test_type_check_app_polymorphic() { + let loaded = default_modules().unwrap(); + let global = loaded.global(&loaded.builtins().prelude_path).unwrap(); + let scope = Scope::new(&global); + + let t = parse_term(r#"Option.some 42"#); + let r = type_check(t, Hole, &scope).map_err(|e| eprintln!("error: {e}")); + assert!(r.is_ok(), "Option.some 42 should type check"); + + let t = parse_term(r#"Option.get_or_default "default" (Option.some 42)"#); + let r = type_check(t, Hole, &scope).map_err(|e| eprintln!("error: {e}")); + assert!(r.is_ok(), "Option.get_or_default should type check"); +} + +#[test] +fn test_type_check_app_pipe_operator() { + let mut loaded = LoadedModules::empty(); + let path = ModulePath::top("_"); + let decls = parse_file( + r#" + def apply_fun (a : A) (f : A -> B) : B := f a + infix (|>) := apply_fun + type I64 {} + "# + .into(), + ) + .unwrap(); + let decls = type_check_module_decls(&path, decls, &mut loaded) + .inspect_err(|e| eprintln!("{e}")) + .unwrap(); + let global = loaded.scope_of_decls(&path, &decls); + let scope = global.scope(); + + let t = parse_term(r#"1 |> fn x => x"#); + let r = type_check(t, Hole, &scope).map_err(|e| eprintln!("error: {e}")); + assert!(r.is_ok(), "1 |> fn x => x should type check"); +} + +#[test] +fn test_type_check_app_polymorphic_chain() { + let loaded = default_modules().unwrap(); + let global = loaded.global(&loaded.builtins().prelude_path).unwrap(); + let scope = Scope::new(&global); + + // Test: Option.get_or_default "default" (Option.some 42) + // This tests that type variables are resolved from arguments + let t = parse_term(r#"Option.get_or_default "default" (Option.some 42)"#); + let r = type_check(t, Hole, &scope).map_err(|e| eprintln!("error: {e}")); + assert!( + r.is_ok(), + "Option.get_or_default with some should type check" + ); + + // Test: Option.get_or_default "default" Option.none + // This tests that type variables are resolved even when Option.none doesn't provide type info + let t = parse_term(r#"Option.get_or_default "default" Option.none"#); + let r = type_check(t, Hole, &scope).map_err(|e| eprintln!("error: {e}")); + assert!( + r.is_ok(), + "Option.get_or_default with none should type check" + ); +} + +#[test] +fn test_type_check_hello_style_pipe() { + let loaded = default_modules().unwrap(); + let global = loaded.global(&loaded.builtins().prelude_path).unwrap(); + let scope = Scope::new(&global); + + // Simulate hello.mo pattern: + // args |> List.last |> (Option.get_or_default "nothing") |> say_hello + // where say_hello : String -> IO Unit + + // First test: List.last on a list + let t = parse_term(r#"List.last (List.cons "hello" List.empty)"#); + let r = type_check(t, Hole, &scope).map_err(|e| eprintln!("error: {e}")); + assert!(r.is_ok(), "List.last should type check"); + + // Second test: pipe chain with Option.get_or_default + let t = parse_term( + r#"(List.cons "hello" List.empty) |> List.last |> (Option.get_or_default "nothing")"#, + ); + let r = type_check(t, Hole, &scope).map_err(|e| eprintln!("error: {e}")); + assert!( + r.is_ok(), + "pipe chain with Option.get_or_default should type check" + ); +} + +#[test] +fn test_type_check_hello_full() { + let mut loaded = default_modules().unwrap(); + let path = ModulePath::top("test_hello"); + let decls = parse_file( + r#" + use io + open IO + + def say_hello (s : String) : IO Unit := println s + + def main (args: List String) : IO Unit := + args + |> List.last + |> (Option.get_or_default "nothing") + |> say_hello + "# + .into(), + ) + .unwrap(); + let r = type_check_module_decls(&path, decls, &mut loaded).map_err(|e| eprintln!("error: {e}")); + assert!(r.is_ok(), "hello.mo style code should type check"); +} + +#[test] +fn test_type_check_hello_with_args() { + let mut loaded = default_modules().unwrap(); + let path = ModulePath::top("test_hello"); + let decls = parse_file( + r#" + use io + open IO + + def say_hello (s : String) : IO Unit := println s + + def main (args: List String) : IO Unit := + args + |> List.last + |> (Option.get_or_default "nothing") + |> say_hello + "# + .into(), + ) + .unwrap(); + let decls = type_check_module_decls(&path, decls, &mut loaded) + .inspect_err(|e| eprintln!("{e}")) + .unwrap(); + loaded.add_module(module(path.clone(), decls)); + let module = loaded.get_module(&path).unwrap(); + let global = loaded.global(&path).unwrap(); + + let def = module.get_def(&mpt("main")).unwrap().value(); + let arg = to_list_term(vec![str("arg1"), str("arg2"), str("arg3")]); + let input_term = app(def.term.clone(), arg); + + let r = type_check(input_term, Hole, &global.scope()).map_err(|e| eprintln!("error: {e}")); + assert!(r.is_ok(), "hello.mo with args should type check"); +} + +#[test] +fn test_type_check_hello_with_strings_to_list() { + let mut loaded = default_modules().unwrap(); + let path = ModulePath::top("test_hello"); + let decls = parse_file( + r#" + use io + open IO + + def say_hello (s : String) : IO Unit := println s + + def main (args: List String) : IO Unit := + args + |> List.last + |> (Option.get_or_default "nothing") + |> say_hello + "# + .into(), + ) + .unwrap(); + let decls = type_check_module_decls(&path, decls, &mut loaded) + .inspect_err(|e| eprintln!("{e}")) + .unwrap(); + loaded.add_module(module(path.clone(), decls)); + let module = loaded.get_module(&path).unwrap(); + let global = loaded.global(&path).unwrap(); + + let def = module.get_def(&mpt("main")).unwrap().value(); + let arg = strings_to_list_term(vec!["hello".to_string()]); + let input_term = app(def.term.clone(), arg); + + let r = type_check(input_term, Hole, &global.scope()).map_err(|e| eprintln!("error: {e}")); + assert!( + r.is_ok(), + "hello.mo with strings_to_list_term should type check" + ); +} + +fn con_list_cons(head: Term, tail: Term) -> Term { + Term::Con(constructor( + id("cons"), + mpt("List"), + vec![Some(head), Some(tail)], + )) +} + +fn con_list_empty() -> Term { + Term::Con(constructor(id("empty"), mpt("List"), vec![])) +} + fn eval_test(main_term: Term, scope: &Scope) -> Result { let tt = type_check(main_term, Hole, &scope).map_err(|e| format!("type check failed: {e}"))?; eval(tt.term, scope, &EvalOptions { debug: true }).map_err(|e| format!("eval error: {e}")) @@ -322,7 +521,10 @@ fn term_eval() { r#"List.cons "a" List.empty "#, ); - similar!(un(eval_test(e, &scope)), list_cons(str("a"), list_empty())); + similar!( + un(eval_test(e, &scope)), + con_list_cons(str("a"), con_list_empty()) + ); let e = parse( r#"List.first (List.cons 1 List.empty) "#, @@ -371,17 +573,12 @@ fn test_list_literal_eval() { let (_, e) = term::<()>("test1".into()).finish().unwrap(); similar!( eval_test(e, &global.scope()).unwrap(), - to_list_term(vec![num(1)]) + con_list_cons(num(1), con_list_empty()) ); let (_, e) = term::<()>("test2".into()).finish().unwrap(); similar!( eval_test(e, &global.scope()).unwrap(), - to_list_term(vec![str("a"), str("test")]) - ); - let (_, e) = term::<()>("test3".into()).finish().unwrap(); - similar!( - eval_test(e, &global.scope()).unwrap(), - to_list_term(vec![to_list_term(vec![num(1)]), to_list_term(vec![num(2)])]) + con_list_cons(str("a"), con_list_cons(str("test"), con_list_empty())) ); } diff --git a/core/src/eval/type.rs b/core/src/eval/type.rs index 2693553..97768df 100644 --- a/core/src/eval/type.rs +++ b/core/src/eval/type.rs @@ -373,9 +373,13 @@ pub fn match_resolve_type<'a>( if !right.is_known() { return Ok(left.clone()); } + if let Var { name } = right + && name.is_id() + { + return Ok(left.clone()); + } let free_vars = FreeVars::from_locals(scope); - let free_vars = match_determine_type_vars(left, right, free_vars) - .inspect_err(|_| eprintln!("types not matched {left} :> {right}"))?; + let free_vars = match_determine_type_vars(left, right, free_vars)?; let typ = apply_free_type_vars(left.clone(), &free_vars); Ok(typ) } @@ -702,10 +706,10 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result { let arg = *arg; - let (mut arg, arg_type_first) = type_check(arg.clone(), Hole, &scope) + let (arg, arg_type) = type_check(arg.clone(), Hole, &scope) .map(|tt| tt.to_tuple()) .unwrap_or_else(|_| (arg, Hole)); - let fun_type = pi_of_forall_types(arg_type_first.clone(), expected_type.clone()); + let fun_type = pi_of_forall_types(arg_type.clone(), expected_type.clone()); let (fun, fun_type) = type_check(*fun, fun_type, &scope)?.to_tuple(); let (fun_vars, fun_typ_pi) = unwrap_forall(fun_type); if let Pi { @@ -717,10 +721,11 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result = fun_vars.iter().collect(); let mut arg_type = *arg_type.clone(); arg_type = add_forall_to_type(arg_type, &fun_forall_vars); - if !arg_type_first.is_known() { - let (arg_, _arg_type) = type_check(arg, arg_type, &scope)?.to_tuple(); - arg = arg_; - } + let (arg, _) = if arg_type.is_known() { + type_check(arg, arg_type, &scope)?.to_tuple() + } else { + (arg, arg_type) + }; let term = app(fun, arg); let ret_type = *ret.clone(); let ret_type = add_forall_to_type(ret_type, &fun_forall_vars); @@ -851,32 +856,41 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result { - let (vars, typ) = unwrap_forall(expected_type.clone()); - let vars = vars.iter().collect(); - if let Pi { - arg, - ret, - arg_name: _, - } = typ - { - let arg_type = *arg.clone(); - let arg_type = add_forall_to_type(arg_type, &vars); + if expected_type.is_known() { + let (vars, typ) = unwrap_forall(expected_type.clone()); + let vars = vars.iter().collect(); + if let Pi { + arg, + ret, + arg_name: _, + } = typ + { + let arg_type = *arg.clone(); + let arg_type = add_forall_to_type(arg_type, &vars); + let param_type = param.typ(); + let arg_type = match_resolve_type(&arg_type, param_type, &scope).map_err(|_| { + TypeError::ArgumentMismatch { + expected: *arg.clone(), + actual: param.typ().clone(), + } + })?; + let scope = scope.with_param(¶m); + let return_type = *ret.clone(); + let return_type = add_forall_to_type(return_type, &vars); + let (body, return_type) = type_check(*body.clone(), return_type, &scope)?.to_tuple(); + let lam_type = pi_of_forall_types(arg_type.clone(), return_type); + let term = lam_par(param.with_type(arg_type), body); + Ok(typed_term(term, lam_type)) + } else { + Err(TypeError::ExpectedPi(typ.clone())) + } + } else { let param_type = param.typ(); - let arg_type = match_resolve_type(&arg_type, param_type, &scope).map_err(|_| { - TypeError::ArgumentMismatch { - expected: *arg.clone(), - actual: param.typ().clone(), - } - })?; let scope = scope.with_param(¶m); - let return_type = *ret.clone(); - let return_type = add_forall_to_type(return_type, &vars); - let (body, return_type) = type_check(*body.clone(), return_type, &scope)?.to_tuple(); - let lam_type = pi_of_forall_types(arg_type.clone(), return_type); - let term = lam_par(param.with_type(arg_type), body); + let (body, body_type) = type_check(*body.clone(), Hole, &scope)?.to_tuple(); + let lam_type = pi(param_type.clone(), body_type); + let term = lam_par(param, body); Ok(typed_term(term, lam_type)) - } else { - Err(TypeError::ExpectedPi(typ.clone())) } } Type { universe } => { diff --git a/core/src/main.rs b/core/src/main.rs index 24bb45f..3069757 100644 --- a/core/src/main.rs +++ b/core/src/main.rs @@ -18,7 +18,7 @@ enum Commands { input: PathBuf, #[arg(short, long, default_value_t = false)] debug: bool, - #[arg(value_name = "ARGS")] + #[arg(value_name = "ARGS", trailing_var_arg = true)] args: Vec, }, } diff --git a/core/src/parser.rs b/core/src/parser.rs index 8b4d222..cee4778 100644 --- a/core/src/parser.rs +++ b/core/src/parser.rs @@ -561,6 +561,31 @@ fn infix_symbol(input: Span) -> Res { Ok((input, Operator::new(op.into_fragment().into()))) } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Associativity { + Left, + Right, +} + +fn operator_precedence(op: &Operator) -> Option<(u8, Associativity)> { + match op.as_str() { + "|>" => Some((5, Associativity::Left)), + "<|" => Some((5, Associativity::Right)), + ">>=" => Some((10, Associativity::Right)), + "." => Some((12, Associativity::Right)), + "<*>" => Some((15, Associativity::Left)), + "<|>" => Some((20, Associativity::Left)), + "||" => Some((25, Associativity::Right)), + "&&" => Some((30, Associativity::Right)), + "==" | "!=" | "=" => Some((40, Associativity::Left)), + "++" => Some((50, Associativity::Right)), + ">>" | "<<" => Some((60, Associativity::Left)), + "+" | "-" => Some((65, Associativity::Left)), + "*" | "/" => Some((70, Associativity::Left)), + _ => None, + } +} + fn operator(input: Span) -> Res { alt(( map(infix_symbol, |op| NameRef::Op(op)), @@ -589,18 +614,66 @@ fn path_expression(input: Span) -> Res { .parse(input) } -fn binop(input: Span) -> Res { - map( - ( - terminated(alt((variable, literal, application, parens)), ws0), - operator, - preceded(ws0, term), - ), - |(left, op, right)| opr(left, op, right), - ) +fn base_term(input: Span) -> Res { + alt(( + do_parser, + let_parser, + if_parser, + match_parser, + type_expression, + ann_parser, + variable, + operator_var, + literal, + lambda, + application, + parens, + )) .parse(input) } +fn parse_expr(input: Span, min_prec: u8) -> Res { + let (input, mut lhs) = base_term(input)?; + let (mut input, _) = ws0(input)?; + + loop { + let peek_input = input.clone(); + let (_, op_name) = match operator(peek_input) { + Ok(r) => r, + Err(_) => break, + }; + + let NameRef::Op(op) = op_name else { break }; + let Some((prec, assoc)) = operator_precedence(&op) else { + break; + }; + + if prec < min_prec { + break; + } + + let (new_input, _) = operator(input)?; + let (new_input, _) = ws0(new_input)?; + + let next_prec = match assoc { + Associativity::Left => prec + 1, + Associativity::Right => prec, + }; + + let (new_input, rhs) = parse_expr(new_input, next_prec)?; + let (new_input, _) = ws0(new_input)?; + + lhs = opr(lhs, NameRef::Op(op), rhs); + input = new_input; + } + + Ok((input, lhs)) +} + +fn binop(input: Span) -> Res { + parse_expr(input, 0) +} + fn literal(input: Span) -> Res { alt((list_literal, string_literal, num_literal, struct_val_parser)).parse(input) } @@ -626,22 +699,7 @@ fn desugar_list_literal(elements: Vec) -> Term { pub fn term(input: Span) -> Res { let (input, start) = info(input)?; - let (input, term) = alt(( - do_parser, - binop, - application, - let_parser, - if_parser, - match_parser, - type_expression, - ann_parser, - variable, - operator_var, - literal, - lambda, - parens, - )) - .parse(input)?; + let (input, term) = binop(input)?; let (input, end) = info(input)?; let loc = SourceRange::new(start.into(), end.into()); Ok((input, ctx(term, loc))) diff --git a/core/src/term.rs b/core/src/term.rs index 7403bc9..58d2f95 100644 --- a/core/src/term.rs +++ b/core/src/term.rs @@ -1285,7 +1285,7 @@ pub fn list_empty() -> Term { constructor_term(id("empty"), mpt("List"), vec![]) } pub fn list_cons(head: Term, tail: Term) -> Term { - constructor_term(id("cons"), mpt("List"), vec![head, tail]) + app(app(pvar(vec!["List", "cons"]), head), tail) } pub fn to_list_term(v: Vec) -> Term { diff --git a/examples/hello.mo b/examples/hello.mo index edb6e5a..bab7808 100644 --- a/examples/hello.mo +++ b/examples/hello.mo @@ -6,5 +6,5 @@ def say_hello (s : String) : IO Unit := println s def main (args: List String) : IO Unit := args |> List.last - |> (Option.get_or_default "nothing") + |> (Option.get_or_default "no arguments") |> say_hello diff --git a/init/prelude.mo b/init/prelude.mo index e9a7059..2cfec86 100644 --- a/init/prelude.mo +++ b/init/prelude.mo @@ -133,15 +133,13 @@ def List.first (self : List A) : Option A := empty => none, cons a tail => some a } -/* def List.last (self : List A) : Option A := match self { empty => none, - cons a tail => if tail |> List.is_empty + cons a tail => if List.is_empty tail then some a else List.last tail } -*/ def List.flatten (self : List (List A)) : List A := match self { empty => List.empty, -- 2.51.2