diff --git a/core/src/eval.rs b/core/src/eval.rs index 19887a5..1b6e064 100644 --- a/core/src/eval.rs +++ b/core/src/eval.rs @@ -99,6 +99,9 @@ pub fn eval(mut main_term: Term, scope: &Scope, options: &EvalOptions) -> Result { let mut term = *case.value.clone(); for (ide, arg) in case.args.clone().into_iter().zip(args) { + if ide.as_str() == "_" { + continue; + } if let Some(arg) = arg { term = substitute(term, &Id(ide), arg); } else { @@ -128,6 +131,13 @@ pub fn eval(mut main_term: Term, scope: &Scope, options: &EvalOptions) -> Result "struct literal reached evaluator without being desugared".to_string(), )); } + Lit { + value: Literal::StructUpdate { .. }, + } => { + return Err(err( + "struct update reached evaluator without being desugared".to_string(), + )); + } _ => break, }; if options.debug { @@ -300,6 +310,17 @@ pub fn apply_dot_macro_recursive(term: Term) -> Term { value: Literal::StructLit { fields }, } } + Lit { + value: Literal::StructUpdate { base, fields }, + } => { + let fields = fields + .into_iter() + .map(|(k, v)| (k, apply_dot_macro_recursive(v))) + .collect(); + Term::Lit { + value: Literal::StructUpdate { base, fields }, + } + } Ctx { loc, term } => { let term = apply_dot_macro_recursive(*term); Ctx { @@ -441,6 +462,17 @@ fn substitute(term: Term, nref: &NameRef, new_term: &Term) -> Term { args, num_args, }) + } else if let Id(_) = nref { + let args = args + .iter() + .map(|a| a.as_ref().map(|a| substitute(a.clone(), nref, new_term))) + .collect(); + Con(Constructor { + name: name.clone(), + typ_name: typ_name.clone(), + args, + num_args, + }) } else { term } @@ -480,6 +512,17 @@ fn substitute(term: Term, nref: &NameRef, new_term: &Term) -> Term { value: Literal::StructLit { fields }, } } + Lit { + value: Literal::StructUpdate { base, fields }, + } => { + let fields = fields + .into_iter() + .map(|(k, v)| (k, substitute(v, nref, new_term))) + .collect(); + Term::Lit { + value: Literal::StructUpdate { base, fields }, + } + } Lit { value: _ } => term, Ctx { loc, term } => Ctx { loc, diff --git a/core/src/eval/type.rs b/core/src/eval/type.rs index 301abf2..f5f7a03 100644 --- a/core/src/eval/type.rs +++ b/core/src/eval/type.rs @@ -1147,22 +1147,32 @@ fn type_check_with_env( .constructors .first() .ok_or_else(|| generic_terr("Structs must have at least one constructor".to_string()))?; - // TODO default values - if fields.len() != mk_cons.params.len() { + if fields.len() > mk_cons.params.len() { return Err(generic_terr(format!( - "too few args in struct {:?} {:?}", + "too many fields in struct {:?} {:?}", fields, mk_cons.params ))); } for Param { name, typ, .. } in mk_cons.params.iter() { - let term = fields.get(name).ok_or_else(|| MissingField(name.clone()))?; - type_check_with_env(term.clone(), *typ.clone(), &scope, usage, track_usage)?; + if let Some(term) = fields.get(name) { + type_check_with_env(term.clone(), *typ.clone(), &scope, usage, track_usage)?; + } else if !ind.defaults.contains_key(name) { + return Err(MissingField(name.clone())); + } } let args: Vec> = mk_cons .params .iter() - .map(|p| Some(fields.get(&p.name).unwrap().clone())) - .collect(); + .map(|p| { + Ok(Some( + fields + .get(&p.name) + .cloned() + .or_else(|| ind.defaults.get(&p.name).cloned()) + .ok_or_else(|| MissingField(p.name.clone()))?, + )) + }) + .collect::, TypeError>>()?; let con = Term::Con(Constructor { name: id("mk"), typ_name: struct_name.clone(), @@ -1174,6 +1184,69 @@ fn type_check_with_env( Ok(typed_term(term, expected_type)) } } + Lit { + value: Literal::StructUpdate { base, fields }, + } => { + // Desugar { id with field := val, ... } to: + // match id { mk orig_fields => mk new_fields } + let base_term = Term::Var { + name: NameRef::Id(base), + }; + let (base, base_type) = + type_check_with_env(base_term, Hole, &scope, usage, track_usage)?.to_tuple(); + if let Some((ind_name, _ind_args)) = extract_first_name(&base_type) { + let ind = scope.find_inductive(&ind_name)?; + let mk_cons = ind + .constructors + .first() + .ok_or_else(|| generic_terr("Struct update requires a struct type".to_string()))?; + // Generate fresh pattern variables + let fresh_names: Map = mk_cons + .params + .iter() + .map(|p| { + let fresh = p.name.rename(); + (p.name.clone(), fresh) + }) + .collect(); + let pat_args: Vec = mk_cons + .params + .iter() + .map(|p| fresh_names.get(&p.name).unwrap().clone()) + .collect(); + // Build body with fresh variable references (will be bound by match pattern) + let mut body_args: Vec> = Vec::new(); + for param in mk_cons.params.iter() { + let fresh = fresh_names.get(¶m.name).unwrap(); + let val: Term = if let Some(override_term) = fields.get(¶m.name) { + override_term.clone() + } else { + Term::Var { + name: NameRef::Id(fresh.clone()), + } + }; + body_args.push(Some(val)); + } + let body = Term::Con(Constructor { + name: id("mk"), + typ_name: ind_name.clone(), + args: body_args, + num_args: mk_cons.params.len(), + }); + let match_case = case(id("mk"), pat_args, body); + let match_expr = match_term(base, vec![match_case]); + return type_check_with_env( + match_expr, + expected_type.clone(), + &scope, + usage, + track_usage, + ); + } + Err(generic_terr(format!( + "Struct update requires an inductive type, found {base_type}" + ))) + } Lit { value: Literal::Match { ref value, @@ -1204,6 +1277,9 @@ fn type_check_with_env( }); } for (name, param) in mcase.args.iter().zip(ind_cons.params.iter()) { + if name.as_str() == "_" { + continue; + } let typ = substitute_params(*param.typ.clone(), ind_params, &ind_args); scope = add_params_to_scope(ind_params, &ind_args, scope); scope = scope.with_type_owned(name, typ); @@ -1219,7 +1295,9 @@ fn type_check_with_env( )?; // Remove pattern variables from usage tracking after each branch for (name, _) in mcase.args.iter().zip(ind_cons.params.iter()) { - usage.remove(name); + if name.as_str() != "_" { + usage.remove(name); + } } if let Ok(typ) = match_resolve_type(&branch_t, t.typ(), &scope) { branch_t = typ; @@ -1420,10 +1498,11 @@ fn type_check_with_env( inductive.params().iter().map(|_| Hole).collect(), ); let lam_types: Vec = arg_res.into_iter().map(|tt| tt.to_tuple().1).collect(); - let cons_type = if !lam_types.is_empty() { - pi_typs(lam_types, ind_type) - } else { + let num_present = args.iter().filter(|a| a.is_some()).count(); + let cons_type = if lam_types.is_empty() || num_present == args.len() { ind_type + } else { + pi_typs(lam_types, ind_type) }; let cons_type = match_resolve_type(&cons_type, &expected_type, &scope)?; Ok(typed_term(term.clone(), cons_type)) diff --git a/core/src/lib.rs b/core/src/lib.rs index 2413847..57fb2de 100644 --- a/core/src/lib.rs +++ b/core/src/lib.rs @@ -250,8 +250,12 @@ pub fn run_tests(input: PathBuf, options: EvalOptions) -> Result<(), String> { let path: ModulePath = input.into(); let mut loaded = default_modules().map_err(|e| format!("{e}"))?; let test_path = ModulePath::new(vec![id("std"), id("test")]); - loaded = load_module_files(&test_path, loaded).map_err(|e| format!("{e}"))?; - loaded = load_module_files(&path, loaded).map_err(|e| format!("{e}"))?; + if loaded.get_module(&test_path).is_none() { + loaded = load_module_files(&test_path, loaded).map_err(|e| format!("{e}"))?; + } + if loaded.get_module(&path).is_none() { + loaded = load_module_files(&path, loaded).map_err(|e| format!("{e}"))?; + } let module = loaded .get_module(&path) .ok_or_else(|| format!("Module {path} not loaded"))?; diff --git a/core/src/parser.rs b/core/src/parser.rs index fb86be5..f660594 100644 --- a/core/src/parser.rs +++ b/core/src/parser.rs @@ -58,7 +58,7 @@ pub fn set_res_extra(res: Res, extra: Y) -> Res(input: Span) -> Res { string_literal, float_literal, num_literal, - struct_val_parser, + struct_or_update_parser, )) .parse(input) } @@ -1232,22 +1232,49 @@ fn struct_val_field_parser(input: Span) -> Res<(Identifier, Term), Ok((input, (name, value))) } -fn struct_val_parser(input: Span) -> Res { - map( - delimited( - (char('{'), ws0), - many0(terminated( - struct_val_field_parser, - (ws0, opt(char(',')), ws0), - )), - (ws0, char('}')), - ), - |fields| Term::Lit { - value: Literal::StructLit { +fn parse_struct_update(input: Span) -> Res { + let (input, id) = identifier(input)?; + let (input, _) = ws0(input)?; + let (input, _) = tag("with")(input)?; + let (input, _) = ws0(input)?; + let (input, fields) = many0(terminated( + struct_val_field_parser, + (ws0, opt(char(',')), ws0), + )) + .parse(input)?; + let (input, _) = ws0(input)?; + let (input, _) = char('}')(input)?; + Ok(( + input, + Term::Lit { + value: Literal::StructUpdate { + base: id, fields: fields.into_iter().collect(), }, }, - ) + )) +} + +fn struct_or_update_parser(input: Span) -> Res { + let (input, _) = char('{')(input)?; + let (input, _) = ws0(input)?; + alt(( + parse_struct_update, + map( + pair( + many0(terminated( + struct_val_field_parser, + (ws0, opt(char(',')), ws0), + )), + preceded(ws0, char('}')), + ), + |(fields, _)| Term::Lit { + value: Literal::StructLit { + fields: fields.into_iter().collect(), + }, + }, + ), + )) .parse(input) } diff --git a/core/src/parser/test.rs b/core/src/parser/test.rs index f8d7a84..ab649c4 100644 --- a/core/src/parser/test.rs +++ b/core/src/parser/test.rs @@ -474,15 +474,15 @@ fn test_let() { #[test] fn test_struct_val() { - let struct_val_parser = |s: &'static str| struct_val_parser::<()>(s.into()); - let (_, r) = struct_val_parser(r#"{}"#.into()).unwrap(); + let p = |s: &'static str| struct_or_update_parser::<()>(s.into()); + let (_, r) = p(r#"{}"#.into()).unwrap(); similar!( r, Term::Lit { value: Literal::StructLit { fields: Map::new() } } ); - let (_, r) = struct_val_parser(r#"{a := b}"#.into()).unwrap(); + let (_, r) = p(r#"{a := b}"#.into()).unwrap(); similar!( r, Term::Lit { @@ -491,7 +491,7 @@ fn test_struct_val() { } } ); - let (_, r) = struct_val_parser(r#"{a := b, b:={c:=0},}"#.into()).unwrap(); + let (_, r) = p(r#"{a := b, b:={c:=0},}"#.into()).unwrap(); similar!( r, Term::Lit { diff --git a/core/src/term.rs b/core/src/term.rs index fabe08f..7df4a95 100644 --- a/core/src/term.rs +++ b/core/src/term.rs @@ -291,13 +291,13 @@ pub fn stru( attributes: Vec, ) -> Inductive { let typ = params_to_inductive_type(¶ms, type0()); + let defaults: Map = fields + .iter() + .filter_map(|d| d.default_value.clone().map(|v| (d.name.clone(), v))) + .collect(); let con_params: Vec = fields .into_iter() - .map(|d| { - param_with_mult( - d.name, d.typ, d.mult, // TODO default value - ) - }) + .map(|d| param_with_mult(d.name, d.typ, d.mult)) .collect(); let struct_type = mpvar(name.clone()); let con_typ = if con_params.is_empty() { @@ -324,6 +324,7 @@ pub fn stru( variant: InductiveVariant::Struct, term, attributes, + defaults, } } @@ -417,6 +418,7 @@ pub fn inductive( constructors, term, attributes, + defaults: Map::new(), } } @@ -479,6 +481,7 @@ pub fn class( typ, term, attributes, + defaults: Map::new(), } } @@ -610,6 +613,7 @@ pub struct Inductive { typ: Term, pub(crate) constructors: Vec, pub attributes: Vec, + pub defaults: Map, } pub trait AsVarRef { @@ -971,6 +975,10 @@ pub enum Literal { StructLit { fields: Map, }, + StructUpdate { + base: Identifier, + fields: Map, + }, } pub fn match_term(value: Term, cases: Vec) -> Term { @@ -1020,6 +1028,14 @@ impl Display for Literal { .join(", "); write!(f, "{{ {fields_str} }}") } + Literal::StructUpdate { base, fields } => { + let fields_str = fields + .iter() + .map(|(k, v)| format!("{k} := {v}")) + .collect::>() + .join(", "); + write!(f, "{{ {base} with {fields_str} }}") + } } } } @@ -1267,6 +1283,7 @@ impl Term { Literal::Match { .. } => "match", Literal::If { .. } => "if", Literal::StructLit { .. } => "struct_lit", + Literal::StructUpdate { .. } => "struct_update", }, Ntv { native: _ } => "ntv", Con(_) => "con", diff --git a/core/src/term/test.rs b/core/src/term/test.rs index bcdcbe4..c89e623 100644 --- a/core/src/term/test.rs +++ b/core/src/term/test.rs @@ -187,7 +187,7 @@ impl Similar for MatchCase { impl Similar for Literal { fn similar(&self, other: &Self) -> bool { - use Literal::{If, Match, StructLit}; + use Literal::{If, Match, StructLit, StructUpdate}; match (self, other) { ( Match { @@ -217,6 +217,22 @@ impl Similar for Literal { .iter() .all(|(k, v)| f2.get(k).map(|v2| v.similar(v2)).unwrap_or(false)) } + ( + StructUpdate { + base: b1, + fields: f1, + }, + StructUpdate { + base: b2, + fields: f2, + }, + ) => { + b1 == b2 + && f1.len() == f2.len() + && f1 + .iter() + .all(|(k, v): (&Identifier, &Term)| f2.get(k).map(|v2| v.similar(v2)).unwrap_or(false)) + } _ => self == other, } } diff --git a/init/tests.mo b/init/tests.mo index 4c9fbaa..892162f 100644 --- a/init/tests.mo +++ b/init/tests.mo @@ -201,6 +201,14 @@ def test_struct_construct_and_match : Bool := mk x y => true } + + +@[test] +def test_eq_in_plain_match : Bool := + match List.cons 5 List.empty { + cons x _ => x == 5 + } + @[test] def test_struct_eq_in_match : Bool := let pt : Point := { x := 1, y := 2 } in @@ -209,7 +217,28 @@ def test_struct_eq_in_match : Bool := } @[test] -def test_eq_in_plain_match : Bool := - match List.cons 5 List.empty { - cons x _ => x == 5 +def test_struct_wildcard : Bool := + let pt : Point := { x := 10, y := 20 } in + match pt { + mk x _ => x == 10 + } + +struct Rect { + w: I64, + h: I64 := 100, +} + +@[test] +def test_struct_default_value : Bool := + let r : Rect := { w := 50 } in + match r { + mk w h => h == 100 + } + +@[test] +def test_struct_update_syntax : Bool := + let p1 : Point := { x := 1, y := 2 } in + let p2 : Point := { p1 with x := 10 } in + match p2 { + mk x y => x == 10 && y == 2 }