While inferring parameter type from usages, check if it is compared with None (PY-27231)

This commit is contained in:
Semyon Proshev
2018-07-16 21:04:15 +03:00
parent 3b109fae50
commit 358edc489f
2 changed files with 86 additions and 7 deletions
@@ -304,9 +304,9 @@ public class PyNamedParameterImpl extends PyBaseElementImpl<PyNamedParameterStub
}
}
if (context.maySwitchToAST(this)) {
final Set<String> attributes = collectUsedAttributes(context);
if (!attributes.isEmpty()) {
return new PyStructuralType(attributes, true);
final PyType typeFromUsages = getTypeFromUsages(context);
if (typeFromUsages != null) {
return typeFromUsages;
}
}
}
@@ -319,12 +319,15 @@ public class PyNamedParameterImpl extends PyBaseElementImpl<PyNamedParameterStub
return new PyElementPresentation(this);
}
@NotNull
private Set<String> collectUsedAttributes(@NotNull final TypeEvalContext context) {
final Set<String> result = new LinkedHashSet<>();
@Nullable
private PyType getTypeFromUsages(@NotNull TypeEvalContext context) {
final Set<String> usedAttributes = new LinkedHashSet<>();
final ScopeOwner owner = ScopeUtil.getScopeOwner(this);
final String name = getName();
final Ref<Boolean> parameterWasReassigned = Ref.create(false);
final Ref<Boolean> noneComparison = Ref.create(false);
if (owner != null && name != null) {
owner.accept(new PyRecursiveElementVisitor() {
@@ -411,6 +414,21 @@ public class PyNamedParameterImpl extends PyBaseElementImpl<PyNamedParameterStub
}
}
@Override
public void visitPyBinaryExpression(PyBinaryExpression node) {
super.visitPyBinaryExpression(node);
if (noneComparison.get() || !node.isOperator(PyNames.IS) && !node.isOperator("isnot")) return;
final PyExpression lhs = node.getLeftExpression();
final PyExpression rhs = node.getRightExpression();
if (isReferenceToParameter(lhs) ^ isReferenceToParameter(rhs) &&
(lhs != null && context.getType(lhs) instanceof PyNoneType) ^ (rhs != null && context.getType(rhs) instanceof PyNoneType)) {
noneComparison.set(true);
}
}
@Contract("null -> false")
private boolean isReferenceToParameter(@Nullable PsiElement element) {
if (element == null) return false;
@@ -419,7 +437,13 @@ public class PyNamedParameterImpl extends PyBaseElementImpl<PyNamedParameterStub
}
});
}
return result;
if (!usedAttributes.isEmpty()) {
final PyStructuralType structuralType = new PyStructuralType(usedAttributes, true);
return noneComparison.get() ? PyUnionType.union(structuralType, PyNoneType.INSTANCE) : structuralType;
}
return null;
}
@NotNull
@@ -579,6 +579,61 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase {
doTest();
}
// PY-27231
public void testStructuralAndNone() {
doTestByText("def func11(value):\n" +
" if value is not None and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"def func12(value):\n" +
" if None is not value and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"def func21(value):\n" +
" if value is None and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"def func22(value):\n" +
" if None is value and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"func11(None)\n" +
"func12(None)\n" +
"func21(None)\n" +
"func22(None)\n" +
"\n" +
"\n" +
"def func31(value):\n" +
" if value and None and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"def func32(value):\n" +
" if value is value and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"def func33(value):\n" +
" if None is None and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"def func34(value):\n" +
" a = 2\n" +
" if a is a and value != 1:\n" +
" pass\n" +
"\n" +
"\n" +
"func31(<warning descr=\"Expected type '{__ne__}', got 'None' instead\">None</warning>)\n" +
"func32(<warning descr=\"Expected type '{__ne__}', got 'None' instead\">None</warning>)\n" +
"func33(<warning descr=\"Expected type '{__ne__}', got 'None' instead\">None</warning>)\n" +
"func34(<warning descr=\"Expected type '{__ne__}', got 'None' instead\">None</warning>)");
}
// PY-29704
public void testPassingAbstractMethodResult() {
doTestByText("import abc\n" +