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:
Sebastian Graf 2026-02-19 08:33:30 +01:00 committed by GitHub
parent 63f7776390
commit 2c68952694
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 79 additions and 68 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)