fixed PY-3126 "Chained comparison" doesn't react if one of the sides is flipped

This commit is contained in:
Ekaterina Tuzova
2011-06-01 14:41:18 +04:00
parent 09d552e6af
commit 283b3b7f87
10 changed files with 158 additions and 32 deletions
@@ -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) {
@@ -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;
}
}
}
@@ -0,0 +1,2 @@
if <warning descr="Simplify chained comparison">b >= <caret>a > e and b < c</warning>:
print "q"
@@ -0,0 +1,2 @@
if e < a <= b < c:
print "q"
@@ -0,0 +1,2 @@
if <warning descr="Simplify chained comparison">b >= <caret>a > e and c < b</warning>:
print "q"
@@ -0,0 +1,2 @@
if e < a <= b > c:
print "q"
@@ -0,0 +1,2 @@
if <warning descr="Simplify chained comparison">e >= <caret>a > b and c < b</warning>:
print "q"
@@ -0,0 +1,2 @@
if e >= a > b > c:
print "q"
@@ -1,5 +1,5 @@
<warning descr="Simplify chained comparison">a < b and b < c</warning>
<warning descr="Simplify chained comparison"><warning descr="Simplify chained comparison">b > a and b < c and c < d</warning> and d < e</warning>
<warning descr="Simplify chained comparison">b > a and b < c</warning>
<warning descr="Simplify chained comparison"><warning descr="Simplify chained comparison">a < b and b < c</warning> and c < d</warning>
if <warning descr="Simplify chained comparison">a < b and b < c</warning>:
pass
@@ -9,3 +9,16 @@ q = <warning descr="Simplify chained comparison">a < c and c < d < e</warning>
result = <warning descr="Simplify chained comparison">a < c and c == 4</warning>
q = <warning descr="Simplify chained comparison">a < b < c and c <= d</warning>
q = <warning descr="Simplify chained comparison">a >= b >= c and c > d</warning>
#PY-3126
if <warning descr="Simplify chained comparison">b > a and b < c</warning>:
print ("a")
if <warning descr="Simplify chained comparison">c > a and b < c</warning>:
print("b")
if <warning descr="Simplify chained comparison">a > c and b < c</warning>:
print("d")
if <warning descr="Simplify chained comparison">b >= a > e and b < c</warning>:
print "q"
@@ -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);