diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java
index be2502b27cf5..114ad867e10d 100644
--- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java
+++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java
@@ -39,6 +39,8 @@ import org.jetbrains.annotations.Nullable;
import java.util.*;
+import static com.jetbrains.python.psi.PyUtil.as;
+
/**
* @author vlan
*/
@@ -110,22 +112,24 @@ public class PyTypeCheckerInspection extends PyInspection {
private PyType getExpectedReturnType(@NotNull PyFunction function) {
final PyType returnType = myTypeEvalContext.getReturnType(function);
- if (returnType instanceof PyCollectionType) {
- final PyCollectionType genericType = (PyCollectionType)returnType;
- if (PyNames.FAKE_COROUTINE.equals(genericType.getName())) {
+ final PyCollectionType genericType = as(returnType, PyCollectionType.class);
+ final PyClassType classType = as(returnType, PyClassType.class);
+
+ if (function.isAsync()) {
+ if (genericType != null && PyNames.FAKE_COROUTINE.equals(genericType.getName())) {
return genericType.getIteratedItemType();
}
- else if (function.isGenerator()) {
- if (PyNames.FAKE_GENERATOR.equals(genericType.getName()) ||
- genericType instanceof PyClassType && "typing.Generator".equals(((PyClassType)genericType).getClassQName())) {
- // Generator's type is parametrized as [YieldType, SendType, ReturnType]
- return ContainerUtil.getOrElse(genericType.getElementTypes(myTypeEvalContext), 2, null);
- }
- else {
- // Assume that any other return type annotation for a generator cannot contain its return type
- return null;
- }
+ // Async generators are not allowed to return anything anyway
+ return null;
+ }
+ else if (function.isGenerator()) {
+ if (genericType != null && classType != null &&
+ (PyNames.FAKE_GENERATOR.equals(genericType.getName()) || "typing.Generator".equals(classType.getClassQName()))) {
+ // Generator's type is parametrized as [YieldType, SendType, ReturnType]
+ return ContainerUtil.getOrElse(genericType.getElementTypes(myTypeEvalContext), 2, null);
}
+ // Assume that any other return type annotation for a generator cannot contain its return type
+ return null;
}
return returnType;
diff --git a/python/testData/inspections/PyTypeCheckerInspection/AsyncGeneratorAnnotatedToReturnAsyncIterable.py b/python/testData/inspections/PyTypeCheckerInspection/AsyncGeneratorAnnotatedToReturnAsyncIterable.py
new file mode 100644
index 000000000000..a9ea51f79360
--- /dev/null
+++ b/python/testData/inspections/PyTypeCheckerInspection/AsyncGeneratorAnnotatedToReturnAsyncIterable.py
@@ -0,0 +1,16 @@
+from typing import AsyncIterable
+
+
+async def g1() -> AsyncIterable[int]:
+ yield 42
+
+async def g2() -> AsyncIterable[int]:
+ yield 42
+ return None
+
+async def g3() -> AsyncIterable:
+ yield 42
+
+async def g4() -> AsyncIterable:
+ yield 42
+ return None
\ No newline at end of file
diff --git a/python/testData/inspections/PyTypeCheckerInspection/GeneratorAnnotatedToReturnIterable.py b/python/testData/inspections/PyTypeCheckerInspection/GeneratorAnnotatedToReturnIterable.py
index 62ac2b327f1f..45964b84209c 100644
--- a/python/testData/inspections/PyTypeCheckerInspection/GeneratorAnnotatedToReturnIterable.py
+++ b/python/testData/inspections/PyTypeCheckerInspection/GeneratorAnnotatedToReturnIterable.py
@@ -8,4 +8,13 @@ def g1() -> Iterable[int]:
def g2() -> Iterable[int]:
yield 42
- return None
\ No newline at end of file
+ return None
+
+
+def g3() -> Iterable:
+ yield 42
+
+
+def g4() -> Iterable:
+ yield 42
+ return None
diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java
index 554a3387219e..fff8fdb3369a 100644
--- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java
+++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java
@@ -179,11 +179,16 @@ public class Py3TypeCheckerInspectionTest extends PyTestCase {
doTest();
}
- // PY-20657
+ // PY-20657, PY-21916
public void testGeneratorAnnotatedToReturnIterable() {
doTest();
}
+ // PY-20657, PY-21916
+ public void testAsyncGeneratorAnnotatedToReturnAsyncIterable() {
+ doTest();
+ }
+
// PY-21083
public void testFloatFromhex() {
doTest();