chore: refactor match elaborator to be used from the do elaborator (#12451)
This PR provides the necessary hooks for the new do elaborator to call into the let and match elaborator. The `do match` elaborator needs access to a couple of functions from the term `match` elaborator to implement its own `elabMatchAlt`. In particular, `withEqs`, `withPatternVars` and `checkNumPatterns` need to be exposed. Furthermore, I think it makes sense to share `instantiateAltLHSs`.
This commit is contained in:
parent
63f7776390
commit
2c68952694
4 changed files with 79 additions and 68 deletions
|
|
@ -574,7 +574,7 @@ private def checkMatchAltPatternCounts (matchAlts : Syntax) (numDiscrs : Nat) (e
|
|||
let sepPats (pats : List Syntax) := MessageData.joinSep (pats.map toMessageData) ", "
|
||||
let maxDiscrs? ← forallTelescopeReducing expectedType fun xs e =>
|
||||
if e.getAppFn.isMVar then pure none else pure (some xs.size)
|
||||
let matchAltViews := matchAlts[0].getArgs.filterMap getMatchAlt
|
||||
let matchAltViews := matchAlts[0].getArgs.filterMap (getMatchAlt `term)
|
||||
let numPatternsStr (n : Nat) := s!"{n} {if n == 1 then "pattern" else "patterns"}"
|
||||
if h : matchAltViews.size > 0 then
|
||||
if let some maxDiscrs := maxDiscrs? then
|
||||
|
|
|
|||
|
|
@ -70,15 +70,17 @@ structure Discr where
|
|||
h? : Option Syntax := none
|
||||
deriving Inhabited
|
||||
|
||||
abbrev TermMatchAltView := MatchAltView `term
|
||||
|
||||
structure ElabMatchTypeAndDiscrsResult where
|
||||
discrs : Array Discr
|
||||
matchType : Expr
|
||||
/-- `true` when performing dependent elimination. We use this to decide whether we optimize the "match unit" case.
|
||||
See `isMatchUnit?`. -/
|
||||
isDep : Bool
|
||||
alts : Array MatchAltView
|
||||
alts : Array TermMatchAltView
|
||||
|
||||
private partial def elabMatchTypeAndDiscrs (discrStxs : Array Syntax) (matchOptMotive : Syntax) (matchAltViews : Array MatchAltView) (expectedType : Expr)
|
||||
private partial def elabMatchTypeAndDiscrs (discrStxs : Array Syntax) (matchOptMotive : Syntax) (matchAltViews : Array TermMatchAltView) (expectedType : Expr)
|
||||
: TermElabM ElabMatchTypeAndDiscrsResult := do
|
||||
if matchOptMotive.isNone then
|
||||
elabDiscrs 0 #[]
|
||||
|
|
@ -149,7 +151,7 @@ where
|
|||
else
|
||||
return { discrs, alts := matchAltViews, isDep := false, matchType := expectedType }
|
||||
|
||||
def expandMacrosInPatterns (matchAlts : Array MatchAltView) : MacroM (Array MatchAltView) := do
|
||||
def expandMacrosInPatterns (matchAlts : Array (MatchAltView k)) : MacroM (Array (MatchAltView k)) := do
|
||||
matchAlts.mapM fun matchAlt => do
|
||||
let patterns ← matchAlt.patterns.mapM expandMacros
|
||||
pure { matchAlt with patterns := patterns }
|
||||
|
|
@ -159,20 +161,20 @@ private def getMatchGeneralizing? : Syntax → Option Bool
|
|||
| `(match (generalizing := false) $[$motive]? $_discrs,* with $_alts:matchAlt*) => some false
|
||||
| _ => none
|
||||
|
||||
/-- Given the `stx` of a single match alternative, return a corresponding `MatchAltView`. -/
|
||||
def getMatchAlt : Syntax → Option MatchAltView
|
||||
/-- Given the `stx` of a single match alternative, return a corresponding `TermMatchAltView`. -/
|
||||
def getMatchAlt (k : SyntaxNodeKinds) : Syntax → Option (MatchAltView k)
|
||||
| alt@`(matchAltExpr| | $patterns,* => $rhs) => some {
|
||||
ref := alt,
|
||||
patterns := patterns,
|
||||
lhs := alt[1], -- this is the ref `$patterns,*`
|
||||
rhs := rhs
|
||||
rhs := ⟨rhs⟩
|
||||
}
|
||||
| _ => none
|
||||
|
||||
/-- Given `stx` a match-expression, return its alternatives. -/
|
||||
def getMatchAlts : Syntax → Array MatchAltView
|
||||
def getMatchAlts (k : SyntaxNodeKinds) : Syntax → Array (MatchAltView k)
|
||||
| `(match $[$gen]? $[$motive]? $_discrs,* with $alts:matchAlt*) =>
|
||||
alts.filterMap getMatchAlt
|
||||
alts.filterMap (getMatchAlt k)
|
||||
| _ => #[]
|
||||
|
||||
@[builtin_term_elab inaccessible] def elabInaccessible : TermElab := fun stx expectedType? => do
|
||||
|
|
@ -194,7 +196,7 @@ open Lean.Elab.Term.Quotation in
|
|||
structure PatternVarDecl where
|
||||
fvarId : FVarId
|
||||
|
||||
private partial def withPatternVars {α} (pVars : Array PatternVar) (k : Array PatternVarDecl → TermElabM α) : TermElabM α :=
|
||||
partial def withPatternVars {α} (pVars : Array PatternVar) (k : Array PatternVarDecl → TermElabM α) : TermElabM α :=
|
||||
let rec loop (i : Nat) (decls : Array PatternVarDecl) (userNames : Array Name) := do
|
||||
if h : i < pVars.size then
|
||||
let type ← mkFreshTypeMVar
|
||||
|
|
@ -344,28 +346,31 @@ private partial def eraseIndices (type : Expr) : MetaM Expr := do
|
|||
let (newIndices, _, _) ← forallMetaTelescopeReducing resultType (some (args.size - info.numParams))
|
||||
return mkAppN result newIndices
|
||||
|
||||
private def withPatternElabConfig (x : TermElabM α) : TermElabM α :=
|
||||
def withPatternElabConfig (x : TermElabM α) : TermElabM α :=
|
||||
withoutErrToSorry <| withReader (fun ctx => { ctx with inPattern := true }) <| x
|
||||
|
||||
open Meta.Match (throwIncorrectNumberOfPatternsAt logIncorrectNumberOfPatternsAt)
|
||||
|
||||
/-- Checks the number of patterns and appropriately adds holes if needed. -/
|
||||
def checkNumPatterns (numDiscrs : Nat) (patternStxs : Array Syntax) : TermElabM (Array Syntax) := do
|
||||
if patternStxs.size < numDiscrs then
|
||||
-- If there are too few patterns, log the error but continue elaborating with holes for the missing patterns
|
||||
logIncorrectNumberOfPatternsAt (← getRef) "Not enough" numDiscrs patternStxs.size patternStxs.toList
|
||||
let numHoles := numDiscrs - patternStxs.size
|
||||
let mut extraStxs := Array.emptyWithCapacity numHoles
|
||||
for _ in *...numHoles do
|
||||
extraStxs := extraStxs.push (← `(_))
|
||||
return patternStxs ++ extraStxs
|
||||
else if patternStxs.size > numDiscrs then
|
||||
throwIncorrectNumberOfPatternsAt (← getRef) "Too many" numDiscrs patternStxs.size patternStxs.toList
|
||||
return patternStxs
|
||||
|
||||
private def elabPatterns (patternStxs : Array Syntax) (numDiscrs : Nat) (matchType : Expr) : ExceptT PatternElabException TermElabM (Array Expr × Expr) :=
|
||||
withReader (fun ctx => { ctx with implicitLambda := false }) do
|
||||
let origMatchType := matchType
|
||||
let mut patterns := #[]
|
||||
let mut matchType := matchType
|
||||
let mut patternStxs := patternStxs
|
||||
if patternStxs.size < numDiscrs then
|
||||
-- If there are too few patterns, log the error but continue elaborating with holes for the missing patterns
|
||||
logIncorrectNumberOfPatternsAt (← getRef) "Not enough" numDiscrs patternStxs.size patternStxs.toList
|
||||
let numHoles := numDiscrs - patternStxs.size
|
||||
let mut extraStxs := Array.emptyWithCapacity numHoles
|
||||
for _ in *...numHoles do
|
||||
extraStxs := extraStxs.push (← `(_))
|
||||
patternStxs := patternStxs ++ extraStxs
|
||||
else if patternStxs.size > numDiscrs then
|
||||
throwIncorrectNumberOfPatternsAt (← getRef) "Too many" numDiscrs patternStxs.size patternStxs.toList
|
||||
|
||||
let patternStxs ← checkNumPatterns numDiscrs patternStxs
|
||||
for h : idx in *...patternStxs.size do
|
||||
let patternStx := patternStxs[idx]
|
||||
matchType ← whnf matchType
|
||||
|
|
@ -789,7 +794,7 @@ private def withToClear (toClear : Array FVarId) (type : Expr) (k : TermElabM α
|
|||
Generate equalities `h : discr = pattern` for discriminants annotated with `h :`.
|
||||
We use these equalities to elaborate the right-hand-side of a `match` alternative.
|
||||
-/
|
||||
private def withEqs (discrs : Array Discr) (patterns : List Pattern) (k : Array Expr → TermElabM α) : TermElabM α := do
|
||||
def withEqs (discrs : Array Discr) (patterns : List Pattern) (k : Array Expr → TermElabM α) : TermElabM α := do
|
||||
go 0 patterns #[]
|
||||
where
|
||||
go (i : Nat) (ps : List Pattern) (eqs : Array Expr) : TermElabM α := do
|
||||
|
|
@ -812,7 +817,7 @@ where
|
|||
The array `toClear` contains variables that must be cleared before elaborating the `rhs` because
|
||||
they have been generalized/refined.
|
||||
-/
|
||||
private def elabMatchAltView (discrs : Array Discr) (alt : MatchAltView) (matchType : Expr) (toClear : Array FVarId) : ExceptT PatternElabException TermElabM (AltLHS × Expr) := withRef alt.ref do
|
||||
private def elabMatchAltView (discrs : Array Discr) (alt : TermMatchAltView) (matchType : Expr) (toClear : Array FVarId) : ExceptT PatternElabException TermElabM (AltLHS × Expr) := withRef alt.ref do
|
||||
let (patternVars, alt) ← collectPatternVars alt
|
||||
trace[Elab.match] "patternVars: {patternVars}"
|
||||
withPatternVars patternVars fun patternVarDecls => do
|
||||
|
|
@ -857,7 +862,7 @@ structure GeneralizeResult where
|
|||
/-- `FVarId`s of the variables that have been generalized. We store them to clear after in each branch. -/
|
||||
toClear : Array FVarId := #[]
|
||||
matchType : Expr
|
||||
altViews : Array MatchAltView
|
||||
altViews : Array TermMatchAltView
|
||||
refined : Bool := false
|
||||
|
||||
/--
|
||||
|
|
@ -871,7 +876,7 @@ structure GeneralizeResult where
|
|||
but they become inaccessible since they are shadowed by the patterns variables. We assume this is ok since
|
||||
this is the exact behavior users would get if they had written it by hand. Recall there is no `clear` in term mode.
|
||||
-/
|
||||
private def generalize (discrs : Array Discr) (matchType : Expr) (altViews : Array MatchAltView) (generalizing? : Option Bool) : TermElabM GeneralizeResult := do
|
||||
private def generalize (discrs : Array Discr) (matchType : Expr) (altViews : Array TermMatchAltView) (generalizing? : Option Bool) : TermElabM GeneralizeResult := do
|
||||
let gen := if let some g := generalizing? then g else true
|
||||
if !gen then
|
||||
return { discrs, matchType, altViews }
|
||||
|
|
@ -911,13 +916,13 @@ private def generalize (discrs : Array Discr) (matchType : Expr) (altViews : Arr
|
|||
return { discrs, matchType, altViews }
|
||||
|
||||
|
||||
private partial def elabMatchAltViews (generalizing? : Option Bool) (discrs : Array Discr) (matchType : Expr) (altViews : Array MatchAltView) : TermElabM (Array Discr × Expr × Array (AltLHS × Expr) × Bool) := do
|
||||
private partial def elabMatchAltViews (generalizing? : Option Bool) (discrs : Array Discr) (matchType : Expr) (altViews : Array TermMatchAltView) : TermElabM (Array Discr × Expr × Array (AltLHS × Expr) × Bool) := do
|
||||
loop discrs #[] matchType altViews none
|
||||
where
|
||||
/--
|
||||
"Discriminant refinement" main loop.
|
||||
`first?` contains the first error message we found before updated the `discrs`. -/
|
||||
loop (discrs : Array Discr) (toClear : Array FVarId) (matchType : Expr) (altViews : Array MatchAltView) (first? : Option (SavedState × Exception))
|
||||
loop (discrs : Array Discr) (toClear : Array FVarId) (matchType : Expr) (altViews : Array TermMatchAltView) (first? : Option (SavedState × Exception))
|
||||
: TermElabM (Array Discr × Expr × Array (AltLHS × Expr) × Bool) := do
|
||||
let s ← saveState
|
||||
let { discrs := discrs', toClear := toClear', matchType := matchType', altViews := altViews', refined } ← generalize discrs matchType altViews generalizing?
|
||||
|
|
@ -1063,7 +1068,44 @@ private def isMatchUnit? (altLHSS : List Match.AltLHS) (rhss : Array Expr) : Met
|
|||
| _ => return none
|
||||
| _ => return none
|
||||
|
||||
private def elabMatchAux (generalizing? : Option Bool) (discrStxs : Array Syntax) (altViews : Array MatchAltView) (matchOptMotive : Syntax) (expectedType : Expr)
|
||||
/--
|
||||
Ensure that there are no metavariables in free variables and patterns of the `lhss`.
|
||||
-/
|
||||
def instantiateAltLHSs (lhss : Array AltLHS) : TermElabM (List AltLHS) :=
|
||||
lhss.toList.mapM fun altLHS => do
|
||||
let altLHS ← Match.instantiateAltLHSMVars altLHS
|
||||
/- Remark: we try to postpone before throwing an error.
|
||||
The combinator `commitIfDidNotPostpone` ensures we backtrack any updates that have been performed.
|
||||
The quick-check `waitExpectedTypeAndDiscrs` minimizes the number of scenarios where we have to postpone here.
|
||||
Here is an example that passes the `waitExpectedTypeAndDiscrs` test, but postpones here.
|
||||
```
|
||||
def bad (ps : Array (Nat × Nat)) : Array (Nat × Nat) :=
|
||||
(ps.filter fun (p : Prod _ _) =>
|
||||
match p with
|
||||
| (x, y) => x == 0)
|
||||
++
|
||||
ps
|
||||
```
|
||||
When we try to elaborate `fun (p : Prod _ _) => ...` for the first time, we haven't propagated the type of `ps` yet
|
||||
because `Array.filter` has type `{α : Type u_1} → (α → Bool) → (as : Array α) → optParam Nat 0 → optParam Nat (Array.size as) → Array α`
|
||||
However, the partial type annotation `(p : Prod _ _)` makes sure we succeed at the quick-check `waitExpectedTypeAndDiscrs`.
|
||||
-/
|
||||
withRef altLHS.ref do
|
||||
for d in altLHS.fvarDecls do
|
||||
if d.hasExprMVar then
|
||||
Term.tryPostpone
|
||||
withExistingLocalDecls altLHS.fvarDecls do
|
||||
Term.runPendingTacticsAt d.type
|
||||
if (← instantiateMVars d.type).hasExprMVar then
|
||||
Term.throwMVarError m!"Invalid match expression: The type of pattern variable '{d.toExpr}' contains metavariables:{indentExpr d.type}"
|
||||
for p in altLHS.patterns do
|
||||
if (← Match.instantiatePatternMVars p).hasExprMVar then
|
||||
Term.tryPostpone
|
||||
withExistingLocalDecls altLHS.fvarDecls do
|
||||
Term.throwMVarError m!"Invalid match expression: This pattern contains metavariables:{indentExpr (← p.toExpr)}"
|
||||
pure altLHS
|
||||
|
||||
private def elabMatchAux (generalizing? : Option Bool) (discrStxs : Array Syntax) (altViews : Array TermMatchAltView) (matchOptMotive : Syntax) (expectedType : Expr)
|
||||
: TermElabM Expr := do
|
||||
let mut generalizing? := generalizing?
|
||||
if !matchOptMotive.isNone then
|
||||
|
|
@ -1098,38 +1140,7 @@ private def elabMatchAux (generalizing? : Option Bool) (discrStxs : Array Syntax
|
|||
synthesizeSyntheticMVarsUsingDefault
|
||||
let rhss := alts.map Prod.snd
|
||||
let matchType ← instantiateMVars matchType
|
||||
let altLHSS ← alts.toList.mapM fun alt => do
|
||||
let altLHS ← Match.instantiateAltLHSMVars alt.1
|
||||
/- Remark: we try to postpone before throwing an error.
|
||||
The combinator `commitIfDidNotPostpone` ensures we backtrack any updates that have been performed.
|
||||
The quick-check `waitExpectedTypeAndDiscrs` minimizes the number of scenarios where we have to postpone here.
|
||||
Here is an example that passes the `waitExpectedTypeAndDiscrs` test, but postpones here.
|
||||
```
|
||||
def bad (ps : Array (Nat × Nat)) : Array (Nat × Nat) :=
|
||||
(ps.filter fun (p : Prod _ _) =>
|
||||
match p with
|
||||
| (x, y) => x == 0)
|
||||
++
|
||||
ps
|
||||
```
|
||||
When we try to elaborate `fun (p : Prod _ _) => ...` for the first time, we haven't propagated the type of `ps` yet
|
||||
because `Array.filter` has type `{α : Type u_1} → (α → Bool) → (as : Array α) → optParam Nat 0 → optParam Nat (Array.size as) → Array α`
|
||||
However, the partial type annotation `(p : Prod _ _)` makes sure we succeed at the quick-check `waitExpectedTypeAndDiscrs`.
|
||||
-/
|
||||
withRef altLHS.ref do
|
||||
for d in altLHS.fvarDecls do
|
||||
if d.hasExprMVar then
|
||||
tryPostpone
|
||||
withExistingLocalDecls altLHS.fvarDecls do
|
||||
runPendingTacticsAt d.type
|
||||
if (← instantiateMVars d.type).hasExprMVar then
|
||||
throwMVarError m!"Invalid match expression: The type of pattern variable '{d.toExpr}' contains metavariables:{indentExpr d.type}"
|
||||
for p in altLHS.patterns do
|
||||
if (← Match.instantiatePatternMVars p).hasExprMVar then
|
||||
tryPostpone
|
||||
withExistingLocalDecls altLHS.fvarDecls do
|
||||
throwMVarError m!"Invalid match expression: This pattern contains metavariables:{indentExpr (← p.toExpr)}"
|
||||
pure altLHS
|
||||
let altLHSS ← instantiateAltLHSs (alts.map Prod.fst)
|
||||
return (discrs, matchType, altLHSS, isDep, rhss)
|
||||
if let some r ← if isDep then pure none else isMatchUnit? altLHSS rhss then
|
||||
return r
|
||||
|
|
@ -1247,11 +1258,11 @@ private def elabMatchCore (stx : Syntax) (expectedType? : Option Expr) : TermEla
|
|||
let expectedType ← waitExpectedTypeAndDiscrs stx expectedType?
|
||||
let discrStxs := getDiscrs stx
|
||||
let gen? := getMatchGeneralizing? stx
|
||||
let altViews := getMatchAlts stx
|
||||
let altViews := getMatchAlts _ stx
|
||||
let matchOptMotive := getMatchOptMotive stx
|
||||
elabMatchAux gen? discrStxs altViews matchOptMotive expectedType
|
||||
|
||||
private def isPatternVar (stx : Syntax) : TermElabM Bool := do
|
||||
def isPatternVar (stx : Syntax) : TermElabM Bool := do
|
||||
match (← resolveId? stx "pattern") with
|
||||
| none => return isAtomicIdent stx
|
||||
| some f => match f with
|
||||
|
|
|
|||
|
|
@ -21,11 +21,11 @@ def «match» := leading_parser:leadPrec "match " >> sepBy1 matchDiscr ", " >> o
|
|||
```
|
||||
-/
|
||||
|
||||
structure MatchAltView where
|
||||
structure MatchAltView (k : SyntaxNodeKinds) where
|
||||
ref : Syntax
|
||||
patterns : Array Syntax
|
||||
lhs : Syntax
|
||||
rhs : Syntax
|
||||
rhs : TSyntax k
|
||||
deriving Inhabited
|
||||
|
||||
end Lean.Elab.Term
|
||||
|
|
|
|||
|
|
@ -415,7 +415,7 @@ where
|
|||
else
|
||||
throwCtorExpected (some fId)
|
||||
|
||||
def main (alt : MatchAltView) : M MatchAltView := do
|
||||
def main (alt : MatchAltView k) : M (MatchAltView k) := do
|
||||
let patterns ← alt.patterns.mapM fun p => do
|
||||
trace[Elab.match] "collecting variables at pattern: {p}"
|
||||
collect p
|
||||
|
|
@ -427,7 +427,7 @@ end CollectPatternVars
|
|||
Collect pattern variables occurring in the `match`-alternative object views.
|
||||
It also returns the updated views.
|
||||
-/
|
||||
def collectPatternVars (alt : MatchAltView) : TermElabM (Array PatternVar × MatchAltView) := do
|
||||
def collectPatternVars (alt : MatchAltView k) : TermElabM (Array PatternVar × MatchAltView k) := do
|
||||
let (alt, s) ← (CollectPatternVars.main alt).run {}
|
||||
return (s.vars, alt)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue