From aa41ca3616615a4ce48816b88ca6e64429f554a2 Mon Sep 17 00:00:00 2001 From: Mikhail Golubev Date: Fri, 3 Feb 2017 12:05:50 +0300 Subject: [PATCH] Don't wrap types returned by type providers in Coroutine and AsyncGenerator It's their responsibility to provide the final return type for a function. Only inferred (i.e. when there's no additional type hints) return value types should be postprocessed like that in order to get not the immediate type of values in return statements but the properly parametrized special type from typing for genetators and coroutines. Otherwise, we might end up returning something like Coroutine[Coroutine[int]] for coroutines annotated using .pyi stubs, because we don't take into account that PyiTypeProvider already wraps the return type from the stub's annotation in typing.Coroutine. --- .../typing/PyTypingTypeProvider.java | 21 +++++++-- .../python/psi/impl/PyFunctionImpl.java | 43 +++++++++---------- .../pyi/type/coroutineType/CoroutineType.py | 4 ++ .../pyi/type/coroutineType/CoroutineType.pyi | 2 + .../com/jetbrains/python/pyi/PyiTypeTest.java | 4 ++ 5 files changed, 47 insertions(+), 27 deletions(-) create mode 100644 python/testData/pyi/type/coroutineType/CoroutineType.py create mode 100644 python/testData/pyi/type/coroutineType/CoroutineType.pyi diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index 1c374a75478e..907b8ef65be0 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -26,6 +26,7 @@ import com.intellij.psi.PsiPolyVariantReference; import com.intellij.psi.util.CachedValueProvider; import com.intellij.psi.util.CachedValuesManager; import com.intellij.psi.util.PsiTreeUtil; +import com.intellij.psi.util.QualifiedName; import com.intellij.util.containers.ContainerUtil; import com.intellij.util.containers.HashSet; import com.jetbrains.python.PyCustomType; @@ -36,6 +37,7 @@ import com.jetbrains.python.codeInsight.functionTypeComments.psi.PyParameterType import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyExpressionCodeFragmentImpl; import com.jetbrains.python.psi.resolve.PyResolveContext; +import com.jetbrains.python.psi.resolve.PyResolveImportUtil; import com.jetbrains.python.psi.types.*; import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; @@ -239,11 +241,15 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { final PyExpression value = getReturnTypeAnnotation(function); if (value != null) { final Ref typeRef = getType(value, new Context(context)); - final PyType type = Ref.deref(typeRef); - if (isInit && type instanceof PyNoneType) { - return null; + if (typeRef != null) { + if (isInit && typeRef.get() instanceof PyNoneType) { + return null; + } + if (function.isAsync() && function.isAsyncAllowed() && !function.isGenerator()) { + return Ref.create(wrapInCoroutineType(typeRef.get(), callable)); + } + return typeRef; } - return typeRef; } } return null; @@ -716,6 +722,13 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return null; } + @Nullable + public static PyType wrapInCoroutineType(@Nullable PyType returnType, @NotNull PsiElement resolveAnchor) { + final PyClass coroutine = as(PyResolveImportUtil.resolveTopLevelMember(QualifiedName.fromDottedString(COROUTINE), + PyResolveImportUtil.fromFoothold(resolveAnchor)), PyClass.class); + return coroutine != null ? new PyCollectionTypeImpl(coroutine, false, Arrays.asList(null, null, returnType)) : null; + } + private static class Context { @NotNull private final TypeEvalContext myContext; @NotNull private final Set myCache = new HashSet<>(); diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index 917259b7398c..6d3fddefa781 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -38,10 +38,10 @@ import com.intellij.util.containers.JBIterable; import com.jetbrains.python.PyElementTypes; import com.jetbrains.python.PyNames; import com.jetbrains.python.PyTokenTypes; -import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.codeInsight.controlflow.ControlFlowCache; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.documentation.docstrings.DocStringUtil; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.resolve.PyResolveImportUtil; @@ -187,12 +187,6 @@ public class PyFunctionImpl extends PyBaseElementImpl implements @Nullable @Override public PyType getReturnType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) { - final PyType type = getReturnType(context); - return isAsync() && isAsyncAllowed() ? createCoroutineType(type) : type; - } - - @Nullable - private PyType getReturnType(@NotNull TypeEvalContext context) { for (PyTypeProvider typeProvider : Extensions.getExtensions(PyTypeProvider.EP_NAME)) { final Ref returnTypeRef = typeProvider.getReturnType(this, context); if (returnTypeRef != null) { @@ -200,15 +194,21 @@ public class PyFunctionImpl extends PyBaseElementImpl implements } } + PyType inferredType = null; if (context.allowReturnTypes(this)) { final Ref yieldTypeRef = getYieldStatementType(context); if (yieldTypeRef != null) { - return yieldTypeRef.get(); + inferredType = yieldTypeRef.get(); } - return getReturnStatementType(context); + else { + inferredType = getReturnStatementType(context); + } } - return null; + if (isAsync() && isAsyncAllowed()) { + return createAsyncType(inferredType); + } + return inferredType; } @Nullable @@ -383,12 +383,11 @@ public class PyFunctionImpl extends PyBaseElementImpl implements } @Nullable - private PyType createCoroutineType(@Nullable PyType returnType) { - if (isGenerator()) { - // Re-wrap typing.Generator[Y, S, R] into typing.AsyncGenerator[Y, Any] - if (returnType instanceof PyCollectionType) { - final PyClassLikeType classType = as(returnType, PyClassLikeType.class); - if (classType != null && PyTypingTypeProvider.GENERATOR.equals(classType.getClassQName())) { + private PyType createAsyncType(@Nullable PyType returnType) { + if (returnType instanceof PyCollectionType) { + final PyClassLikeType classType = as(returnType, PyClassLikeType.class); + if (classType != null) { + if (PyTypingTypeProvider.GENERATOR.equals(classType.getClassQName())) { final QualifiedName asyncGeneratorName = QualifiedName.fromDottedString(PyTypingTypeProvider.ASYNC_GENERATOR); final PsiElement resolvedGenerator = PyResolveImportUtil.resolveTopLevelMember(asyncGeneratorName, PyResolveImportUtil.fromFoothold(this)); @@ -397,16 +396,14 @@ public class PyFunctionImpl extends PyBaseElementImpl implements return new PyCollectionTypeImpl(asyncGenerator, false, Arrays.asList(((PyCollectionType)returnType).getIteratedItemType(), null)); } + else { + return null; + } } } - // Leave the type as is - return returnType; - } - else { - final PyClass coroutine = as(PyResolveImportUtil.resolveTopLevelMember(QualifiedName.fromDottedString(PyTypingTypeProvider.COROUTINE), - PyResolveImportUtil.fromFoothold(this)), PyClass.class); - return coroutine != null ? new PyCollectionTypeImpl(coroutine, false, Arrays.asList(null, null, returnType)) : null; } + + return PyTypingTypeProvider.wrapInCoroutineType(returnType, this); } public PyFunction asMethod() { diff --git a/python/testData/pyi/type/coroutineType/CoroutineType.py b/python/testData/pyi/type/coroutineType/CoroutineType.py new file mode 100644 index 000000000000..5eb61e524dc4 --- /dev/null +++ b/python/testData/pyi/type/coroutineType/CoroutineType.py @@ -0,0 +1,4 @@ +async def f(): + return 42 + +coroutine = f() \ No newline at end of file diff --git a/python/testData/pyi/type/coroutineType/CoroutineType.pyi b/python/testData/pyi/type/coroutineType/CoroutineType.pyi new file mode 100644 index 000000000000..34d671d403e0 --- /dev/null +++ b/python/testData/pyi/type/coroutineType/CoroutineType.pyi @@ -0,0 +1,2 @@ +async def f() -> int: + ... \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/pyi/PyiTypeTest.java b/python/testSrc/com/jetbrains/python/pyi/PyiTypeTest.java index eddca9007514..0349ee5cee7f 100644 --- a/python/testSrc/com/jetbrains/python/pyi/PyiTypeTest.java +++ b/python/testSrc/com/jetbrains/python/pyi/PyiTypeTest.java @@ -101,6 +101,10 @@ public class PyiTypeTest extends PyTestCase { doTest("int"); } + public void testCoroutineType() { + doTest("Coroutine[Any, Any, int]"); + } + public void testPyiOnPythonPath() { addPyiStubsToContentRoot(myFixture); doTest("int");