Add issubclass support into PyTypeAssertionEvaluator.

This commit is contained in:
Semyon Proshev
2017-06-20 11:34:50 +03:00
parent b6482f44a1
commit 0db3c6ef81
3 changed files with 42 additions and 12 deletions
@@ -112,6 +112,7 @@ public class PyNames {
public static final String ISINSTANCE = "isinstance";
public static final String ASSERT_IS_INSTANCE = "assertIsInstance";
public static final String HAS_ATTR = "hasattr";
public static final String ISSUBCLASS = "issubclass";
public static final String DOC = "__doc__";
public static final String DOCFORMAT = "__docformat__";
@@ -62,13 +62,13 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
@Override
public void visitPyCallExpression(PyCallExpression node) {
if (node.isCalleeText(PyNames.ISINSTANCE) || node.isCalleeText(PyNames.ASSERT_IS_INSTANCE)) {
if (node.isCalleeText(PyNames.ISINSTANCE, PyNames.ASSERT_IS_INSTANCE)) {
final PyExpression[] args = node.getArguments();
if (args.length == 2 && args[0] instanceof PyReferenceExpression) {
final PyReferenceExpression target = (PyReferenceExpression)args[0];
final PyExpression typeElement = args[1];
pushAssertion(target, myPositive, context -> context.getType(typeElement));
pushAssertion(target, myPositive, false, context -> context.getType(typeElement));
}
}
else if (node.isCalleeText(PyNames.CALLABLE_BUILTIN)) {
@@ -76,7 +76,16 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
if (args.length == 1 && args[0] instanceof PyReferenceExpression) {
final PyReferenceExpression target = (PyReferenceExpression)args[0];
pushAssertion(target, myPositive, context -> PyTypeParser.getTypeByName(target, "collections." + PyNames.CALLABLE, context));
pushAssertion(target, myPositive, false, context -> PyTypeParser.getTypeByName(target, "collections." + PyNames.CALLABLE, context));
}
}
else if (node.isCalleeText(PyNames.ISSUBCLASS)) {
final PyExpression[] args = node.getArguments();
if (args.length == 2 && args[0] instanceof PyReferenceExpression) {
final PyReferenceExpression target = (PyReferenceExpression)args[0];
final PyExpression typeElement = args[1];
pushAssertion(target, myPositive, true, context -> context.getType(typeElement));
}
}
}
@@ -86,7 +95,7 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
if (myPositive && (isIfReferenceStatement(node) || isIfReferenceConditionalStatement(node) || isIfNotReferenceStatement(node))) {
// we could not suggest `None` because it could be a reference to an empty collection
// so we could push only non-`None` assertions
pushAssertion(node, !myPositive, context -> PyNoneType.INSTANCE);
pushAssertion(node, !myPositive, false, context -> PyNoneType.INSTANCE);
return;
}
@@ -111,12 +120,12 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
final PyReferenceExpression target = (PyReferenceExpression)(rightIsNone ? lhs : rhs);
if (node.isOperator(PyNames.IS)) {
pushAssertion(target, myPositive, context -> PyNoneType.INSTANCE);
pushAssertion(target, myPositive, false, context -> PyNoneType.INSTANCE);
return;
}
if (node.isOperator("isnot")) {
pushAssertion(target, !myPositive, context -> PyNoneType.INSTANCE);
pushAssertion(target, !myPositive, false, context -> PyNoneType.INSTANCE);
return;
}
}
@@ -143,8 +152,9 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
private static PyType createAssertionType(@Nullable PyType initial,
@Nullable PyType suggested,
boolean positive,
boolean transformToDefinition,
@NotNull TypeEvalContext context) {
final PyType transformedType = transformTypeFromAssertion(suggested);
final PyType transformedType = transformTypeFromAssertion(suggested, transformToDefinition);
if (positive) {
if (!(initial instanceof PyUnionType) &&
!PyTypeChecker.isUnknown(initial, context) &&
@@ -163,29 +173,31 @@ public class PyTypeAssertionEvaluator extends PyRecursiveElementVisitor {
}
@Nullable
private static PyType transformTypeFromAssertion(@Nullable PyType type) {
private static PyType transformTypeFromAssertion(@Nullable PyType type, boolean transformToDefinition) {
if (type instanceof PyTupleType) {
final List<PyType> members = new ArrayList<>();
final PyTupleType tupleType = (PyTupleType)type;
final int count = tupleType.getElementCount();
for (int i = 0; i < count; i++) {
members.add(transformTypeFromAssertion(tupleType.getElementType(i)));
members.add(transformTypeFromAssertion(tupleType.getElementType(i), transformToDefinition));
}
return PyUnionType.union(members);
}
else if (type instanceof PyClassType) {
return ((PyClassType)type).toInstance();
else if (type instanceof PyInstantiableType) {
final PyInstantiableType instantiableType = (PyInstantiableType)type;
return transformToDefinition ? instantiableType.toClass() : instantiableType.toInstance();
}
return type;
}
private void pushAssertion(@NotNull PyReferenceExpression target,
boolean positive,
boolean transformToDefinition,
@NotNull Function<TypeEvalContext, PyType> suggestedType) {
final InstructionTypeCallback typeCallback = new InstructionTypeCallback() {
@Override
public PyType getType(TypeEvalContext context, @Nullable PsiElement anchor) {
return createAssertionType(context.getType(target), suggestedType.apply(context), positive, context);
return createAssertionType(context.getType(target), suggestedType.apply(context), positive, transformToDefinition, context);
}
};
@@ -1915,6 +1915,23 @@ public class PyTypeTest extends PyTestCase {
"expr = resort");
}
public void testIsSubclass() {
doTest("Type[A]",
"class A: pass\n" +
"def foo(cls):\n" +
" if issubclass(cls, A):\n" +
" expr = cls");
}
public void testIsSubclassWithTupleOfTypeObjects() {
doTest("Type[Union[A, B]]",
"class A: pass\n" +
"class B: pass\n" +
"def foo(cls):\n" +
" if issubclass(cls, (A, B)):\n" +
" expr = cls");
}
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());