diff --git a/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java b/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java index df9348bbb18d..7dc151b23741 100644 --- a/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyChainedComparisonsInspection.java @@ -51,6 +51,9 @@ public class PyChainedComparisonsInspection extends PyInspection { } private static class Visitor extends PyInspectionVisitor { + /** + * @see ChainedComparisonsQuickFix#ChainedComparisonsQuickFix(boolean, boolean, boolean) + */ boolean myIsLeft; boolean myIsRight; PyElementType myOperator; @@ -61,22 +64,22 @@ public class PyChainedComparisonsInspection extends PyInspection { } @Override - public void visitPyBinaryExpression(final PyBinaryExpression node){ + public void visitPyBinaryExpression(final PyBinaryExpression node) { myIsLeft = false; myIsRight = false; myOperator = null; getInnerRight = false; - PyExpression leftExpression = node.getLeftExpression(); - PyExpression rightExpression = node.getRightExpression(); + final PyExpression leftExpression = node.getLeftExpression(); + final PyExpression rightExpression = node.getRightExpression(); if (leftExpression instanceof PyBinaryExpression && - rightExpression instanceof PyBinaryExpression) { + rightExpression instanceof PyBinaryExpression) { if (node.getOperator() == PyTokenTypes.AND_KEYWORD) { if (isRightSimplified((PyBinaryExpression)leftExpression, (PyBinaryExpression)rightExpression) || - isLeftSimplified((PyBinaryExpression)leftExpression, (PyBinaryExpression)rightExpression)) - registerProblem(node, "Simplify chained comparison", new ChainedComparisonsQuickFix(myIsLeft, myIsRight, - getInnerRight)); + isLeftSimplified((PyBinaryExpression)leftExpression, (PyBinaryExpression)rightExpression)) { + registerProblem(node, "Simplify chained comparison", new ChainedComparisonsQuickFix(myIsLeft, myIsRight, getInnerRight)); + } } } } @@ -92,9 +95,10 @@ public class PyChainedComparisonsInspection extends PyInspection { } if (leftRight instanceof PyBinaryExpression && - PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)leftRight).getOperator())){ - if (isRightSimplified((PyBinaryExpression)leftRight, rightExpression)) + PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)leftRight).getOperator())) { + if (isRightSimplified((PyBinaryExpression)leftRight, rightExpression)) { return true; + } } myOperator = leftExpression.getOperator(); @@ -118,11 +122,12 @@ public class PyChainedComparisonsInspection extends PyInspection { } private static boolean isOpposite(final PyElementType op1, final PyElementType op2) { - if ((op1 == PyTokenTypes.GT || op1 == PyTokenTypes.GE) && (op2 == PyTokenTypes.LT || op2 == PyTokenTypes.LE)) + if ((op1 == PyTokenTypes.GT || op1 == PyTokenTypes.GE) && (op2 == PyTokenTypes.LT || op2 == PyTokenTypes.LE)) { return true; - if ((op2 == PyTokenTypes.GT || op2 == PyTokenTypes.GE) && (op1 == PyTokenTypes.LT || op1 == PyTokenTypes.LE)) + } + if ((op2 == PyTokenTypes.GT || op2 == PyTokenTypes.GE) && (op1 == PyTokenTypes.LT || op1 == PyTokenTypes.LE)) { return true; - + } return false; } @@ -139,9 +144,10 @@ public class PyChainedComparisonsInspection extends PyInspection { } if (leftLeft instanceof PyBinaryExpression && - PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)leftLeft).getOperator())){ - if (isLeftSimplified((PyBinaryExpression)leftLeft, rightExpression)) + PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)leftLeft).getOperator())) { + if (isLeftSimplified((PyBinaryExpression)leftLeft, rightExpression)) { return true; + } } myOperator = leftExpression.getOperator(); @@ -165,13 +171,14 @@ public class PyChainedComparisonsInspection extends PyInspection { private PyExpression getLeftExpression(PyBinaryExpression expression, boolean isRight) { PyExpression result = expression; - while (result instanceof PyBinaryExpression && ( - PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()) - || PyTokenTypes.EQUALITY_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()))) { + while (result instanceof PyBinaryExpression && + (PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()) || + PyTokenTypes.EQUALITY_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()))) { final boolean opposite = isOpposite(((PyBinaryExpression)result).getOperator(), myOperator); - if ((isRight && opposite) || (!isRight && !opposite)) + if ((isRight && opposite) || (!isRight && !opposite)) { break; + } result = ((PyBinaryExpression)result).getLeftExpression(); } return result; @@ -180,13 +187,14 @@ public class PyChainedComparisonsInspection extends PyInspection { @Nullable private PyExpression getSmallestRight(PyBinaryExpression expression, boolean isRight) { PyExpression result = expression; - while (result instanceof PyBinaryExpression && ( - PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()) - || PyTokenTypes.EQUALITY_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()))) { + while (result instanceof PyBinaryExpression && + (PyTokenTypes.RELATIONAL_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()) || + PyTokenTypes.EQUALITY_OPERATIONS.contains(((PyBinaryExpression)result).getOperator()))) { final boolean opposite = isOpposite(((PyBinaryExpression)result).getOperator(), myOperator); - if ((isRight && !opposite) || (!isRight && opposite)) + if ((isRight && !opposite) || (!isRight && opposite)) { break; + } result = ((PyBinaryExpression)result).getRightExpression(); } return result; diff --git a/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java b/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java index d09bc4e0a52c..a5643afff75b 100644 --- a/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java +++ b/python/src/com/jetbrains/python/inspections/quickfix/ChainedComparisonsQuickFix.java @@ -21,10 +21,15 @@ import com.intellij.openapi.project.Project; import com.intellij.psi.PsiElement; import com.jetbrains.python.PyBundle; import com.jetbrains.python.PyTokenTypes; -import com.jetbrains.python.psi.*; +import com.jetbrains.python.psi.PyBinaryExpression; +import com.jetbrains.python.psi.PyElementGenerator; +import com.jetbrains.python.psi.PyExpression; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; +import static com.intellij.util.ObjectUtils.assertNotNull; +import static com.jetbrains.python.psi.PyUtil.as; + /** * User: catherine * @@ -36,6 +41,13 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { boolean myIsRightLeft; boolean getInnerRight; + /** + * @param isLeft true if common expression is on the left hand side of the left comparison + * @param isRight true if common expression is on the left hand side of the right comparison + * @param getInner whether left comparison is deeper in PSI tree than the right comparison. + * E.g. in {@code foo and x > 1 and x < 3} expressions {@code x > 1} and {@code x < 3} are targets for simplification but + * because of associativity of {@code and} operator they are not siblings: {@code (foo and x > 1) and x < 3} + */ public ChainedComparisonsQuickFix(boolean isLeft, boolean isRight, boolean getInner) { myIsLeftLeft = isLeft; myIsRightLeft = isRight; @@ -53,26 +65,23 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { } public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - PsiElement expression = descriptor.getPsiElement(); + final PyBinaryExpression expression = as(descriptor.getPsiElement(), PyBinaryExpression.class); if (expression != null && expression.isWritable()) { - if (expression instanceof PyBinaryExpression) { - PyExpression leftExpression = ((PyBinaryExpression)expression).getLeftExpression(); - PyExpression rightExpression = ((PyBinaryExpression)expression).getRightExpression(); - if (rightExpression instanceof PyBinaryExpression && leftExpression instanceof PyBinaryExpression) { - if (((PyBinaryExpression)expression).getOperator() == PyTokenTypes.AND_KEYWORD) { - if (getInnerRight && ((PyBinaryExpression)leftExpression).getRightExpression() instanceof PyBinaryExpression - && PyTokenTypes.AND_KEYWORD == ((PyBinaryExpression)leftExpression).getOperator()) { - leftExpression = ((PyBinaryExpression)leftExpression).getRightExpression(); - } - checkOperator((PyBinaryExpression)leftExpression, (PyBinaryExpression)rightExpression, project); - } + 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) { + leftExpression = (PyBinaryExpression)leftExpression.getRightExpression(); } + checkOperator(leftExpression, rightExpression, project); } } } - private void checkOperator(final PyBinaryExpression leftExpression, - final PyBinaryExpression rightExpression, final Project project) { + private void checkOperator(@NotNull PyBinaryExpression leftExpression, + @NotNull PyBinaryExpression rightExpression, + @NotNull Project project) { final PyElementGenerator elementGenerator = PyElementGenerator.getInstance(project); if (myIsLeftLeft) { final PyExpression newLeftExpression = invertExpression(leftExpression, elementGenerator); @@ -80,14 +89,14 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { if (myIsRightLeft) { final PsiElement operator = getLeftestOperator(rightExpression); final PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( - operator.getText(), newLeftExpression, getLargeRightExpression(rightExpression, project)); + operator.getText(), newLeftExpression, getLargeRightExpression(rightExpression, project)); leftExpression.replace(binaryExpression); rightExpression.delete(); } else { - final String operator = invertOperator(rightExpression.getPsiOperator()); + final String operator = invertOperator(assertNotNull(rightExpression.getPsiOperator())); final PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( - operator, newLeftExpression, rightExpression.getLeftExpression()); + operator, newLeftExpression, rightExpression.getLeftExpression()); leftExpression.replace(binaryExpression); rightExpression.delete(); } @@ -96,25 +105,26 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { if (myIsRightLeft) { final PsiElement operator = getLeftestOperator(rightExpression); final PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( - operator.getText(), leftExpression, getLargeRightExpression(rightExpression, project)); + operator.getText(), leftExpression, getLargeRightExpression(rightExpression, project)); leftExpression.replace(binaryExpression); rightExpression.delete(); } else { PyExpression expression = rightExpression.getLeftExpression(); - if (expression instanceof PyBinaryExpression) + if (expression instanceof PyBinaryExpression) { expression = invertExpression((PyBinaryExpression)expression, elementGenerator); - final String operator = invertOperator(rightExpression.getPsiOperator()); + } + final String operator = invertOperator(assertNotNull(rightExpression.getPsiOperator())); final PyBinaryExpression binaryExpression = elementGenerator.createBinaryExpression( - operator, leftExpression, expression); + operator, leftExpression, expression); leftExpression.replace(binaryExpression); rightExpression.delete(); } } - } - private PsiElement getLeftestOperator(PyBinaryExpression expression) { + @NotNull + private static PsiElement getLeftestOperator(@NotNull PyBinaryExpression expression) { PsiElement op = expression.getPsiOperator(); while (expression.getLeftExpression() instanceof PyBinaryExpression) { expression = (PyBinaryExpression)expression.getLeftExpression(); @@ -124,45 +134,47 @@ public class ChainedComparisonsQuickFix implements LocalQuickFix { return op; } - private PyExpression invertExpression(PyBinaryExpression leftExpression, PyElementGenerator elementGenerator) { + @NotNull + private static PyExpression invertExpression(@NotNull PyBinaryExpression leftExpression, @NotNull PyElementGenerator elementGenerator) { final PsiElement operator = leftExpression.getPsiOperator(); final PyExpression right = leftExpression.getRightExpression(); PyExpression left = leftExpression.getLeftExpression(); - if (left instanceof PyBinaryExpression){ + if (left instanceof PyBinaryExpression) { left = invertExpression((PyBinaryExpression)left, elementGenerator); } - final String newOperator = invertOperator(operator); - return elementGenerator.createBinaryExpression( - newOperator, right, left); + final String newOperator = invertOperator(assertNotNull(operator)); + return elementGenerator.createBinaryExpression(newOperator, right, left); } - private String invertOperator(PsiElement op) { - if (op.getText().equals(">")) + @NotNull + private static String invertOperator(@NotNull PsiElement op) { + if (op.getText().equals(">")) { return "<"; - if (op.getText().equals("<")) + } + if (op.getText().equals("<")) { return ">"; - if (op.getText().equals(">=")) + } + if (op.getText().equals(">=")) { return "<="; - if (op.getText().equals("<=")) + } + if (op.getText().equals("<=")) { return ">="; + } return op.getText(); } @Nullable - static private PyExpression getLargeRightExpression(PyBinaryExpression expression, Project project) { + static private PyExpression getLargeRightExpression(@NotNull PyBinaryExpression expression, @NotNull Project project) { final PyElementGenerator elementGenerator = PyElementGenerator.getInstance(project); PyExpression left = expression.getLeftExpression(); PyExpression right = expression.getRightExpression(); PsiElement operator = expression.getPsiOperator(); while (left instanceof PyBinaryExpression) { assert operator != null; - right = elementGenerator.createBinaryExpression(operator.getText(), - ((PyBinaryExpression)left).getRightExpression(), - right); + right = elementGenerator.createBinaryExpression(operator.getText(), ((PyBinaryExpression)left).getRightExpression(), right); operator = ((PyBinaryExpression)left).getPsiOperator(); left = ((PyBinaryExpression)left).getLeftExpression(); } return right; } - }