mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Initial implementation of lists type inference (PY-1182)
This commit is contained in:
committed by
Elizaveta Shashkova
parent
adf85192a7
commit
cedb7bda67
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user