From 91044050b06ee31c8d13328be44f51a3f89dfbc4 Mon Sep 17 00:00:00 2001 From: Elizaveta Shashkova Date: Fri, 22 Dec 2017 12:50:28 +0300 Subject: [PATCH] Support in-place modifications for dict, list and set (PY-1182) --- python/src/META-INF/python-core-common.xml | 1 + .../src/com/jetbrains/python/psi/PyUtil.java | 84 ---- .../python/psi/impl/PyBuiltinCache.java | 94 +--- ...CollectionTypeByModificationsProvider.java | 98 ++++ .../python/psi/types/PyCollectionTypeUtil.kt | 468 ++++++++++++++++++ .../PyTypeCheckerInspection/SetMethods.py | 17 +- .../com/jetbrains/python/PyTypeTest.java | 327 +++++++++++- 7 files changed, 907 insertions(+), 182 deletions(-) create mode 100644 python/src/com/jetbrains/python/psi/types/PyCollectionTypeByModificationsProvider.java create mode 100644 python/src/com/jetbrains/python/psi/types/PyCollectionTypeUtil.kt 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