diff --git a/python/src/com/jetbrains/python/psi/PyUtil.java b/python/src/com/jetbrains/python/psi/PyUtil.java index e0ccc25c53e0..d268d6229fc4 100644 --- a/python/src/com/jetbrains/python/psi/PyUtil.java +++ b/python/src/com/jetbrains/python/psi/PyUtil.java @@ -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> modifications = findModifications(expr, context); + final Set types = new LinkedHashSet<>(); + for (Pair 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> 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> myModifications; + private final TypeEvalContext myTypeEvalContext; + + private static final Set 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> result() { + return myModifications; + } + } } diff --git a/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java b/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java index 23672a29d81f..0403209f59e9 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java +++ b/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java @@ -170,7 +170,7 @@ public class PyBuiltinCache { @NotNull private static List 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); } diff --git a/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py b/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py index 1ff0f01d078d..4f367df931bb 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py +++ b/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py @@ -6,4 +6,4 @@ xs.remove(object()) ys = ['green', 'eggs'] -ys.extend(xs) +ys.extend(xs) diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index 7d21037b0276..61417f44fb3e 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -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",