From 8abae98db4f6d107d5549051dd3e3166844e7268 Mon Sep 17 00:00:00 2001 From: Lada Gagina Date: Thu, 6 Jun 2019 13:03:06 +0300 Subject: [PATCH] IDEA-CR-52805: PY-36008 Add support of TypedDict TypedDict was introduced in https://www.python.org/dev/peps/pep-0589 GitOrigin-RevId: cca138735c7927302214accde298b2f5aae20b48 --- .../src/com/jetbrains/python/PyNames.java | 9 +- .../src/META-INF/python-psi-impl.xml | 2 + .../PyTypedDictCompletionContributor.kt | 31 ++ .../PyTypedDictOverridingTypeProvider.kt | 19 + .../typing/PyTypedDictTypeProvider.kt | 374 ++++++++++++++++++ .../typing/PyTypingTypeProvider.java | 18 +- .../documentation/PyTypeModelBuilder.java | 48 ++- .../python/psi/PyFileElementType.java | 2 +- .../src/com/jetbrains/python/psi/PyUtil.java | 2 +- .../python/psi/impl/PyClassImpl.java | 2 +- .../impl/PySubscriptionExpressionImpl.java | 10 +- .../psi/impl/stubs/PyTypedDictStubImpl.kt | 146 +++++++ .../psi/impl/stubs/PyTypedDictStubType.kt | 20 + .../python/psi/stubs/PyTypedDictStub.kt | 21 + .../python/psi/types/PyTypeChecker.java | 12 + .../python/psi/types/PyTypedDictType.kt | 214 ++++++++++ .../completion/PythonCompletionTest.java | 26 ++ .../PyTypedDictInspection.html | 5 + python/src/META-INF/python-core-common.xml | 3 + .../PyClassHasNoInitInspection.java | 2 + .../inspections/PyStringFormatInspection.java | 2 +- .../inspections/PyTypeCheckerInspection.java | 24 +- .../inspections/PyTypedDictInspection.kt | 328 +++++++++++++++ .../PyArgumentListInspection/typedDict.py | 44 +++ .../PyClassHasNoInitInspection/typedDict.py | 5 + .../TypedDictConsistency.py | 74 ++++ .../com/jetbrains/python/Py3TypeTest.java | 30 ++ .../com/jetbrains/python/PyTypeTest.java | 15 + .../PyArgumentListInspectionTest.java | 5 + .../PyClassHasNoInitInspectionTest.java | 5 + .../PyTypeCheckerInspectionTest.java | 90 +++++ .../PyTypedDictInspectionTest.java | 230 +++++++++++ .../PyUnresolvedReferencesInspectionTest.java | 20 + 33 files changed, 1819 insertions(+), 19 deletions(-) create mode 100644 python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyTypedDictCompletionContributor.kt create mode 100644 python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictOverridingTypeProvider.kt create mode 100644 python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt create mode 100644 python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubImpl.kt create mode 100644 python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubType.kt create mode 100644 python/python-psi-impl/src/com/jetbrains/python/psi/stubs/PyTypedDictStub.kt create mode 100644 python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt create mode 100644 python/resources/inspectionDescriptions/PyTypedDictInspection.html create mode 100644 python/src/com/jetbrains/python/inspections/PyTypedDictInspection.kt create mode 100644 python/testData/inspections/PyArgumentListInspection/typedDict.py create mode 100644 python/testData/inspections/PyClassHasNoInitInspection/typedDict.py create mode 100644 python/testData/inspections/PyTypeCheckerInspection/TypedDictConsistency.py create mode 100644 python/testSrc/com/jetbrains/python/inspections/PyTypedDictInspectionTest.java diff --git a/python/python-psi-api/src/com/jetbrains/python/PyNames.java b/python/python-psi-api/src/com/jetbrains/python/PyNames.java index acfdc953065b..eeb5b5cbe7cc 100644 --- a/python/python-psi-api/src/com/jetbrains/python/PyNames.java +++ b/python/python-psi-api/src/com/jetbrains/python/PyNames.java @@ -79,7 +79,7 @@ public class PyNames { } public static final String INIT = "__init__"; - public static final String DICT = "__dict__"; + public static final String DUNDER_DICT = "__dict__"; public static final String DOT_PY = ".py"; public static final String DOT_PYI = ".pyi"; public static final String INIT_DOT_PY = INIT + DOT_PY; @@ -209,6 +209,8 @@ public class PyNames { public static final String COLLECTIONS_NAMEDTUPLE_PY2 = COLLECTIONS + "." + NAMEDTUPLE; public static final String COLLECTIONS_NAMEDTUPLE_PY3 = COLLECTIONS + "." + INIT + "." + NAMEDTUPLE; + public static final String TYPED_DICT = "TypedDict"; + public static final String FORMAT = "format"; public static final String ABSTRACTMETHOD = "abstractmethod"; @@ -220,11 +222,16 @@ public class PyNames { public static final String TUPLE = "tuple"; public static final String SET = "set"; public static final String SLICE = "slice"; + public static final String DICT = "dict"; public static final String KEYS = "keys"; public static final String APPEND = "append"; public static final String EXTEND = "extend"; public static final String UPDATE = "update"; + public static final String CLEAR = "clear"; + public static final String POP = "pop"; + public static final String POPITEM = "popitem"; + public static final String SETDEFAULT = "setdefault"; public static final String PASS = "pass"; diff --git a/python/python-psi-impl/src/META-INF/python-psi-impl.xml b/python/python-psi-impl/src/META-INF/python-psi-impl.xml index 49df542bebdc..d55611616be9 100644 --- a/python/python-psi-impl/src/META-INF/python-psi-impl.xml +++ b/python/python-psi-impl/src/META-INF/python-psi-impl.xml @@ -107,6 +107,8 @@ + + diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyTypedDictCompletionContributor.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyTypedDictCompletionContributor.kt new file mode 100644 index 000000000000..f3c0cf8df356 --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyTypedDictCompletionContributor.kt @@ -0,0 +1,31 @@ +/* + * Copyright 2000-2017 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.codeInsight.completion + +import com.intellij.codeInsight.completion.* +import com.intellij.codeInsight.lookup.LookupElementBuilder +import com.intellij.patterns.PlatformPatterns +import com.intellij.psi.util.PsiTreeUtil +import com.intellij.util.ProcessingContext +import com.jetbrains.python.psi.PyKeywordArgument + +class PyTypedDictCompletionContributor : CompletionContributor() { + + override fun handleAutoCompletionPossibility(context: AutoCompletionContext): AutoCompletionDecision = autoInsertSingleItem(context) + + init { + extend(CompletionType.BASIC, PlatformPatterns.psiElement().inside(PyKeywordArgument::class.java), TotalityValueProvider) + } + + private object TotalityValueProvider : CompletionProvider() { + + override fun addCompletions(parameters: CompletionParameters, context: ProcessingContext, result: CompletionResultSet) { + val keywordArgument = PsiTreeUtil.getParentOfType(parameters.position, PyKeywordArgument::class.java) + if (keywordArgument != null && keywordArgument.keyword == "total") { + result.addElement(LookupElementBuilder.create("True").bold()) + result.addElement(LookupElementBuilder.create("False").bold()) + } + } + } +} diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictOverridingTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictOverridingTypeProvider.kt new file mode 100644 index 000000000000..3c45ee0bafe8 --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictOverridingTypeProvider.kt @@ -0,0 +1,19 @@ +// Copyright 2000-2019 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.codeInsight.typing + +import com.intellij.openapi.util.Ref +import com.intellij.psi.PsiElement +import com.jetbrains.python.psi.impl.PyOverridingTypeProvider +import com.jetbrains.python.psi.types.PyType +import com.jetbrains.python.psi.types.PyTypeProviderBase +import com.jetbrains.python.psi.types.PyTypeUtil +import com.jetbrains.python.psi.types.TypeEvalContext + +class PyTypedDictOverridingTypeProvider : PyTypeProviderBase(), PyOverridingTypeProvider { + + override fun getReferenceType(referenceTarget: PsiElement, context: TypeEvalContext, anchor: PsiElement?): Ref? { + val type = PyTypedDictTypeProvider.getTypedDictTypeForResolvedCallee(referenceTarget, context) + + return PyTypeUtil.notNullToRef(type) + } +} \ No newline at end of file diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt new file mode 100644 index 000000000000..aaf80e366eec --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt @@ -0,0 +1,374 @@ +// Copyright 2000-2019 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.codeInsight.typing + +import com.intellij.openapi.util.Ref +import com.intellij.psi.PsiElement +import com.intellij.psi.util.QualifiedName +import com.jetbrains.python.PyNames +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.* +import com.jetbrains.python.psi.* +import com.jetbrains.python.psi.impl.PyBuiltinCache +import com.jetbrains.python.psi.impl.PyCallExpressionNavigator +import com.jetbrains.python.psi.impl.stubs.PyClassElementType +import com.jetbrains.python.psi.impl.stubs.PyTypedDictStubImpl +import com.jetbrains.python.psi.resolve.PyResolveContext +import com.jetbrains.python.psi.stubs.PyTypedDictStub +import com.jetbrains.python.psi.types.* +import java.util.* +import java.util.stream.Collectors + +typealias TDFields = LinkedHashMap + +class PyTypedDictTypeProvider : PyTypeProviderBase() { + override fun getReferenceExpressionType(referenceExpression: PyReferenceExpression, context: TypeEvalContext): PyType? { + return getTypedDictTypeForCallee(referenceExpression, context) + } + + override fun getReferenceType(referenceTarget: PsiElement, context: TypeEvalContext, anchor: PsiElement?): Ref? { + return PyTypeUtil.notNullToRef(getTypedDictTypeForResolvedCallee(referenceTarget, context)) + } + + companion object { + val nameIsTypedDict = { name: String? -> name == TYPED_DICT || name == TYPED_DICT_EXT } + + fun isTypingTypedDictInheritor(cls: PyClass, context: TypeEvalContext): Boolean { + val isTypingTD = { type: PyClassLikeType? -> + type is PyTypedDictType || nameIsTypedDict(type?.classQName) + } + val ancestors = cls.getAncestorTypes(context) + + if (ancestors.any(isTypingTD)) return true + + val hasTDAsSuperclass = hasTypedDictAsSuperclass(cls, context) + val hasTDAncestors = ancestors.filterIsInstance() + .any { hasTypedDictAsSuperclass(it.pyClass, context) } + return hasTDAsSuperclass || hasTDAncestors + } + + private fun hasTypedDictAsSuperclass(cls: PyClass, context: TypeEvalContext): Boolean { + when { + context.maySwitchToAST(cls) -> return cls.superClassExpressions.any { superClassExpr -> + resolveToQualifiedNames(superClassExpr, context).any(nameIsTypedDict) + } + cls.stub != null -> { + return containsTypedDictQName(cls.stub.superClasses) + } + else -> return containsTypedDictQName(PyClassElementType.getSuperClassQNames(cls)) + } + } + + private fun containsTypedDictQName(map: Map): Boolean { + return map.any { name -> name == QualifiedName.fromDottedString(TYPED_DICT) || name == QualifiedName.fromDottedString(TYPED_DICT_EXT) } + } + + fun getTypedDictTypeForResolvedCallee(referenceTarget: PsiElement, context: TypeEvalContext): PyTypedDictType? { + return when (referenceTarget) { + is PyClass -> getTypedDictTypeForTypingTDInheritorAsCallee(referenceTarget, context) + is PyTargetExpression -> getTypedDictTypeForTarget(referenceTarget, context) + else -> null + } + } + + private fun getTypingTypedDictTypeForResolvedCallee(referenceTarget: PyClass, context: TypeEvalContext): PyTypedDictType? { + return getTypedDictTypeForTypingTDInheritorAsCallee(referenceTarget, context, true) + } + + private fun getTypedDictTypeForCallee(referenceExpression: PyReferenceExpression, context: TypeEvalContext): PyType? { + if (PyCallExpressionNavigator.getPyCallExpressionByCallee(referenceExpression) == null) return null + + val resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context) + val resolveResults = referenceExpression.getReference(resolveContext).multiResolve(false) + + for (element in PyUtil.filterTopPriorityResults(resolveResults)) { + if (element is PyTargetExpression) { + val result = getTypedDictTypeForTarget(element, context) + if (result != null) { + return result + } + } + + if (element is PyClass) { + val result = getTypedDictTypeForTypingTDInheritorAsCallee(element, context) + if (result != null) { + return result + } + } + + if (element is PyTypedElement) { + val type = context.getType(element) + if (type is PyClassType) { + if (isTypingTypedDictInheritor(type.pyClass, context)) { + return getTypedDictTypeForTypingTDInheritorAsCallee(type.pyClass, context) + } + } + } + + if (resolveToQualifiedNames(referenceExpression, context).contains(TYPED_DICT)) { + val parameters = mutableListOf() + + val builtinCache = PyBuiltinCache.getInstance(referenceExpression) + val languageLevel = LanguageLevel.forElement(referenceExpression) + val generator = PyElementGenerator.getInstance(referenceExpression.project) + + parameters.add(PyCallableParameterImpl.nonPsi("name", builtinCache.getStringType(languageLevel))) + parameters.add(PyCallableParameterImpl.nonPsi("fields", builtinCache.dictType)) + parameters.add( + PyCallableParameterImpl.nonPsi("total", builtinCache.boolType, generator.createExpressionFromText(languageLevel, "True"))) + + return PyCallableTypeImpl(parameters, null) + } + } + + return null + } + + private fun getTypedDictTypeForTypingTDInheritorAsCallee(cls: PyClass, context: TypeEvalContext): PyTypedDictType? { + return getTypedDictTypeForTypingTDInheritorAsCallee(cls, context, false) + } + + private fun getTypedDictTypeForTypingTDInheritorAsCallee(cls: PyClass, context: TypeEvalContext, isInstance: Boolean): PyTypedDictType? { + if (isTypingTypedDictInheritor(cls, context)) { + val ancestors = cls.getAncestorTypes(context).filterIsInstance() + val name = cls.name ?: return null + val fields = collectFields(cls, context) + val overallFields = mutableMapOf() + overallFields.putAll(fields.filter { it.value.isRequired }) + overallFields.putAll(fields.filter { !it.value.isRequired }) + + val dictClass = PyBuiltinCache.getInstance(cls).dictType?.pyClass + if (dictClass == null) return null + + return PyTypedDictType(name, + TDFields(overallFields), + false, + dictClass, + if (isInstance) PyTypedDictType.DefinitionLevel.INSTANCE else PyTypedDictType.DefinitionLevel.NEW_TYPE, + ancestors) + } + + return null + } + + private fun collectFields(cls: PyClass, context: TypeEvalContext): TDFields { + val fields = mutableMapOf() + fields.putAll(collectTypingTDInheritorFields(cls, context)) + val ancestors = cls.getAncestorTypes(context) + ancestors.forEach { if (it is PyTypedDictType) fields.putAll(it.fields) } + return TDFields(fields) + } + + private fun collectTypingTDInheritorFields(cls: PyClass, context: TypeEvalContext): TDFields { + val type = cls.getType(context) + if (type is PyTypedDictType) { + return TDFields(type.fields) + } + + val fields = mutableListOf() + + cls.processClassLevelDeclarations { element, _ -> + if (element is PyTargetExpression && element.annotationValue != null) { + fields.add(element) + } + + true + } + + val argumentList: PyArgumentList? = cls.children.filterIsInstance().firstOrNull() + val totalityValue = argumentList?.getKeywordArgument("total") + val fieldsRequired = if (totalityValue != null && totalityValue.valueExpression is PyBoolLiteralExpression) + (totalityValue.valueExpression as PyBoolLiteralExpression).value + else true + + val toTDFields = Collectors.toMap( + { it.name }, + { field -> PyTypedDictType.FieldTypeAndTotality(context.getType(field), fieldsRequired) }, + { _, v2 -> v2 }, + { TDFields() }) + + return fields.stream().collect(toTDFields) + } + + private fun getTypedDictTypeForTarget(target: PyTargetExpression, context: TypeEvalContext): PyTypedDictType? { + val stub = target.stub + + return if (stub != null) { + getTypedDictTypeFromStub(target, + stub.getCustomStub(PyTypedDictStub::class.java), + context) + } + else getTypedDictTypeFromAST(target, context) + } + + fun getTypedDictTypeForResolvedElement(resolved: PsiElement, context: TypeEvalContext): PyType? { + if (resolved is PyClass && isTypingTypedDictInheritor(resolved, context)) { + return getTypingTypedDictTypeForResolvedCallee(resolved, context) + } + if (resolved is PyCallExpression) { + val callee = resolved.callee + if (callee is PyReferenceExpression && isTypedDict(callee, context)) { + val type = getTypedDictTypeFromAST(resolved, context) + if (type != null) { + return type + } + } + } + + return null + } + + private fun getTypedDictTypeFromAST(expression: PyCallExpression, context: TypeEvalContext): PyTypedDictType? { + return if (context.maySwitchToAST(expression)) { + getTypedDictTypeFromStub(expression, PyTypedDictStubImpl.create(expression), context) + } + else null + } + + private fun getTypedDictTypeFromAST(expression: PyTargetExpression, context: TypeEvalContext): PyTypedDictType? { + return if (context.maySwitchToAST(expression)) { + getTypedDictTypeFromStub(expression, PyTypedDictStubImpl.create(expression), context) + } + else null + } + + private fun getTypedDictTypeFromStub(referenceTarget: PsiElement, + stub: PyTypedDictStub?, + context: TypeEvalContext): PyTypedDictType? { + if (stub == null) return null + + val dictClass = PyBuiltinCache.getInstance(referenceTarget).dictType?.pyClass + if (dictClass == null) return null + val fields = stub.fields + val total = stub.isTotal + val typedDictFields = parseTypedDictFields(referenceTarget, fields, context, total) + + return PyTypedDictType(stub.name, + typedDictFields, + false, + dictClass, + PyTypedDictType.DefinitionLevel.NEW_TYPE, + listOf(), + referenceTarget as? PyTargetExpression) + } + + private fun parseTypedDictFields(anchor: PsiElement, + fields: Map>, + context: TypeEvalContext, + total: Boolean): TDFields { + val result = TDFields() + for ((name, type) in fields) { + result[name] = parseTypedDictField(anchor, type.orElse(null), context, total) + } + return result + } + + private fun parseTypedDictField(anchor: PsiElement, + type: String?, + context: TypeEvalContext, + total: Boolean): PyTypedDictType.FieldTypeAndTotality { + if (type == null) return PyTypedDictType.FieldTypeAndTotality(null) + + val pyType = Ref.deref(getStringBasedType(type, anchor, context)) + return PyTypedDictType.FieldTypeAndTotality(pyType, total) + } + + /** + * If [expected] type is `typing.TypedDict[...]`, + * then tries to infer `typing.TypedDict[...]` for [expression], + * otherwise returns type inferred by [context]. + */ + fun promoteToTypedDict(expression: PyExpression, expected: PyType?, context: TypeEvalContext): PyType? { + if (expected is PyTypedDictType) { + return fromValue(expression, context) ?: context.getType(expression) + } + else { + return context.getType(expression) + } + } + + /** + * Tries to construct TypedDict type for a value that could be considered as TypedDict and downcasted to `typing.TypedDict[...]` type. + */ + private fun fromValue(expression: PyExpression, context: TypeEvalContext): PyType? = newInstance(expression, context) + + private fun newInstance(expression: PyExpression, context: TypeEvalContext): PyType? { + return when (expression) { + is PyTupleExpression -> { + val elements = expression.elements + val classes = elements.mapNotNull { toTypedDictType(it, context) } + if (elements.size == classes.size) PyUnionType.union(classes) else null + } + else -> toTypedDictType(expression, context) + } + } + + private fun toTypedDictType(expression: PyExpression, context: TypeEvalContext): PyType? { + if (expression is PyNoneLiteralExpression && + !expression.isEllipsis || + expression is PyReferenceExpression && + expression.name == PyNames.NONE && + LanguageLevel.forElement(expression).isPython2) return PyNoneType.INSTANCE + + if (expression is PyDictLiteralExpression) { + val fields = getTypingTDFieldsFromDictLiteral(expression, context) + if (fields != null) { + val dictClass = PyBuiltinCache.getInstance(expression).dictType?.pyClass + if (dictClass == null) return null + return PyTypedDictType("TypedDict", fields, true, dictClass, + PyTypedDictType.DefinitionLevel.INSTANCE, + listOf()) + } + } else if (expression is PyCallExpression) { + val resolvedQualifiedNames = if (expression.callee != null) resolveToQualifiedNames(expression.callee!!, context) else return null + if (resolvedQualifiedNames.any { it == PyNames.DICT }) { + val arguments = expression.arguments + if (arguments.size > 1) { + val fields = getTypingTDFieldsFromDictKeywordArguments(arguments, context) + if (fields != null) { + val dictClass = PyBuiltinCache.getInstance(expression).dictType?.pyClass + if (dictClass == null) return null + return PyTypedDictType("TypedDict", fields, true, dictClass, + PyTypedDictType.DefinitionLevel.INSTANCE, + listOf()) + } + } + } + } + return null + } + + private fun getTypingTDFieldsFromDictLiteral(dictLiteral: PyDictLiteralExpression, context: TypeEvalContext): TDFields? { + val fields = LinkedHashMap() + + dictLiteral.elements.forEach { + val name: PyExpression = it.key + val value: PyExpression? = it.value + + if (name !is PyStringLiteralExpression) return null + + fields[name.stringValue] = value + } + + return typedDictFieldsFromKeysAndValues(fields, context) + } + + private fun getTypingTDFieldsFromDictKeywordArguments(keywordArguments: Array, context: TypeEvalContext): TDFields? { + val fields = LinkedHashMap() + + keywordArguments.forEach { + if (it !is PyKeywordArgument || it.keyword == null) return null + fields[it.keyword!!] = it.valueExpression + } + + return typedDictFieldsFromKeysAndValues(fields, context) + } + + private fun typedDictFieldsFromKeysAndValues(fields: Map, context: TypeEvalContext): TDFields? { + val result = TDFields() + for ((name, type) in fields) { + result[name] = if (type != null) PyTypedDictType.FieldTypeAndTotality(context.getType(type)) + else PyTypedDictType.FieldTypeAndTotality(null) + } + return result + } + } +} \ No newline at end of file diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index 0eafdc887598..1a1b9172cb0b 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -69,6 +69,8 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { public static final String ASYNC_GENERATOR = "typing.AsyncGenerator"; public static final String COROUTINE = "typing.Coroutine"; public static final String NAMEDTUPLE = "typing.NamedTuple"; + public static final String TYPED_DICT = "typing.TypedDict"; + public static final String TYPED_DICT_EXT = "typing_extensions.TypedDict"; public static final String GENERIC = "typing.Generic"; public static final String PROTOCOL = "typing.Protocol"; public static final String PROTOCOL_EXT = "typing_extensions.Protocol"; @@ -76,6 +78,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { public static final String ANY = "typing.Any"; public static final String NEW_TYPE = "typing.NewType"; public static final String CALLABLE = "typing.Callable"; + public static final String MAPPING = "typing.Mapping"; private static final String LIST = "typing.List"; private static final String DICT = "typing.Dict"; private static final String DEFAULT_DICT = "typing.DefaultDict"; @@ -164,6 +167,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { .add(NO_RETURN) .add(FINAL, FINAL_EXT) .add(LITERAL, LITERAL_EXT) + .add(TYPED_DICT, TYPED_DICT_EXT) .build(); @Nullable @@ -854,6 +858,10 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return aliasedType; } } + final PyType typedDictType = PyTypedDictTypeProvider.Companion.getTypedDictTypeForResolvedElement(resolved, context.getTypeContext()); + if (typedDictType != null) { + return Ref.create(typedDictType); + } final Ref classType = getClassType(resolved, context.getTypeContext()); if (classType != null) { return classType; @@ -1000,6 +1008,11 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return null; } + public static boolean isTypedDict(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + Collection qualifiedNames = resolveToQualifiedNames(expression, context); + return qualifiedNames.stream().anyMatch(name -> TYPED_DICT.equals(name) || TYPED_DICT_EXT.equals(name)); + } + public static boolean isFinal(@NotNull PyDecoratable decoratable, @NotNull TypeEvalContext context) { return ContainerUtil.exists(PyKnownDecoratorUtil.getKnownDecorators(decoratable, context), d -> d == TYPING_FINAL || d == TYPING_FINAL_EXT); @@ -1009,7 +1022,8 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return PyUtil.getParameterizedCachedValue(owner, context, p -> isFinalImpl(owner, p)); } - private static boolean isFinalImpl(@NotNull T owner, @NotNull TypeEvalContext context) { + private static boolean isFinalImpl(@NotNull T owner, + @NotNull TypeEvalContext context) { final PyExpression annotation = getAnnotationValue(owner, context); if (annotation instanceof PySubscriptionExpression) { return eventuallyResolvesToFinal(((PySubscriptionExpression)annotation).getOperand(), context); @@ -1455,7 +1469,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @NotNull - private static Collection resolveToQualifiedNames(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { + public static Collection resolveToQualifiedNames(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { final Set names = Sets.newLinkedHashSet(); for (PsiElement resolved : tryResolving(expression, context)) { final String name = getQualifiedName(resolved); diff --git a/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java b/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java index fde45b6aefc7..4513c7e1f0d5 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java +++ b/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java @@ -151,6 +151,19 @@ public class PyTypeModelBuilder { } } + static class InferredTypedDictType extends TypeModel { + private final List members; + + InferredTypedDictType(List members) { + this.members = members; + } + + @Override + void accept(@NotNull TypeVisitor visitor) { + visitor.typedDict(this); + } + } + static class FunctionType extends TypeModel { @NotNull private final TypeModel returnType; @Nullable private final Collection parameters; @@ -180,7 +193,7 @@ public class PyTypeModelBuilder { visitor.param(this); } } - + static class ClassObjectType extends TypeModel { private final TypeModel classType; @@ -192,8 +205,8 @@ public class PyTypeModelBuilder { void accept(@NotNull TypeVisitor visitor) { visitor.classObject(this); } - } - + } + static class GenericType extends TypeModel { private final String name; @@ -226,7 +239,14 @@ public class PyTypeModelBuilder { myVisited.put(type, null); //mark as evaluating TypeModel result = null; - if (type instanceof PyInstantiableType && ((PyInstantiableType)type).isDefinition()) { + if (type instanceof PyTypedDictType) { + if (((PyTypedDictType)type).isInferred()) { + result = new InferredTypedDictType(Collections.singletonList(build(((PyTypedDictType)type).getValuesType(), true))); + } else { + result = NamedType.nameOrAny(type); + } + } + else if (type instanceof PyInstantiableType && ((PyInstantiableType)type).isDefinition()) { final PyInstantiableType instanceType = ((PyInstantiableType)type).toInstance(); // Special case: render Type[type] as just type if (type instanceof PyClassType && instanceType.equals(PyBuiltinCache.getInstance(((PyClassType)type).getPyClass()).getTypeType())) { @@ -343,6 +363,8 @@ public class PyTypeModelBuilder { void tuple(TupleType type); + void typedDict(InferredTypedDictType type); + void classObject(ClassObjectType type); void genericType(GenericType type); @@ -593,6 +615,24 @@ public class PyTypeModelBuilder { } } + @Override + public void typedDict(InferredTypedDictType type) { + add("Dict[str, "); + boolean first = true; + if (!type.members.isEmpty()) { + for (TypeModel member : type.members) { + if (!first) { + add(", "); + } + else { + first = false; + } + member.accept(this); + } + add("]"); + } + } + @Override public void classObject(ClassObjectType type) { add("Type["); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/PyFileElementType.java b/python/python-psi-impl/src/com/jetbrains/python/psi/PyFileElementType.java index 13353ea4a80e..239dbc64a24f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/PyFileElementType.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/PyFileElementType.java @@ -62,7 +62,7 @@ public class PyFileElementType extends IStubFileElementType { @Override public int getStubVersion() { // Don't forget to update versions of indexes that use the updated stub-based elements - return 77; + return 78; } @Nullable diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java b/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java index c2b88c55c6f1..8ddfe7344b10 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/PyUtil.java @@ -230,7 +230,7 @@ public class PyUtil { return new PyClassTypeImpl(((PyClassType)qualifierType).getPyClass(), true); // always as class, never instance } } - else if (PyNames.DICT.equals(attr_name)) { + else if (PyNames.DUNDER_DICT.equals(attr_name)) { PyType qualifierType = context.getType(qualifier); if (qualifierType instanceof PyClassType && ((PyClassType)qualifierType).isDefinition()) { return PyBuiltinCache.getInstance(ref).getDictType(); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassImpl.java index c8aff196d509..6adb29b66ba2 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyClassImpl.java @@ -271,7 +271,7 @@ public class PyClassImpl extends PyBaseElementImpl implements PyCla if (!cls.isNewStyleClass(contextToUse)) return null; final List ownSlots = cls.getOwnSlots(); - if (ownSlots == null || ownSlots.contains(PyNames.DICT)) { + if (ownSlots == null || ownSlots.contains(PyNames.DUNDER_DICT)) { return null; } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java index fa78cdc6ef24..d5ac42517ee6 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java @@ -79,11 +79,17 @@ public class PySubscriptionExpressionImpl extends PyElementImpl implements PySub if (type instanceof PyTupleType) { final PyTupleType tupleType = (PyTupleType)type; return Optional - .ofNullable(new PyEvaluator().evaluate(indexExpression)) - .map(value -> PyUtil.as(value, Integer.class)) + .ofNullable(PyEvaluator.evaluate(indexExpression, Integer.class)) .map(tupleType::getElementType) .orElse(null); } + if (type instanceof PyTypedDictType) { + final PyTypedDictType typedDictType = (PyTypedDictType)type; + return Optional + .ofNullable(PyEvaluator.evaluate(indexExpression, String.class)) + .map(typedDictType::getElementType) + .orElse(null); + } for (PsiElement resolved : PyUtil.multiResolveTopPriority(reference)) { PyType res = null; if (resolved instanceof PyCallable) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubImpl.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubImpl.kt new file mode 100644 index 000000000000..e5e26283477c --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubImpl.kt @@ -0,0 +1,146 @@ +// Copyright 2000-2019 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.impl.stubs + +import com.intellij.openapi.util.Pair +import com.intellij.psi.PsiElement +import com.intellij.psi.stubs.StubInputStream +import com.intellij.psi.stubs.StubOutputStream +import com.intellij.psi.util.QualifiedName +import com.jetbrains.python.codeInsight.typing.PyTypedDictTypeProvider +import com.jetbrains.python.psi.* +import com.jetbrains.python.psi.impl.PyEvaluator +import com.jetbrains.python.psi.impl.PyPsiUtils +import com.jetbrains.python.psi.resolve.PyResolveUtil +import com.jetbrains.python.psi.stubs.PyTypedDictStub +import java.io.IOException +import java.util.* + +class PyTypedDictStubImpl private constructor(private val myCalleeName: QualifiedName?, + override val name: String, + override val fields: LinkedHashMap>, + override val isTotal: Boolean = true) : PyTypedDictStub { + + override fun getTypeClass(): Class> { + return PyTypedDictStubType::class.java + } + + @Throws(IOException::class) + override fun serialize(stream: StubOutputStream) { + stream.writeName(myCalleeName?.toString()) + stream.writeName(name) + stream.writeVarInt(fields.size) + + for ((key, value) in fields) { + stream.writeName(key) + stream.writeName(value.orElse(null)) + } + } + + override fun getCalleeName(): QualifiedName? { + return myCalleeName + } + + companion object { + + fun create(expression: PyTargetExpression): PyTypedDictStub? { + val assignedValue = expression.findAssignedValue() + + return if (assignedValue is PyCallExpression) create(assignedValue) else null + } + + fun create(expression: PyCallExpression): PyTypedDictStub? { + val calleeReference = expression.callee as? PyReferenceExpression ?: return null + + val calleeName = getCalleeName(calleeReference) + + if (calleeName != null) { + val name = PyResolveUtil.resolveStrArgument(expression, 0, "name") ?: return null + + val fieldsAndTotality = resolveTypingTDFields(expression) + + if (fieldsAndTotality?.first != null && fieldsAndTotality.second != null) { + return PyTypedDictStubImpl(calleeName, name, fieldsAndTotality.first, fieldsAndTotality.second) + } + } + + return null + } + + @Throws(IOException::class) + fun deserialize(stream: StubInputStream): PyTypedDictStub? { + val calleeName = stream.readNameString() + val name = stream.readNameString() + val fields = deserializeFields(stream, stream.readVarInt()) + + return if (calleeName == null || name == null) { + null + } + else PyTypedDictStubImpl(QualifiedName.fromDottedString(calleeName), name, fields) + } + + private fun getCalleeName(referenceExpression: PyReferenceExpression): QualifiedName? { + val calleeName = PyPsiUtils.asQualifiedName(referenceExpression) ?: return null + + for (name in PyResolveUtil.resolveImportedElementQNameLocally(referenceExpression).map { it.toString() }) { + if (PyTypedDictTypeProvider.nameIsTypedDict(name)) { + return calleeName + } + } + + return null + } + + @Throws(IOException::class) + private fun deserializeFields(stream: StubInputStream, fieldsSize: Int): LinkedHashMap> { + val fields = LinkedHashMap>(fieldsSize) + + for (i in 0 until fieldsSize) { + val name = stream.readNameString() + val type = stream.readNameString() + + if (name != null) { + fields[name] = Optional.ofNullable(type) + } + } + + return fields + } + + private fun resolveTypingTDFields(callExpression: PyCallExpression): Pair>, Boolean>? { + // SUPPORTED CASES: + + // fields = {"x": str, "y": int} + // Movie = TypedDict(..., fields) + + // Movie = TypedDict(..., {'name': str, 'year': int}, total=False) + + val secondArgument = PyPsiUtils.flattenParens(callExpression.getArgument(1, PyExpression::class.java)) + + val resolvedFields = if (secondArgument is PyReferenceExpression) PyResolveUtil.fullResolveLocally(secondArgument) else secondArgument + return if (resolvedFields !is PySequenceExpression) null + else Pair.create(getTypingTDFieldsFromIterable(resolvedFields), + PyEvaluator.evaluateAsBoolean(callExpression.getKeywordArgument("total"), true)) + } + + private fun getTypingTDFieldsFromIterable(fields: PySequenceExpression): LinkedHashMap>? { + val result = LinkedHashMap>() + + fields.elements.forEach { + if (it !is PyKeyValueExpression) return null + + val name: PyExpression = it.key + val type: PyExpression? = it.value + + if (name !is PyStringLiteralExpression) return null + + result[name.stringValue] = Optional.ofNullable(textIfPresent(type)) + } + + return result + } + + private fun textIfPresent(element: PsiElement?): String? { + return element?.text + } + } +} diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubType.kt new file mode 100644 index 000000000000..4620c9cc8185 --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/stubs/PyTypedDictStubType.kt @@ -0,0 +1,20 @@ +// Copyright 2000-2019 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.impl.stubs + +import com.intellij.psi.stubs.StubInputStream +import com.jetbrains.python.psi.PyTargetExpression +import com.jetbrains.python.psi.stubs.PyTypedDictStub + +import java.io.IOException + +class PyTypedDictStubType : CustomTargetExpressionStubType() { + + override fun createStub(psi: PyTargetExpression): PyTypedDictStub? { + return PyTypedDictStubImpl.create(psi) + } + + @Throws(IOException::class) + override fun deserializeStub(stream: StubInputStream): PyTypedDictStub? { + return PyTypedDictStubImpl.deserialize(stream) + } +} diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/stubs/PyTypedDictStub.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/stubs/PyTypedDictStub.kt new file mode 100644 index 000000000000..1a7fe89c1df7 --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/stubs/PyTypedDictStub.kt @@ -0,0 +1,21 @@ +// Copyright 2000-2019 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.stubs + +import com.jetbrains.python.psi.impl.stubs.CustomTargetExpressionStub +import java.util.Optional + +interface PyTypedDictStub : CustomTargetExpressionStub { + + /** + * @return TypedDict's name. + */ + val name: String + + /** + * @return keys' names and their values' types. + * Iteration order repeats the declaration order. + */ + val fields: Map> + + val isTotal: Boolean +} diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 2b9d63472531..4b549ceee973 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -281,6 +281,18 @@ public class PyTypeChecker { return Optional.of(actual instanceof PyLiteralType && PyLiteralType.Companion.match((PyLiteralType)expected, (PyLiteralType)actual)); } + if (actual instanceof PyTypedDictType) { + if (!((PyTypedDictType)actual).isInferred()) { + Optional match = PyTypedDictType.Companion.checkStructuralCompatibility(expected, (PyTypedDictType)actual, context.context); + if (match.isPresent()) { + return match; + } + } + if (expected instanceof PyTypedDictType) { + return Optional.of(PyTypedDictType.Companion.match((PyTypedDictType)expected, (PyTypedDictType)actual, context.context)); + } + } + final PyClass superClass = expected.getPyClass(); final PyClass subClass = actual.getPyClass(); final boolean matchClasses = matchClasses(superClass, subClass, context.context); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt new file mode 100644 index 000000000000..818db32cc65d --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypedDictType.kt @@ -0,0 +1,214 @@ +// Copyright 2000-2019 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.jetbrains.python.PyNames +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.psi.* +import com.jetbrains.python.psi.impl.PyBoolLiteralExpressionImpl +import com.jetbrains.python.psi.impl.PyBuiltinCache +import com.jetbrains.python.psi.impl.PyPassStatementImpl +import one.util.streamex.StreamEx +import java.util.* + +class PyTypedDictType @JvmOverloads constructor(private val name: String, + val fields: LinkedHashMap, + private val inferred: Boolean, + private val dictClass: PyClass, + private val definitionLevel: DefinitionLevel, + private val ancestors: List, + private val targetExpression: PyTargetExpression? = null) : PyClassTypeImpl(dictClass, + definitionLevel != DefinitionLevel.INSTANCE), PyCollectionType { + override fun getElementTypes(): List { + return listOf(PyBuiltinCache.getInstance(dictClass).strType, getValuesType()) + } + + override fun getIteratedItemType(): PyType? { + return PyBuiltinCache.getInstance(dictClass).strType + } + + fun getValuesType(): PyType? { + return PyUnionType.union(fields.map { it.value.type }) + } + + fun getElementType(key: String): PyType? { + return fields[key]?.type + } + + override fun getCallType(context: TypeEvalContext, callSite: PyCallSiteExpression): PyType? { + if (definitionLevel == DefinitionLevel.NEW_TYPE) { + return toInstance() + } + + return null + } + + override fun isDefinition(): Boolean { + return definitionLevel == DefinitionLevel.NEW_TYPE + } + + override fun toInstance(): PyClassType { + return if (definitionLevel == DefinitionLevel.NEW_TYPE) + PyTypedDictType(name, fields, inferred, dictClass, + DefinitionLevel.INSTANCE, ancestors, + targetExpression) + else + this + } + + override fun toClass(): PyClassLikeType { + return if (definitionLevel == DefinitionLevel.INSTANCE) + PyTypedDictType(name, fields, inferred, dictClass, + DefinitionLevel.NEW_TYPE, ancestors, + targetExpression) + else + this + } + + override fun getName(): String? { + return name + } + + override fun isBuiltin(): Boolean { + return false + } + + override fun isCallable(): Boolean { + return definitionLevel != DefinitionLevel.INSTANCE + } + + override fun getParameters(context: TypeEvalContext): List? { + val elementGenerator = PyElementGenerator.getInstance(dictClass.project) + val psi = PyCallableParameterImpl.psi(elementGenerator.createSingleStarParameter()) + val ellipsis = elementGenerator.createEllipsis() + return if (isCallable) + listOf(psi) + fields.map { + if (it.value.isRequired) PyCallableParameterImpl.nonPsi(it.key, it.value.type) + else { + PyCallableParameterImpl.nonPsi(it.key, it.value.type, ellipsis) + } + } + else + null + } + + private fun getKeysToValueTypes(): Map { + return fields.mapValues { it.value.type } + } + + override fun toString(): String { + return "PyTypedDictType: $name" + } + + override fun equals(other: Any?): Boolean { + if (other === this) return true + if (other == null || javaClass != other.javaClass) return false + + val otherTypedDict = other as? PyTypedDictType ?: return false + return name == otherTypedDict.name + && fields == otherTypedDict.fields + && inferred == otherTypedDict.inferred + && definitionLevel == otherTypedDict.definitionLevel + && ancestors == otherTypedDict.ancestors + && targetExpression == otherTypedDict.targetExpression + } + + override fun hashCode(): Int { + return Objects.hash(super.hashCode(), name, fields, inferred, definitionLevel, ancestors, targetExpression) + } + + enum class DefinitionLevel { + NEW_TYPE, + INSTANCE + } + + /** + * Is this an actual TypedDict type or something that is inferred to match the expected TypedDict type (e.g. [PyDictLiteralExpression]) + */ + fun isInferred(): Boolean { + return inferred + } + + class FieldTypeAndTotality @JvmOverloads constructor(val type: PyType?, val isRequired: Boolean = true) + + companion object { + + /** + * [actual] matches [expected] if: + * * all required keys from [expected] are present in [actual] + * * all keys from [actual] are present in [expected] + * * each key has the same value type in [expected] and [actual] + */ + fun match(expected: PyTypedDictType, actual: PyTypedDictType, context: TypeEvalContext): Boolean { + val mandatoryArguments = expected.fields.filterValues { it.isRequired }.mapValues { it.value.type } + val actualArguments = actual.getKeysToValueTypes() + val expectedArguments = expected.getKeysToValueTypes() + + return match(mandatoryArguments, expectedArguments, actualArguments, context) + } + + fun match(expected: PyTypedDictType, actual: PyDictLiteralExpression, context: TypeEvalContext): Boolean { + if (actual.elements.any { it.key !is PyStringLiteralExpression }) return false + val mandatoryArguments = expected.fields.filter { it.value.isRequired }.map { it.key to it.value.type }.toMap() + val actualArguments = actual.elements.map { + (it.key as PyStringLiteralExpression).stringValue to if (it.value != null) context.getType(it.value!!) else null + }.toMap() + val expectedArguments = expected.getKeysToValueTypes() + + return match(mandatoryArguments, expectedArguments, actualArguments, context) + } + + private fun match(mandatoryArguments: Map, + expectedArguments: Map, + actualArguments: Map, + context: TypeEvalContext): Boolean { + if (!actualArguments.keys.containsAll(mandatoryArguments.keys)) return false + + actualArguments.forEach { + if (!expectedArguments.containsKey(it.key)) { + return false + } + val matchResult: Boolean = strictUnionMatch(expectedArguments[it.key], it.value, context) + if (!matchResult) { + return false + } + } + + return true + } + + private fun strictUnionMatch(expected: PyType?, actual: PyType?, context: TypeEvalContext): Boolean { + if (actual is PyUnionType) { + return StreamEx.of(actual.members).allMatch { type -> PyTypeChecker.match(expected, type, context) } + } + + return PyTypeChecker.match(expected, actual, context) + } + + fun checkStructuralCompatibility(expected: PyType?, actual: PyTypedDictType, context: TypeEvalContext): Optional { + if (expected is PyCollectionType && PyTypingTypeProvider.MAPPING == expected.classQName) { + val builtinCache = PyBuiltinCache.getInstance(actual.dictClass) + val elementTypes = expected.elementTypes + return Optional.of(elementTypes.size == 2 + && builtinCache.strType == elementTypes[0] + && (elementTypes[1] == null || PyNames.OBJECT == elementTypes[1].name)) + } + + if (expected !is PyTypedDictType) return Optional.empty() + + expected.fields.forEach { + val expectedTypeAndTotality = it.value + + if (!actual.fields.containsKey(it.key)) return Optional.of(false) + + val actualTypeAndTotality = actual.fields[it.key] + if (actualTypeAndTotality == null + || !strictUnionMatch(expectedTypeAndTotality.type, actualTypeAndTotality.type, context) + || !strictUnionMatch(actualTypeAndTotality.type, expectedTypeAndTotality.type, context) + || expectedTypeAndTotality.isRequired.xor(actualTypeAndTotality.isRequired)) { + return Optional.of(false) + } + } + return Optional.of(true) + } + } +} diff --git a/python/python-psi-impl/test/com/jetbrains/python/completion/PythonCompletionTest.java b/python/python-psi-impl/test/com/jetbrains/python/completion/PythonCompletionTest.java index 21961fb866bf..ff573c158b18 100644 --- a/python/python-psi-impl/test/com/jetbrains/python/completion/PythonCompletionTest.java +++ b/python/python-psi-impl/test/com/jetbrains/python/completion/PythonCompletionTest.java @@ -1580,6 +1580,32 @@ public class PythonCompletionTest extends PyTestCase { ); } + // PY-36008 + public void testTypedDictHasDictMethods() { + final List suggested = doTestByText("from typing import TypedDict\n" + + "class A(TypedDict):\n" + + " pass\n" + + "A()."); + + assertNotNull(suggested); + assertContainsElements(suggested, "update", "clear", "pop", "popitem", "setdefault"); + } + + // PY-36008 + public void testTypedDictDefinition() { + final List suggested = doTestByText("from typing import TypedDict\n" + + "class A(TypedDict, total=):\n"); + + final List suggestedInAlternativeSyntax = doTestByText("from typing import TypedDict\n" + + "A = TypedDict('A', {}, total=):\n"); + + assertNotNull(suggested); + assertContainsElements(suggested, "True", "False"); + + assertNotNull(suggestedInAlternativeSyntax); + assertContainsElements(suggestedInAlternativeSyntax, "True", "False"); + } + private void assertNoVariantsInExtendedCompletion() { myFixture.copyDirectoryToProject(getTestName(true), ""); myFixture.configureByFile("a.py"); diff --git a/python/resources/inspectionDescriptions/PyTypedDictInspection.html b/python/resources/inspectionDescriptions/PyTypedDictInspection.html new file mode 100644 index 000000000000..665f49a5f956 --- /dev/null +++ b/python/resources/inspectionDescriptions/PyTypedDictInspection.html @@ -0,0 +1,5 @@ + + +This inspection detects invalid definition and usage of TypedDict. + + \ No newline at end of file diff --git a/python/src/META-INF/python-core-common.xml b/python/src/META-INF/python-core-common.xml index 3f27bc977f3c..71d0bbbbd6ce 100644 --- a/python/src/META-INF/python-core-common.xml +++ b/python/src/META-INF/python-core-common.xml @@ -99,6 +99,8 @@ implementationClass="com.jetbrains.python.codeInsight.completion.PyStructuralTypeAttributesCompletionContributor"/> + @@ -434,6 +436,7 @@ + diff --git a/python/src/com/jetbrains/python/inspections/PyClassHasNoInitInspection.java b/python/src/com/jetbrains/python/inspections/PyClassHasNoInitInspection.java index 43b22126c49b..764cf9544909 100644 --- a/python/src/com/jetbrains/python/inspections/PyClassHasNoInitInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyClassHasNoInitInspection.java @@ -22,6 +22,7 @@ import com.intellij.psi.PsiElementVisitor; import com.intellij.psi.util.PsiTreeUtil; import com.jetbrains.python.PyBundle; import com.jetbrains.python.PyNames; +import com.jetbrains.python.codeInsight.typing.PyTypedDictTypeProvider; import com.jetbrains.python.inspections.quickfix.AddMethodQuickFix; import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyFunction; @@ -66,6 +67,7 @@ public class PyClassHasNoInitInspection extends PyInspection { return; } final List types = node.getAncestorTypes(myTypeEvalContext); + if (PyTypedDictTypeProvider.Companion.isTypingTypedDictInheritor(node, myTypeEvalContext)) return; for (PyClassLikeType type : types) { if (type == null) return; final String qName = type.getClassQName(); diff --git a/python/src/com/jetbrains/python/inspections/PyStringFormatInspection.java b/python/src/com/jetbrains/python/inspections/PyStringFormatInspection.java index 8ab0b7d31b93..50315bca9334 100644 --- a/python/src/com/jetbrains/python/inspections/PyStringFormatInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyStringFormatInspection.java @@ -111,7 +111,7 @@ public class PyStringFormatInspection extends PyInspection { return 1; } else if (rightExpression instanceof PyReferenceExpression) { - if (PyNames.DICT.equals(rightExpression.getName())) return -1; + if (PyNames.DUNDER_DICT.equals(rightExpression.getName())) return -1; final List resolveResults = ((PyReferenceExpression)rightExpression).multiFollowAssignmentsChain(resolveContext); diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index 8943a001ed83..d6941eef0c8b 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -11,6 +11,8 @@ import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; +import com.jetbrains.python.codeInsight.typing.PyTypedDictTypeProvider; +import com.jetbrains.python.psi.types.PyTypedDictType; import com.jetbrains.python.codeInsight.typing.PyProtocolsKt; import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.documentation.PythonDocumentationProvider; @@ -112,7 +114,7 @@ public class PyTypeCheckerInspection extends PyInspection { final PyExpression value = node.findAssignedValue(); if (value == null) return; final PyType expected = myTypeEvalContext.getType(node); - final PyType actual = PyLiteralType.Companion.promoteToLiteral(value, expected, myTypeEvalContext); + final PyType actual = tryPromotingType(value, expected); if (!PyTypeChecker.match(expected, actual, myTypeEvalContext)) { registerProblem(value, String.format("Expected type '%s', got '%s' instead", PythonDocumentationProvider.getTypeName(expected, myTypeEvalContext), @@ -120,6 +122,13 @@ public class PyTypeCheckerInspection extends PyInspection { } } + @Nullable + private PyType tryPromotingType(@NotNull PyExpression value, @Nullable PyType expected) { + final PyType promotedToLiteral = PyLiteralType.Companion.promoteToLiteral(value, expected, myTypeEvalContext); + if (promotedToLiteral instanceof PyLiteralType) return promotedToLiteral; + return PyTypedDictTypeProvider.Companion.promoteToTypedDict(value, expected, myTypeEvalContext); + } + @Override public void visitPyFunction(PyFunction node) { final PyAnnotation annotation = node.getAnnotation(); @@ -215,7 +224,7 @@ public class PyTypeCheckerInspection extends PyInspection { final PyCallableParameter parameter = entry.getValue(); final PyType expected = parameter.getArgumentType(myTypeEvalContext); final PyType actual = PyLiteralType.Companion.promoteToLiteral(argument, expected, myTypeEvalContext); - final boolean matched = matchParameterAndArgument(expected, actual, substitutions); + final boolean matched = matchParameterAndArgument(expected, actual, argument, substitutions); result.add(new AnalyzeArgumentResult(argument, expected, substituteGenerics(expected, substitutions), actual, matched)); } final PyCallableParameter positionalContainer = getMappedPositionalContainer(mappedParameters); @@ -238,7 +247,7 @@ public class PyTypeCheckerInspection extends PyInspection { // For an expected type with generics we have to match all the actual types against it in order to do proper generic unification if (PyTypeChecker.hasGenerics(expected, myTypeEvalContext)) { final PyType actual = PyUnionType.union(ContainerUtil.map(arguments, myTypeEvalContext::getType)); - final boolean matched = matchParameterAndArgument(expected, actual, substitutions); + final boolean matched = matchParameterAndArgument(expected, actual, null, substitutions); return ContainerUtil.map(arguments, argument -> new AnalyzeArgumentResult(argument, expected, expectedWithSubstitutions, actual, matched)); } @@ -247,7 +256,7 @@ public class PyTypeCheckerInspection extends PyInspection { arguments, argument -> { final PyType actual = myTypeEvalContext.getType(argument); - final boolean matched = matchParameterAndArgument(expected, actual, substitutions); + final boolean matched = matchParameterAndArgument(expected, actual, argument, substitutions); return new AnalyzeArgumentResult(argument, expected, expectedWithSubstitutions, actual, matched); } ); @@ -256,9 +265,12 @@ public class PyTypeCheckerInspection extends PyInspection { private boolean matchParameterAndArgument(@Nullable PyType parameterType, @Nullable PyType argumentType, + @Nullable PyExpression argument, @NotNull Map substitutions) { - return PyTypeChecker.match(parameterType, argumentType, myTypeEvalContext, substitutions) && - !PyProtocolsKt.matchingProtocolDefinitions(parameterType, argumentType, myTypeEvalContext); + return parameterType instanceof PyTypedDictType && argument instanceof PyDictLiteralExpression + ? PyTypedDictType.Companion.match((PyTypedDictType)parameterType, (PyDictLiteralExpression)argument, myTypeEvalContext) + : (PyTypeChecker.match(parameterType, argumentType, myTypeEvalContext, substitutions) && + !PyProtocolsKt.matchingProtocolDefinitions(parameterType, argumentType, myTypeEvalContext)); } @Nullable diff --git a/python/src/com/jetbrains/python/inspections/PyTypedDictInspection.kt b/python/src/com/jetbrains/python/inspections/PyTypedDictInspection.kt new file mode 100644 index 000000000000..e74b4a2cc3d0 --- /dev/null +++ b/python/src/com/jetbrains/python/inspections/PyTypedDictInspection.kt @@ -0,0 +1,328 @@ +// Copyright 2000-2019 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.inspections + +import com.intellij.codeInspection.LocalInspectionToolSession +import com.intellij.codeInspection.ProblemHighlightType +import com.intellij.codeInspection.ProblemsHolder +import com.intellij.openapi.util.Ref +import com.intellij.psi.PsiElement +import com.intellij.psi.PsiElementVisitor +import com.intellij.psi.PsiNameIdentifierOwner +import com.jetbrains.python.PyNames +import com.jetbrains.python.codeInsight.typing.PyTypedDictTypeProvider +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.documentation.PythonDocumentationProvider +import com.jetbrains.python.psi.* +import com.jetbrains.python.psi.impl.PyExpressionStatementImpl +import com.jetbrains.python.psi.impl.PyPsiUtils +import com.jetbrains.python.psi.types.PyType +import com.jetbrains.python.psi.types.PyTypeChecker +import com.jetbrains.python.psi.types.PyTypedDictType + +class PyTypedDictInspection : PyInspection() { + + override fun buildVisitor(holder: ProblemsHolder, + isOnTheFly: Boolean, + session: LocalInspectionToolSession): PsiElementVisitor { + return Visitor(holder, session) + } + + private class Visitor(holder: ProblemsHolder, session: LocalInspectionToolSession) : PyInspectionVisitor(holder, session) { + + override fun visitPySubscriptionExpression(node: PySubscriptionExpression) { + super.visitPySubscriptionExpression(node) + + val rootOperand = node.rootOperand + val rootOperandType = myTypeEvalContext.getType(rootOperand) + if (rootOperandType !is PyTypedDictType) return + + val indexExpression = node.indexExpression + val indexExprValue = getIndexExpressionAsString(indexExpression) + if (indexExprValue == null) { + registerProblem(indexExpression, "TypedDict key type must be string") + return + } + + if (!rootOperandType.fields.containsKey(indexExprValue)) { + registerProblem(indexExpression, String.format("TypedDict '%s' cannot have key '%s'", rootOperandType.name, indexExprValue)) + } + } + + override fun visitPyTargetExpression(node: PyTargetExpression?) { + super.visitPyTargetExpression(node) + if (node == null) return + + if (node.hasAssignedValue()) { + val value = node.findAssignedValue() + if (value is PyCallExpression && value.callee != null && + PyTypingTypeProvider.resolveToQualifiedNames(value.callee!!, myTypeEvalContext).any { + PyTypedDictTypeProvider.nameIsTypedDict(it) + }) { + if (value.arguments.isNotEmpty() && node.name != (value.arguments[0] as? PyStringLiteralExpression)?.stringValue) { + registerProblem(value.arguments[0], "First argument has to match the variable name") + } + } + } + } + + override fun visitPyArgumentList(node: PyArgumentList) { + super.visitPyArgumentList(node) + + if (node.parent is PyClass && PyTypedDictTypeProvider.isTypingTypedDictInheritor(node.parent as PyClass, myTypeEvalContext)) { + val arguments = node.arguments + for (argument in arguments) { + val type = myTypeEvalContext.getType(argument) + if (argument !is PyKeywordArgument && !PyTypingTypeProvider.isTypedDict(argument, + myTypeEvalContext) && type !is PyTypedDictType) { + registerProblem(argument, "TypedDict cannot inherit from a non-TypedDict base class") + } + if (argument is PyKeywordArgument && argument.keyword == "total" && !checkValidTotality(argument.valueExpression)) { + registerProblem(argument.valueExpression, "Value of 'total' must be True or False") + } + } + } + else if (node.callExpression != null) { + val callee = node.callExpression!!.callee + if (callee != null && PyTypingTypeProvider.resolveToQualifiedNames(callee, myTypeEvalContext).any { + PyTypedDictTypeProvider.nameIsTypedDict(it) + }) { + val totality = node.getKeywordArgument("total")?.valueExpression + if (!checkValidTotality(totality)) { + registerProblem(totality, "Value of 'total' must be True or False") + } + val fields = node.arguments.filterIsInstance().firstOrNull() + if (fields == null) return + + fields.elements.forEach { + if (it !is PyKeyValueExpression) return + + if (it.value !is PyReferenceExpression) { + registerProblem(it.value, "Value must be a type") + } + val name = it.value?.name + if (name != null) { + val type = Ref.deref(PyTypingTypeProvider.getStringBasedType(name, it, myTypeEvalContext)) + if (type == null && !PyTypingTypeProvider.resolveToQualifiedNames(it.value!!, myTypeEvalContext).any { qualifiedName -> + PyTypingTypeProvider.ANY == qualifiedName + }) { + registerProblem(it.value, "Value must be a type") + } + } + } + } + } + } + + override fun visitPyClass(node: PyClass) { + super.visitPyClass(node) + + if (LanguageLevel.forElement(node).isAtLeast(LanguageLevel.PYTHON36) && + PyTypedDictTypeProvider.isTypingTypedDictInheritor(node, myTypeEvalContext)) { + if (node.metaClassExpression != null) { + registerProblem((node.metaClassExpression as PyExpression).parent, "Specifying a metaclass is not allowed in TypedDict") + } + + val ancestorsFields = mutableMapOf() + val typedDictAncestors = node.getAncestorTypes(myTypeEvalContext).filterIsInstance() + typedDictAncestors.forEach { typedDict -> + typedDict.fields.forEach { field -> + val key = field.key + val value = field.value + if (ancestorsFields.containsKey(key) && ancestorsFields[key] != value) { + registerProblem(node.superClassExpressionList, + "Cannot overwrite TypedDict field \'$key\' while merging") + } + else { + ancestorsFields[key] = value + } + } + } + + typedDictAncestors + .flatMap { it.fields.entries } + .map { it.key to it.value } + .toMap(ancestorsFields) + + val statements = node.statementList.statements + if (statements.size == 1) { + val statement = statements[0] + if (statement !is PyTypeDeclarationStatement + && statement !is PyPassStatement + && !isDocString(statement)) { + registerProblem(tryGetNameIdentifier(statement), "Invalid statement in TypedDict definition; expected 'field_name: field_type'", + ProblemHighlightType.WEAK_WARNING) + return + } + } + + node.processClassLevelDeclarations { element, _ -> + if (element !is PyTargetExpression) { + registerProblem(tryGetNameIdentifier(element), "Invalid statement in TypedDict definition; expected 'field_name: field_type'", + ProblemHighlightType.WEAK_WARNING) + return@processClassLevelDeclarations true + } + if (element.hasAssignedValue()) { + registerProblem(element.findAssignedValue(), "Right hand side values are not supported in TypedDict") + return@processClassLevelDeclarations true + } + if (ancestorsFields.containsKey(element.name)) { + registerProblem(element, "Cannot overwrite TypedDict field") + return@processClassLevelDeclarations true + } + true + } + } + } + + private fun tryGetNameIdentifier(element: PsiElement): PsiElement { + return if (element is PsiNameIdentifierOwner) element.nameIdentifier ?: element else element + } + + private fun isDocString(statement: PyStatement): Boolean { + return statement is PyExpressionStatementImpl + && (statement as PyExpressionStatement).expression is PyStringLiteralExpression + && (statement.expression as PyStringLiteralExpression).isDocString + } + + override fun visitPyDelStatement(node: PyDelStatement) { + super.visitPyDelStatement(node) + + for (target in node.targets) { + for (expr in PyUtil.flattenedParensAndTuples(target)) { + if (expr !is PySubscriptionExpression) return + val rootOp = expr.rootOperand + val type = myTypeEvalContext.getType(rootOp) + if (type is PyTypedDictType) { + val index = getIndexExpressionAsString(expr.indexExpression) + if (index == null || !type.fields.containsKey(index)) return + if (type.fields[index]!!.isRequired) { + registerProblem(expr.indexExpression, "Key '$index' of TypedDict '${type.name}' cannot be deleted") + } + } + } + } + } + + override fun visitPyCallExpression(node: PyCallExpression) { + super.visitPyCallExpression(node) + + val callee = node.callee + if (callee !is PyReferenceExpression || callee.qualifier == null) return + + val nodeType = myTypeEvalContext.getType(callee.qualifier!!) + if (nodeType !is PyTypedDictType) return + val arguments = node.arguments + + if (PyNames.UPDATE == callee.name) { + inspectUpdateSequenceArgument( + if (arguments.size == 1 && arguments[0] is PySequenceExpression) (arguments[0] as PySequenceExpression).elements else arguments, + nodeType) + } + + if (PyNames.CLEAR == callee.name || PyNames.POPITEM == callee.name) { + if (nodeType.fields.any { it.value.isRequired }) { + registerProblem(callee.nameElement?.psi, "This operation might break TypedDict consistency", + ProblemHighlightType.WEAK_WARNING) + } + } + + if (PyNames.POP == callee.name) { + val key = if (arguments.isNotEmpty()) getIndexExpressionAsString(arguments[0]) else null + if (key != null && nodeType.fields.containsKey(key) && nodeType.fields[key]!!.isRequired) { + registerProblem(callee.nameElement?.psi, "Key '$key' of TypedDict '${nodeType.name}' cannot be deleted") + } + } + + if (PyNames.SETDEFAULT == callee.name) { + val key = if (arguments.isNotEmpty()) getIndexExpressionAsString(arguments[0]) else null + if (key != null && nodeType.fields.containsKey(key) && !nodeType.fields[key]!!.isRequired) { + if (node.arguments.size > 1) { + val valueType = myTypeEvalContext.getType(arguments[1]) + if (nodeType.fields[key]!!.type != valueType) { + registerProblem(arguments[1], String.format("Expected type '%s', got '%s' instead", + PythonDocumentationProvider.getTypeName(nodeType.fields[key]!!.type, + myTypeEvalContext), + PythonDocumentationProvider.getTypeName(valueType, myTypeEvalContext))) + } + } + } + } + } + + private fun checkValidTotality(totalityExpression: PyExpression?): Boolean { + val languageLevel = (totalityExpression?.containingFile as? PyFile)?.languageLevel ?: return false + if (languageLevel.isAtLeast(LanguageLevel.PYTHON38)) return totalityExpression is PyBoolLiteralExpression + else return listOf(PyNames.TRUE, PyNames.FALSE).contains(totalityExpression.text) + } + + private fun inspectUpdateSequenceArgument(sequenceElements: Array, typedDictType: PyTypedDictType) { + sequenceElements.forEach { + var key: PsiElement? = null + var keyAsString: String? = null + var value: PyExpression? = null + + if (it is PyKeyValueExpression && it.key is PyStringLiteralExpression) { + key = it.key + keyAsString = (it.key as PyStringLiteralExpression).stringValue + value = it.value + } + else if (it is PyParenthesizedExpression) { + var expression: PyExpression? = it + while (expression is PyParenthesizedExpression) { + expression = PyPsiUtils.flattenParens(expression) + } + if (expression == null) return@forEach + + if (expression is PyTupleExpression && expression.elements.size == 2 && expression.elements[0] is PyStringLiteralExpression) { + key = expression.elements[0] + keyAsString = (expression.elements[0] as PyStringLiteralExpression).stringValue + value = expression.elements[1] + } + } + else if (it is PyKeywordArgument && it.valueExpression != null) { + key = it.keywordNode?.psi + keyAsString = it.keyword + value = it.valueExpression + } + else return@forEach + + val fields = typedDictType.fields + if (value == null) { + return@forEach + } + if (keyAsString == null) { + registerProblem(key, "Cannot add a non-string key to TypedDict ${typedDictType.name}") + return@forEach + } + if (!fields.containsKey(keyAsString)) { + registerProblem(key, "TypedDict ${typedDictType.name} cannot have key ${keyAsString}") + return@forEach + } + val valueType = myTypeEvalContext.getType(value) + if (!PyTypeChecker.match(fields[keyAsString]?.type, valueType, myTypeEvalContext)) { + registerProblem(value, String.format("Expected type '%s', got '%s' instead", + PythonDocumentationProvider.getTypeName(fields[keyAsString]!!.type, myTypeEvalContext), + PythonDocumentationProvider.getTypeName(valueType, myTypeEvalContext))) + return@forEach + } + } + } + + private fun getIndexExpressionAsString(indexExpression: PyExpression?): String? { + if (indexExpression is PyStringLiteralExpression) { + return indexExpression.stringValue + } + + var index = indexExpression + var target = indexExpression?.reference?.resolve() + while (index is PyReferenceExpression && target is PyTargetExpression) { + index = target.findAssignedValue() + target = index?.reference?.resolve() + } + if (index is PyStringLiteralExpression) + return index.stringValue + + return null + } + } +} \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/typedDict.py b/python/testData/inspections/PyArgumentListInspection/typedDict.py new file mode 100644 index 000000000000..7994ac6b8bfd --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/typedDict.py @@ -0,0 +1,44 @@ +from typing import TypedDict + + +class X(TypedDict): + x: int + + +class Y(TypedDict, total=False): + y: str + + +class XYZ(X, Y): + z: bool + + +xyz = XYZ(z=True) + +x = X() +x.clear() +x.setdefault() + +x1: X = {'x': 42} +x1.clear() +x1.setdefault() + + +class Employee(TypedDict): + name: str + id: int + + +class Employee2(Employee, total=False): + director: str + + +em = Employee2(name='str') +em2 = Employee2("str", id=2) + + +Movie = TypedDict('Movie', {'name': str, 'year': int}, total=False) +Movie2 = TypedDict('Movie', {'name': str, 'year': int}) +Movie3 = TypedDict(3) +movie = Movie() +movie2 = Movie2() diff --git a/python/testData/inspections/PyClassHasNoInitInspection/typedDict.py b/python/testData/inspections/PyClassHasNoInitInspection/typedDict.py new file mode 100644 index 000000000000..7d959fb85c99 --- /dev/null +++ b/python/testData/inspections/PyClassHasNoInitInspection/typedDict.py @@ -0,0 +1,5 @@ +from typing import TypedDict + + +class X(TypedDict, total=False): + x: str diff --git a/python/testData/inspections/PyTypeCheckerInspection/TypedDictConsistency.py b/python/testData/inspections/PyTypeCheckerInspection/TypedDictConsistency.py new file mode 100644 index 000000000000..c62566657a6a --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/TypedDictConsistency.py @@ -0,0 +1,74 @@ +from typing import TypedDict, Optional, Union, Mapping, Any + + +class A(TypedDict): + x: Optional[int] +class B(TypedDict): + x: Optional[int] +def f(a: A) -> None: + a['x'] = None + +b: B = B(x=0) +f(b) + + +class C(TypedDict): + x: Union[int, str] +c: C = C(x = '0') +f(c) + + +class D(TypedDict): + x: int +def bar(a: A) -> None: + a['x'] = None +d: D = {'x': 0} +bar(d) + + +class E(TypedDict): + x: int +def f(d: Mapping[str, object]) -> None: + print(d) +def g(d: Mapping[str, Any]) -> None: + print(d) +def h(d: Mapping[str, int]) -> None: + print(d) +e: E = E(x=1) +f(e) +g(e) +h(e) + + +class A1(TypedDict, total=False): + x: int + y: int +class B1(TypedDict, total=False): + x: int +class C1(TypedDict, total=False): + x: int + y: str +def f1(a: A1) -> None: + a['y'] = 1 +def g1(b: B1) -> None: + f1(b) + + +class A2(TypedDict, total=False): + x: int +class B2(TypedDict): + x: int +def f2(a: A2) -> None: + del a['x'] +b: B2 = {'x': 0} +f2(b) + + + +class A3(TypedDict): + x: str +class B3(TypedDict): + x: str + y: str +a: A3 = B3(x = '', y = '') +b: B3 = A3(x = '') \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index a10a9f437175..aab5c006b648 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -974,6 +974,16 @@ public class Py3TypeTest extends PyTestCase { "\n" + " def __post_init__(self, a, b):\n" + " expr = a"); + } + ); + } + + // PY-28506 + public void testDataclassPostInitInheritedParameter2() { + runWithLanguageLevel( + LanguageLevel.PYTHON37, + () -> { + myFixture.copyDirectoryToProject(TEST_DIRECTORY + "DataclassPostInitParameter", ""); // both are dataclasses, base with enabled `init` doTest("Any", @@ -989,6 +999,16 @@ public class Py3TypeTest extends PyTestCase { "\n" + " def __post_init__(self, a, b):\n" + " expr = a"); + } + ); + } + + // PY-28506 + public void testDataclassPostInitInheritedParameter3() { + runWithLanguageLevel( + LanguageLevel.PYTHON37, + () -> { + myFixture.copyDirectoryToProject(TEST_DIRECTORY + "DataclassPostInitParameter", ""); // both are dataclasses, derived with enabled `init` doTest("int", @@ -1004,6 +1024,16 @@ public class Py3TypeTest extends PyTestCase { "\n" + " def __post_init__(self, a, b):\n" + " expr = a"); + } + ); + } + + // PY-28506 + public void testDataclassPostInitInheritedParameter4() { + runWithLanguageLevel( + LanguageLevel.PYTHON37, + () -> { + myFixture.copyDirectoryToProject(TEST_DIRECTORY + "DataclassPostInitParameter", ""); // both are dataclasses with disabled `init` doTest("Any", diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index a9b66bb70e10..f2c3d3dac06f 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -3576,6 +3576,21 @@ public class PyTypeTest extends PyTestCase { "expr = mytime.now()"); } + // PY-36008 + public void testTypedDict() { + runWithLanguageLevel( + LanguageLevel.getLatest(), + () -> { + doTest("A", + "from typing import TypedDict\n" + + "class A(TypedDict):\n" + + " x: int\n" + + "a: A = {'x': 42}\n" + + "expr = a"); + } + ); + } + private static List getTypeEvalContexts(@NotNull PyExpression element) { return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(), TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing()); diff --git a/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java index ea8cfb34e33f..5c4927826218 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java @@ -376,4 +376,9 @@ public class PyArgumentListInspectionTest extends PyInspectionTestCase { public void testPositionalOnlyParameters() { runWithLanguageLevel(LanguageLevel.PYTHON38, this::doTest); } + + // PY-36008 + public void testTypedDict() { + runWithLanguageLevel(LanguageLevel.PYTHON38, this::doTest); + } } diff --git a/python/testSrc/com/jetbrains/python/inspections/PyClassHasNoInitInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyClassHasNoInitInspectionTest.java index 9c7789e5262b..499ae4bd5398 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyClassHasNoInitInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyClassHasNoInitInspectionTest.java @@ -58,6 +58,11 @@ public class PyClassHasNoInitInspectionTest extends PyInspectionTestCase { runWithLanguageLevel(LanguageLevel.PYTHON34, this::doMultiFileTest); } + // PY-36008 + public void testTypedDict() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> doTest()); + } + @NotNull @Override protected Class getInspectionClass() { diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java index 7c921c9aa341..eadd670d925c 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java @@ -905,4 +905,94 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { "c: Literal[\"22\"] = f\"2{two}\"") ); } + + // PY-36008 + public void testTypedDictDefinitionAlternativeSyntax() { + doTestByText("from typing import TypedDict\n" + + "\n" + + "Movie = TypedDict('Movie', {'name': str, 'year': int}, total=False)\n" + + "movie = {'name': 'Blade Runner', 'lo': 1234} # type: Movie\n" + + "movie['year'] = '1984'\n"); + } + + // PY-36008 + public void testTypedDictAsArgument() { + runWithLanguageLevel( + LanguageLevel.PYTHON36, + () -> doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int\n" + + "def record_movie(movie: Movie) -> None: ...\n" + + "record_movie({'name': 'Blade Runner', 'year': 1982})\n" + + "record_movie({'name': 1984})") + ); + } + + // PY-36008 + public void testTypedDictSubscriptionAsArgument() { + runWithLanguageLevel( + LanguageLevel.PYTHON36, + () -> doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int\n" + + "m1: Movie = dict(name='Alien', year=1979)\n" + + "m2 = Movie(name='Garden State', year=2004)\n" + + "def foo(p: int):\n" + + " pass\n" + + "foo(m2[\"year\"])\n" + + "foo(m2[\"name\"])\n" + + "foo(m1[\"name\"])") + ); + } + + // PY-36008 + public void testTypedDictSubscriptionAssignment() { + runWithLanguageLevel( + LanguageLevel.PYTHON36, + () -> doTestByText("from typing import TypedDict\n" + + "class X(TypedDict):\n" + + " x: int\n" + + "f = X(x=12)\n" + + "f['x'] = 13\n" + + "f['x'] = '14'\n" + + "g: X = {'x': 12}\n" + + "g['x'] = 13\n" + + "g['x'] = '14'")); + } + + // PY-36008 + public void testTypedDictAssignment() { + runWithLanguageLevel( + LanguageLevel.PYTHON36, + () -> doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int\n" + + "m1: Movie = dict(name='Alien', year=1979)\n" + + "m2: Movie = dict(name='Alien', year='1979')\n" + + "m3: Movie = typing.cast(Movie, dict(zip(['name', 'year'], ['Alien', 1979])))\n" + + "m4: Movie = {'name': 'Alien', 'year': '1979'}\n" + + "m5 = Movie(name='Garden State', year=2004)")); + } + + // PY-36008 + public void testTypedDictDefinition() { + runWithLanguageLevel( + LanguageLevel.PYTHON36, + () -> doTestByText("from typing import TypedDict\n" + + "class Employee(TypedDict):\n" + + " name: str\n" + + " id: int\n" + + "class Employee2(Employee, total=False):\n" + + " director: str\n" + + "em = Employee2(name='John Dorian', id=1234, director=3)\n" + + "Movie = TypedDict(3, [1, 2, 3])")); + } + + // PY-36008 + public void testTypedDictConsistency() { + runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest); + } } diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypedDictInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypedDictInspectionTest.java new file mode 100644 index 000000000000..6ac5a9a84923 --- /dev/null +++ b/python/testSrc/com/jetbrains/python/inspections/PyTypedDictInspectionTest.java @@ -0,0 +1,230 @@ +// Copyright 2000-2019 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.inspections; + +import com.jetbrains.python.fixtures.PyInspectionTestCase; +import com.jetbrains.python.psi.LanguageLevel; +import org.jetbrains.annotations.NotNull; + +public class PyTypedDictInspectionTest extends PyInspectionTestCase { + + public void testClassBasedSyntax() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict\n" + + "\n" + + "class Employee(TypedDict):\n" + + " name: str\n" + + " id: int\n" + + "\n" + + "\n" + + "class Employee2(Employee, total=False):\n" + + " director: str\n" + + "\n" + + "\n" + + "em = Employee2(name='John Dorian', id=1234)\n" + + "em['director'] = 'Robert Kelso'\n" + + "em['slave'] = 'Doug Murphy'\n")); + } + + public void testAlternativeSyntax() { + doTestByText("from typing import TypedDict\n" + + "\n" + + "Movie = TypedDict('Movie', {'name': str, 'year': int}, total=False)\n" + + "movie = Movie(name='Blade Runner')\n" + + "movie['based_on_book'] = True\n" + + "movie['year'] = 1984"); + } + + public void testMetaclass() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict\n" + + "\n" + + "class Movie(TypedDict, metaclass=Meta):\n" + + " name: str")); + } + + public void testExtraClassDeclarations() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " def my_method(self):\n" + + " pass\n" + + " class Horror:\n" + + " def __init__(self):\n" + + " ...")); + } + + public void testInitializer() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int = 42")); + } + + public void testPass() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " ...\n" + + "class HorrorMovie(TypedDict):\n" + + " pass")); + } + + public void testNonTypedDictAsSuperclass() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict, NamedTuple\n" + + "class Bastard:\n" + + " pass\n" + + "class X(TypedDict):\n" + + " x: int\n" + + "class Y(TypedDict):\n" + + " y: str\n" + + "class XYZ(X, Bastard):\n" + + " z: bool\n" + + "class MyNT(NamedTuple):\n" + + " a: str")); + } + + public void testIncorrectTotalityValue() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict\n" + + "class X(TypedDict, total=1):\n" + + " x: int\n" + + "Movie = TypedDict(\"Movie\", {}, total=2)")); + } + + public void testNameAndVariableNameDoNotMatch() { + doTestByText("from typing import TypedDict\n" + + "Movie2 = TypedDict('Movie', {'name': str, 'year': int}, total=False)"); + } + + public void testKeyTypes() { + doTestByText("from typing import TypedDict, Any\n" + + "Movie = TypedDict('Movie', {'name': Any, 'year': 2}, total=False)"); + } + + public void testIncorrectKeyValue() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict\n" + + "class A(TypedDict):\n" + + " x: int\n" + + "a: A\n" + + "a = A(x=2)\n" + + "a['x'] = 1\n" + + "a['new'] = 10\n" + + "a[1] = 2")); + } + + public void testFinalKey() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict, Final\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int\n" + + "YEAR: Final = 'year'\n" + + "m = Movie(name='Alien', year=1979)\n" + + "years_since_epoch = m[YEAR] - 1970")); + } + + public void testStringVariableAsKey() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict, Final\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int\n" + + "year = 'year'\n" + + "year2 = year\n" + + "m = Movie(name='Alien', year=1979)\n" + + "years_since_epoch = m[year2] - 1970\n" + + "year = 42\n" + + "print(m[year])")); + } + + public void testDelStatement() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int\n" + + "class HorrorMovie(Movie, total=False):\n" + + " based_on_book: bool\n" + + "year = 'year'\n" + + "year2 = year\n" + + "m = HorrorMovie(name='Alien', year=1979)\n" + + "del (m['based_on_book'], m['name'])\n" + + "del m[year2], m['based_on_book']")); + } + + public void testDictModificationMethods() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: int\n" + + "class Horror(Movie, total=False):\n" + + " based_on_book: bool\n" + + "m = Horror(name='Alien', year=1979)\n" + + "m.clear()\n" + + "name = 'name'\n" + + "m.pop('based_on_book')\n" + + "m.pop('year')\n" + + "m.popitem()\n" + + "m.setdefault('based_on_book', 42)")); + } + + public void testUpdateMethods() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict, Optional\n" + + "class Movie(TypedDict):\n" + + " name: str\n" + + " year: Optional[int]\n" + + "class Horror(Movie, total=False):\n" + + " based_on_book: bool\n" + + "m = Horror(name='Alien', year=1979)\n" + + "d={'name':'Garden State', 'year':2004}\n" + + "m.update(d)\n" + + "m.update({'name':'Garden State', 'year':'2004', 'based_on': 'book'})\n" + + "m.update(name=1984, year=1984, based_on_book='yes')\n" + + "m.update([('name',1984), ('year',None)])")); + } + + public void testDocString() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict\n" + + "class Cinema(TypedDict):\n" + + " \"\"\"\n" + + " It's doc string\n" + + " \"\"\"")); + } + + public void testFieldOverwrittenByInheritance() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict\n" + + "class X(TypedDict):\n" + + " y: int\n" + + "class Y(TypedDict):\n" + + " y: str\n" + + "class XYZ(X, Y):\n" + + " y: bool")); + } + + public void testIncorrectTypedDictArguments() { + runWithLanguageLevel(LanguageLevel.PYTHON38, () -> + doTestByText("from typing import TypedDict\n" + + "c = TypedDict(\"c\", [1, 2, 3])")); + } + + @NotNull + @Override + protected Class getInspectionClass() { + return PyTypedDictInspection.class; + } +} \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java index e2a2857904fa..9ea637650d71 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java @@ -797,6 +797,26 @@ public class PyUnresolvedReferencesInspectionTest extends PyInspectionTestCase { ); } + // PY-36008 + public void testTypedDict() { + runWithLanguageLevel( + LanguageLevel.PYTHON38, + () -> doTestByText("from typing import TypedDict\n" + + "class X(TypedDict):\n" + + " x: str\n" + + "x = X(x='str')\n" + + "x.clear()\n" + + "x['x'] = 'rts'\n" + + "x.clea()\n" + + "x.x()\n" + + "x1: X = {'x1': 'str'}\n" + + "x1['x1'] = 'rts'\n" + + "x1.clear()\n" + + "x1.clea()\n" + + "x1.x()") + ); + } + @NotNull @Override protected Class getInspectionClass() {