Support in-place modifications for dict, list and set (PY-1182)

This commit is contained in:
Elizaveta Shashkova
2018-02-05 19:30:26 +03:00
parent cedb7bda67
commit 91044050b0
7 changed files with 907 additions and 182 deletions
@@ -708,6 +708,7 @@
<typeProvider implementation="com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider"/>
<typeProvider implementation="com.jetbrains.python.codeInsight.stdlib.PyStdlibOverridingTypeProvider"/>
<typeProvider implementation="com.jetbrains.python.codeInsight.stdlib.PyDataclassesTypeProvider"/>
<typeProvider implementation="com.jetbrains.python.psi.types.PyCollectionTypeByModificationsProvider"/>
<pyModuleMembersProvider implementation="com.jetbrains.python.codeInsight.stdlib.PyStdlibModuleMembersProvider"/>
<documentationLinkProvider implementation="com.jetbrains.python.codeInsight.stdlib.PyStdlibDocumentationLinkProvider"/>
<canonicalPathProvider implementation="com.jetbrains.python.codeInsight.stdlib.PyStdlibCanonicalPathProvider"/>
@@ -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<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;
}
}
}
@@ -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<PyType> 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<PyType> getDictElementTypes(@NotNull PyExpression[] elements, @NotNull TypeEvalContext context) {
final int maxAnalyzedElements = Math.min(MAX_ANALYZED_ELEMENTS_OF_LITERALS, elements.length);
final List<PyType> keyTypes = new ArrayList<>();
final List<PyType> valueTypes = new ArrayList<>();
StreamEx
.of(elements, 0, maxAnalyzedElements)
.map(element -> as(context.getType(element), PyTupleType.class))
.forEach(
tupleType -> {
if (tupleType != null) {
final List<PyType> 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;
@@ -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<PyType> 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<PyExpression> arguments = callSite.getArguments(null);
List<PyType> argumentTypes = getTypesFromConstructorArguments(context, arguments);
PyTargetExpression element = (PyTargetExpression)target;
ScopeOwner owner = ScopeUtil.getScopeOwner(element);
if (owner != null) {
final List<PyType> 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<PyType> getTypesFromConstructorArguments(@NotNull TypeEvalContext context,
@NotNull List<PyExpression> arguments) {
List<PyType> argumentTypes = new ArrayList<>();
if (arguments.size() == 1 && arguments.get(0) != null) {
PyType type = context.getType(arguments.get(0));
if (type instanceof PyCollectionType) {
List<PyType> elementTypes = ((PyCollectionType)type).getElementTypes();
argumentTypes.addAll(elementTypes);
}
else {
argumentTypes.add(type);
}
}
return argumentTypes;
}
@NotNull
private static List<PyType> extractTypesForDict(@NotNull List<PyType> argumentTypes, @NotNull List<PyType> 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;
}
}
@@ -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<PyType?> {
return if (sequence is PyListLiteralExpression || sequence is PySetLiteralExpression) {
listOf(getListOrSetIteratedValueType(sequence, context, true))
}
else if (sequence is PyDictLiteralExpression) {
getDictElementTypesWithModifications(sequence, context)
}
else {
listOf<PyType?>(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<PyType> {
val elements = sequence.elements
val maxAnalyzedElements = Math.min(MAX_ANALYZED_ELEMENTS_OF_LITERALS, elements.size)
val keyTypes = ArrayList<PyType?>()
val valueTypes = ArrayList<PyType?>()
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<PyType>(PyUnionType.union(keyTypes), PyUnionType.union(valueTypes))
}
private fun getDictElementTypesWithModifications(sequence: PySequenceExpression,
context: TypeEvalContext): List<PyType> {
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<PyType>(keyType, valueType)
}
private fun getCollectionTypeByModifications(sequence: PySequenceExpression, context: TypeEvalContext): List<PyType?> {
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<PyType?> {
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<PyExpression>, argumentIndex: Int, typeEvalContext: TypeEvalContext): PyType? {
return if (argumentIndex < arguments.size)
typeEvalContext.getType(arguments[argumentIndex])
else
null
}
private fun getTypeByModifications(node: PyCallExpression,
modificationMethods: Map<String, (Array<PyExpression>) -> List<PyType?>>,
element: PsiElement,
typeEvalContext: TypeEvalContext): MutableList<PyType?>? {
val valueTypes = mutableListOf<PyType?>()
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<PyType?>, List<PyType?>>? {
var parent = node.parent
val keyTypes = ArrayList<PyType?>()
val valueTypes = ArrayList<PyType?>()
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<PyType?>
abstract fun initMethods(): Map<String, (Array<PyExpression>) -> List<PyType?>>
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<String, (Array<PyExpression>) -> List<PyType?>>
private val valueTypes: MutableList<PyType?>
override var isModificationExist = false
override val result: List<PyType?>
get() = if (isModificationExist) valueTypes else emptyList()
init {
modificationMethods = initMethods()
valueTypes = mutableListOf()
}
override fun initMethods(): Map<String, (Array<PyExpression>) -> List<PyType?>> {
val modificationMethods = HashMap<String, (Array<PyExpression>) -> List<PyType?>>()
modificationMethods.put("append", { arguments: Array<PyExpression> -> listOf(getTypeForArgument(arguments, 0, myTypeEvalContext)) })
modificationMethods.put("index", { arguments: Array<PyExpression> -> listOf(getTypeForArgument(arguments, 0, myTypeEvalContext)) })
modificationMethods.put("insert", { arguments: Array<PyExpression> -> listOf(getTypeForArgument(arguments, 1, myTypeEvalContext)) })
modificationMethods.put("extend", { arguments: Array<PyExpression> ->
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<String, (Array<PyExpression>) -> List<PyType?>>
private val keyTypes: MutableList<PyType?>
private val valueTypes: MutableList<PyType?>
override val result: List<PyType>
get() = if (isModificationExist) Arrays.asList<PyType>(PyUnionType.union(keyTypes), PyUnionType.union(valueTypes))
else emptyList()
init {
modificationMethods = initMethods()
keyTypes = mutableListOf()
valueTypes = mutableListOf()
}
override fun initMethods(): Map<String, (Array<PyExpression>) -> List<PyType?>> {
val modificationMethods = HashMap<String, (Array<PyExpression>) -> List<PyType?>>()
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<String, (Array<PyExpression>) -> List<PyType?>>
private val valueTypes: MutableList<PyType?>
override val result: List<PyType?>
get() = if (isModificationExist) valueTypes else emptyList()
init {
modificationMethods = initMethods()
valueTypes = ArrayList()
}
override fun initMethods(): Map<String, (Array<PyExpression>) -> List<PyType?>> {
val modificationMethods = HashMap<String, (Array<PyExpression>) -> List<PyType?>>()
modificationMethods.put("add", { arguments -> listOf(getTypeForArgument(arguments, 0, myTypeEvalContext)) })
modificationMethods.put("update", { arguments ->
val types = ArrayList<PyType?>()
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)
}
}
}
}
@@ -1,9 +1,12 @@
xs = set([1, 2, 3])
'foo' + <warning descr="Expected type 'AnyStr', got 'int' instead">xs.pop()</warning>
xs.discard(<weak_warning descr="Expected type 'int' (matched generic type '_T'), got 'str' instead">'foo'</weak_warning>)
xs.remove(<weak_warning descr="Expected type 'int' (matched generic type '_T'), got 'str' instead">'bar'</weak_warning>)
xs.add(<weak_warning descr="Expected type 'int' (matched generic type '_T'), got 'object' instead">object()</weak_warning>)
def foo(xs, ys):
"""
:type xs: set of int
:type ys: set of string
"""
'foo' + <warning descr="Expected type 'AnyStr', got 'int' instead">xs.pop()</warning>
xs.discard(<weak_warning descr="Expected type 'int' (matched generic type '_T'), got 'str' instead">'foo'</weak_warning>)
xs.remove(<weak_warning descr="Expected type 'int' (matched generic type '_T'), got 'str' instead">'bar'</weak_warning>)
xs.add(<weak_warning descr="Expected type 'int' (matched generic type '_T'), got 'object' instead">object()</weak_warning>)
ys = ['green', 'eggs']
ys.extend(xs)
ys.extend(xs)
@@ -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