Improve "is None" and "is not None" processing in PyTypeAssertionEvaluator

This commit is contained in:
Semyon Proshev
2016-09-15 16:30:12 +03:00
parent da153b0935
commit 7d7ead47a3
2 changed files with 73 additions and 15 deletions
@@ -115,21 +115,39 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
final PyExpression lhs = node.getLeftExpression();
final PyExpression rhs = node.getRightExpression();
if (node.isOperator("isnot")) {
if (lhs instanceof PyReferenceExpression && rhs instanceof PyReferenceExpression) {
final PyReferenceExpression target = (PyReferenceExpression)lhs;
if (PyNames.NONE.equals(rhs.getName())) {
final boolean positive = myPositive;
pushAssertion(target, new InstructionTypeCallback() {
@Override
public PyType getType(TypeEvalContext context, @Nullable PsiElement anchor) {
final List<PyType> types = new ArrayList<>();
types.add(PyNoneType.INSTANCE);
return createAssertionType(context.getType(target), types, !positive, context);
}
});
return;
}
if (lhs instanceof PyReferenceExpression && rhs instanceof PyReferenceExpression) {
final boolean leftIsNone = PyNames.NONE.equals(lhs.getName());
final boolean rightIsNone = PyNames.NONE.equals(rhs.getName());
if (leftIsNone && rightIsNone) {
return;
}
final PyReferenceExpression target = (PyReferenceExpression)(rightIsNone ? lhs : rhs);
final boolean positive = myPositive;
if (node.isOperator(PyNames.IS)) {
pushAssertion(target, new InstructionTypeCallback() {
@Override
public PyType getType(TypeEvalContext context, @Nullable PsiElement anchor) {
final List<PyType> types = new ArrayList<>();
types.add(PyNoneType.INSTANCE);
return createAssertionType(context.getType(target), types, positive, context);
}
});
return;
}
if (node.isOperator("isnot")) {
pushAssertion(target, new InstructionTypeCallback() {
@Override
public PyType getType(TypeEvalContext context, @Nullable PsiElement anchor) {
final List<PyType> types = new ArrayList<>();
types.add(PyNoneType.INSTANCE);
return createAssertionType(context.getType(target), types, !positive, context);
}
});
return;
}
}
@@ -1181,6 +1181,46 @@ public class PyTypeTest extends PyTestCase {
" print(expr)");
}
public void testIsNotNone() {
doTest("int",
"def test_1(self, c):\n" +
" x = 1 if c else None\n" +
" if x is not None:\n" +
" expr = x\n");
doTest("int",
"def test_1(self, c):\n" +
" x = 1 if c else None\n" +
" if None is not x:\n" +
" expr = x\n");
doTest("int",
"def test_1(self, c):\n" +
" x = 1 if c else None\n" +
" if not x is None:\n" +
" expr = x\n");
doTest("int",
"def test_1(self, c):\n" +
" x = 1 if c else None\n" +
" if not None is x:\n" +
" expr = x\n");
}
public void testIsNone() {
doTest("None",
"def test_1(self, c):\n" +
" x = 1 if c else None\n" +
" if x is None:\n" +
" expr = x\n");
doTest("None",
"def test_1(self, c):\n" +
" x = 1 if c else None\n" +
" if None is x:\n" +
" expr = x\n");
}
private static List<TypeEvalContext> getTypeEvalContexts(@NotNull PyExpression element) {
return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(),
TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());