Fixed type inference for list comprehensions (PY-7021)

This commit is contained in:
Andrey Vlasovskikh
2012-08-16 18:47:54 +04:00
parent f026238360
commit e14a7aa834
2 changed files with 28 additions and 1 deletions
@@ -1,11 +1,15 @@
package com.jetbrains.python.psi.impl;
import com.intellij.lang.ASTNode;
import com.jetbrains.python.psi.PyClass;
import com.jetbrains.python.psi.PyElementVisitor;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.PyListCompExpression;
import com.jetbrains.python.psi.types.PyCollectionTypeImpl;
import com.jetbrains.python.psi.types.PyType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
/**
* @author yole
@@ -20,7 +24,16 @@ public class PyListCompExpressionImpl extends PyComprehensionElementImpl impleme
pyVisitor.visitPyListCompExpression(this);
}
@Nullable
@Override
public PyType getType(@NotNull TypeEvalContext context) {
return PyBuiltinCache.getInstance(this).getListType();
final PyExpression resultExpr = getResultExpression();
final PyBuiltinCache cache = PyBuiltinCache.getInstance(this);
final PyClass list = cache.getClass("list");
if (resultExpr != null && list != null) {
final PyType elementType = resultExpr.getType(context);
return new PyCollectionTypeImpl(list, false, elementType);
}
return cache.getListType();
}
}
@@ -441,6 +441,20 @@ public class PyTypeTest extends PyTestCase {
"expr = f()\n");
}
// PY-7020
public void testListComprehensionType() {
final PyExpression expr = parseExpr("expr = [str(x) for x in range(10)]\n");
final TypeEvalContext context = TypeEvalContext.slow().withTracing();
final PyType type = expr.getType(context);
assertNotNull(type);
assertInstanceOf(type, PyCollectionType.class);
assertEquals(type.getName(), "list");
final PyCollectionType collectionType = (PyCollectionType)type;
final PyType elementType = collectionType.getElementType(context);
assertNotNull(elementType);
assertEquals(elementType.getName(), "str");
}
private PyExpression parseExpr(String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
return myFixture.findElementByText("expr", PyExpression.class);