Initial implementation of lists type inference (PY-1182)

This commit is contained in:
Lada Gagina
2018-02-05 19:30:26 +03:00
committed by Elizaveta Shashkova
parent adf85192a7
commit cedb7bda67
4 changed files with 135 additions and 4 deletions
@@ -2,6 +2,7 @@
package com.jetbrains.python.psi;
import com.google.common.collect.Collections2;
import com.google.common.collect.ImmutableSet;
import com.google.common.collect.Maps;
import com.intellij.codeInsight.FileModificationService;
import com.intellij.codeInsight.completion.PrioritizedLookupElement;
@@ -34,6 +35,7 @@ import com.intellij.openapi.roots.ModuleRootManager;
import com.intellij.openapi.ui.MessageType;
import com.intellij.openapi.ui.popup.Balloon;
import com.intellij.openapi.ui.popup.JBPopupFactory;
import com.intellij.openapi.util.Pair;
import com.intellij.openapi.util.TextRange;
import com.intellij.openapi.util.io.FileUtil;
import com.intellij.openapi.util.io.FileUtilRt;
@@ -1988,4 +1990,86 @@ public class PyUtil {
return ret;
}
}
@Nullable
public static PyType getCollectionTypeByModifications(@Nullable PsiElement parent, @NotNull TypeEvalContext context) {
if (parent instanceof PyAssignmentStatement) {
final PyExpression[] targets = ((PyAssignmentStatement)parent).getTargets();
if (targets.length == 1 && targets[0] != null) {
final PyExpression expr = targets[0];
final List<Pair<String, PyType>> modifications = findModifications(expr, context);
final Set<PyType> types = new LinkedHashSet<>();
for (Pair<String, PyType> modification : modifications) {
final String funcName = modification.getFirst();
final PyType argType = modification.getSecond();
if (funcName.equals("extend")) {
if (argType != null && argType instanceof PyCollectionType) {
final PyType argElemType = PyUnionType.union(((PyCollectionType)argType).getElementTypes());
types.add(argElemType);
}
}
else {
types.add(argType);
}
}
return PyUnionType.union(types);
}
}
return null;
}
@NotNull
private static List<Pair<String, PyType>> findModifications(@NotNull PsiElement element, TypeEvalContext context) {
final CollectionTypeVisitor visitor = new CollectionTypeVisitor(element, context);
ScopeOwner owner = ScopeUtil.getScopeOwner(element);
if (owner != null) {
owner.accept(visitor);
}
return visitor.result();
}
private static class CollectionTypeVisitor extends PyRecursiveElementVisitor {
private final PsiElement myElement;
private final List<Pair<String, PyType>> myModifications;
private final TypeEvalContext myTypeEvalContext;
private static final Set<String> SEQUENCE_MODIFICATION_METHODS = ImmutableSet.of(
"append",
"extend",
"insert",
"index"
);
public CollectionTypeVisitor(@NotNull PsiElement element, @NotNull TypeEvalContext context) {
myElement = element;
myTypeEvalContext = context;
myModifications = new ArrayList<>();
}
@Override
public void visitPyCallExpression(PyCallExpression node) {
final PyExpression callee = node.getCallee();
if (callee instanceof PyQualifiedExpression) {
final PyExpression qualifier = ((PyQualifiedExpression)callee).getQualifier();
final String funcName = ((PyQualifiedExpression)callee).getReferencedName();
if (qualifier != null) {
final PsiReference reference = qualifier.getReference();
if (SEQUENCE_MODIFICATION_METHODS.contains(funcName) && reference != null && reference.isReferenceTo(myElement)) {
PyExpression[] arguments = node.getArguments();
if (arguments.length == 1 && arguments[0] != null) {
myModifications.add(Pair.create(funcName, myTypeEvalContext.getType(arguments[0])));
}
if (arguments.length == 2) { // insert(pos, item)
myModifications.add(Pair.create(funcName, myTypeEvalContext.getType(arguments[1])));
}
}
}
}
}
@NotNull
public List<Pair<String, PyType>> result() {
return myModifications;
}
}
}
@@ -170,7 +170,7 @@ public class PyBuiltinCache {
@NotNull
private static List<PyType> getSequenceElementTypes(@NotNull PySequenceExpression sequence, @NotNull TypeEvalContext context) {
if (sequence instanceof PyListLiteralExpression || sequence instanceof PySetLiteralExpression) {
return Collections.singletonList(getListOrSetIteratedValueType(sequence.getElements(), context));
return Collections.singletonList(getListOrSetIteratedValueType(sequence.getElements(), context, sequence.getParent()));
}
else if (sequence instanceof PyDictLiteralExpression) {
return getDictElementTypes(sequence.getElements(), context);
@@ -181,14 +181,24 @@ public class PyBuiltinCache {
}
@Nullable
private static PyType getListOrSetIteratedValueType(@NotNull PyExpression[] elements, @NotNull TypeEvalContext context) {
private static PyType getListOrSetIteratedValueType(@NotNull PyExpression[] elements, @NotNull TypeEvalContext context,
@Nullable PsiElement parent) {
final int maxAnalyzedElements = Math.min(MAX_ANALYZED_ELEMENTS_OF_LITERALS, elements.length);
final PyType analyzedElementsType = StreamEx
PyType analyzedElementsType = StreamEx
.of(elements, 0, maxAnalyzedElements)
.map(context::getType)
.toListAndThen(PyUnionType::union);
PyType typeByModifications = PyUtil.getCollectionTypeByModifications(parent, context);
if (analyzedElementsType == null) {
analyzedElementsType = typeByModifications;
}
else {
if (typeByModifications != null) {
analyzedElementsType = PyUnionType.union(analyzedElementsType, typeByModifications);
}
}
if (elements.length > maxAnalyzedElements) {
return PyUnionType.createWeakType(analyzedElementsType);
}
@@ -6,4 +6,4 @@ xs.remove(<weak_warning descr="Expected type 'int' (matched generic type '_T'),
xs.add(<weak_warning descr="Expected type 'int' (matched generic type '_T'), got 'object' instead">object()</weak_warning>)
ys = ['green', 'eggs']
ys.extend(<weak_warning descr="Expected type 'Iterable[str]' (matched generic type 'Iterable[_T]'), got 'Set[int]' instead">xs</weak_warning>)
ys.extend(xs)
@@ -1228,6 +1228,43 @@ public class PyTypeTest extends PyTestCase {
TypeEvalContext.userInitiated(expr.getProject(), expr.getContainingFile()));
}
// PY-1182
public void testCollectionType() {
doTest("List[int]",
"def f():\n" +
" expr = []\n" +
" expr.append(42)\n" +
" expr.append(0)"
);
doTest("List[Union[str, int]]",
"def f():\n" +
" expr = []\n" +
" expr.append('a')\n" +
" expr.append(1)"
);
doTest("List[Union[Union[int, str], Any]]",
"expr = [3, 4, 5, 6, 7, 8, 9, 1, 2, 3, 4]\n" +
"expr.append('a')"
);
doTest("List[Union[int, str]]",
"expr = [1, 2]\n" +
"expr.append(42)\n" +
"expr.extend(['a']"
);
doTest("List[Union[int, str, None]]",
"expr = []\n" +
"expr.extend([1, 'a', None])"
);
doTest("List[int]",
"expr = []\n" +
"expr.index(42)");
}
// PY-20063
public void testIteratedSetElement() {
doTest("int",