diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PyCaseClause.java b/python/python-psi-api/src/com/jetbrains/python/psi/PyCaseClause.java index 036170f28dee..302ceb42d4cc 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/PyCaseClause.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PyCaseClause.java @@ -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); } diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/PySequencePattern.kt b/python/python-psi-api/src/com/jetbrains/python/psi/PySequencePattern.kt index 0228670eb2ef..0a09b9f9ffcb 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/PySequencePattern.kt +++ b/python/python-psi-api/src/com/jetbrains/python/psi/PySequencePattern.kt @@ -6,5 +6,5 @@ import com.jetbrains.python.ast.findChildrenByClass interface PySequencePattern : PyAstSequencePattern, PyPattern { val elements: List - get() = findChildrenByClass(PyPattern::class.java).toList() + get() = findChildrenByClass(PyPattern::class.java).asList() } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java index d7c9a3006027..ace9a3231e40 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java @@ -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 clauses = matchStatement.getCaseClauses(); + if (!clauses.isEmpty()) { + return clauses.getLast().getSubjectTypeAfter(context); } - - return subjectType; + return null; }); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCaseClauseImpl.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCaseClauseImpl.kt index fb106c5c1cb3..0a86d04b9d4d 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCaseClauseImpl.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCaseClauseImpl.kt @@ -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() } } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java index 28ae17ad3a05..9bcdd8f0a17d 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassPatternImpl.java @@ -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 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) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyDoubleStarPatternImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyDoubleStarPatternImpl.java index 64fe3beae3c1..b6cf339d4bcd 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyDoubleStarPatternImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyDoubleStarPatternImpl.java @@ -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); } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyMappingPatternImpl.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyMappingPatternImpl.kt index 9714cb75f7e7..86383cb017dc 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyMappingPatternImpl.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyMappingPatternImpl.kt @@ -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 + get() = findChildrenByClass(PyPattern::class.java).asList() + override fun acceptPyVisitor(pyVisitor: PyElementVisitor) { pyVisitor.visitPyMappingPattern(this) } - override fun getComponents(): List = findChildrenByClass(PyKeyValuePattern::class.java).toList() + override fun getComponents(): List = 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() val valueTypes = mutableListOf() - for (it in components) { + for (it in elements.filterIsInstance()) { 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 diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyOrPatternImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyOrPatternImpl.java index ca5f3ba159aa..df2a126f3cbe 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyOrPatternImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyOrPatternImpl.java @@ -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)); + } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySequencePatternImpl.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySequencePatternImpl.kt index 8c9f5a849e17..61b1eb843e5f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySequencePatternImpl.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySequencePatternImpl.kt @@ -31,7 +31,7 @@ class PySequencePatternImpl(astNode: ASTNode?) : PyElementImpl(astNode), PySeque override fun getComponents(): List = 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 { 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? { diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySingleStarPatternImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySingleStarPatternImpl.java index 7156f7fe8281..1dbf720c9aca 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySingleStarPatternImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySingleStarPatternImpl.java @@ -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); } } diff --git a/python/testSrc/com/jetbrains/python/PyPatternTypeTest.java b/python/testSrc/com/jetbrains/python/PyPatternTypeTest.java index 001ea3bbb29c..d2c16ecf6e2a 100644 --- a/python/testSrc/com/jetbrains/python/PyPatternTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyPatternTypeTest.java @@ -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) + """); + } } \ No newline at end of file