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 extends PyInspection> 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 extends PyInspection> 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 extends PyInspection> getInspectionClass() {