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
+ """)
}