From 307a1ea89adb7436ef67f56b6fe0d4019533601c Mon Sep 17 00:00:00 2001 From: Adam Dangoor Date: Thu, 17 Sep 2026 09:53:47 +0100 Subject: [PATCH] Narrow stored match patterns to concrete node types --- mypy/nativeparse.py | 4 ++-- mypy/nodes.py | 6 +++--- mypy/patterns.py | 38 +++++++++++++++++++++++++------------- mypy/treetransform.py | 6 +++--- 4 files changed, 33 insertions(+), 21 deletions(-) diff --git a/mypy/nativeparse.py b/mypy/nativeparse.py index 11252e346620..2535f8b6888f 100644 --- a/mypy/nativeparse.py +++ b/mypy/nativeparse.py @@ -130,9 +130,9 @@ from mypy.patterns import ( AsPattern, ClassPattern, + ConcretePattern, MappingPattern, OrPattern, - Pattern, SequencePattern, SingletonPattern, StarredPattern, @@ -1186,7 +1186,7 @@ def read_call_type(state: State, data: ReadBuffer) -> Type: return call_arg -def read_pattern(state: State, data: ReadBuffer) -> Pattern: +def read_pattern(state: State, data: ReadBuffer) -> ConcretePattern: tag = read_tag(data) if tag == nodes.AS_PATTERN: has_pattern = read_bool(data) diff --git a/mypy/nodes.py b/mypy/nodes.py index e64d86f0ab15..dc86550cebe4 100644 --- a/mypy/nodes.py +++ b/mypy/nodes.py @@ -80,7 +80,7 @@ from mypy.visitor import ExpressionVisitor, NodeVisitor, StatementVisitor if TYPE_CHECKING: - from mypy.patterns import Pattern + from mypy.patterns import ConcretePattern @unique @@ -2200,14 +2200,14 @@ class MatchStmt(Statement): subject: Expression subject_dummy: NameExpr | None - patterns: list[Pattern] + patterns: list[ConcretePattern] guards: list[Expression | None] bodies: list[Block] def __init__( self, subject: Expression, - patterns: list[Pattern], + patterns: list[ConcretePattern], guards: list[Expression | None], bodies: list[Block], ) -> None: diff --git a/mypy/patterns.py b/mypy/patterns.py index a01bf6acc876..2bd623579982 100644 --- a/mypy/patterns.py +++ b/mypy/patterns.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import TypeVar +from typing import TypeAlias, TypeVar from mypy_extensions import trait @@ -30,10 +30,10 @@ class AsPattern(Pattern): # If pattern is None this is a capture pattern. If name and pattern are both none this is a # wildcard pattern. # Only name being None should not happen but also won't break anything. - pattern: Pattern | None + pattern: ConcretePattern | None name: NameExpr | None - def __init__(self, pattern: Pattern | None, name: NameExpr | None) -> None: + def __init__(self, pattern: ConcretePattern | None, name: NameExpr | None) -> None: super().__init__() self.pattern = pattern self.name = name @@ -45,9 +45,9 @@ def accept(self, visitor: PatternVisitor[T]) -> T: class OrPattern(Pattern): """The pattern | | ...""" - patterns: list[Pattern] + patterns: list[ConcretePattern] - def __init__(self, patterns: list[Pattern]) -> None: + def __init__(self, patterns: list[ConcretePattern]) -> None: super().__init__() self.patterns = patterns @@ -83,9 +83,9 @@ def accept(self, visitor: PatternVisitor[T]) -> T: class SequencePattern(Pattern): """The pattern [, ...]""" - patterns: list[Pattern] + patterns: list[ConcretePattern] - def __init__(self, patterns: list[Pattern]) -> None: + def __init__(self, patterns: list[ConcretePattern]) -> None: super().__init__() self.patterns = patterns @@ -108,11 +108,11 @@ def accept(self, visitor: PatternVisitor[T]) -> T: class MappingPattern(Pattern): keys: list[Expression] - values: list[Pattern] + values: list[ConcretePattern] rest: NameExpr | None def __init__( - self, keys: list[Expression], values: list[Pattern], rest: NameExpr | None + self, keys: list[Expression], values: list[ConcretePattern], rest: NameExpr | None ) -> None: super().__init__() assert len(keys) == len(values) @@ -128,16 +128,16 @@ class ClassPattern(Pattern): """The pattern Cls(...)""" class_ref: RefExpr - positionals: list[Pattern] + positionals: list[ConcretePattern] keyword_keys: list[str] - keyword_values: list[Pattern] + keyword_values: list[ConcretePattern] def __init__( self, class_ref: RefExpr, - positionals: list[Pattern], + positionals: list[ConcretePattern], keyword_keys: list[str], - keyword_values: list[Pattern], + keyword_values: list[ConcretePattern], ) -> None: super().__init__() assert len(keyword_keys) == len(keyword_values) @@ -148,3 +148,15 @@ def __init__( def accept(self, visitor: PatternVisitor[T]) -> T: return visitor.visit_class_pattern(self) + + +ConcretePattern: TypeAlias = ( + AsPattern + | OrPattern + | ValuePattern + | SingletonPattern + | SequencePattern + | StarredPattern + | MappingPattern + | ClassPattern +) diff --git a/mypy/treetransform.py b/mypy/treetransform.py index c5e1ad44ea4c..f33cba1e3c7a 100644 --- a/mypy/treetransform.py +++ b/mypy/treetransform.py @@ -97,9 +97,9 @@ from mypy.patterns import ( AsPattern, ClassPattern, + ConcretePattern, MappingPattern, OrPattern, - Pattern, SequencePattern, SingletonPattern, StarredPattern, @@ -735,9 +735,9 @@ def stmt(self, stmt: Statement) -> Statement: new.set_line(stmt) return new - def pattern(self, pattern: Pattern) -> Pattern: + def pattern(self, pattern: ConcretePattern) -> ConcretePattern: new = pattern.accept(self) - assert isinstance(new, Pattern) + assert isinstance(new, ConcretePattern) new.set_line(pattern) return new