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");