diff --git a/python/psi-api/src/com/jetbrains/python/PyNames.java b/python/psi-api/src/com/jetbrains/python/PyNames.java index 3e30b0286642..4d36979ca649 100644 --- a/python/psi-api/src/com/jetbrains/python/PyNames.java +++ b/python/psi-api/src/com/jetbrains/python/PyNames.java @@ -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__"; diff --git a/python/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java b/python/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java index 3175748a4d4e..c26aceb635ac 100644 --- a/python/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java +++ b/python/src/com/jetbrains/python/codeInsight/controlflow/PyTypeAssertionEvaluator.java @@ -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 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 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); } }; diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index 225294b3d442..cc79ec302c00 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -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 getTypeEvalContexts(@NotNull PyExpression element) { return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(), TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());