From 8c745a5c229b4a0e87b07c0618f824d1b7c467cd Mon Sep 17 00:00:00 2001 From: "Aleksandr.Govenko" Date: Fri, 17 Oct 2025 14:55:29 +0200 Subject: [PATCH] PY-83339 Narrow `X | None` type to `X` after `assert x` (cherry picked from commit a49863065bf26f615100b21f423a650cbebbcab3) IJ-MR-178978 GitOrigin-RevId: 4bf02cab1e94a2ab1aaa70cb0b4578688cea88b9 --- .../controlflow/PyTypeAssertionEvaluator.java | 24 ++++++------------- .../codeInsight/controlflow/WithAssert.txt | 15 ++++++------ .../com/jetbrains/python/Py3TypeTest.java | 9 +++++++ 3 files changed, 24 insertions(+), 24 deletions(-) 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 5428520403eb..597a4389013a 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 @@ -62,11 +62,7 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor { } private void visitExpressionInCondition(@NotNull PyExpression node) { - if (myPositive && ( - isIfReferenceStatement(node) || - isIfReferenceConditionalStatement(node) || - isBinaryExpressionPart(node) - )) { + if (myPositive && isReferenceInTruthyCondition(node)) { // TODO: we can actually check if the class defines __bool__ or __len__, and use it to exclude the type // we could not suggest `None` because it could be a reference to an empty collection // so we could push only non-`None` assertions @@ -403,19 +399,13 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor { return null; } - private static boolean isIfReferenceStatement(@NotNull PyExpression node) { - return skipNotAndParens(node) instanceof PyConditionalStatementPart; - } - - private static boolean isIfReferenceConditionalStatement(@NotNull PyExpression node) { + private static boolean isReferenceInTruthyCondition(@NotNull PyExpression node) { final PsiElement parent = skipNotAndParens(node); - return parent instanceof PyConditionalExpression cond && - PsiTreeUtil.isAncestor(cond.getCondition(), node, false); - } - - private static boolean isBinaryExpressionPart(@NotNull PyExpression node) { - return skipNotAndParens(node) instanceof PyBinaryExpression binExpr && - (binExpr.isOperator(PyNames.AND) || binExpr.isOperator(PyNames.OR)); + if (parent instanceof PyConditionalStatementPart) return true; + if (parent instanceof PyConditionalExpression cond && PsiTreeUtil.isAncestor(cond.getCondition(), node, false)) return true; + if (parent instanceof PyBinaryExpression binExpr && (binExpr.isOperator(PyNames.AND) || binExpr.isOperator(PyNames.OR))) return true; + if (parent instanceof PyAssertStatement) return true; + return false; } static class Assertion { diff --git a/python/testData/codeInsight/controlflow/WithAssert.txt b/python/testData/codeInsight/controlflow/WithAssert.txt index ea7a016b21b6..7487a7d65f14 100644 --- a/python/testData/codeInsight/controlflow/WithAssert.txt +++ b/python/testData/codeInsight/controlflow/WithAssert.txt @@ -1,15 +1,16 @@ 0(1) element: null 1(2) element: PyWithStatement 2(4) READ ACCESS: context_manager -3(13) exit context manager: context_manager +3(14) exit context manager: context_manager 4(5,3) element: PyAssertStatement 5(6,3) READ ACCESS: True 6(7,3) READ ACCESS: f 7(8,3) element: PyCallExpression: f 8(9,10,3) READ ACCESS: True -9(11,3) element: null. Condition: True:false -10(12,3) element: null. Condition: True:true -11(14,3) raise: PyAssertStatement -12(13,3) element: PyPrintStatement -13(14) element: PyPrintStatement -14() element: null \ No newline at end of file +9(12,3) element: null. Condition: True:false +10(11) element: null. Condition: True:true +11(3,13) ASSERTTYPE ACCESS: True +12(15,3) raise: PyAssertStatement +13(14,3) element: PyPrintStatement +14(15) element: PyPrintStatement +15() element: null \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index bfe87aee2b7c..eb11ce694b84 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -4170,6 +4170,15 @@ public class Py3TypeTest extends PyTestCase { """); } + @TestFor(issues="PY-83339") + public void testAssertNarrowsOptionalAfterAssert() { + doTest("int", """ + def foo(param: int | None): + assert param + expr = param + """); + } + private void doTest(final String expectedType, final String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class);