PY-87777 python: async def should return a types.CoroutineType, not a typing.Coroutine

(cherry picked from commit 7992f6b1eef198e35771fa59e3587656e252b24d)

GitOrigin-RevId: db84f2dc3a485670f02f9aba039007c769c03ed9
This commit is contained in:
Tatiana Ber
2026-03-12 12:22:02 +00:00
committed by intellij-monorepo-bot
parent bc23d305d4
commit cc77134fc9
10 changed files with 48 additions and 29 deletions
@@ -109,6 +109,7 @@ object PyNames {
const val FUNCTION: String = "function"
const val TYPES_FUNCTION_TYPE: String = "types.FunctionType"
const val TYPES_COROUTINE_TYPE: String = "types.CoroutineType"
const val TYPES_METHOD_TYPE: String = "types.UnboundMethodType"
const val FUTURE_MODULE: String = "__future__"
@@ -221,7 +221,7 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext<Context?>() {
if (typeRef != null) {
// Do not use toAsyncIfNeeded, as it also converts Generators. Here we do not need it.
if (callable.isAsync && callable.isAsyncAllowed && !callable.isGenerator) {
return Ref(wrapInCoroutineType(typeRef.get(), callable))
return Ref(typeRef.get().wrapInCoroutineType(callable))
}
return typeRef
}
@@ -2803,7 +2803,7 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext<Context?>() {
fun toAsyncIfNeeded(function: PyFunction, returnType: PyType?): PyType? {
if (function.isAsync && function.isAsyncAllowed) {
if (!function.isGenerator) {
return wrapInCoroutineType(returnType, function)
return returnType.wrapInCoroutineType(function)
}
val desc = GeneratorTypeDescriptor.fromGenerator(returnType)
if (desc != null) {
@@ -2828,9 +2828,17 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext<Context?>() {
}
}
private fun wrapInCoroutineType(returnType: PyType?, resolveAnchor: PsiElement): PyType? {
val coroutine = PyPsiFacade.getInstance(resolveAnchor.project).createClassByQName(COROUTINE, resolveAnchor)
return if (coroutine != null) PyCollectionTypeImpl(coroutine, false, listOf(null, null, returnType)) else null
private fun PyType?.wrapInCoroutineType(anchor: PsiElement): PyType? {
val facade = PyPsiFacade.getInstance(anchor.project)
val targetClass =
facade.createClassByQName(PyNames.TYPES_COROUTINE_TYPE, anchor)
?: facade.createClassByQName(COROUTINE, anchor)
?: return null
return PyCollectionTypeImpl(
targetClass,
false,
listOf(null, null, this)
)
}
@JvmStatic
@@ -2851,18 +2859,15 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext<Context?>() {
@JvmStatic
fun unwrapCoroutineReturnType(coroutineType: PyType?): Ref<PyType?>? {
val genericType = PyUtil.`as`(coroutineType, PyCollectionType::class.java)
if (coroutineType !is PyCollectionType) return null
val qName = coroutineType.classQName
if (genericType != null) {
val qName = genericType.classQName
if (AWAITABLE == qName) {
return Ref(coroutineType.elementTypes.getOrNull(0))
}
if (AWAITABLE == qName) {
return Ref(ContainerUtil.getOrElse<PyType?>(genericType.elementTypes, 0, null))
}
if (COROUTINE == qName) {
return Ref(ContainerUtil.getOrElse<PyType?>(genericType.elementTypes, 2, null))
}
if (qName in arrayOf(COROUTINE, PyNames.TYPES_COROUTINE_TYPE)) {
return Ref(coroutineType.elementTypes.getOrNull(2))
}
return null
@@ -86,8 +86,7 @@ class PySuspiciousBooleanConditionInspection : PyInspection() {
val type = myTypeEvalContext.getType(this) ?: return
// Check if the type is a coroutine
// TODO: use `CoroutineType` instead
if (type is PyClassType && type.classQName == PyTypingTypeProvider.COROUTINE) {
if (type is PyClassType && (type.classQName == PyTypingTypeProvider.COROUTINE || type.classQName == PyNames.TYPES_COROUTINE_TYPE)) {
registerProblem(this, PyPsiBundle.message("INSP.suspicious.boolean.condition.coroutine"), PyAddAwaitQuickFix())
}
}
@@ -20,6 +20,7 @@ import com.intellij.util.ProcessingContext
import com.intellij.util.containers.ContainerUtil
import com.intellij.util.containers.toArray
import com.jetbrains.python.BaseReference
import com.jetbrains.python.PyNames
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider
import com.jetbrains.python.psi.PyArgumentList
import com.jetbrains.python.psi.PyCallExpression
@@ -145,7 +146,7 @@ class PyTextFixtureTypeProvider : PyTypeProviderBase() {
if (ArrayUtil.contains(qName, "typing.Awaitable", PyTypingTypeProvider.GENERATOR)) {
return Ref.create(ContainerUtil.getOrElse(genericType.getElementTypes(), 0, null))
}
if (PyTypingTypeProvider.COROUTINE == qName) {
if (PyTypingTypeProvider.COROUTINE == qName || PyNames.TYPES_COROUTINE_TYPE == qName) {
return Ref.create(ContainerUtil.getOrElse(genericType.getElementTypes(), 2, null))
}
}
@@ -1,8 +1,9 @@
from typing import Any, Coroutine
from types import CoroutineType
from typing import Any
async def bar() -> int:
return 42
def foo(x, y) -> Coroutine[Any, Any, int]:
def foo(x, y) -> CoroutineType[Any, Any, int]:
return bar()
@@ -1 +1 @@
<html><body><div class="bottom"><icon src="AllIcons.Nodes.Package"/>&nbsp;<code><a href="psi_element://#module#AsyncFunctionQuickDoc">AsyncFunctionQuickDoc</a></code></div><div class="definition"><pre><span style="color:#000080;font-weight:bold;">async </span><span style="color:#000080;font-weight:bold;">def </span><span style="color:#000000;">func</span><span style="">(</span><span style="">)</span> -&gt; <span style="color:#000000;">Coroutine<span style="">[</span>Any<span style="">, </span>Any<span style="">, </span><span style="color:#000080;font-weight:bold;"><a href="psi_element://#typename#types.NoneType">None</a></span><span style="">]</span></span></pre></div></body></html>
<html><body><div class="bottom"><icon src="AllIcons.Nodes.Package"/>&nbsp;<code><a href="psi_element://#module#AsyncFunctionQuickDoc">AsyncFunctionQuickDoc</a></code></div><div class="definition"><pre><span style="color:#000080;font-weight:bold;">async </span><span style="color:#000080;font-weight:bold;">def </span><span style="color:#000000;">func</span><span style="">(</span><span style="">)</span> -&gt; <span style="color:#000000;">CoroutineType<span style="">[</span>Any<span style="">, </span>Any<span style="">, </span><span style="color:#000080;font-weight:bold;"><a href="psi_element://#typename#types.NoneType">None</a></span><span style="">]</span></span></pre></div></body></html>
@@ -1 +1 @@
<span style="color:#000080;font-weight:bold;">async </span><span style="color:#000080;font-weight:bold;">def </span><span style="color:#000000;">func</span><span style="">(</span><span style="">)</span> -&gt; <span style="color:#000000;">Coroutine<span style="">[</span>Any<span style="">, </span>Any<span style="">, </span><span style="color:#000080;font-weight:bold;"><a href="psi_element://#typename#types.NoneType">None</a></span><span style="">]</span></span>
<span style="color:#000080;font-weight:bold;">async </span><span style="color:#000080;font-weight:bold;">def </span><span style="color:#000000;">func</span><span style="">(</span><span style="">)</span> -&gt; <span style="color:#000000;">CoroutineType<span style="">[</span>Any<span style="">, </span>Any<span style="">, </span><span style="color:#000080;font-weight:bold;"><a href="psi_element://#typename#types.NoneType">None</a></span><span style="">]</span></span>
@@ -531,7 +531,7 @@ public class Py3TypeTest extends PyTestCase {
}
public void testAsyncDefReturnType() {
doTest("Coroutine[Any, Any, int]",
doTest("CoroutineType[Any, Any, int]",
"""
async def foo(x):
await x
@@ -565,6 +565,18 @@ public class Py3TypeTest extends PyTestCase {
""");
}
public void testAwaitOnTypingCoroutineAnnotation() {
doTest("int",
"""
from typing import Any, Coroutine
x: Coroutine[Any, Any, int]
async def bar():
expr = await x
""");
}
// PY-16987
public void testNoTypeInGoogleDocstringParamAnnotation() {
doTest("int", """
@@ -1553,7 +1565,7 @@ public class Py3TypeTest extends PyTestCase {
// PY-24067
public void testAsyncFunctionReturnTypeInDocstring() {
doTest("Coroutine[Any, Any, int]",
doTest("CoroutineType[Any, Any, int]",
"""
async def f():
""\"
@@ -1565,7 +1577,7 @@ public class Py3TypeTest extends PyTestCase {
// PY-27518
public void testAsyncFunctionReturnTypeInNumpyDocstring() {
doTest("Coroutine[Any, Any, int]",
doTest("CoroutineType[Any, Any, int]",
"""
async def f():
""\"
@@ -1592,7 +1604,7 @@ public class Py3TypeTest extends PyTestCase {
// PY-26643
public void testReplaceSelfInCoroutine() {
doTest("Coroutine[Any, Any, B]",
doTest("CoroutineType[Any, Any, B]",
"""
class A:
async def foo(self):
@@ -1152,7 +1152,7 @@ public class PyTypingTest extends PyTestCase {
}
public void testCoroutineReturnsGenerator() {
doTest("Coroutine[Any, Any, Generator[int, Any, Any]]",
doTest("CoroutineType[Any, Any, Generator[int, Any, Any]]",
"""
from typing import Generator
@@ -6271,7 +6271,7 @@ public class PyTypingTest extends PyTestCase {
// PY-36416
public void testReturnTypeOfNonAnnotatedAsyncOverride() {
doTest("Coroutine[Any, Any, str]", """
doTest("CoroutineType[Any, Any, str]", """
class Base:
async def get(self) -> str:
...
@@ -97,7 +97,7 @@ public class PyiTypeTest extends PyTestCase {
}
public void testCoroutineType() {
doTest("Coroutine[Any, Any, int]");
doTest("CoroutineType[Any, Any, int]");
}
public void testPyiOnPythonPath() {