Update PyFunctionImpl.getYieldStatementType to correctly handle collections without type arguments

This commit is contained in:
Semyon Proshev
2016-10-13 21:28:17 +03:00
committed by Semyon Proshev
parent c5ed4e2659
commit a93b24d7e6
2 changed files with 43 additions and 5 deletions
@@ -324,11 +324,19 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> 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<PyType> 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<PyType> 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);
@@ -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" +