diff --git a/.gitignore b/.gitignore index 5b217c8..b8815fb 100644 --- a/.gitignore +++ b/.gitignore @@ -12,3 +12,6 @@ devenv.local.yaml # pre-commit .pre-commit-config.yaml docs/book + +# Plans directory for implementation planning +plans/ diff --git a/core/src/eval.rs b/core/src/eval.rs index 153bc2f..068a365 100644 --- a/core/src/eval.rs +++ b/core/src/eval.rs @@ -151,7 +151,7 @@ pub fn recognize_bool(value: &Term) -> Result { fn substitute_lam(param: Par, body: Term, arg: &Term) -> Term { match param { Par::P(param) => substitute(body, &Id(param.name.clone()), arg), - Par::I { typ: _ } => substitute(body, &Index(1), arg), + Par::I { typ: _, .. } => substitute(body, &Index(1), arg), } } fn native_apply_arg(mut native: Native, index: usize, term: Term) -> Option { @@ -264,6 +264,7 @@ fn substitute(term: Term, nref: &NameRef, new_term: &Term) -> Term { let new_param = Param { name: new_name.clone(), typ: param.typ.clone(), + mult: param.mult.clone(), }; let new_body = rename_variable(*body, new_name, old_name.clone()); let term = substitute(new_body, nref, new_term); @@ -273,7 +274,7 @@ fn substitute(term: Term, nref: &NameRef, new_term: &Term) -> Term { lam(param.clone(), term) } } - Par::I { typ } => { + Par::I { typ, .. } => { if let Index(i) = nref { let term = substitute(*body, &Index(i + 1), new_term); lam_index(*typ, term) @@ -290,7 +291,9 @@ fn substitute(term: Term, nref: &NameRef, new_term: &Term) -> Term { term.clone() } } - Pi { arg, ret, arg_name } => { + Pi { + arg, ret, arg_name, .. + } => { let arg = substitute(*arg, nref, new_term); let ret = substitute(*ret, nref, new_term); pi_name(arg_name, arg, ret) diff --git a/core/src/eval/type.rs b/core/src/eval/type.rs index 1564024..553b340 100644 --- a/core/src/eval/type.rs +++ b/core/src/eval/type.rs @@ -4,7 +4,7 @@ use crate::{ Map, Set, empty_set, set_of, term::{ Ann, ClassDefRef, Decl, Def, Identifier, Inductive, InductiveVariant, Instance, InstanceKey, - Literal, ModulePath, NameRef, Named, NumSuffix, SourceContext, + Literal, ModulePath, Multiplicity, NameRef, Named, NumSuffix, SourceContext, Term::{Forall, Hole, Pi}, TypeConstraint, Typed, TypedTerm, VarRef, app, bvar, ctx, forall, lam_par, module::{LoadedModules, names_of_decls}, @@ -68,6 +68,10 @@ pub enum TypeError { target: &'static str, }, Many(Vec), + // Linear type errors + LinearUsedMultipleTimes(Identifier), + LinearUnused(Identifier), + AffineUsedMultipleTimes(Identifier), } impl From for TypeError { @@ -139,6 +143,15 @@ impl Display for TypeError { TypeError::Overflow { value, target } => { write!(f, "Integer overflow: {value} does not fit in {target}") } + TypeError::LinearUsedMultipleTimes(id) => { + write!(f, "Linear variable '{}' used more than once", id) + } + TypeError::LinearUnused(id) => { + write!(f, "Linear variable '{}' must be used exactly once", id) + } + TypeError::AffineUsedMultipleTimes(id) => { + write!(f, "Affine variable '{}' used more than once", id) + } } } } @@ -182,6 +195,62 @@ impl From for TypeError { } } +/// Tracks variable usage counts for linear type checking (compile-time only) +#[derive(Debug, Clone)] +pub struct UsageEnv { + usages: Map, +} + +impl UsageEnv { + pub fn new() -> Self { + UsageEnv { usages: Map::new() } + } + + /// Register a new variable with its multiplicity + pub fn register(&mut self, name: Identifier, mult: Multiplicity) { + self.usages.insert(name, (mult, 0)); + } + + /// Check if a variable can be used (based on its multiplicity) + pub fn check_usage(&self, name: &Identifier) -> Result<(), TypeError> { + if let Some((mult, count)) = self.usages.get(name) { + match mult { + Multiplicity::Linear => { + if *count >= 1 { + return Err(TypeError::LinearUsedMultipleTimes(name.clone())); + } + } + Multiplicity::Affine => { + if *count >= 1 { + return Err(TypeError::AffineUsedMultipleTimes(name.clone())); + } + } + Multiplicity::Many => { + // Always ok + } + } + } + Ok(()) + } + + /// Mark a variable as used (increment usage count) + pub fn mark_used(&mut self, name: &Identifier) { + if let Some((_, count)) = self.usages.get_mut(name) { + *count += 1; + } + } + + /// Verify all linear variables were used exactly once + pub fn verify_linear_usage(&self) -> Result<(), TypeError> { + for (name, (mult, count)) in &self.usages { + if *mult == Multiplicity::Linear && *count != 1 { + return Err(TypeError::LinearUnused(name.clone())); + } + } + Ok(()) + } +} + pub fn derive_instance_key(class_def: &ClassDefRef, typ: &Term) -> Result { use FreeVar::*; use InstanceError::*; @@ -275,7 +344,7 @@ pub fn type_check_instance<'a>( } for (param, arg) in class.params.iter().zip(instance.args.iter()) { - type_check(arg.clone(), *param.typ.clone(), &scope)?; + type_check_with_env(arg.clone(), *param.typ.clone(), &scope)?; } let cons = class @@ -287,7 +356,7 @@ pub fn type_check_instance<'a>( if let Some(impl_def) = instance.impls_map.get_mut(¶m.name) { let class_def_type = param.typ(); let typ = match_resolve_type(class_def_type, &impl_def.typ, &scope)?; - let (term, _) = type_check(impl_def.term.clone(), typ.clone(), &scope)?.to_tuple(); + let (term, _) = type_check_with_env(impl_def.term.clone(), typ.clone(), &scope)?.to_tuple(); impl_def.term = term; // Wrap type with forall bindings for type variables impl_def.typ = wrap_with_foralls(typ, &type_vars); @@ -567,11 +636,13 @@ fn match_resolve_type_inner<'a>( arg: a_arg, ret: a_ret, arg_name, + .. }, Pi { arg: b_arg, ret: b_ret, arg_name: _, + .. }, ) => { if let Some(name) = arg_name { @@ -758,7 +829,9 @@ fn extract_first_name(term: &Term) -> Option<(ModulePath, Vec)> { /// Find unknown identifiers in a type pub fn free_vars(typ: &Term, known_names: &Set<&ModulePath>) -> Set { match typ { - Pi { arg, ret, arg_name } => { + Pi { + arg, ret, arg_name, .. + } => { let mut a = free_vars(arg, known_names); if let Some(name) = arg_name { let mut known_names = known_names.clone(); @@ -950,9 +1023,20 @@ fn convert_int_literal(value: i64, suffix: NumSuffix) -> Result /// Check and compute the Type of a Term /// Resolves type classes +/// Public wrapper that creates a fresh UsageEnv if not present pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result { + type_check_with_env(term, expected_type, scope) +} + +/// Internal type check with UsageEnv for linear type tracking +/// Uses the UsageEnv stored in the Scope +fn type_check_with_env( + term: Term, + expected_type: Term, + scope: &Scope, +) -> Result { use TypeError::*; - let scope = add_forall_to_scope(&expected_type, scope.clone()); + let mut scope = add_forall_to_scope(&expected_type, scope.clone()); match term { App { fun, arg } => { // Try desugaring method calls with args (x.fun arg -> A.fun arg x) @@ -966,23 +1050,24 @@ 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); let (arg, _) = if arg_type.is_known() { - type_check(arg, arg_type, &scope)?.to_tuple() + type_check_with_env(arg, arg_type, &scope)?.to_tuple() } else { (arg, arg_type) }; @@ -1013,9 +1098,9 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result Result { - let con = type_check(*value.clone(), Hole, &scope)?; + let con = type_check_with_env(*value.clone(), Hole, &scope)?; if let Some((ind_name, ind_args)) = extract_first_name(con.typ()) { let ind = scope.find_inductive(&ind_name)?; @@ -1057,7 +1142,7 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result Result { - let b = type_check(*value, var("Bool"), &scope)?; - let t1 = type_check(*then, expected_type.clone(), &scope)?; - let t2 = type_check(*els, expected_type.clone(), &scope)?; + let b = type_check_with_env(*value, var("Bool"), &scope)?; + let t1 = type_check_with_env(*then, expected_type.clone(), &scope)?; + let t2 = type_check_with_env(*els, expected_type.clone(), &scope)?; if let Ok(typ) = match_resolve_type(t1.typ(), t2.typ(), &scope) { let new_term = Lit { value: Literal::If { @@ -1138,9 +1223,23 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result { + // Register parameter in usage environment + let mult = param.multiplicity(); + let param_name = match ¶m { + Par::P(p) => Some(p.name.clone()), + Par::I { .. } => None, // Anonymous implicit param + }; + if let Some(ref name) = param_name { + scope.usage_env_mut().register(name.clone(), mult.clone()); + } if expected_type.is_known() { let (vars, typ) = unwrap_forall(expected_type.clone()); let vars = vars.iter().collect(); @@ -1148,6 +1247,7 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result Result Result Result Result { - let mut tt = type_check(*term, expected_type.clone(), &scope) + let mut tt = type_check_with_env(*term, expected_type.clone(), &scope) .map_err(|err| t_context(err, None, loc.clone()))?; *tt.mut_term() = ctx(tt.term().clone(), loc.clone()); @@ -1242,10 +1347,11 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result { if expected_type.is_type() { - let _arg = type_check(*arg.clone(), type0(), &scope)?; - let _ret = type_check(*ret.clone(), type0(), &scope)?; + let _arg = type_check_with_env(*arg.clone(), type0(), &scope)?; + let _ret = type_check_with_env(*ret.clone(), type0(), &scope)?; Ok(typed_term(term.clone(), type0())) } else { Err(TypeError::ExpectedType(expected_type.clone())) @@ -1254,7 +1360,7 @@ pub fn type_check(term: Term, expected_type: Term, scope: &Scope) -> Result Ok(typed_term(term, type0())), Hole => Ok(typed_term(term, expected_type)), Ann { term, typ } => { - let tt = type_check(*term, *typ, &scope)?; + let tt = type_check_with_env(*term, *typ, &scope)?; let typ = match_resolve_type(tt.typ(), &expected_type, &scope)?; Ok(typed_term(tt.term, typ)) } @@ -1372,6 +1478,7 @@ pub fn pi_to_vec(mut typ: Term) -> (Vec, Term) { arg, ret, arg_name: _, + .. } = typ { res.push(*arg); diff --git a/core/src/parser.rs b/core/src/parser.rs index c75a46d..0576775 100644 --- a/core/src/parser.rs +++ b/core/src/parser.rs @@ -10,15 +10,15 @@ use crate::{ parser::error::{ParseError, ReplParserError}, term::{ AttrArg, Attribute, ClassDef, Decl, Def, Documentation, Identifier, InductConstructor, - Inductive, Infix, Instance, LetVar, Literal, MatchCase, ModulePath, NameRef, NumSuffix, Open, - Operator, Param, SourceContext, SourceRange, StructField, + Inductive, Infix, Instance, LetVar, Literal, MatchCase, ModulePath, Multiplicity, NameRef, + NumSuffix, Open, Operator, Param, SourceContext, SourceRange, StructField, Term::{self, Hole, Var}, TypeConstraint, Use, app, apps, case, class, class_def, ctx, def, def_with_native, float_suffix, forall, foralls, id, if_term, induct_constructor, inductive, infix, instance, - lam, lams, lets, map_term, match_term, + ivar, lam, lams, lets, map_term, match_term, module::ParsedModule, - mpvar, num_suffix, opr, param, pi_name, pi_typs, pvar, stru, stru_field, ty, type_constraint, - var_id, + mpvar, num_suffix, opr, param, param_with_mult, pi_name, pi_typs, pvar, stru, stru_field, + type_constraint, var_id, }, }; use locate::{LocatedSpan, info}; @@ -31,7 +31,7 @@ use nom::{ }, combinator::{eof, map, not, opt, recognize, success, verify}, multi::{fold_many0, many0, many1}, - sequence::{delimited, preceded, separated_pair, terminated}, + sequence::{delimited, pair, preceded, separated_pair, terminated}, }; use string::parse_string; @@ -297,14 +297,27 @@ fn float_suffix_parser(input: Span) -> Res { Ok((input, suffix)) } +/// Parse multiplicity prefix: "!" = Linear, "?" = Affine, none = Many +fn multiplicity_prefix(input: Span) -> Res { + map(opt(alt((char('!'), char('?')))), |m| match m { + Some('!') => Multiplicity::Linear, + Some('?') => Multiplicity::Affine, + _ => Multiplicity::Many, + }) + .parse(input) +} + fn lam_param(input: Span) -> Res { alt(( map(identifier, |i| param(i, Hole)), delimited( (char('('), ws0), map( - separated_pair(identifier, ws0, opt_type_annotation), - |(name, typ)| param(name, typ), + pair( + multiplicity_prefix, + separated_pair(identifier, ws0, opt_type_annotation), + ), + |(mult, (name, typ))| param_with_mult(name, typ, mult), ), (ws0, char(')')), ), @@ -314,7 +327,7 @@ fn lam_param(input: Span) -> Res { fn cons_param(input: Span) -> Res, X> { alt(( - map(identifier, |t| vec![param(id(""), ty(t))]), + map(identifier, |t| vec![param(id(""), ivar(t))]), delimited( (char('('), ws0), alt(( @@ -345,8 +358,16 @@ fn implicit_param(input: Span) -> Res, X> { delimited( (char('{'), ws0), map( - separated_pair(many1(terminated(identifier, ws0)), ws0, type_annotation), - |(ids, typ)| ids.into_iter().map(|i| param(i, typ.clone())).collect(), + pair( + multiplicity_prefix, + separated_pair(many1(terminated(identifier, ws0)), ws0, type_annotation), + ), + |(mult, (ids, typ))| { + ids + .into_iter() + .map(|i| param_with_mult(i, typ.clone(), mult.clone())) + .collect() + }, ), (ws0, char('}')), ) @@ -369,8 +390,16 @@ fn def_param(input: Span) -> Res, X> { delimited( (char('('), ws0), map( - separated_pair(many1(terminated(identifier, ws0)), ws0, type_annotation), - |(ids, typ)| ids.into_iter().map(|i| param(i, typ.clone())).collect(), + pair( + multiplicity_prefix, + separated_pair(many1(terminated(identifier, ws0)), ws0, type_annotation), + ), + |(mult, (ids, typ))| { + ids + .into_iter() + .map(|i| param_with_mult(i, typ.clone(), mult.clone())) + .collect() + }, ), (ws0, char(')')), ) @@ -876,12 +905,7 @@ fn opt_attributes(input: Span) -> Res, X> { many0(attribute_parser).parse(input) } -#[cfg(test)] fn def_parser(input: Span) -> Res { - def_with_attrs_parser(input) -} - -fn def_with_attrs_parser(input: Span) -> Res { let (input, attrs) = opt_attributes(input)?; let (input, _) = ws0(input)?; let (input, _) = tag("def")(input)?; @@ -1045,7 +1069,7 @@ fn class_parser(input: Span) -> Res { fn instance_inner_parser(input: Span) -> Res> { delimited( (char('{'), ws0), - many1(delimited(ws0, def_with_attrs_parser, ws0)), + many1(delimited(ws0, def_parser, ws0)), (ws0, char('}')), ) .parse(input) @@ -1263,7 +1287,7 @@ fn decl_parser(input: Span) -> Res> { let (input, decl) = alt(( map(use_parser, Decl::Use), map(open_parser, Decl::Open), - map(def_with_attrs_parser, Decl::Def), + map(def_parser, Decl::Def), map(class_parser, Decl::Type), map(instance_parser, Decl::Ins), map(struct_parser, Decl::Type), diff --git a/core/src/parser/test.rs b/core/src/parser/test.rs index b803df0..64b9828 100644 --- a/core/src/parser/test.rs +++ b/core/src/parser/test.rs @@ -698,10 +698,12 @@ fn test_native() { let expected_term = Term::Lam { param: Par::I { typ: Box::new(typ("I64")), + mult: Multiplicity::Many, }, body: Box::new(Term::Lam { param: Par::I { typ: Box::new(typ("I64")), + mult: Multiplicity::Many, }, body: Box::new(native_term), }), @@ -1278,10 +1280,12 @@ fn test_native_with_named_arg() { let expected_term = Term::Lam { param: Par::I { typ: Box::new(typ("I64")), + mult: Multiplicity::Many, }, body: Box::new(Term::Lam { param: Par::I { typ: Box::new(typ("I64")), + mult: Multiplicity::Many, }, body: Box::new(native_term), }), diff --git a/core/src/term.rs b/core/src/term.rs index 9208e50..810c19b 100644 --- a/core/src/term.rs +++ b/core/src/term.rs @@ -191,6 +191,7 @@ pub fn pi_var(name: Identifier, arg: Term, ret: Term) -> Term { arg_name: Some(name), arg: Box::new(arg), ret: Box::new(ret), + mult: Multiplicity::default(), } } @@ -201,6 +202,7 @@ pub fn type_count_args(typ: &Term) -> u8 { arg: _, ret, arg_name: _, + .. } = t { t = ret; @@ -619,17 +621,51 @@ impl Named for Inductive { } } +/// Multiplicity for linear type system (Rust-style syntax) +/// - Many: unrestricted (default), can be used any number of times +/// - Linear: must be used exactly once (syntax: !x) +/// - Affine: can be used 0 or 1 time (syntax: ?x) +#[derive(Debug, Clone, PartialEq, Hash, Eq, PartialOrd, Ord)] +pub enum Multiplicity { + Many, // ω - unrestricted (default) + Linear, // 1 - must use exactly once + Affine, // ≤1 - can use 0 or 1 time +} + +impl Default for Multiplicity { + fn default() -> Self { + Multiplicity::Many + } +} + +impl Display for Multiplicity { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Multiplicity::Many => write!(f, "ω"), + Multiplicity::Linear => write!(f, "!"), + Multiplicity::Affine => write!(f, "?"), + } + } +} + #[derive(Debug, Clone, PartialEq, Hash, Eq, PartialOrd, Ord)] pub enum Par { - P(Param), - I { typ: Box }, + P(Param), // explicit param (has mult in Param) + I { typ: Box, mult: Multiplicity }, // implicit param with multiplicity } impl Par { pub fn typ(&self) -> &Term { match self { Par::P(param) => ¶m.typ, - Par::I { typ } => typ, + Par::I { typ, .. } => typ, + } + } + + pub fn multiplicity(&self) -> &Multiplicity { + match self { + Par::P(param) => ¶m.mult, + Par::I { mult, .. } => mult, } } @@ -637,7 +673,10 @@ impl Par { use Par::*; match self { P(param_) => P(param(param_.name.clone(), *param_.typ.clone())), - I { typ: _ } => I { typ: Box::new(typ) }, + I { typ: _, mult } => I { + typ: Box::new(typ), + mult: mult.clone(), + }, } } } @@ -646,7 +685,7 @@ impl Display for Par { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { Par::P(p) => write!(f, "{p}"), - Par::I { typ } => write!(f, "(' : {typ})"), + Par::I { typ, .. } => write!(f, "(' : {typ})"), } } } @@ -654,6 +693,7 @@ impl Display for Par { pub struct Param { pub name: Identifier, pub typ: Box, + pub mult: Multiplicity, // NEW: Many (default), Linear (!), or Affine (?) } impl Typed for Param { @@ -665,9 +705,9 @@ impl Typed for Param { impl Display for Param { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { if self.typ.is_known() { - write!(f, "({} : {})", self.name, self.typ) + write!(f, "({}{} : {})", self.mult, self.name, self.typ) } else { - write!(f, "{}", self.name) + write!(f, "{}{}", self.mult, self.name) } } } @@ -680,6 +720,15 @@ pub fn param(name: Identifier, typ: Term) -> Param { Param { name, typ: Box::new(typ), + mult: Multiplicity::default(), + } +} + +pub fn param_with_mult(name: Identifier, typ: Term, mult: Multiplicity) -> Param { + Param { + name, + typ: Box::new(typ), + mult, } } @@ -692,7 +741,7 @@ pub fn mpvar(name: ModulePath) -> Term { name: NameRef::P(name), } } -pub fn ty(name: Identifier) -> Term { +pub fn ivar(name: Identifier) -> Term { Var { name: NameRef::Id(name), } @@ -702,7 +751,7 @@ pub fn mpv(s: &str) -> Term { } pub fn typ(s: &str) -> Term { - ty(id(s)) + ivar(id(s)) } pub fn forall(param: Param, body: Term) -> Term { @@ -737,6 +786,7 @@ pub fn pi(arg: Term, ret: Term) -> Term { arg_name: None, arg: Box::new(arg), ret: Box::new(ret), + mult: Multiplicity::default(), } } pub fn pi_name(arg_name: Option, arg: Term, ret: Term) -> Term { @@ -744,6 +794,7 @@ pub fn pi_name(arg_name: Option, arg: Term, ret: Term) -> Term { arg_name, arg: Box::new(arg), ret: Box::new(ret), + mult: Multiplicity::default(), } } @@ -1069,6 +1120,7 @@ pub enum Term { arg_name: Option, arg: Box, ret: Box, + mult: Multiplicity, // NEW: Many, Linear (!), or Affine (?) }, /// Variable Var { @@ -1137,6 +1189,7 @@ impl Term { arg, ret: _, arg_name: _, + .. } => arg.is_forall(), _ => false, } @@ -1149,6 +1202,7 @@ impl Term { arg: _, ret: _, arg_name: _, + .. } => true, Forall { name: _, @@ -1217,6 +1271,7 @@ impl Term { arg: _, ret: _, arg_name: _, + .. } => "pi", Prop => "prop", Type { universe: _ } => "type", @@ -1329,7 +1384,9 @@ impl Display for Term { } } Ctx { loc: _, term } => write!(f, "{term}"), - Pi { arg, ret, arg_name } => { + Pi { + arg, ret, arg_name, .. + } => { if let Some(name) = arg_name { write!(f, "({name} : {arg} -> {ret})") } else { @@ -1417,7 +1474,10 @@ pub fn apps(fun: Term, args: Vec) -> Term { pub fn lam_index(typ: Term, body: Term) -> Term { Term::Lam { - param: Par::I { typ: Box::new(typ) }, + param: Par::I { + typ: Box::new(typ), + mult: Multiplicity::default(), + }, body: Box::new(body), } } @@ -1427,7 +1487,10 @@ pub fn lam_indecies(params: Vec, body: Term) -> Term { let mut body = body; for param in params.into_iter().rev() { body = Term::Lam { - param: Par::I { typ: param.typ }, + param: Par::I { + typ: param.typ, + mult: Multiplicity::default(), + }, body: Box::new(body), } } diff --git a/core/src/term/module.rs b/core/src/term/module.rs index 58b93c6..e108658 100644 --- a/core/src/term/module.rs +++ b/core/src/term/module.rs @@ -4,7 +4,7 @@ pub mod test; use super::*; use crate::Set; use crate::eval::native::{NativeFun, load_native_funs}; -use crate::eval::r#type::{TypeError, derive_instance_key, type_check_module_decls}; +use crate::eval::r#type::{TypeError, UsageEnv, derive_instance_key, type_check_module_decls}; use crate::term::{Inductive, Instance, InstanceKey, ModulePath, SourceContext, Term}; use crate::{ parser::parse_file, @@ -982,16 +982,37 @@ impl<'a> GlobalScope<'a> { pub enum Scope<'a> { Top { global: &'a GlobalScope<'a>, + usage_env: UsageEnv, }, Sub { local: LocalVar<'a>, parent: Box>, + usage_env: UsageEnv, }, } impl<'a> Scope<'a> { pub fn new(global: &'a GlobalScope<'a>) -> Scope<'a> { - Scope::Top { global } + Scope::Top { + global, + usage_env: UsageEnv::new(), + } + } + + /// Get a reference to the usage environment + pub fn usage_env(&self) -> &UsageEnv { + match self { + Scope::Top { usage_env, .. } => usage_env, + Scope::Sub { usage_env, .. } => usage_env, + } + } + + /// Get a mutable reference to the usage environment + pub fn usage_env_mut(&mut self) -> &mut UsageEnv { + match self { + Scope::Top { usage_env, .. } => usage_env, + Scope::Sub { usage_env, .. } => usage_env, + } } /// Extract term of NameRef pub fn resolve_name(&self, nref: &NameRef) -> Result<&Term, ScopeError> { @@ -1028,8 +1049,8 @@ impl<'a> Scope<'a> { pub fn find_local(&'a self, local_name: &Identifier) -> Option<&'a LocalVar<'a>> { match self { - Scope::Top { global: _ } => None, - Scope::Sub { local, parent } => { + Scope::Top { global: _, .. } => None, + Scope::Sub { local, parent, .. } => { if let Some(name) = local.name() && name == local_name { @@ -1044,8 +1065,8 @@ impl<'a> Scope<'a> { /// Named local variables pub fn locals(&'a self) -> Map<&'a Identifier, &'a LocalVar<'a>> { match self { - Scope::Top { global: _ } => Map::new(), - Scope::Sub { local, parent } => { + Scope::Top { global: _, .. } => Map::new(), + Scope::Sub { local, parent, .. } => { let mut loc = parent.locals(); if let Some(name) = local.name() { loc.insert(name, local); @@ -1056,8 +1077,8 @@ impl<'a> Scope<'a> { } pub fn local_foralls(&'a self) -> Map<&'a Identifier, &'a LocalVar<'a>> { match self { - Scope::Top { global: _ } => Map::new(), - Scope::Sub { local, parent } => { + Scope::Top { global: _, .. } => Map::new(), + Scope::Sub { local, parent, .. } => { let mut loc = parent.local_foralls(); if let LocalVar::Forall { .. } = local && let Some(name) = local.name() @@ -1083,7 +1104,7 @@ impl<'a> Scope<'a> { ) -> Result, ScopeError> { use Scope::{Sub, Top}; match self { - Sub { local, parent } => { + Sub { local, parent, .. } => { if let Some(name) = nref.as_id() && let Some(local_name) = local.name() && name == local_name @@ -1093,7 +1114,7 @@ impl<'a> Scope<'a> { parent.find_var_ref_of(nref, given_type) } } - Top { global } => { + Top { global, .. } => { let def = global.find_any_name_ref(nref, given_type)?; Ok(def) } @@ -1101,40 +1122,62 @@ impl<'a> Scope<'a> { } pub fn with_param(&self, param: &'a Par) -> Scope<'a> { + let mult = param.multiplicity(); + let mut usage_env = self.usage_env().clone(); match param { - Par::P(param) => self.with_local_var(¶m.name, param.typ.as_ref()), - Par::I { typ } => self.with_local_index_var(typ.as_ref()), + Par::P(param) => { + usage_env.register(param.name.clone(), mult.clone()); + Scope::Sub { + local: local_var(¶m.name, param.typ.as_ref()), + parent: Box::new(self.clone()), + usage_env, + } + } + Par::I { typ, .. } => { + // Anonymous implicit param - no name to register + Scope::Sub { + local: local_index_var(typ.as_ref()), + parent: Box::new(self.clone()), + usage_env, + } + } } } pub fn with_local_var(&self, name: &'a Identifier, typ: &'a Term) -> Scope<'a> { Scope::Sub { local: local_var(name, typ), parent: Box::new(self.clone()), + usage_env: self.usage_env().clone(), } } pub fn with_forall(&self, name: &'a Identifier, typ: &'a Term) -> Scope<'a> { Scope::Sub { local: local_forall(name, typ), parent: Box::new(self.clone()), + usage_env: self.usage_env().clone(), } } pub fn with_local_index_var(&self, typ: &'a Term) -> Scope<'a> { Scope::Sub { local: local_index_var(typ), parent: Box::new(self.clone()), + usage_env: self.usage_env().clone(), } } pub fn with_type_owned(&self, name: &'a Identifier, typ: Term) -> Scope<'a> { Scope::Sub { local: local_var_owned(name, typ), parent: Box::new(self.clone()), + usage_env: self.usage_env().clone(), } } pub fn global(&self) -> &GlobalScope<'a> { match self { - Scope::Top { global } => global, - Scope::Sub { local: _, parent } => parent.global(), + Scope::Top { global, .. } => global, + Scope::Sub { + local: _, parent, .. + } => parent.global(), } } } diff --git a/core/src/term/test.rs b/core/src/term/test.rs index e1b29df..9065cfa 100644 --- a/core/src/term/test.rs +++ b/core/src/term/test.rs @@ -91,7 +91,7 @@ impl Similar for Par { fn similar(&self, other: &Par) -> bool { match (self, other) { (Par::P(param), Par::P(o_param)) => param.similar(o_param), - (Par::I { typ: t1 }, Par::I { typ: t2 }) => t1.similar(t2), + (Par::I { typ: t1, .. }, Par::I { typ: t2, .. }) => t1.similar(t2), _ => false, } } @@ -254,13 +254,15 @@ impl Similar for Term { arg: a1, ret: r1, arg_name: n1, + mult: m1, }, Pi { arg: a2, ret: r2, arg_name: n2, + mult: m2, }, - ) => (*a1).similar(&**a2) && (*r1).similar(&**r2) && n1 == n2, + ) => (*a1).similar(&**a2) && (*r1).similar(&**r2) && n1 == n2 && m1 == m2, (Var { name: n1 }, Var { name: n2 }) => n1 == n2, ( Forall { diff --git a/llvm-codegen/src/compiler.rs b/llvm-codegen/src/compiler.rs index 0b00f0b..dfec585 100644 --- a/llvm-codegen/src/compiler.rs +++ b/llvm-codegen/src/compiler.rs @@ -286,7 +286,7 @@ mod tests { use std::fs; use std::sync::atomic::{AtomicU64, Ordering}; - use monad_core::term::{Decl, def, id, lams, mpt, num, param, type0}; + use monad_core::term::{Decl, def, id, lams, mpt, num, param, pi, type0}; static TEST_COUNTER: AtomicU64 = AtomicU64::new(0); @@ -305,11 +305,7 @@ mod tests { } else { let mut typ = type0(); for pt in param_types.into_iter().rev() { - typ = Term::Pi { - arg_name: None, - arg: Box::new(pt), - ret: Box::new(typ), - }; + typ = pi(pt, typ); } typ }; diff --git a/llvm-codegen/tests/codegen_tests.rs b/llvm-codegen/tests/codegen_tests.rs index 3c47aa2..9ed8f81 100644 --- a/llvm-codegen/tests/codegen_tests.rs +++ b/llvm-codegen/tests/codegen_tests.rs @@ -1,4 +1,4 @@ -use monad_core::term::{Decl, Param, Term, def, id, lam, lams, mpt, num, param, type0}; +use monad_core::term::{Decl, Param, Term, def, id, lam, lams, mpt, num, param, pi, type0}; use monad_llvm_codegen::compile_decls; @@ -9,11 +9,7 @@ fn make_def(name: &str, params: Vec, body: Term) -> Decl { } else { let mut typ = type0(); for pt in param_types.into_iter().rev() { - typ = Term::Pi { - arg_name: None, - arg: Box::new(pt), - ret: Box::new(typ), - }; + typ = pi(pt, typ); } typ };