diff --git a/python/src/com/jetbrains/python/actions/ChainedComparisonsQuickFix.java b/python/src/com/jetbrains/python/actions/ChainedComparisonsQuickFix.java index a9b806ce26d1..e0845c08ceaf 100644 --- a/python/src/com/jetbrains/python/actions/ChainedComparisonsQuickFix.java +++ b/python/src/com/jetbrains/python/actions/ChainedComparisonsQuickFix.java @@ -16,8 +16,11 @@ import org.jetbrains.annotations.NotNull; * For instance, a < b and b < c --> a < b < c */ public class ChainedComparisonsQuickFix implements LocalQuickFix { - - public ChainedComparisonsQuickFix() { + boolean myIsLeftLeft; + boolean myIsRightLeft; + public ChainedComparisonsQuickFix(boolean isLeft, boolean isRight) { + myIsLeftLeft = isLeft; + myIsRightLeft = isRight; } @NotNull @@ -45,33 +48,69 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { } } - static private void checkOperator(PyBinaryExpression leftExpression, + private void checkOperator(PyBinaryExpression leftExpression, PyBinaryExpression rightExpression, Project project) { - if (leftExpression.getRightExpression() instanceof PyBinaryExpression) { - checkOperator((PyBinaryExpression)leftExpression.getRightExpression(), rightExpression, project); - } - else if (/*leftExpression.getOperator() == rightExpression.getOperator() && */ - PyTokenTypes.RELATIONAL_OPERATIONS.contains(leftExpression.getOperator())) { - PyExpression leftRight = leftExpression.getRightExpression(); - if (leftRight != null) { - if (leftRight.getText().equals(getSmallLeftExpression(rightExpression).getText())) { - PyElementGenerator elementGenerator = PyElementGenerator.getInstance(project); - PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( - (rightExpression).getPsiOperator().getText(), leftExpression, - getLargeRightExpression(rightExpression, project)); - leftExpression.replace(binaryExpression); - rightExpression.delete(); - } + PyElementGenerator elementGenerator = PyElementGenerator.getInstance(project); + if (myIsLeftLeft) { + PyExpression newLeftExpression = invertExpression(leftExpression, elementGenerator); + + if (myIsRightLeft) { + PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( + (rightExpression).getPsiOperator().getText(), newLeftExpression, + getLargeRightExpression(rightExpression, project)); + leftExpression.replace(binaryExpression); + rightExpression.delete(); + } + else { + PsiElement op = rightExpression.getPsiOperator(); + String newOp = invertOperator(op); + PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( + newOp, newLeftExpression, rightExpression.getLeftExpression()); + leftExpression.replace(binaryExpression); + rightExpression.delete(); } } + else { + if (myIsRightLeft) { + PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( + (rightExpression).getPsiOperator().getText(), leftExpression, getLargeRightExpression(rightExpression, project)); + leftExpression.replace(binaryExpression); + rightExpression.delete(); + } + else { + PsiElement op = rightExpression.getPsiOperator(); + String newOp = invertOperator(op); + PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( + newOp, leftExpression, rightExpression.getLeftExpression()); + leftExpression.replace(binaryExpression); + rightExpression.delete(); + } + } + } - static private PyExpression getSmallLeftExpression(PyBinaryExpression expression) { - PyExpression result = expression; - while (result instanceof PyBinaryExpression) { - result = ((PyBinaryExpression)result).getLeftExpression(); + private PyExpression invertExpression(PyBinaryExpression leftExpression, PyElementGenerator elementGenerator) { + PsiElement op = leftExpression.getPsiOperator(); + PyExpression right = leftExpression.getRightExpression(); + PyExpression left = leftExpression.getLeftExpression(); + if (left instanceof PyBinaryExpression){ + left = invertExpression((PyBinaryExpression)left, elementGenerator); } - return result; + String newOp = invertOperator(op); + return elementGenerator.createBinaryExpression( + newOp, right, left); + } + + private String invertOperator(PsiElement op) { + if (op.getText().equals(">")) + return "<"; + if (op.getText().equals("<")) + return ">"; + if (op.getText().equals(">=")) + return "<="; + if (op.getText().equals("<=")) + return ">="; + return op.getText(); } static private PyExpression getLargeRightExpression(PyBinaryExpression expression, Project project) { diff --git a/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java b/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java index 6cdeb12f12c3..2cfe14bdac3d 100644 --- a/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java @@ -30,6 +30,8 @@ public class PyChainedComparisonsInspection extends PyInspection { } private static class Visitor extends PyInspectionVisitor { + boolean myIsLeft; + boolean myIsRight; public Visitor(final ProblemsHolder holder) { super(holder); @@ -43,24 +45,61 @@ public class PyChainedComparisonsInspection extends PyInspection { if (leftExpression instanceof PyBinaryExpression && rightExpression instanceof PyBinaryExpression) { if (node.getOperator() == PyTokenTypes.AND_KEYWORD) { - if (checkOperator((PyBinaryExpression)leftExpression, (PyBinaryExpression)rightExpression)) - registerProblem(node, "Simplify chained comparison", new ChainedComparisonsQuickFix()); + if (isRightSimplified((PyBinaryExpression)leftExpression, (PyBinaryExpression)rightExpression) || + isLeftSimplified((PyBinaryExpression)leftExpression, (PyBinaryExpression)rightExpression)) + registerProblem(node, "Simplify chained comparison", new ChainedComparisonsQuickFix(myIsLeft, myIsRight)); } } } - static private boolean checkOperator(PyBinaryExpression leftExpression,PyBinaryExpression rightExpression) { - if (leftExpression.getRightExpression() instanceof PyBinaryExpression) { - if (checkOperator((PyBinaryExpression)leftExpression.getRightExpression(), rightExpression)) + private boolean isRightSimplified(PyBinaryExpression leftExpression, PyBinaryExpression rightExpression) { + if (leftExpression.getRightExpression() instanceof PyBinaryExpression && + PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)leftExpression.getRightExpression()).getOperator())){ + if (isRightSimplified((PyBinaryExpression)leftExpression.getRightExpression(), rightExpression)) return true; } - if (/*leftExpression.getOperator() == rightExpression.getOperator() && */ - PyTokenTypes.RELATIONAL_OPERATIONS.contains(leftExpression.getOperator())) { + if (PyTokenTypes.RELATIONAL_OPERATIONS.contains(leftExpression.getOperator())) { PyExpression leftRight = leftExpression.getRightExpression(); if (leftRight != null) { - if (leftRight.getText().equals(getLeftExpression(rightExpression).getText())) + if (leftRight.getText().equals(getLeftExpression(rightExpression).getText())) { + myIsLeft = false; + myIsRight = true; return true; + } + + PyExpression right = getSmallestRight(rightExpression); + if (right != null && leftRight.getText().equals(right.getText())) { + myIsLeft = false; + myIsRight = false; + return true; + } + } + } + return false; + } + + private boolean isLeftSimplified(PyBinaryExpression leftExpression, PyBinaryExpression rightExpression) { + if (leftExpression.getLeftExpression() instanceof PyBinaryExpression && + PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)leftExpression.getLeftExpression()).getOperator())){ + if (isLeftSimplified((PyBinaryExpression)leftExpression.getLeftExpression(), rightExpression)) + return true; + } + + if (PyTokenTypes.RELATIONAL_OPERATIONS.contains(leftExpression.getOperator())) { + PyExpression leftRight = leftExpression.getLeftExpression(); + if (leftRight != null) { + if (leftRight.getText().equals(getLeftExpression(rightExpression).getText())) { + myIsLeft = true; + myIsRight = true; + return true; + } + PyExpression right = getSmallestRight(rightExpression); + if (right != null && leftRight.getText().equals(right.getText())) { + myIsLeft = true; + myIsRight = false; + return true; + } } } return false; @@ -73,5 +112,13 @@ public class PyChainedComparisonsInspection extends PyInspection { } return result; } + + static private PyExpression getSmallestRight(PyBinaryExpression expression) { + PyExpression result = expression; + while (result instanceof PyBinaryExpression) { + result = ((PyBinaryExpression)result).getRightExpression(); + } + return result; + } } } diff --git a/python/testData/inspections/ChainedComparison1.py b/python/testData/inspections/ChainedComparison1.py new file mode 100644 index 000000000000..0a166712908a --- /dev/null +++ b/python/testData/inspections/ChainedComparison1.py @@ -0,0 +1,2 @@ +if b >= a > e and b < c: + print "q" \ No newline at end of file diff --git a/python/testData/inspections/ChainedComparison1_after.py b/python/testData/inspections/ChainedComparison1_after.py new file mode 100644 index 000000000000..060bbda1968f --- /dev/null +++ b/python/testData/inspections/ChainedComparison1_after.py @@ -0,0 +1,2 @@ +if e < a <= b < c: + print "q" \ No newline at end of file diff --git a/python/testData/inspections/ChainedComparison2.py b/python/testData/inspections/ChainedComparison2.py new file mode 100644 index 000000000000..7ad69f265bb8 --- /dev/null +++ b/python/testData/inspections/ChainedComparison2.py @@ -0,0 +1,2 @@ +if b >= a > e and c < b: + print "q" \ No newline at end of file diff --git a/python/testData/inspections/ChainedComparison2_after.py b/python/testData/inspections/ChainedComparison2_after.py new file mode 100644 index 000000000000..5059778a064d --- /dev/null +++ b/python/testData/inspections/ChainedComparison2_after.py @@ -0,0 +1,2 @@ +if e < a <= b > c: + print "q" \ No newline at end of file diff --git a/python/testData/inspections/ChainedComparison3.py b/python/testData/inspections/ChainedComparison3.py new file mode 100644 index 000000000000..3c12548e2f0a --- /dev/null +++ b/python/testData/inspections/ChainedComparison3.py @@ -0,0 +1,2 @@ +if e >= a > b and c < b: + print "q" \ No newline at end of file diff --git a/python/testData/inspections/ChainedComparison3_after.py b/python/testData/inspections/ChainedComparison3_after.py new file mode 100644 index 000000000000..f5f2e205d6bf --- /dev/null +++ b/python/testData/inspections/ChainedComparison3_after.py @@ -0,0 +1,2 @@ +if e >= a > b > c: + print "q" \ No newline at end of file diff --git a/python/testData/inspections/PyChainedComparisonsInspection/test.py b/python/testData/inspections/PyChainedComparisonsInspection/test.py index 027fb1ec9f7b..b19bdfdab483 100644 --- a/python/testData/inspections/PyChainedComparisonsInspection/test.py +++ b/python/testData/inspections/PyChainedComparisonsInspection/test.py @@ -1,5 +1,5 @@ a < b and b < c -b > a and b < c and c < d and d < e +b > a and b < c a < b and b < c and c < d if a < b and b < c: pass @@ -9,3 +9,16 @@ q = a < c and c < d < e result = a < c and c == 4 q = a < b < c and c <= d q = a >= b >= c and c > d + +#PY-3126 +if b > a and b < c: + print ("a") + +if c > a and b < c: + print("b") + +if a > c and b < c: + print("d") + +if b >= a > e and b < c: + print "q" \ 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 4710bce9e428..e9def65d7fa4 100644 --- a/python/testSrc/com/jetbrains/python/PyQuickFixTest.java +++ b/python/testSrc/com/jetbrains/python/PyQuickFixTest.java @@ -211,6 +211,21 @@ public class PyQuickFixTest extends PyLightFixtureTestCase { PyBundle.message("QFIX.chained.comparison"), true, true); } + public void testChainedComparison1() { // PY-3126 + doInspectionTest("ChainedComparison1.py", PyChainedComparisonsInspection.class, + PyBundle.message("QFIX.chained.comparison"), true, true); + } + + public void testChainedComparison2() { // PY-3126 + doInspectionTest("ChainedComparison2.py", PyChainedComparisonsInspection.class, + PyBundle.message("QFIX.chained.comparison"), true, true); + } + + public void testChainedComparison3() { // PY-3126 + doInspectionTest("ChainedComparison3.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);