diff --git a/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java b/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java index a5643afff75b..715d619d425f 100644 --- a/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java +++ b/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java @@ -23,6 +23,7 @@ import com.jetbrains.python.PyBundle; import com.jetbrains.python.PyTokenTypes; import com.jetbrains.python.psi.PyBinaryExpression; import com.jetbrains.python.psi.PyElementGenerator; +import com.jetbrains.python.psi.PyElementType; import com.jetbrains.python.psi.PyExpression; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -37,6 +38,7 @@ import static com.jetbrains.python.psi.PyUtil.as; * For instance, a < b and b < c --> a < b < c */ public class ChainedComparisonsQuickFix implements LocalQuickFix { + boolean myIsLeftLeft; boolean myIsRightLeft; boolean getInnerRight; @@ -69,12 +71,11 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { if (expression != null && expression.isWritable()) { final PyBinaryExpression rightExpression = as(expression.getRightExpression(), PyBinaryExpression.class); PyBinaryExpression leftExpression = as(expression.getLeftExpression(), PyBinaryExpression.class); - if (rightExpression != null && leftExpression != null && expression.getOperator() == PyTokenTypes.AND_KEYWORD) { - if (getInnerRight && leftExpression.getRightExpression() instanceof PyBinaryExpression && - leftExpression.getOperator() == PyTokenTypes.AND_KEYWORD) { + if (rightExpression != null && leftExpression != null && isLogicalAndExpression(expression)) { + if (getInnerRight && leftExpression.getRightExpression() instanceof PyBinaryExpression && isLogicalAndExpression(leftExpression)) { leftExpression = (PyBinaryExpression)leftExpression.getRightExpression(); } - checkOperator(leftExpression, rightExpression, project); + checkOperator(assertNotNull(leftExpression), rightExpression, project); } } } @@ -126,7 +127,7 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { @NotNull private static PsiElement getLeftestOperator(@NotNull PyBinaryExpression expression) { PsiElement op = expression.getPsiOperator(); - while (expression.getLeftExpression() instanceof PyBinaryExpression) { + while (isComparisonExpression(expression.getLeftExpression())) { expression = (PyBinaryExpression)expression.getLeftExpression(); op = expression.getPsiOperator(); } @@ -139,7 +140,7 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { final PsiElement operator = leftExpression.getPsiOperator(); final PyExpression right = leftExpression.getRightExpression(); PyExpression left = leftExpression.getLeftExpression(); - if (left instanceof PyBinaryExpression) { + if (isComparisonExpression(left)) { left = invertExpression((PyBinaryExpression)left, elementGenerator); } final String newOperator = invertOperator(assertNotNull(operator)); @@ -169,7 +170,7 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { PyExpression left = expression.getLeftExpression(); PyExpression right = expression.getRightExpression(); PsiElement operator = expression.getPsiOperator(); - while (left instanceof PyBinaryExpression) { + while (isComparisonExpression(left)) { assert operator != null; right = elementGenerator.createBinaryExpression(operator.getText(), ((PyBinaryExpression)left).getRightExpression(), right); operator = ((PyBinaryExpression)left).getPsiOperator(); @@ -177,4 +178,16 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { } return right; } + + private static boolean isComparisonExpression(@Nullable PyExpression expression) { + if (!(expression instanceof PyBinaryExpression)) { + return false; + } + final PyElementType operator = ((PyBinaryExpression)expression).getOperator(); + return PyTokenTypes.RELATIONAL_OPERATIONS.contains(operator) || PyTokenTypes.EQUALITY_OPERATIONS.contains(operator); + } + + private static boolean isLogicalAndExpression(@Nullable PyExpression expression) { + return expression instanceof PyBinaryExpression && ((PyBinaryExpression)expression).getOperator() == PyTokenTypes.AND_KEYWORD; + } } diff --git a/python/testData/inspections/ChainedComparisonWithCommonBinaryExpression.py b/python/testData/inspections/ChainedComparisonWithCommonBinaryExpression.py new file mode 100644 index 000000000000..fad398623788 --- /dev/null +++ b/python/testData/inspections/ChainedComparisonWithCommonBinaryExpression.py @@ -0,0 +1,2 @@ +if 3 < x + 1 and x + 1 < 10: + pass \ No newline at end of file diff --git a/python/testData/inspections/ChainedComparisonWithCommonBinaryExpression_after.py b/python/testData/inspections/ChainedComparisonWithCommonBinaryExpression_after.py new file mode 100644 index 000000000000..cf13465ac226 --- /dev/null +++ b/python/testData/inspections/ChainedComparisonWithCommonBinaryExpression_after.py @@ -0,0 +1,2 @@ +if 3 < x + 1 < 10: + pass \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/PyQuickFixTest.java b/python/testSrc/com/jetbrains/python/PyQuickFixTest.java index 811eb0f679a3..8582e692de50 100644 --- a/python/testSrc/com/jetbrains/python/PyQuickFixTest.java +++ b/python/testSrc/com/jetbrains/python/PyQuickFixTest.java @@ -251,6 +251,12 @@ public class PyQuickFixTest extends PyTestCase { PyBundle.message("QFIX.chained.comparison"), true, true); } + // PY-14002 + public void testChainedComparisonWithCommonBinaryExpression() { + doInspectionTest("ChainedComparisonWithCommonBinaryExpression.py", PyChainedComparisonsInspection.class, + PyBundle.message("QFIX.chained.comparison"), true, true); + } + public void testStatementEffect() { // PY-1362, PY-2585 doInspectionTest("StatementEffect.py", PyStatementEffectInspection.class, PyBundle.message("QFIX.statement.effect"), true, true);