diff --git a/python/python-ast/src/com/jetbrains/python/ast/PyAstBinaryExpression.java b/python/python-ast/src/com/jetbrains/python/ast/PyAstBinaryExpression.java index 211083b45208..fd2b70b1254e 100644 --- a/python/python-ast/src/com/jetbrains/python/ast/PyAstBinaryExpression.java +++ b/python/python-ast/src/com/jetbrains/python/ast/PyAstBinaryExpression.java @@ -44,6 +44,11 @@ public interface PyAstBinaryExpression extends PyAstQualifiedExpression, PyAstCa return null; } + default boolean isShortCircuit() { + var operator = getOperator(); + return operator == PyTokenTypes.AND_KEYWORD || operator == PyTokenTypes.OR_KEYWORD; + } + default boolean isOperator(String chars) { ASTNode child = getNode().getFirstChildNode(); StringBuilder buf = new StringBuilder(); diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspection.kt index f13a30ca9a2c..3d91278e3b4f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspection.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspection.kt @@ -11,6 +11,7 @@ import com.intellij.psi.PsiElementVisitor import com.jetbrains.python.PyNames import com.jetbrains.python.PyPsiBundle import com.jetbrains.python.PyTokenTypes +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider import com.jetbrains.python.psi.LanguageLevel import com.jetbrains.python.psi.PyAssertStatement import com.jetbrains.python.psi.PyBinaryExpression @@ -20,7 +21,8 @@ import com.jetbrains.python.psi.PyExpression import com.jetbrains.python.psi.PyIfStatement import com.jetbrains.python.psi.PyPrefixExpression import com.jetbrains.python.psi.PyWhileStatement -import com.jetbrains.python.psi.types.PyABCUtil +import com.jetbrains.python.psi.impl.PyPsiUtils +import com.jetbrains.python.psi.types.PyClassType import com.jetbrains.python.psi.types.TypeEvalContext @@ -33,63 +35,60 @@ class PySuspiciousBooleanConditionInspection : PyInspection() { private class Visitor(holder: ProblemsHolder, context: TypeEvalContext) : PyInspectionVisitor(holder, context) { override fun visitPyIfStatement(node: PyIfStatement) { - super.visitPyIfStatement(node) - checkBooleanContext(node.ifPart.condition) - node.elifParts.forEach { checkBooleanContext(it.condition) } + node.ifPart.condition.checkBooleanContext() + node.elifParts.forEach { it.condition.checkBooleanContext() } } override fun visitPyWhileStatement(node: PyWhileStatement) { - super.visitPyWhileStatement(node) - checkBooleanContext(node.whilePart.condition) + node.whilePart.condition.checkBooleanContext() } override fun visitPyConditionalExpression(node: PyConditionalExpression) { - super.visitPyConditionalExpression(node) - checkBooleanContext(node.condition) + node.condition.checkBooleanContext() } override fun visitPyBinaryExpression(node: PyBinaryExpression) { - super.visitPyBinaryExpression(node) - // Check operands of 'and' and 'or' boolean operators - if (node.operator == PyTokenTypes.AND_KEYWORD || node.operator == PyTokenTypes.OR_KEYWORD) { - checkExpression(node.leftExpression) - checkExpression(node.rightExpression) + if (node.isShortCircuit) { + node.leftExpression.checkBooleanContext() + node.rightExpression.checkBooleanContext() } } override fun visitPyPrefixExpression(node: PyPrefixExpression) { - super.visitPyPrefixExpression(node) - // Check operand of 'not' operator if (node.operator == PyTokenTypes.NOT_KEYWORD) { - checkExpression(node.operand) + node.operand.checkBooleanContext() } } override fun visitPyAssertStatement(node: PyAssertStatement) { - super.visitPyAssertStatement(node) - // Check the first argument (the condition being asserted) - val arguments = node.arguments - if (arguments.isNotEmpty()) { - checkExpression(arguments[0]) + node.arguments.firstOrNull()?.checkBooleanContext() + } + + private fun PyExpression?.checkBooleanContext() { + if (this == null) { + return } + val node = PyPsiUtils.flattenParens(this) + if (node is PyBinaryExpression && node.isShortCircuit) { + return + } + this.checkExpression() } - private fun checkBooleanContext(condition: PyExpression?) { - condition ?: return - checkExpression(condition) - } - - private fun checkExpression(expr: PyExpression?) { - expr ?: return + private fun PyExpression?.checkExpression() { + if (this == null) { + return + } // Don't flag awaited expressions - they're already being awaited - if (expr is PyPrefixExpression && expr.operator == PyTokenTypes.AWAIT_KEYWORD) return + if (this is PyPrefixExpression && this.operator == PyTokenTypes.AWAIT_KEYWORD) return - val type = myTypeEvalContext.getType(expr) ?: return + val type = myTypeEvalContext.getType(this) ?: return - // Check if the type is a coroutine/awaitable - if (PyABCUtil.isSubtype(type, PyNames.AWAITABLE, myTypeEvalContext)) { - registerProblem(expr, PyPsiBundle.message("INSP.suspicious.boolean.condition.coroutine"), PyAddAwaitQuickFix()) + // Check if the type is a coroutine + // TODO: use `CoroutineType` instead + if (type is PyClassType && type.classQName == PyTypingTypeProvider.COROUTINE) { + registerProblem(this, PyPsiBundle.message("INSP.suspicious.boolean.condition.coroutine"), PyAddAwaitQuickFix()) } } } diff --git a/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition.py b/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition.py index 0cefd70ac6ef..01deb21834a1 100644 --- a/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition.py +++ b/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition.py @@ -3,5 +3,5 @@ async def f() -> bool: async def main(): - if f(): + if f() or f(): print("hi") diff --git a/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition_after.py b/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition_after.py index 50e4302e7924..05ab7d788459 100644 --- a/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition_after.py +++ b/python/testData/quickFixes/PySuspiciousBooleanConditionQuickFixTest/addAwaitInIfCondition_after.py @@ -3,5 +3,5 @@ async def f() -> bool: async def main(): - if await f(): + if await f() or f(): print("hi") diff --git a/python/testSrc/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspectionTest.kt b/python/testSrc/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspectionTest.kt index c7972e2693d1..9565e8dbcc29 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspectionTest.kt +++ b/python/testSrc/com/jetbrains/python/inspections/PySuspiciousBooleanConditionInspectionTest.kt @@ -65,6 +65,24 @@ class PySuspiciousBooleanConditionInspectionTest : PyInspectionTestCase() { print("hi") """.trimIndent()) + fun `test coroutine in assert statement`() = doTestByText(""" + async def check() -> bool: + return True + + + async def main(): + assert check() + """.trimIndent()) + + fun `test coroutine with 'not' operator`() = doTestByText(""" + async def f() -> bool: + return True + + + async def main(): + assert not f() + """.trimIndent()) + fun `test coroutine variable with 'and' condition`() = doTestByText(""" async def f() -> bool: return True @@ -72,10 +90,35 @@ class PySuspiciousBooleanConditionInspectionTest : PyInspectionTestCase() { async def main(): coro = f() - if bool() and coro: - print("hi") + assert bool() and coro """.trimIndent()) + fun `test coroutine with 'or' condition`() = doTestByText(""" + async def f() -> bool: + return True + + + async def main(): + assert ( + f() + or f() + or f() + ) + """.trimIndent()) + + fun `test coroutine with bare expression`() = doTestByText(""" + async def f() -> bool: + return True + + + async def main(): + a = ( + f() + or f() + ) + a = not f() + """.trimIndent()) + fun `test awaited coroutine no warning`() = doTestByText(""" async def f() -> bool: return True @@ -96,32 +139,9 @@ class PySuspiciousBooleanConditionInspectionTest : PyInspectionTestCase() { print("hi") """.trimIndent()) - fun `test coroutine with 'or' condition`() = doTestByText(""" - async def f() -> bool: - return True - - - async def main(): - if f() or bool(): - print("hi") - """.trimIndent()) - - fun `test coroutine with 'not' operator`() = doTestByText(""" - async def f() -> bool: - return True - - - async def main(): - if not f(): - print("hi") - """.trimIndent()) - - fun `test coroutine in assert statement`() = doTestByText(""" - async def check() -> bool: - return True - - - async def main(): - assert check() - """.trimIndent()) + fun `test shape type no warning`() = doTestByText(""" + def f(val): + val[0] + assert val + """) }