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);