From 23aed093d54583d2d70ea3c2f28834d33005694d Mon Sep 17 00:00:00 2001 From: Jonathan Cubides Date: Thu, 13 Aug 2026 11:37:10 -0500 Subject: [PATCH] refactor(typing): centralize scoped environments --- EventB/Typing/Check.lean | 43 +++++++++++++++------------- EventB/Typing/Infer.lean | 61 ++++++++++++++++++++++++---------------- 2 files changed, 60 insertions(+), 44 deletions(-) diff --git a/EventB/Typing/Check.lean b/EventB/Typing/Check.lean index 98ce708..3da6106 100644 --- a/EventB/Typing/Check.lean +++ b/EventB/Typing/Check.lean @@ -78,7 +78,7 @@ def componentTheoryRoots (p : Project) (name : String) : List String := /-- Declare the identifiers a component introduces, then feed every predicate it states to the checker. Errors are collected rather than thrown: one unsupported guard should cost that guard's constraints, not the whole file's types. -/ -private def addComponent (c : Component) : StateT St (Except String) (List String) := do +private def addComponent (c : Component) : M (List String) := do let mut errs : List String := [] -- Carrier sets and constants first, so axioms can refer to them in any order. for s in childrenOf c.elem "carrierSet" do @@ -100,36 +100,37 @@ private def addComponent (c : Component) : StateT St (Except String) (List Strin errs := errs ++ (← runPredicate f) -- Each event's parameters are scoped to that event. for ev in childrenOf c.elem "event" do - let saved := (← get).env - for prm in childrenOf ev "parameter" do - if let some n := attrOf prm "identifier" then bind n (← fresh) - let withParams := (← get).env - for g in childrenOf ev "guard" do - if let some f := attrOf g "predicate" then - errs := errs ++ (← runPredicate f) - for act in childrenOf ev "action" do - if let some f := attrOf act "assignment" then - errs := errs ++ (← runPredicate f) - for w in childrenOf ev "witness" do - if let some f := attrOf w "predicate" then - errs := errs ++ (← runPredicate f) + let (eventErrors, bound) ← withEnvBindings do + let mut eventErrors : List String := [] + for prm in childrenOf ev "parameter" do + if let some n := attrOf prm "identifier" then bind n (← fresh) + for g in childrenOf ev "guard" do + if let some f := attrOf g "predicate" then + eventErrors := eventErrors ++ (← runPredicate f) + for act in childrenOf ev "action" do + if let some f := attrOf act "assignment" then + eventErrors := eventErrors ++ (← runPredicate f) + for w in childrenOf ev "witness" do + if let some f := attrOf w "predicate" then + eventErrors := eventErrors ++ (← runPredicate f) + return eventErrors + errs := errs ++ eventErrors -- Parameters leave the environment so a later event cannot see them, but they are -- kept in `params` because the `.bpo` records their types alongside the variables. - let bound := withParams.take (withParams.length - saved.length) - modify fun s => { s with env := saved, params := s.params ++ bound } + modify fun s => { s with params := s.params ++ bound } return errs where /-- Reuse the existing type if the name is already declared, so a refinement does not discard what the abstract machine established. -/ - freshFor (n : String) : StateT St (Except String) Ty := do + freshFor (n : String) : M Ty := do match ← lookup? n with | some t => return t | none => fresh - declare (n : String) (t : Ty) : StateT St (Except String) Unit := do + declare (n : String) (t : Ty) : M Unit := do match ← lookup? n with | some _ => return () | none => bind n t - runPredicate (f : String) : StateT St (Except String) (List String) := do + runPredicate (f : String) : M (List String) := do match Formula.parse f with | .error e => return [s!"parse: {EventB.Error.render e}"] | .ok term => @@ -186,6 +187,10 @@ def inferTerm (env : List (String × Ty)) (t : Term) : Except EventB.Error Ty := /-! Self-checks. The corpus pins the common cases; these pin the shapes it happens not to contain, and the printer conventions the `.bpo` comparison depends on. -/ +#guard match (withEnvBindings (bind "x" .int)).run {} with + | .ok ((_, [ ("x", .int) ]), state) => state.env.isEmpty + | _ => false + /-- `given` are identifiers with a known type, `unknown` are the ones inference has to work out. Metavariables must come from `fresh` so the substitution has a slot for them. -/ private def inferOne (given : List (String × Ty)) (unknown : List String) diff --git a/EventB/Typing/Infer.lean b/EventB/Typing/Infer.lean index 4af63bf..af64e12 100644 --- a/EventB/Typing/Infer.lean +++ b/EventB/Typing/Infer.lean @@ -121,6 +121,22 @@ def lookup? (name : String) : M (Option Ty) := do def bind (name : String) (t : Ty) : M Unit := modify fun s => { s with env := (name, t) :: s.env } +/-- Run a typing action in a lexical environment and restore that environment afterward. -/ +def withEnv {α : Type} (action : M α) : M α := do + let saved := (← get).env + let value ← action + modify fun s => { s with env := saved } + return value + +/-- Run a scoped action, returning the bindings it introduced after restoring the environment. -/ +def withEnvBindings {α : Type} (action : M α) : M (α × List (String × Ty)) := do + let saved := (← get).env + let value ← action + let current := (← get).env + let bound := current.take (current.length - saved.length) + modify fun s => { s with env := saved } + return (value, bound) + /-- A relation `ℙ(A×B)`, returning the two sides. -/ private def asRelation (t : Ty) : M (Ty × Ty) := do let a ← fresh @@ -165,11 +181,10 @@ def checkPred (t : Term) : M Unit := do | .id "⊤" | .id "⊥" => return () | .pre "¬" p => checkPred p | .bind k pat body => - if k == "∀" || k == "∃" then do - let saved := (← get).env + if k == "∀" || k == "∃" then + withEnv do bindPattern pat checkPred body - modify fun s => { s with env := saved } else throw s!"binder {k} is not a predicate" | .bin o a b => if connectives.contains o then do checkPred a; checkPred b @@ -325,29 +340,25 @@ termination_by sizeOf f + sizeOf a decreasing_by all_goals simp_all +arith [termSizePos, Term.id.sizeOf_spec] -def inferBinder (k : String) (pat body : Term) : M Ty := do - let saved := (← get).env +def inferBinder (k : String) (pat body : Term) : M Ty := withEnv do bindPattern pat - let result ← do - match k, body with - | "λ", .bin "∣" p e => do - checkPred p - return .pow (.prod (← patternType pat) (← inferExpr e)) - | "{", .bin "∣" p e => do - checkPred p - return .pow (← inferExpr e) - | "{", p => do - -- `{x · P}` with no expression part means the bound variables themselves. - checkPred p - return .pow (← patternType pat) - | "⋃", .bin "∣" p e | "⋂", .bin "∣" p e => do - checkPred p - let t ← inferExpr e - let _ ← asSet t - return t - | k, _ => throw s!"binder {k} is not an expression" - modify fun s => { s with env := saved } - return result + match k, body with + | "λ", .bin "∣" p e => do + checkPred p + return .pow (.prod (← patternType pat) (← inferExpr e)) + | "{", .bin "∣" p e => do + checkPred p + return .pow (← inferExpr e) + | "{", p => do + -- `{x · P}` with no expression part means the bound variables themselves. + checkPred p + return .pow (← patternType pat) + | "⋃", .bin "∣" p e | "⋂", .bin "∣" p e => do + checkPred p + let t ← inferExpr e + let _ ← asSet t + return t + | k, _ => throw s!"binder {k} is not an expression" termination_by sizeOf pat + sizeOf body decreasing_by