diff --git a/python/src/META-INF/python-core-common.xml b/python/src/META-INF/python-core-common.xml
index 0674591086d1..4a5e192e4c23 100644
--- a/python/src/META-INF/python-core-common.xml
+++ b/python/src/META-INF/python-core-common.xml
@@ -708,6 +708,7 @@
+
diff --git a/python/src/com/jetbrains/python/psi/PyUtil.java b/python/src/com/jetbrains/python/psi/PyUtil.java
index d268d6229fc4..e0ccc25c53e0 100644
--- a/python/src/com/jetbrains/python/psi/PyUtil.java
+++ b/python/src/com/jetbrains/python/psi/PyUtil.java
@@ -2,7 +2,6 @@
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;
@@ -35,7 +34,6 @@ 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;
@@ -1990,86 +1988,4 @@ 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 0403209f59e9..623f5c7d00ee 100644
--- a/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java
+++ b/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java
@@ -34,12 +34,13 @@ import com.jetbrains.python.psi.resolve.PyResolveImportUtil;
import com.jetbrains.python.psi.resolve.PythonSdkPathCache;
import com.jetbrains.python.psi.types.*;
import com.jetbrains.python.sdk.PythonSdkType;
-import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NonNls;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
-import java.util.*;
+import java.util.HashMap;
+import java.util.List;
+import java.util.Map;
import static com.jetbrains.python.psi.PyUtil.as;
@@ -53,8 +54,6 @@ public class PyBuiltinCache {
private static final PyBuiltinCache DUD_INSTANCE = new PyBuiltinCache(null, null);
- private static final int MAX_ANALYZED_ELEMENTS_OF_LITERALS = 10; /* performance */
-
/**
* Stores the most often used types, returned by getNNNType().
*/
@@ -162,96 +161,11 @@ public class PyBuiltinCache {
public PyType createLiteralCollectionType(final PySequenceExpression sequence, final String name, @NotNull TypeEvalContext context) {
final PyClass cls = getClass(name);
if (cls != null) {
- return new PyCollectionTypeImpl(cls, false, getSequenceElementTypes(sequence, context));
+ return new PyCollectionTypeImpl(cls, false, PyCollectionTypeUtil.INSTANCE.getTypeByModifications(sequence, context));
}
return null;
}
- @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, sequence.getParent()));
- }
- else if (sequence instanceof PyDictLiteralExpression) {
- return getDictElementTypes(sequence.getElements(), context);
- }
- else {
- return Collections.singletonList(null);
- }
- }
-
- @Nullable
- 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);
-
- 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);
- }
- else {
- return analyzedElementsType;
- }
- }
-
- @NotNull
- private static List getDictElementTypes(@NotNull PyExpression[] elements, @NotNull TypeEvalContext context) {
- final int maxAnalyzedElements = Math.min(MAX_ANALYZED_ELEMENTS_OF_LITERALS, elements.length);
-
- final List keyTypes = new ArrayList<>();
- final List valueTypes = new ArrayList<>();
-
- StreamEx
- .of(elements, 0, maxAnalyzedElements)
- .map(element -> as(context.getType(element), PyTupleType.class))
- .forEach(
- tupleType -> {
- if (tupleType != null) {
- final List tupleElementTypes = tupleType.getElementTypes();
-
- if (tupleType.isHomogeneous()) {
- final PyType keyAndValueType = tupleType.getIteratedItemType();
-
- keyTypes.add(keyAndValueType);
- valueTypes.add(keyAndValueType);
- }
- else if (tupleElementTypes.size() == 2) {
- keyTypes.add(tupleElementTypes.get(0));
- valueTypes.add(tupleElementTypes.get(1));
- }
- else {
- keyTypes.add(null);
- valueTypes.add(null);
- }
- }
- else {
- keyTypes.add(null);
- valueTypes.add(null);
- }
- }
- );
-
- if (elements.length > maxAnalyzedElements) {
- keyTypes.add(null);
- valueTypes.add(null);
- }
-
- return Arrays.asList(PyUnionType.union(keyTypes), PyUnionType.union(valueTypes));
- }
-
@Nullable
public PyFile getBuiltinsFile() {
return myBuiltinsFile;
diff --git a/python/src/com/jetbrains/python/psi/types/PyCollectionTypeByModificationsProvider.java b/python/src/com/jetbrains/python/psi/types/PyCollectionTypeByModificationsProvider.java
new file mode 100644
index 000000000000..145b1054f1cb
--- /dev/null
+++ b/python/src/com/jetbrains/python/psi/types/PyCollectionTypeByModificationsProvider.java
@@ -0,0 +1,98 @@
+/*
+ * Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
+ */
+package com.jetbrains.python.psi.types;
+
+import com.intellij.openapi.util.Ref;
+import com.jetbrains.python.codeInsight.controlflow.ScopeOwner;
+import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
+import com.jetbrains.python.psi.*;
+import com.jetbrains.python.psi.impl.PyOverridingTypeProvider;
+import org.jetbrains.annotations.NotNull;
+import org.jetbrains.annotations.Nullable;
+
+import java.util.ArrayList;
+import java.util.Collections;
+import java.util.List;
+
+final public class PyCollectionTypeByModificationsProvider extends PyTypeProviderBase implements PyOverridingTypeProvider {
+
+ @Nullable
+ @Override
+ public Ref getCallType(@NotNull PyFunction function, @Nullable PyCallSiteExpression callSite, @NotNull TypeEvalContext context) {
+ String qualifiedName = function.getQualifiedName();
+ if (qualifiedName != null && PyCollectionTypeUtil.INSTANCE.getCOLLECTION_CONSTRUCTORS().contains(qualifiedName)) {
+ if (callSite == null) {
+ return null;
+ }
+
+ PyExpression target = PyCollectionTypeUtil.INSTANCE.getTargetForValueInAssignment(callSite);
+ if (target instanceof PyTargetExpression) {
+ List arguments = callSite.getArguments(null);
+ List argumentTypes = getTypesFromConstructorArguments(context, arguments);
+
+ PyTargetExpression element = (PyTargetExpression)target;
+ ScopeOwner owner = ScopeUtil.getScopeOwner(element);
+ if (owner != null) {
+ final List typesByModifications = PyCollectionTypeUtil.INSTANCE
+ .getCollectionTypeByModifications(qualifiedName, element, context);
+ if (!typesByModifications.isEmpty()) {
+ if (qualifiedName.equals(PyCollectionTypeUtil.INSTANCE.getDICT_CONSTRUCTOR())) {
+ argumentTypes = extractTypesForDict(argumentTypes, typesByModifications);
+ }
+ else {
+ argumentTypes.addAll(typesByModifications);
+ argumentTypes = Collections.singletonList(PyUnionType.union(argumentTypes));
+ }
+
+ final PyClass cls = function.getContainingClass();
+ if (cls != null) {
+ return Ref.create(new PyCollectionTypeImpl(cls, false, argumentTypes));
+ }
+ }
+ }
+ }
+ }
+ return null;
+ }
+
+ @NotNull
+ private static List getTypesFromConstructorArguments(@NotNull TypeEvalContext context,
+ @NotNull List arguments) {
+ List argumentTypes = new ArrayList<>();
+ if (arguments.size() == 1 && arguments.get(0) != null) {
+ PyType type = context.getType(arguments.get(0));
+ if (type instanceof PyCollectionType) {
+ List elementTypes = ((PyCollectionType)type).getElementTypes();
+ argumentTypes.addAll(elementTypes);
+ }
+ else {
+ argumentTypes.add(type);
+ }
+ }
+ return argumentTypes;
+ }
+
+ @NotNull
+ private static List extractTypesForDict(@NotNull List argumentTypes, @NotNull List typesByModifications) {
+ if (argumentTypes.size() == 1) {
+ if (argumentTypes.get(0) instanceof PyTupleType) {
+ PyTupleType tuple = (PyTupleType)argumentTypes.get(0);
+ argumentTypes = tuple.getElementTypes();
+ }
+ else if (argumentTypes.get(0) == null) {
+ argumentTypes.add(null);
+ }
+ }
+ if (typesByModifications.size() == 2) {
+ if (argumentTypes.size() == 2) {
+ argumentTypes.set(0, PyUnionType.union(argumentTypes.get(0), typesByModifications.get(0)));
+ argumentTypes.set(1, PyUnionType.union(argumentTypes.get(1), typesByModifications.get(1)));
+ }
+ else {
+ argumentTypes = typesByModifications;
+ }
+ }
+ return argumentTypes;
+ }
+}
diff --git a/python/src/com/jetbrains/python/psi/types/PyCollectionTypeUtil.kt b/python/src/com/jetbrains/python/psi/types/PyCollectionTypeUtil.kt
new file mode 100644
index 000000000000..c4eedcaf1199
--- /dev/null
+++ b/python/src/com/jetbrains/python/psi/types/PyCollectionTypeUtil.kt
@@ -0,0 +1,468 @@
+/*
+ * Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
+ */
+package com.jetbrains.python.psi.types
+
+import com.intellij.openapi.util.Pair
+import com.intellij.psi.PsiElement
+import com.intellij.psi.util.PsiTreeUtil
+import com.intellij.util.ArrayUtil
+import com.jetbrains.python.codeInsight.controlflow.ScopeOwner
+import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil
+import com.jetbrains.python.psi.*
+import com.jetbrains.python.psi.PyUtil.`as`
+import com.jetbrains.python.psi.impl.PyBuiltinCache
+import com.jetbrains.python.psi.resolve.PyResolveContext
+import java.util.*
+
+object PyCollectionTypeUtil {
+
+ val DICT_CONSTRUCTOR = "dict.__init__"
+ private val LIST_CONSTRUCTOR = "list.__init__"
+ private val SET_CONSTRUCTOR = "set.__init__"
+
+ val COLLECTION_CONSTRUCTORS: Set<*> = HashSet(Arrays.asList(LIST_CONSTRUCTOR, DICT_CONSTRUCTOR, SET_CONSTRUCTOR))
+
+ private val MAX_ANALYZED_ELEMENTS_OF_LITERALS = 10 /* performance */
+
+ fun getTypeByModifications(sequence: PySequenceExpression, context: TypeEvalContext): List {
+ return if (sequence is PyListLiteralExpression || sequence is PySetLiteralExpression) {
+ listOf(getListOrSetIteratedValueType(sequence, context, true))
+ }
+ else if (sequence is PyDictLiteralExpression) {
+ getDictElementTypesWithModifications(sequence, context)
+ }
+ else {
+ listOf(null)
+ }
+ }
+
+ fun getListOrSetIteratedValueType(sequence: PySequenceExpression, context: TypeEvalContext,
+ withModifications: Boolean): PyType? {
+ val elements = sequence.elements
+ val maxAnalyzedElements = Math.min(MAX_ANALYZED_ELEMENTS_OF_LITERALS, elements.size)
+ var analyzedElementsType = PyUnionType.union(elements
+ .take(maxAnalyzedElements)
+ .map { context.getType(it) })
+ if (withModifications) {
+ val typesByModifications = getCollectionTypeByModifications(sequence, context)
+ if (!typesByModifications.isEmpty()) {
+ val typeByModifications = PyUnionType.union(typesByModifications)
+ analyzedElementsType = if (analyzedElementsType == null) typeByModifications
+ else PyUnionType.union(analyzedElementsType, typeByModifications)
+ }
+ }
+
+ return if (elements.size > maxAnalyzedElements) {
+ PyUnionType.createWeakType(analyzedElementsType)
+ }
+ else {
+ analyzedElementsType
+ }
+ }
+
+ private fun getDictElementTypes(sequence: PySequenceExpression, context: TypeEvalContext): List {
+ val elements = sequence.elements
+ val maxAnalyzedElements = Math.min(MAX_ANALYZED_ELEMENTS_OF_LITERALS, elements.size)
+ val keyTypes = ArrayList()
+ val valueTypes = ArrayList()
+
+ elements
+ .take(maxAnalyzedElements)
+ .map { element -> `as`(context.getType(element), PyTupleType::class.java) }
+ .forEach { tupleType ->
+ if (tupleType != null) {
+ val tupleElementTypes = tupleType.elementTypes
+
+ when {
+ tupleType.isHomogeneous -> {
+ val keyAndValueType = tupleType.iteratedItemType
+ keyTypes.add(keyAndValueType)
+ valueTypes.add(keyAndValueType)
+ }
+ tupleElementTypes.size == 2 -> {
+ keyTypes.add(tupleElementTypes[0])
+ valueTypes.add(tupleElementTypes[1])
+ }
+ else -> {
+ keyTypes.add(null)
+ valueTypes.add(null)
+ }
+ }
+ }
+ else {
+ keyTypes.add(null)
+ valueTypes.add(null)
+ }
+ }
+
+ if (elements.size > maxAnalyzedElements) {
+ keyTypes.add(null)
+ valueTypes.add(null)
+ }
+
+ return Arrays.asList(PyUnionType.union(keyTypes), PyUnionType.union(valueTypes))
+ }
+
+ private fun getDictElementTypesWithModifications(sequence: PySequenceExpression,
+ context: TypeEvalContext): List {
+ val dictTypes = getDictElementTypes(sequence, context)
+ var keyType: PyType? = null
+ var valueType: PyType? = null
+ if (dictTypes.size == 2) {
+ keyType = dictTypes[0]
+ valueType = dictTypes[1]
+ }
+
+ val elements = sequence.elements
+ val typesByModifications = getCollectionTypeByModifications(sequence, context)
+ if (typesByModifications.size == 2) {
+ val keysByModifications = typesByModifications[0]
+ keyType = if (elements.isNotEmpty()) {
+ PyUnionType.union(keyType, keysByModifications)
+ }
+ else {
+ keysByModifications
+ }
+ val valuesByModifications = typesByModifications[1]
+ valueType = if (elements.isNotEmpty()) {
+ PyUnionType.union(valueType, valuesByModifications)
+ }
+ else {
+ valuesByModifications
+ }
+ }
+
+ return Arrays.asList(keyType, valueType)
+ }
+
+ private fun getCollectionTypeByModifications(sequence: PySequenceExpression, context: TypeEvalContext): List {
+ val target = getTargetForValueInAssignment(sequence)
+ if (target != null) {
+ val owner = ScopeUtil.getScopeOwner(target)
+ if (owner != null) {
+ val visitor = getVisitorForSequence(sequence, target, context)
+ if (visitor != null) {
+ owner.accept(visitor)
+ return visitor.result
+ }
+ }
+ }
+ return emptyList()
+ }
+
+ fun getCollectionTypeByModifications(qualifiedName: String, element: PsiElement,
+ context: TypeEvalContext): List {
+ val owner = ScopeUtil.getScopeOwner(element)
+ if (owner != null) {
+ val typeVisitor = getVisitorForQualifiedName(qualifiedName, element, context)
+ if (typeVisitor != null) {
+ owner.accept(typeVisitor)
+ return typeVisitor.result
+ }
+ }
+ return emptyList()
+ }
+
+ private fun getVisitorForSequence(sequence: PySequenceExpression, element: PsiElement,
+ context: TypeEvalContext): PyCollectionTypeVisitor? {
+ return when (sequence) {
+ is PyListLiteralExpression -> PyListTypeVisitor(element, context)
+ is PyDictLiteralExpression -> PyDictTypeVisitor(element, context)
+ is PySetLiteralExpression -> PySetTypeVisitor(element, context)
+ else -> null
+ }
+ }
+
+ private fun getVisitorForQualifiedName(qualifiedName: String, element: PsiElement,
+ context: TypeEvalContext): PyCollectionTypeVisitor? {
+ when (qualifiedName) {
+ LIST_CONSTRUCTOR -> return PyListTypeVisitor(element, context)
+ DICT_CONSTRUCTOR -> return PyDictTypeVisitor(element, context)
+ SET_CONSTRUCTOR -> return PySetTypeVisitor(element, context)
+ }
+ return null
+ }
+
+ fun getTargetForValueInAssignment(value: PyExpression): PyExpression? {
+ val assignmentStatement = PsiTreeUtil.getParentOfType(value, PyAssignmentStatement::class.java, true, ScopeOwner::class.java)
+ assignmentStatement?.targetsToValuesMapping?.filter { it.second === value }?.forEach { return it.first }
+ return null
+ }
+
+ private fun getTypeForArgument(arguments: Array, argumentIndex: Int, typeEvalContext: TypeEvalContext): PyType? {
+ return if (argumentIndex < arguments.size)
+ typeEvalContext.getType(arguments[argumentIndex])
+ else
+ null
+ }
+
+ private fun getTypeByModifications(node: PyCallExpression,
+ modificationMethods: Map) -> List>,
+ element: PsiElement,
+ typeEvalContext: TypeEvalContext): MutableList? {
+ val valueTypes = mutableListOf()
+ var isModificationExist = false
+ val qualifiedExpression = node.callee as? PyQualifiedExpression ?: return null;
+ val funcName = qualifiedExpression.referencedName
+ if (modificationMethods.containsKey(funcName)) {
+ val referenceOwner = qualifiedExpression.qualifier as? PyReferenceOwner ?: return null
+ val resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(typeEvalContext)
+ if (referenceOwner.getReference(resolveContext).isReferenceTo(element)) {
+ isModificationExist = true
+ val function = modificationMethods[funcName]
+ if (function != null) {
+ valueTypes.addAll(function(node.arguments))
+ }
+ }
+ }
+ return if (isModificationExist) valueTypes else null
+ }
+
+ private fun getTypeByModifications(node: PySubscriptionExpression,
+ element: PsiElement,
+ typeEvalContext: TypeEvalContext): Pair, List>? {
+ var parent = node.parent
+ val keyTypes = ArrayList()
+ val valueTypes = ArrayList()
+ var isModificationExist = false
+
+ var tupleParent: PyTupleExpression? = null
+ if (parent is PyTupleExpression) {
+ tupleParent = parent
+ parent = tupleParent.parent
+ }
+
+ if (parent is PyAssignmentStatement) {
+ val assignment = parent
+ val leftExpression = assignment.leftHandSideExpression
+
+ if (tupleParent == null) {
+ if (leftExpression !== node) return null
+ }
+ else {
+ if (leftExpression !== tupleParent || !ArrayUtil.contains(node, *tupleParent.elements)) {
+ return null
+ }
+ }
+
+ val resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(typeEvalContext)
+ val referenceOwner = node.operand as? PyReferenceOwner ?: return null;
+ val reference = referenceOwner.getReference(resolveContext)
+ isModificationExist = if (reference.isReferenceTo(element)) true else return null
+
+ val indexExpression = node.indexExpression
+ if (indexExpression != null) {
+ keyTypes.add(typeEvalContext.getType(indexExpression))
+ }
+
+ var rightValue = assignment.assignedValue
+ if (tupleParent != null && rightValue is PyTupleExpression) {
+ val rightTuple = rightValue as PyTupleExpression?
+ val rightElements = rightTuple!!.elements
+ val indexInAssignment = Arrays.asList(*tupleParent.elements).indexOf(node)
+ if (indexInAssignment < rightElements.size) {
+ rightValue = rightElements[indexInAssignment]
+ }
+ }
+
+ if (rightValue != null) {
+ valueTypes.add(typeEvalContext.getType(rightValue))
+ }
+ }
+
+ return if (isModificationExist) Pair(keyTypes, valueTypes) else null
+ }
+
+ private abstract class PyCollectionTypeVisitor(protected val myElement: PsiElement,
+ protected val myTypeEvalContext: TypeEvalContext) : PyRecursiveElementVisitor() {
+ protected val scopeOwner: ScopeOwner? = ScopeUtil.getScopeOwner(myElement)
+ protected open var isModificationExist = false
+
+ abstract val result: List
+
+ abstract fun initMethods(): Map) -> List>
+
+ override fun visitPyFunction(node: PyFunction) {
+ if (node === scopeOwner) {
+ super.visitPyFunction(node)
+ }
+ // ignore nested functions
+ }
+
+ override fun visitPyClass(node: PyClass) {
+ if (node === scopeOwner) {
+ super.visitPyClass(node)
+ }
+ // ignore nested classes
+ }
+ }
+
+ private class PyListTypeVisitor(element: PsiElement,
+ typeEvalContext: TypeEvalContext) : PyCollectionTypeVisitor(element, typeEvalContext) {
+ private val modificationMethods: Map) -> List>
+ private val valueTypes: MutableList
+ override var isModificationExist = false
+
+ override val result: List
+ get() = if (isModificationExist) valueTypes else emptyList()
+
+ init {
+ modificationMethods = initMethods()
+ valueTypes = mutableListOf()
+ }
+
+ override fun initMethods(): Map) -> List> {
+ val modificationMethods = HashMap) -> List>()
+
+ modificationMethods.put("append", { arguments: Array -> listOf(getTypeForArgument(arguments, 0, myTypeEvalContext)) })
+ modificationMethods.put("index", { arguments: Array -> listOf(getTypeForArgument(arguments, 0, myTypeEvalContext)) })
+ modificationMethods.put("insert", { arguments: Array -> listOf(getTypeForArgument(arguments, 1, myTypeEvalContext)) })
+ modificationMethods.put("extend", { arguments: Array ->
+ val argType = getTypeForArgument(arguments, 0, myTypeEvalContext)
+ if (argType is PyCollectionType) {
+ argType.elementTypes
+ }
+ else {
+ emptyList()
+ }
+ })
+
+ return modificationMethods
+ }
+
+ override fun visitPyCallExpression(node: PyCallExpression) {
+ val types = getTypeByModifications(node, modificationMethods, myElement, myTypeEvalContext)
+ if (types != null) {
+ isModificationExist = true
+ valueTypes.addAll(types)
+ }
+ }
+
+ override fun visitPySubscriptionExpression(node: PySubscriptionExpression) {
+ val types = getTypeByModifications(node, myElement, myTypeEvalContext)
+ if (types != null) {
+ isModificationExist = true
+ valueTypes.addAll(types.second)
+ }
+ }
+ }
+
+ private class PyDictTypeVisitor(element: PsiElement, typeEvalContext: TypeEvalContext) : PyCollectionTypeVisitor(element,
+ typeEvalContext) {
+ private val modificationMethods: Map) -> List>
+ private val keyTypes: MutableList
+ private val valueTypes: MutableList
+
+ override val result: List
+ get() = if (isModificationExist) Arrays.asList(PyUnionType.union(keyTypes), PyUnionType.union(valueTypes))
+ else emptyList()
+
+ init {
+ modificationMethods = initMethods()
+ keyTypes = mutableListOf()
+ valueTypes = mutableListOf()
+ }
+
+ override fun initMethods(): Map) -> List> {
+ val modificationMethods = HashMap) -> List>()
+
+ modificationMethods.put("update", { arguments ->
+ if (arguments.size == 1 && arguments[0] is PyDictLiteralExpression) {
+ val dict = arguments[0] as PyDictLiteralExpression
+ val dictTypes = getDictElementTypes(dict, myTypeEvalContext)
+ if (dictTypes.size == 2) {
+ keyTypes.add(dictTypes[0])
+ valueTypes.add(dictTypes[1])
+ }
+ }
+ else if (arguments.isNotEmpty()) {
+ var keyStrAdded = false
+ for (arg in arguments) {
+ if (arg is PyKeywordArgument) {
+ if (!keyStrAdded) {
+ val strType = PyBuiltinCache.getInstance(myElement).strType
+ if (strType != null) {
+ keyTypes.add(strType)
+ }
+ keyStrAdded = true
+ }
+ val value = PyUtil.peelArgument(arg)
+ if (value != null) {
+ valueTypes.add(myTypeEvalContext.getType(value))
+ }
+ }
+ }
+ }
+ emptyList()
+ })
+
+ return modificationMethods
+ }
+
+ override fun visitPyCallExpression(node: PyCallExpression) {
+ val types = getTypeByModifications(node, modificationMethods, myElement, myTypeEvalContext)
+ if (types != null) {
+ isModificationExist = true
+ valueTypes.addAll(types)
+ }
+ }
+
+ override fun visitPySubscriptionExpression(node: PySubscriptionExpression) {
+ val types = getTypeByModifications(node, myElement, myTypeEvalContext)
+ if (types != null) {
+ isModificationExist = true
+ keyTypes.addAll(types.first)
+ valueTypes.addAll(types.second)
+ }
+ }
+ }
+
+ private class PySetTypeVisitor(element: PsiElement, typeEvalContext: TypeEvalContext) : PyCollectionTypeVisitor(element,
+ typeEvalContext) {
+ private val modificationMethods: Map) -> List>
+ private val valueTypes: MutableList
+
+ override val result: List
+ get() = if (isModificationExist) valueTypes else emptyList()
+
+ init {
+ modificationMethods = initMethods()
+ valueTypes = ArrayList()
+ }
+
+ override fun initMethods(): Map) -> List> {
+ val modificationMethods = HashMap) -> List>()
+ modificationMethods.put("add", { arguments -> listOf(getTypeForArgument(arguments, 0, myTypeEvalContext)) })
+ modificationMethods.put("update", { arguments ->
+ val types = ArrayList()
+ for (argument in arguments) {
+ when (argument) {
+ is PySetLiteralExpression -> types.add(getListOrSetIteratedValueType(argument as PySequenceExpression, myTypeEvalContext, false))
+ is PyListLiteralExpression -> types.add(getListOrSetIteratedValueType(argument as PySequenceExpression, myTypeEvalContext, false))
+ is PyDictLiteralExpression -> types.add(getDictElementTypes(argument as PySequenceExpression, myTypeEvalContext)[0])
+ else -> {
+ val argType = myTypeEvalContext.getType(argument)
+ if (argType is PyCollectionType) {
+ types.addAll(argType.elementTypes)
+ }
+ else {
+ types.add(argType)
+ }
+ }
+ }
+ }
+ return@put types
+ })
+ return modificationMethods
+ }
+
+ override fun visitPyCallExpression(node: PyCallExpression) {
+ val types = getTypeByModifications(node, modificationMethods, myElement, myTypeEvalContext)
+ if (types != null) {
+ isModificationExist = true
+ valueTypes.addAll(types)
+ }
+ }
+ }
+}
diff --git a/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py b/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py
index 4f367df931bb..622534f20cdd 100644
--- a/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py
+++ b/python/testData/inspections/PyTypeCheckerInspection/SetMethods.py
@@ -1,9 +1,12 @@
-xs = set([1, 2, 3])
-'foo' + xs.pop()
-xs.discard('foo')
-xs.remove('bar')
-xs.add(object())
+def foo(xs, ys):
+ """
+ :type xs: set of int
+ :type ys: set of string
+ """
+ 'foo' + xs.pop()
+ xs.discard('foo')
+ xs.remove('bar')
+ xs.add(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 61417f44fb3e..795b9d20066c 100644
--- a/python/testSrc/com/jetbrains/python/PyTypeTest.java
+++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java
@@ -1229,7 +1229,7 @@ public class PyTypeTest extends PyTestCase {
}
// PY-1182
- public void testCollectionType() {
+ public void testListTypeByModifications() {
doTest("List[int]",
"def f():\n" +
" expr = []\n" +
@@ -1263,6 +1263,331 @@ public class PyTypeTest extends PyTestCase {
doTest("List[int]",
"expr = []\n" +
"expr.index(42)");
+
+ doTest("List[int]",
+ "expr = [1, 2, 3]\n");
+
+ doTest("List[Union[int, Any]]",
+ "expr = [1, 2, 3]\n" +
+ "expr.append(var)\n");
+
+ doTest("List[Union[int, str]]",
+ "expr = [1, 2, 3]\n" +
+ "expr[0] = 'a'\n" +
+ "expr[1] = 'b'\n");
+
+ doTest("List[Union[int, str]]",
+ "expr = [1, 2, 3]\n" +
+ "expr[0] = 'a'\n" +
+ "expr[1] = 'b'\n");
+
+ doTest("List[Union[int, str]]",
+ "expr = [1, 2]\n" +
+ "t, expr[1] = 23, 'b'\n");
+
+ doTest("List[Union[int, str]]",
+ "def f():\n" +
+ " expr, b = [1, 2, 3], 'abc'\n" +
+ " expr.append('a')\n"
+ );
+
+ doTest("List[int]",
+ "def f():" +
+ " expr = [1, 2, 3]\n" +
+ " def inner():\n" +
+ " expr.append('a')\n"
+ );
+ }
+
+ // PY-1182
+ public void testListTypeByModificationsConstructor() {
+ doTest("List[str]",
+ "expr = list()\n" +
+ "expr.append('a')\n"
+ );
+
+ doTest("List[Union[str, int]]",
+ "expr = list()\n" +
+ "expr.append('a')\n" +
+ "expr.append(1)\n"
+ );
+
+ doTest("List[Union[int, str]]",
+ "a = list([1, 2, 3])\n" +
+ "a.append('a')\n" +
+ "expr = a\n"
+ );
+
+ doTest("List[Union[str, Any]]",
+ "expr = list()\n" +
+ "expr.append('a')\n" +
+ "expr.append(var)\n"
+ );
+
+ doTest("List[Union[int, str]]",
+ "expr = list([1, 2])\n" +
+ "t, expr[1] = 23, 'b'\n");
+
+ doTest("List[Union[str, Any]]",
+ "expr = list(var)\n" +
+ "expr[0] = 'abc'\n");
+
+ doTest("List[Union[int, str]]",
+ "b, expr = 1, list([1, 2, 3])\n" +
+ "expr.append('a')\n"
+ );
+
+ doTest("List[int]",
+ "def f():" +
+ " expr = list([1, 2, 3])\n" +
+ " def inner():\n" +
+ " expr.append('a')\n"
+ );
+ }
+
+ // PY-1182
+ public void testDictTypeByModifications() {
+ doTest("Dict[str, Union[int, str]]",
+ "def f():\n" +
+ " expr = {'a': 3}\n" +
+ " expr['b'] = \"s\""
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "def f():\n" +
+ " expr = {'a': 3}\n" +
+ " expr['b'] = \"s\""
+ );
+
+ doTest("Dict[str, Union[int, List[int]]]",
+ "def f():\n" +
+ " expr = {}\n" +
+ " expr['a'] = 0\n" +
+ " expr['c'] = [1, 2]"
+ );
+
+ doTest("Dict[str, Union[int, Any]]",
+ "def f():\n" +
+ " expr = {'b': D()}\n" +
+ " expr['a'] = 2\n"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "def f():\n" +
+ " expr = {'a': 3}\n" +
+ " expr['b'], t = \"s\", 12"
+ );
+
+ doTest("Dict[str, Union[int, Any]]",
+ "def f():\n" +
+ " expr = {'a': 3}\n" +
+ " expr['a'] = var\n"
+ );
+
+ doTest("Dict[str, int]",
+ "def f():\n" +
+ " expr = {'a': 3, 'b': 4}\n"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "def f():\n" +
+ " expr = {'a': 3}\n" +
+ " expr.update({'a': 'str'})\n"
+ );
+
+ doTest("Dict[str, Union[int, Any]]",
+ "def f():\n" +
+ " expr = {'a': 3}\n" +
+ " expr.update({'b': var})\n"
+ );
+
+ doTest("Dict[str, int]",
+ "def f():\n" +
+ " expr = {}\n" +
+ " expr.update(a=1, b=2)"
+ );
+
+ doTest("Dict[Union[int, str], Union[int, str]]",
+ "def f():\n" +
+ " expr = {1: '3'}\n" +
+ " expr.update(a=1, b=2)"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "def f():\n" +
+ " expr = {}\n" +
+ " expr['a'] = 23\n" +
+ " expr.update(a='m', b='n')"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "def f():\n" +
+ " b, expr = 23, {'a': 3}\n" +
+ " expr['b'] = 'l'"
+ );
+
+ doTest("Dict[str, int]",
+ "def f():" +
+ " expr = {'a': 1}\n" +
+ " def inner():\n" +
+ " expr['b'] = 'a'\n"
+ );
+ }
+
+ // PY-1182
+ public void testDictTypeByModificationConstructor() {
+ doTest("Dict[str, int]",
+ "expr = dict()\n" +
+ "expr['d'] = 12\n"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "expr = dict({'a': 1, 'b': 2})\n" +
+ "expr['a'] = '12'\n"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "expr = dict(zip(['a', 'b', 'c'], [1, 2, 3]))\n" +
+ "expr['d'] = '12'\n"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "expr = dict(zip(['a', 'b', 'c'], [1, 2, 3]))\n" +
+ "expr['d'] = '12'\n"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "expr = dict([('two', 2), ('one', 1), ('three', 3)])\n" +
+ "expr['d'] = '12'\n"
+ );
+
+ doTest("Dict[str, Union[int, Any]]",
+ "expr = dict({'a': 1, 'b': 2})\n" +
+ "expr['a'] = var\n"
+ );
+
+ doTest("Dict[Union[str, Any], Union[int, Any]]",
+ "expr = dict(var)\n" +
+ "expr.update({'c': 12})\n"
+ );
+
+ doTest("Dict[str, Union[int, str]]",
+ "a, expr = 23, dict({'a': 1})\n" +
+ "expr.update({'c': '34'})\n"
+ );
+
+ doTest("Dict[str, int]",
+ "def f():" +
+ " expr = dict({'a': 1})\n" +
+ " def inner():\n" +
+ " expr['b'] = 'a'\n"
+ );
+ }
+
+ // PY-1182
+ public void testSetTypeByModifications() {
+ doTest("Set[Union[str, int]]",
+ "def f():\n" +
+ " expr = {'abc'}\n" +
+ " expr.add(1)"
+ );
+
+ doTest("Set[Union[int, str]]",
+ "def f():\n" +
+ " expr = {1, 2}\n" +
+ " b = {'abc'}\n" +
+ " expr.update(b)"
+ );
+
+ doTest("Set[Union[int, str]]",
+ "def f():\n" +
+ " expr = {1, 2}\n" +
+ " b = {2, 3}\n" +
+ " expr.update(b, ['a', 'b'], {1, 2})"
+ );
+
+ doTest("Set[str]",
+ "def f():\n" +
+ " expr = {'m', 'n'}\n" +
+ " expr.update({'a': 1, 'b': 2})"
+ );
+
+ doTest("Set[Union[Union[int, str], Any]]",
+ "def f():\n" +
+ " expr = {1, 2}\n" +
+ " b = {'a', 'b'}\n" +
+ " expr.update(b, var)"
+ );
+
+ doTest("Set[str]",
+ "def f():\n" +
+ " expr, var = {'a', 'b'}, 'lala'\n" +
+ " expr.add('b')"
+ );
+
+ doTest("Set[int]",
+ "def f():" +
+ " expr = {1, 2, 3}\n" +
+ " def inner():\n" +
+ " expr.add('a')\n"
+ );
+ }
+
+ // PY-1182
+ public void testSetTypeByModificationsConstructor() {
+ doTest("Set[int]",
+ "def f():\n" +
+ " expr = set()\n" +
+ " expr.add(1)"
+ );
+
+ doTest("Set[Union[int, str]]",
+ "def f():\n" +
+ " expr = set({1, 2})\n" +
+ " expr.add('abc')"
+ );
+
+ doTest("Set[Union[int, str]]",
+ "def f():\n" +
+ " expr = set({1, 2})\n" +
+ " b = {'abc'}\n" +
+ " expr.update(b)"
+ );
+
+ doTest("Set[Union[int, str]]",
+ "def f():\n" +
+ " expr = set({1, 2})\n" +
+ " b = {2, 3}\n" +
+ " expr.update(b, ['a', 'b'], {1, 2})"
+ );
+
+ doTest("Set[Union[str, Any]]",
+ "def f():\n" +
+ " expr = set()\n" +
+ " b = {'a', 'b'}\n" +
+ " expr.update(b, var)"
+ );
+
+ doTest("Set[Union[str, Any]]",
+ "def f():\n" +
+ " expr = set(var)\n" +
+ " b = {'a', 'b'}\n" +
+ " expr.update(b)"
+ );
+
+ doTest("Set[Union[int, str]]",
+ "def f():\n" +
+ " expr, var = set([1, 2, 3]), 'lala'\n" +
+ " expr.add('b')"
+ );
+
+ doTest("Set[int]",
+ "def f():\n" +
+ " expr = set()\n" +
+ " expr.add(1)\n" +
+ " def inner():\n" +
+ " expr.add('a')\n"
+ );
}
// PY-20063