From 26731f5738c4a62844375a95a3099dc3dae592a8 Mon Sep 17 00:00:00 2001 From: abragin Date: Wed, 6 Dec 2017 16:12:21 +0300 Subject: [PATCH] PY-4813 Type hints retrieval from superclasses implemented Type information from superclass method is now available in overriden methods. --- .../src/com/jetbrains/python/PyNames.java | 1 + .../typing/PyTypingTypeProvider.java | 108 +++++++-- .../com/jetbrains/python/Py3TypeTest.java | 227 ++++++++++++++++++ 3 files changed, 319 insertions(+), 17 deletions(-) diff --git a/python/psi-api/src/com/jetbrains/python/PyNames.java b/python/psi-api/src/com/jetbrains/python/PyNames.java index 85778330e6ef..d1554df06c35 100644 --- a/python/psi-api/src/com/jetbrains/python/PyNames.java +++ b/python/psi-api/src/com/jetbrains/python/PyNames.java @@ -108,6 +108,7 @@ public class PyNames { public static final String CLASSMETHOD = "classmethod"; public static final String STATICMETHOD = "staticmethod"; + public static final String OVERLOADMETHOD = "overload"; public static final String PROPERTY = "property"; public static final String SETTER = "setter"; diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index d96087e0ebc7..a61cf4439b7d 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -14,6 +14,7 @@ import com.intellij.psi.util.CachedValuesManager; import com.intellij.psi.util.PsiTreeUtil; import com.intellij.psi.util.QualifiedName; import com.intellij.util.ArrayUtil; +import com.intellij.util.Query; import com.intellij.util.containers.ContainerUtil; import com.intellij.util.containers.HashMap; import com.intellij.util.containers.HashSet; @@ -35,6 +36,7 @@ import com.jetbrains.python.psi.impl.stubs.PyTypingAliasStubType; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.PyResolveImportUtil; import com.jetbrains.python.psi.resolve.PyResolveUtil; +import com.jetbrains.python.psi.search.PySuperMethodsSearch; import com.jetbrains.python.psi.resolve.RatedResolveResult; import com.jetbrains.python.psi.stubs.PyClassStub; import com.jetbrains.python.psi.types.*; @@ -156,23 +158,28 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } final PyFunctionTypeAnnotation annotation = getFunctionTypeAnnotation(func); - if (annotation == null) { - return null; - } - final PyParameterTypeList list = annotation.getParameterTypeList(); - final List paramTypes = list.getParameterTypes(); - if (paramTypes.size() == 1) { - final PyNoneLiteralExpression noneExpr = as(paramTypes.get(0), PyNoneLiteralExpression.class); - if (noneExpr != null && noneExpr.isEllipsis()) { - return Ref.create(); + if (annotation != null) { + PyParameterTypeList list = annotation.getParameterTypeList(); + List paramTypes = list.getParameterTypes(); + if (paramTypes.size() == 1) { + final PyNoneLiteralExpression noneExpr = as(paramTypes.get(0), PyNoneLiteralExpression.class); + if (noneExpr != null && noneExpr.isEllipsis()) { + return Ref.create(); + } + } + final int startOffset = omitFirstParamInTypeComment(func, annotation) ? 1 : 0; + final List funcParams = Arrays.asList(func.getParameterList().getParameters()); + final int i = funcParams.indexOf(param) - startOffset; + if (i >= 0 && i < paramTypes.size()) { + return getParameterTypeFromFunctionComment(paramTypes.get(i), context); } } - final int startOffset = omitFirstParamInTypeComment(func, annotation) ? 1 : 0; - final List funcParams = Arrays.asList(func.getParameterList().getParameters()); - final int i = funcParams.indexOf(param) - startOffset; - if (i >= 0 && i < paramTypes.size()) { - return getParameterTypeFromFunctionComment(paramTypes.get(i), context); + + final Ref typeFromAncestors = getParameterTypeFromSupertype(param, func, context); + if (typeFromAncestors != null) { + return typeFromAncestors; } + return null; } @@ -243,6 +250,56 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return null; } + @Nullable + private static PyFunctionType getSuperFunctionType(@NotNull PyFunction function, @NotNull TypeEvalContext context) { + final Query superMethodSearchQuery = PySuperMethodsSearch.search(function, context); + final PsiElement firstSuperMethod = superMethodSearchQuery.findFirst(); + + if (firstSuperMethod instanceof PyFunction) { + final PyFunction superFunction = (PyFunction)firstSuperMethod; + + if (superFunction.getDecoratorList() != null) { + if (StreamEx.of(superFunction.getDecoratorList().getDecorators()) + .map(PyDecorator::getQualifiedName) + .nonNull() + .anyMatch(decoratorName -> PyNames.OVERLOADMETHOD.equals(decoratorName.toString()))) { + return null; + } + } + + final PyClass superClass = superFunction.getContainingClass(); + if (superClass != null && !PyNames.OBJECT.equals(superClass.getName())) { + return (PyFunctionType)context.getType(superFunction); + } + } + + return null; + } + + @Nullable + private static Ref getParameterTypeFromSupertype( + @NotNull PyNamedParameter param, @NotNull PyFunction func, @NotNull TypeEvalContext context) { + final PyFunctionType superFunctionType = getSuperFunctionType(func, context); + + if (superFunctionType != null) { + final PyFunctionType superFunctionTypeRemovedSelf = superFunctionType.dropSelf(context); + final List parameters = superFunctionTypeRemovedSelf.getParameters(context); + if (parameters != null) { + for (PyCallableParameter parameter : parameters) { + final String parameterName = parameter.getName(); + if (parameterName != null && parameterName.equals(param.getName())) { + final PyType pyType = parameter.getType(context); + if (pyType != null) { + return new Ref<>(pyType); + } + } + } + } + } + + return null; + } + @NotNull private static PyType createTypingGenericType() { return new PyCustomType(GENERIC, null, false); @@ -263,13 +320,19 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { public Ref getReturnType(@NotNull PyCallable callable, @NotNull TypeEvalContext context) { if (callable instanceof PyFunction) { final PyFunction function = (PyFunction)callable; - final PyExpression value = getReturnTypeAnnotation(function, context); - if (value != null) { - final Ref typeRef = getType(value, new Context(context)); + + final PyExpression returnTypeAnnotation = getReturnTypeAnnotation(function, context); + if (returnTypeAnnotation != null) { + final Ref typeRef = getType(returnTypeAnnotation, new Context(context)); if (typeRef != null) { return Ref.create(toAsyncIfNeeded(function, typeRef.get())); } } + + final Ref typeFromSupertype = getReturnTypeFromSupertype(function, context); + if (typeFromSupertype != null) { + return Ref.create(toAsyncIfNeeded(function, typeFromSupertype.get())); + } } return null; } @@ -287,6 +350,17 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return null; } + @Nullable + private static Ref getReturnTypeFromSupertype(@NotNull PyFunction function, @NotNull TypeEvalContext context) { + final PyFunctionType superFunctionType = getSuperFunctionType(function, context); + + if (superFunctionType != null) { + return new Ref<>(context.getReturnType(superFunctionType.getCallable())); + } + + return null; + } + @Nullable private static PyFunctionTypeAnnotation getFunctionTypeAnnotation(@NotNull PyFunction function) { final String comment = function.getTypeCommentAnnotation(); diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index b16be7c033fd..5e9bcb27b986 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -663,6 +663,233 @@ public class Py3TypeTest extends PyTestCase { ); } + // PY-4813 + public void testParameterTypeInferenceInSubclassFromDocstring() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " def test(self, param):\n" + + " \"\"\"\n" + + " :param param:\n" + + " :type param: int\n" + + " \"\"\"\n" + + " pass\n" + + "\n" + + "class Subclass(Base):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInSubclassFromAnnotation() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " def test(self, param: int) -> int: pass\n" + + "\n" + + "class Subclass(Base):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInSubclassHierarchyFromAnnotation1() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " def test(self, param: int) -> int: pass\n" + + "\n" + + "class Base1(Base):\n" + + " pass\n" + + "\n" + + "class Subclass(Base1):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInSubclassHierarchyFromAnnotation2() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base1:\n" + + " def test(self, param: int) -> int: pass\n" + + "\n" + + "class Base2:\n" + + " def test(self, param: str) -> str: pass\n" + + "\n" + + "class Subclass(Base1, Base2):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInSubclassHierarchyFromAnnotation3() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base1:\n" + + " def test(self, param: int) -> int: pass\n" + + "\n" + + "class Base2:\n" + + " def test(self, param: str) -> str: pass\n" + + "\n" + + "class Base3(Base1):\n" + + " pass\n" + + "\n" + + "class Subclass(Base3, Base2):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInSubclassHierarchyFromAnnotation4() { + /* + This behavior mimics C3 MRO used by *new style* classes for Python 2.3 and below + For details see: https://www.python.org/download/releases/2.3/mro/ + Since annotations are supported in Python 3.5 and above, *classic classes* + should not be tested here, but it's NOT the case for docstrings + */ + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base1:\n" + + " def test(self, param: str) -> str: pass\n" + + "\n" + + "class Base2(Base1):\n" + + " def test(self, param: int) -> int: pass\n" + + "\n" + + "class Base3(Base1):\n" + + " pass\n" + + "\n" + + "class Subclass(Base3, Base2):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInSubclassHierarchyFromAnnotation5() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base1:\n" + + " def test(self, param: int) -> int: pass\n" + + "\n" + + "class Base2(Base1):\n" + + " def test(self, param): pass\n" + + "\n" + + "class Subclass(Base2):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInOverloadedMethods() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("Any", + "class Base:\n" + + " @overload\n" + + " def test(self, param: int) -> int: pass\n" + + "\n" + + " @overload\n" + + " def test(self, param: str) -> str: pass\n" + + "\n" + + "class Subclass(Base):\n" + + " def test(self, param):\n" + + " expr = param") + ); + } + + public void testParameterTypeInferenceInSubclassHierarchyInStaticMethods() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " @staticmethod\n" + + " def test(param: int, param1: int) -> int: pass\n" + + "\n" + + "class Subclass(Base):\n" + + " @staticmethod\n" + + " def test(param, param1):\n" + + " expr = param\n" + + "\n") + ); + } + + public void testReturnTypeInferenceInSubclassFromAnnotation() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " def test(self) -> int: pass\n" + + "\n" + + "class Subclass(Base):\n" + + " def test(self): pass\n" + + "\n" + + "expr = Subclass().test()") + ); + } + + public void testReturnTypeInferenceInSubclassHierarchyFromAnnotation() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " def test(self) -> int: pass\n" + + "\n" + + "class Base1(Base):\n" + + " pass\n" + + "\n" + + "class Subclass(Base1):\n" + + " def test(self): pass\n" + + "\n" + + "expr = Subclass().test()") + ); + } + + public void testReturnTypeInferenceInSubclassFromDocstring() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " def test(self):" + + " \"\"\"\n" + + " :rtype: int\n" + + " \"\"\"\n" + + " pass\n" + + "\n" + + "class Subclass(Base):\n" + + " def test(self): pass\n" + + "\n" + + "expr = Subclass().test()") + ); + } + + public void testReturnTypeInferenceInSubclassHierarchyFromDocstring() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> doTest("int", + "class Base:\n" + + " def test(self):" + + " \"\"\"\n" + + " :rtype: int\n" + + " \"\"\"\n" + + " pass\n" + + "\n" + + "class Base1(Base):\n" + + " pass\n" + + "\n" + + "class Subclass(Base1):\n" + + " def test(self): pass\n" + + "\n" + + "expr = Subclass().test()") + ); + } + private void doTest(final String expectedType, final String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class);