From a93b24d7e6be2339d0d12a01899138afe28b53bb Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Tue, 4 Oct 2016 16:48:35 +0300 Subject: [PATCH] Update PyFunctionImpl.getYieldStatementType to correctly handle collections without type arguments --- .../python/psi/impl/PyFunctionImpl.java | 18 +++++++---- .../com/jetbrains/python/Py3TypeTest.java | 30 +++++++++++++++++++ 2 files changed, 43 insertions(+), 5 deletions(-) diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index 985b34884fb1..23a49dc5cd54 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -324,11 +324,19 @@ public class PyFunctionImpl extends PyBaseElementImpl implements public void visitPyYieldExpression(PyYieldExpression node) { final PyExpression expr = node.getExpression(); final PyType type = expr != null ? context.getType(expr) : null; - if (node.isDelegating() && type instanceof PyCollectionType) { - final PyCollectionType collectionType = (PyCollectionType)type; - // TODO: Select the parameter types that matches T in Iterable[T] - final List elementTypes = collectionType.getElementTypes(context); - types.add(elementTypes.isEmpty() ? null : elementTypes.get(0)); + if (node.isDelegating()) { + if (type instanceof PyCollectionType) { + final PyCollectionType collectionType = (PyCollectionType)type; + // TODO: Select the parameter types that matches T in Iterable[T] + final List elementTypes = collectionType.getElementTypes(context); + types.add(elementTypes.isEmpty() ? null : elementTypes.get(0)); + } + else if (ArrayUtil.contains(type, cache.getListType(), cache.getDictType(), cache.getSetType())) { + types.add(null); + } + else { + types.add(type); + } } else { types.add(type); diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index c8e8d7972546..232d6df67327 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -84,6 +84,36 @@ public class Py3TypeTest extends PyTestCase { " pass")); } + public void testYieldFromUnknownList() { + doTest("Any", + "def get_list() -> list:\n" + + " pass\n" + + "def gen()\n" + + " yield from get_list()\n" + + "for expr in gen():" + + " pass"); + } + + public void testYieldFromUnknownDict() { + doTest("Any", + "def get_dict() -> dict:\n" + + " pass\n" + + "def gen()\n" + + " yield from get_dict()\n" + + "for expr in gen():" + + " pass"); + } + + public void testYieldFromUnknownSet() { + doTest("Any", + "def get_set() -> set:\n" + + " pass\n" + + "def gen()\n" + + " yield from get_set()\n" + + "for expr in gen():" + + " pass"); + } + public void testAwaitAwaitable() { runWithLanguageLevel(LanguageLevel.PYTHON35, () -> doTest("int", "class C:\n" +