PY-85990 [python]: unified logic in PyCaseClause and PyTypeAssertionEvaluator about exhaustiveness of match statement

Also, some small fixes here and there to ensure it works correctly

Space-RevId: 90f4b41c0eca9aae4c49841d756f7c96f805fcff

GitOrigin-RevId: 844c13869fb113df46537c4a644dbbfb8ba8c977
This commit is contained in:
Aleksandr.Govenko
2026-02-03 18:14:34 +00:00
committed by intellij-monorepo-bot
parent d01aa9e583
commit da873ac92b
11 changed files with 229 additions and 98 deletions
@@ -2,6 +2,10 @@
package com.jetbrains.python.psi;
import com.jetbrains.python.ast.PyAstCaseClause;
import com.jetbrains.python.psi.types.PyType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import org.jetbrains.annotations.ApiStatus;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
public interface PyCaseClause extends PyAstCaseClause, PyStatementPart {
@@ -14,4 +18,7 @@ public interface PyCaseClause extends PyAstCaseClause, PyStatementPart {
default @Nullable PyExpression getGuardCondition() {
return (PyExpression)PyAstCaseClause.super.getGuardCondition();
}
@ApiStatus.Internal
@Nullable PyType getSubjectTypeAfter(@NotNull TypeEvalContext context);
}
@@ -6,5 +6,5 @@ import com.jetbrains.python.ast.findChildrenByClass
interface PySequencePattern : PyAstSequencePattern, PyPattern {
val elements: List<PyPattern>
get() = findChildrenByClass(PyPattern::class.java).toList()
get() = findChildrenByClass(PyPattern::class.java).asList()
}
@@ -252,18 +252,11 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
if (subject == null) return;
// allowAnyExpr is here because we need negative edges with Never even when subject is not reference expression
pushAssertion(subject, true, true, true, context -> {
PyType subjectType = context.getType(subject);
for (PyCaseClause cs : matchStatement.getCaseClauses()) {
if (cs.getPattern() == null) continue;
if (cs.getGuardCondition() != null) continue;
if (cs.getPattern().isIrrefutable()) {
subjectType = PyNeverType.NEVER;
break;
}
subjectType = Ref.deref(createAssertionType(subjectType, context.getType(cs.getPattern()), false, true, context));
List<PyCaseClause> clauses = matchStatement.getCaseClauses();
if (!clauses.isEmpty()) {
return clauses.getLast().getSubjectTypeAfter(context);
}
return subjectType;
return null;
});
}
@@ -1,12 +1,10 @@
package com.jetbrains.python.psi.impl
import com.intellij.lang.ASTNode
import com.intellij.openapi.util.Ref
import com.intellij.psi.util.PsiTreeUtil
import com.jetbrains.python.codeInsight.controlflow.PyTypeAssertionEvaluator
import com.jetbrains.python.psi.PyCaseClause
import com.jetbrains.python.psi.PyElementVisitor
import com.jetbrains.python.psi.PyMatchStatement
import com.jetbrains.python.psi.PyPattern
import com.jetbrains.python.psi.*
import com.jetbrains.python.psi.types.PyNeverType
import com.jetbrains.python.psi.types.PyType
import com.jetbrains.python.psi.types.TypeEvalContext
@@ -16,20 +14,41 @@ class PyCaseClauseImpl(astNode: ASTNode?) : PyElementImpl(astNode), PyCaseClause
}
override fun getCaptureTypeForChild(pattern: PyPattern, context: TypeEvalContext): PyType? {
val matchStatement = getParent() as? PyMatchStatement ?: return null
val subject = matchStatement.subject ?: return null
return getSubjectTypeBefore(context)
}
var subjectType = context.getType(subject)
for (cs in matchStatement.caseClauses) {
if (cs === this) break
if (cs.pattern == null) continue
if (cs.guardCondition != null && !PyEvaluator.evaluateAsBoolean(cs.guardCondition, false)) continue
if (cs.pattern!!.canExcludePatternType(context)) {
subjectType = Ref.deref(
PyTypeAssertionEvaluator.createAssertionType(subjectType, context.getType(cs.pattern!!), false, true, context))
private fun getSubjectTypeBefore(context: TypeEvalContext): PyType? {
val prevClause = PsiTreeUtil.getPrevSiblingOfType(this, PyCaseClause::class.java)
if (prevClause != null) {
return prevClause.getSubjectTypeAfter(context)
}
else {
val matchStatement = parent as? PyMatchStatement ?: return null
val subject = matchStatement.subject ?: return null
return context.getType(subject)
}
}
override fun getSubjectTypeAfter(context: TypeEvalContext): PyType? {
fun getSubjectTypeAfterNoCache(): PyType? {
val beforeType = getSubjectTypeBefore(context)
val pattern = pattern ?: return beforeType
if (guardCondition != null && !PyEvaluator.evaluateAsBoolean(guardCondition, false)) {
return beforeType
}
// because subject can be 'Any', and then negative narrowing won't help
if (pattern.isIrrefutable) return PyNeverType.NEVER
if (pattern.canExcludePatternType(context)) {
val patternType = context.getType(pattern)
val narrowing = PyTypeAssertionEvaluator.createAssertionType(beforeType, patternType, false, true, context)
if (narrowing != null) {
return narrowing.get()
}
}
return beforeType
}
return subjectType
return PyUtil.getNullableParameterizedCachedValue(this, context) { getSubjectTypeAfterNoCache() }
}
}
@@ -70,7 +70,7 @@ public class PyClassPatternImpl extends PyElementImpl implements PyClassPattern,
if (classQName != null && SPECIAL_BUILTINS.contains(classQName)) {
if (arguments.isEmpty()) return true;
if (arguments.size() > 1) return false;
return arguments.getFirst().canExcludePatternType(context);
return isExhaustive(arguments.getFirst(), context);
}
List<String> matchArgs = getMatchArgs(classType, context);
@@ -81,14 +81,14 @@ public class PyClassPatternImpl extends PyElementImpl implements PyClassPattern,
return false;
}
final var valuePattern = keywordPattern.getValuePattern();
if (valuePattern != null && !canExcludeArgumentPatternType(valuePattern, context)) {
if (valuePattern != null && !isExhaustive(valuePattern, context)) {
return false;
}
}
else {
if (matchArgs == null) return false;
if (i >= matchArgs.size()) return false;
if (!canExcludeArgumentPatternType(member, context)) {
if (!isExhaustive(member, context)) {
return false;
}
}
@@ -100,7 +100,7 @@ public class PyClassPatternImpl extends PyElementImpl implements PyClassPattern,
@Override
public @Nullable PyType getCaptureTypeForChild(@NotNull PyPattern pattern, @NotNull TypeEvalContext context) {
pattern = as(PsiTreeUtil.findFirstParent(pattern, el -> this.getArgumentList() == el.getParent()), PyPattern.class);
pattern = as(PsiTreeUtil.findFirstParent(pattern, el -> getArgumentList() == el.getParent()), PyPattern.class);
if (pattern == null) return null;
if (pattern instanceof PyKeywordPattern keywordPattern) {
@@ -137,15 +137,16 @@ public class PyClassPatternImpl extends PyElementImpl implements PyClassPattern,
}).collect(PyTypeUtil.toUnion());
}
static boolean canExcludeArgumentPatternType(@NotNull PyPattern pattern, @NotNull TypeEvalContext context) {
/**
* Checks if the pattern covers its entire capture type and is itself exhaustive (e.g., not a non-literal value pattern).
*/
public static boolean isExhaustive(@NotNull PyPattern pattern, @NotNull TypeEvalContext context) {
if (!pattern.canExcludePatternType(context)) return false;
final var captureType = PyCaptureContext.getCaptureType(pattern, context);
final var patternType = context.getType(pattern);
// For class pattern arguments, we need to ensure that the argument pattern covers its capture type fully
if (Ref.deref(PyTypeAssertionEvaluator.createAssertionType(captureType, patternType, false, true, context)) instanceof PyNeverType) {
// in case the argument pattern is also class pattern with arguments
return pattern.canExcludePatternType(context);
}
return false;
// For composite pattern components, we need to ensure that the component pattern covers its capture type fully
return Ref.deref(PyTypeAssertionEvaluator.createAssertionType(captureType, patternType, false, true, context)) instanceof PyNeverType;
}
public static @Nullable List<@NotNull String> getMatchArgs(@NotNull PyClassType type, @NotNull TypeEvalContext context) {
@@ -21,6 +21,6 @@ public class PyDoubleStarPatternImpl extends PyElementImpl implements PyDoubleSt
@Override
public @Nullable PyType getType(@NotNull TypeEvalContext context, TypeEvalContext.@NotNull Key key) {
return null;
return PyCaptureContext.getCaptureType(this, context);
}
}
@@ -3,6 +3,7 @@ package com.jetbrains.python.psi.impl
import com.intellij.lang.ASTNode
import com.intellij.psi.PsiListLikeElement
import com.intellij.psi.util.findParentInFile
import com.jetbrains.python.ast.findChildrenByClass
import com.jetbrains.python.psi.PyDoubleStarPattern
import com.jetbrains.python.psi.PyElementVisitor
import com.jetbrains.python.psi.PyKeyValuePattern
@@ -24,18 +25,23 @@ import com.jetbrains.python.psi.types.PyUnionType
import com.jetbrains.python.psi.types.TypeEvalContext
class PyMappingPatternImpl(astNode: ASTNode?) : PyElementImpl(astNode), PyMappingPattern, PyCaptureContext, PsiListLikeElement {
private val elements: List<PyPattern>
get() = findChildrenByClass(PyPattern::class.java).asList()
override fun acceptPyVisitor(pyVisitor: PyElementVisitor) {
pyVisitor.visitPyMappingPattern(this)
}
override fun getComponents(): List<PyKeyValuePattern> = findChildrenByClass(PyKeyValuePattern::class.java).toList()
override fun getComponents(): List<PyPattern> = elements
override fun canExcludePatternType(context: TypeEvalContext): Boolean = false
override fun canExcludePatternType(context: TypeEvalContext): Boolean {
return elements.size == 1 && elements[0] is PyDoubleStarPattern
}
override fun getType(context: TypeEvalContext, key: TypeEvalContext.Key): PyType? {
val keyTypes = mutableListOf<PyType?>()
val valueTypes = mutableListOf<PyType?>()
for (it in components) {
for (it in elements.filterIsInstance<PyKeyValuePattern>()) {
keyTypes.add(context.getType(it.keyPattern))
if (it.valuePattern != null) {
valueTypes.add(context.getType(it.valuePattern!!))
@@ -81,7 +87,7 @@ class PyMappingPatternImpl(astNode: ASTNode?) : PyElementImpl(astNode), PyMappin
private fun PyType?.getValueType(sequenceMember: PyKeyValuePattern, context: TypeEvalContext): PyType? {
if (this is PyTypedDictType) {
val key = sequenceMember.getKeyString(context)
if (key != null) return this.getElementType(key)
if (key != null) return getElementType(key)
}
val mappingType = PyTypeUtil.convertToType(this, "typing.Mapping", sequenceMember, context)
?: return PyNeverType.NEVER
@@ -24,4 +24,9 @@ public class PyOrPatternImpl extends PyElementImpl implements PyOrPattern {
public @Nullable PyType getType(@NotNull TypeEvalContext context, TypeEvalContext.@NotNull Key key) {
return PyUnionType.union(ContainerUtil.map(getAlternatives(), it -> context.getType(it)));
}
@Override
public boolean canExcludePatternType(@NotNull TypeEvalContext context) {
return ContainerUtil.all(getAlternatives(), it -> it.canExcludePatternType(context));
}
}
@@ -31,7 +31,7 @@ class PySequencePatternImpl(astNode: ASTNode?) : PyElementImpl(astNode), PySeque
override fun getComponents(): List<PyPattern> = elements
override fun getType(context: TypeEvalContext, key: TypeEvalContext.Key): PyType? {
val sequenceCaptureType = this.getSequenceCaptureType(context)
val sequenceCaptureType = getSequenceCaptureType(context)
val types = elements.flatMap { pattern ->
when (pattern) {
is PySingleStarPattern -> pattern.getCapturedTypesFromSequenceType(sequenceCaptureType, context)
@@ -51,13 +51,15 @@ class PySequencePatternImpl(astNode: ASTNode?) : PyElementImpl(astNode), PySeque
}
override fun canExcludePatternType(context: TypeEvalContext): Boolean {
val allElementsCoverCapture = elements.all { PyClassPatternImpl.canExcludeArgumentPatternType(it, context) }
val allCapturesOfThisAreHeteroTuples = this.getSequenceCaptureType(context).toList().all { it.isHeterogeneousTuple() }
return allElementsCoverCapture && allCapturesOfThisAreHeteroTuples
if (elements.size == 1 && elements[0] is PySingleStarPattern) return true
val allElementsExhaustive = elements.all { PyClassPatternImpl.isExhaustive(it, context) }
if (!allElementsExhaustive) return false
val captureType = getSequenceCaptureType(context) ?: return false
return captureType.toList().all { type -> type.isHeterogeneousTuple() }
}
override fun getCaptureTypeForChild(pattern: PyPattern, context: TypeEvalContext): PyType? {
val sequenceType = this.getSequenceCaptureType(context) ?: return null
val sequenceType = getSequenceCaptureType(context) ?: return null
// This is done to skip group- and as-patterns
val sequenceMember = pattern.findParentInFile(withSelf = true) { el -> this === el.parent }
@@ -82,7 +84,7 @@ class PySequencePatternImpl(astNode: ASTNode?) : PyElementImpl(astNode), PySeque
return sequence.getElementType(idx)
}
else {
val starSpan = sequence.elementCount - this.elements.size
val starSpan = sequence.elementCount - elements.size
return sequence.getElementType(idx + starSpan)
}
}
@@ -123,7 +125,7 @@ class PySequencePatternImpl(astNode: ASTNode?) : PyElementImpl(astNode), PySeque
private fun PySingleStarPattern.getCapturedTypesFromSequenceType(sequenceType: PyType?, context: TypeEvalContext): List<PyType?> {
if (sequenceType.isHeterogeneousTuple()) {
val sequenceParent = this.parent as? PySequencePattern ?: return listOf()
val sequenceParent = parent as? PySequencePattern ?: return listOf()
val idx = sequenceParent.elements.indexOf(this)
return sequenceType.elementTypes.subList(idx, idx + sequenceType.elementCount - sequenceParent.elements.size + 1)
}
@@ -144,14 +146,14 @@ private fun PyType?.isHeterogeneousTuple(): Boolean {
}
private fun PyTupleType.takeIfSizeMatches(desiredSize: Int, hasStar: Boolean): PyTupleType? {
if (this.elementTypes.any { it is PyUnpackedTupleType }) {
val variadicElementsCount: Int = desiredSize - this.elementTypes.size + 1
if (elementTypes.any { it is PyUnpackedTupleType }) {
val variadicElementsCount: Int = desiredSize - elementTypes.size + 1
if (variadicElementsCount >= 0) {
return this.expandVariadics(variadicElementsCount)
return expandVariadics(variadicElementsCount)
}
}
else {
if (hasStar && desiredSize <= this.elementCount || desiredSize == this.elementCount) {
if (hasStar && desiredSize <= elementCount || desiredSize == elementCount) {
return this
}
}
@@ -159,7 +161,7 @@ private fun PyTupleType.takeIfSizeMatches(desiredSize: Int, hasStar: Boolean): P
}
private fun PyTupleType.expandVariadics(variadicElementCount: Int): PyTupleType {
require(!this.isHomogeneous) { "Supplied tuple must not be homogeneous: $this" }
require(!isHomogeneous) { "Supplied tuple must not be homogeneous: $this" }
require(variadicElementCount >= 0) { "Supplied variadic element count must not be negative: $variadicElementCount" }
val unpackedTupleIndex = elementTypes.indexOfFirst { it is PyUnpackedTupleType }
@@ -173,7 +175,7 @@ private fun PyTupleType.expandVariadics(variadicElementCount: Int): PyTupleType
}
addAll(elementTypes.subList(unpackedTupleIndex + 1, elementTypes.size))
}
return PyTupleType(this.pyClass, adjustedTupleElementTypes, false)
return PyTupleType(pyClass, adjustedTupleElementTypes, false)
}
private fun intersect(type1: PyType?, type2: PyType?, context: TypeEvalContext): PyType? {
@@ -21,6 +21,6 @@ public class PySingleStarPatternImpl extends PyElementImpl implements PySingleSt
@Override
public @Nullable PyType getType(@NotNull TypeEvalContext context, TypeEvalContext.@NotNull Key key) {
return null;
return PyCaptureContext.getCaptureType(this, context);
}
}
@@ -175,29 +175,6 @@ match m:
""");
}
public void testMatchSequencePatternAlreadyNarrowerOuter() {
doTestByText("""
from typing import assert_type
from typing import Sequence
m: Sequence[object]
match m:
case [1, True]:
assert_type(m, Sequence[int | bool])
""");
}
public void testMatchSequencePatternAlreadyNarrowerBoth() {
doTestByText("""
from typing import assert_type
from typing import Sequence
m: Sequence[bool]
match m:
case [1, True]:
assert_type(m, Sequence[bool])
""");
}
public void testMatchSequencePatternNarrowSubjectItems() {
doTestByText("""
@@ -700,23 +677,6 @@ match m:
);
}
public void testMatchClassPatternCapture() {
doTestByText("""
from typing import assert_type
class A:
__match_args__ = ("a", "b")
a: str
b: int
m: A
match m:
case A(i, j):
assert_type(i, str)
assert_type(j, int)
""");
}
public void testMatchClassPatternCaptureDataclass() {
doTestByText("""
@@ -1117,4 +1077,142 @@ match p2:
assert_type(color_val, str)
""");
}
// PY-85990
public void testClassAttributeValuePatternDoesNotNarrow() {
doTestByText("""
from typing import assert_type
class A:
s = "s"
def func(x: str):
match x:
case A.s:
return
assert_type(x, str)
""");
}
// PY-85990
public void testOrPatternWithNonLiteralDoesNotNarrow() {
doTestByText("""
from typing import assert_type
class A:
s = "s"
def func(x: str):
match x:
case "literal" | A.s:
return
assert_type(x, str)
""");
}
// PY-85990
public void testSequencePatternWithNonExhaustiveElementDoesNotNarrow() {
doTestByText("""
from typing import assert_type
class A:
s = "s"
def func(x: list[str]):
match x:
case [A.s]:
return
assert_type(x, list[str])
""");
}
// PY-85990
public void testMappingPatternWithDoubleStarIsExhaustive() {
doTestByText("""
from typing import assert_type, Never
def func(x: dict[str, int]):
match x:
case {**rest}:
assert_type(x, dict[str, int])
assert_type(rest, dict[str, int])
return
assert_type(x, Never)
""");
}
// PY-85990
public void testSequencePatternWithSingleStarIsExhaustive() {
doTestByText("""
from typing import assert_type, Never
def func(x: list[int]):
match x:
case [*rest]:
assert_type(x, list[int])
assert_type(rest, list[int])
return
assert_type(x, Never)
""");
}
// PY-85990
public void testClassPatternWithNonExhaustiveArgumentDoesNotNarrow() {
doTestByText("""
from typing import assert_type
class B:
s = "s"
class A:
__match_args__ = ("a",)
a: str
def func(x: A):
match x:
case A(B.s):
return
assert_type(x, A)
""");
}
// PY-85990
public void testClassPatternWithNonExhaustiveKeywordArgumentDoesNotNarrow() {
doTestByText("""
from typing import assert_type
class B:
s = "s"
class A:
a: str
def func(x: A):
match x:
case A(a=B.s):
return
assert_type(x, A)
""");
}
// PY-85990
public void testSpecialBuiltinWithNonExhaustiveArgumentDoesNotNarrow() {
doTestByText("""
from typing import assert_type
def func(x: str):
match x:
case str("literal"):
return
assert_type(x, str)
""");
}
}