From 39cac66e40ba1579652a668312ce7e2fabb3085b Mon Sep 17 00:00:00 2001 From: Petr Date: Fri, 28 Nov 2025 13:34:41 +0100 Subject: [PATCH] PY-41061 Fix type inference for (async) iteration over objects with both `__iter__` and `__aiter__` GitOrigin-RevId: 0eb20894a1df40b62b45e7ef0166a7ce72cdc753 --- .../psi/impl/PyTargetExpressionImpl.java | 58 +++++++++++-------- .../psi/impl/PyYieldExpressionImpl.java | 2 +- .../com/jetbrains/python/Py3TypeTest.java | 52 +++++++++++++++++ 3 files changed, 86 insertions(+), 26 deletions(-) diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java index 4d1b1d9a9c95..64fb56ddf8dc 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java @@ -218,7 +218,7 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl getIterationType(member, source, anchor, context)); + if (iterableType instanceof PyUnionType) { + return ((PyUnionType)iterableType).map(member -> getIterationType(member, source, anchor, isAsync, context)); } - else if (iterableType != null && PyABCUtil.isSubtype(iterableType, PyNames.ITERABLE, context)) { - final PyFunction iterateMethod = findMethodByName(iterableType, PyNames.ITER, context); - if (iterateMethod != null) { - final PyType iterateReturnType = getContextSensitiveType(iterateMethod, context, source); - return getIteratedItemType(iterateReturnType, source, anchor, context, false); - } - final Ref nextMethodCallType = getNextMethodCallType(iterableType, source, anchor, context, false); - if (nextMethodCallType != null) { - return nextMethodCallType.get(); - } - final PyFunction getItem = findMethodByName(iterableType, PyNames.GETITEM, context); - if (getItem != null) { - return getContextSensitiveType(getItem, context, source); + if (!isAsync) { + if (iterableType != null && PyABCUtil.isSubtype(iterableType, PyNames.ITERABLE, context)) { + final PyFunction iterateMethod = findMethodByName(iterableType, PyNames.ITER, context); + if (iterateMethod != null) { + final PyType iterateReturnType = getContextSensitiveType(iterateMethod, context, source); + return getIteratedItemType(iterateReturnType, source, anchor, context, false); + } + final Ref nextMethodCallType = getNextMethodCallType(iterableType, source, anchor, context, false); + if (nextMethodCallType != null) { + return nextMethodCallType.get(); + } + final PyFunction getItem = findMethodByName(iterableType, PyNames.GETITEM, context); + if (getItem != null) { + return getContextSensitiveType(getItem, context, source); + } } } - else if (iterableType != null && PyABCUtil.isSubtype(iterableType, PyNames.ASYNC_ITERABLE, context)) { - final PyFunction iterateMethod = findMethodByName(iterableType, PyNames.AITER, context); - if (iterateMethod != null) { - final PyType iterateReturnType = getContextSensitiveType(iterateMethod, context, source); - return getIteratedItemType(iterateReturnType, source, anchor, context, true); + else { + if (iterableType != null && PyABCUtil.isSubtype(iterableType, PyNames.ASYNC_ITERABLE, context)) { + final PyFunction iterateMethod = findMethodByName(iterableType, PyNames.AITER, context); + if (iterateMethod != null) { + final PyType iterateReturnType = getContextSensitiveType(iterateMethod, context, source); + return getIteratedItemType(iterateReturnType, source, anchor, context, true); + } } } return null; diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyYieldExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyYieldExpressionImpl.java index 59918bc2226c..8484c9864e33 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyYieldExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyYieldExpressionImpl.java @@ -46,7 +46,7 @@ public class PyYieldExpressionImpl extends PyElementImpl implements PyYieldExpre final PyType type = expr != null ? context.getType(expr) : PyBuiltinCache.getInstance(this).getNoneType(); if (isDelegating()) { - return PyTargetExpressionImpl.getIterationType(type, expr, this, context); + return PyTargetExpressionImpl.getIterationType(type, expr, this, false, context); } return type; } diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index c4a53dd48b12..a5fdd1e6dac8 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -1071,6 +1071,58 @@ public class Py3TypeTest extends PyTestCase { """); } + // PY-41061 + public void testForIterationOverObjectWithIterAndAiter() { + doTest("int", + """ + from collections.abc import Iterator, AsyncIterator + + class MyIterable: + def __iter__(self) -> Iterator[int]: ... + + def __aiter__(self) -> AsyncIterator[str]: ... + + for expr in MyIterable(): ..."""); + + doTest("int", + """ + from collections.abc import Iterator, AsyncIterator + + class MyIterable: + def __iter__(self) -> Iterator[int]: ... + + def __aiter__(self) -> AsyncIterator[str]: ... + + _ = [expr for expr in MyIterable()]"""); + } + + // PY-41061 + public void testAsyncForIterationOverObjectWithIterAndAiter() { + doTest("str", + """ + from collections.abc import Iterator, AsyncIterator + + class MyIterable: + def __iter__(self) -> Iterator[int]: ... + + def __aiter__(self) -> AsyncIterator[str]: ... + + async def foo(): + async for expr in MyIterable(): ..."""); + + doTest("str", + """ + from collections.abc import Iterator, AsyncIterator + + class MyIterable: + def __iter__(self) -> Iterator[int]: ... + + def __aiter__(self) -> AsyncIterator[str]: ... + + async def foo(): + _ = [expr async for expr in MyIterable()]"""); + } + // PY-21655 public void testUsageOfFunctionDecoratedWithAsyncioCoroutine() { myFixture.copyDirectoryToProject(TEST_DIRECTORY + getTestName(false), "");