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