diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyEnumMemberDeclarationProvider.kt b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyEnumMemberDeclarationProvider.kt index 859088223240..84f66e9aa075 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyEnumMemberDeclarationProvider.kt +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyEnumMemberDeclarationProvider.kt @@ -1,12 +1,24 @@ package com.jetbrains.python.psi.types import com.intellij.openapi.extensions.ExtensionPointName +import com.intellij.openapi.util.Ref import com.jetbrains.python.psi.PyClass +import com.jetbrains.python.psi.PyExpression import org.jetbrains.annotations.ApiStatus +/** + * Lets framework-specific enum implementations model how their metaclass/constructor transforms a member declaration + * into the actual member value. For example, Django's `ChoicesType.__new__` treats a trailing `str`/`Promise` element + * of a tuple declaration as the member's label and drops it, so `A = "x", "label"` has value type `str`, not + * `tuple[str, str]`. + */ @ApiStatus.Internal interface PyEnumMemberDeclarationProvider { - fun allowsTupleMemberDeclaration(enumClass: PyClass, context: TypeEvalContext): Boolean + /** + * Returns the member value type produced from [valueExpression] for [enumClass], wrapped in a [Ref] (which may hold + * `null` for an unknown value type), or `null` if this provider does not transform declarations for [enumClass]. + */ + fun getMemberValueType(enumClass: PyClass, valueExpression: PyExpression, context: TypeEvalContext): Ref? companion object { @JvmField diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java index 6488fbb6c453..ba959a8f02d3 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java @@ -143,29 +143,22 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase { return Ref.create(PyBuiltinCache.getInstance(referenceTarget).getStrType()); } else if ("enum.IntEnum.value".equals(name) && anchor instanceof PyReferenceExpression) { - return Ref.create(PyBuiltinCache.getInstance(referenceTarget).getIntType()); + PyType memberValueType = getEnumValueTypeFromQualifier(anchor, context); + return Ref.create(memberValueType != null ? memberValueType : PyBuiltinCache.getInstance(referenceTarget).getIntType()); } else if ("enum.StrEnum.value".equals(name) && anchor instanceof PyReferenceExpression) { - return Ref.create(PyBuiltinCache.getInstance(referenceTarget).getStrType()); + PyType memberValueType = getEnumValueTypeFromQualifier(anchor, context); + return Ref.create(memberValueType != null ? memberValueType : PyBuiltinCache.getInstance(referenceTarget).getStrType()); } else if ((PyNames.TYPE_ENUM_FLAG + ".value").equals(name) && anchor instanceof PyReferenceExpression) { - return Ref.create(PyBuiltinCache.getInstance(referenceTarget).getIntType()); + PyType memberValueType = getEnumValueTypeFromQualifier(anchor, context); + return Ref.create(memberValueType != null ? memberValueType : PyBuiltinCache.getInstance(referenceTarget).getIntType()); } - else if ((PyNames.TYPE_ENUM + ".value").equals(name) && - anchor instanceof PyReferenceExpression anchorExpr && context.maySwitchToAST(anchor)) { - final PyExpression qualifier = anchorExpr.getQualifier(); + else if ((PyNames.TYPE_ENUM + ".value").equals(name) && anchor instanceof PyReferenceExpression) { // An enum value is retrieved programmatically, e.g. MyEnum[name].value, or just type-hinted - if (qualifier != null) { - PyClassType enumType = as(context.getType(qualifier), PyClassType.class); - if (enumType != null) { - PyClass enumClass = enumType.getPyClass(); - if (isCustomEnum(enumClass, context)) { - // If the qualifier is a specific member (e.g. MyEnum.B), use that member's value type; - // otherwise fall back to the union of all members' value types. - String memberName = enumType instanceof PyLiteralType literalType ? literalType.getEnumMemberName() : null; - return Ref.create(getEnumValueType(enumClass, memberName, context)); - } - } + PyType memberValueType = getEnumValueTypeFromQualifier(anchor, context); + if (memberValueType != null) { + return Ref.create(memberValueType); } } else if ("enum.EnumMeta.__members__".equals(name)) { @@ -179,6 +172,21 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase { return null; } + /** + * The value type of the enum member (or enum instance) that qualifies a {@code .value} access at {@code anchor}, or + * {@code null} when the AST is unavailable or the qualifier is not a custom enum. If the qualifier is a specific + * member (e.g. {@code MyEnum.B}), its own value type is used; otherwise the union of all members' value types. + */ + private static @Nullable PyType getEnumValueTypeFromQualifier(@Nullable PsiElement anchor, @NotNull TypeEvalContext context) { + if (!(anchor instanceof PyReferenceExpression anchorExpr) || !context.maySwitchToAST(anchor)) return null; + final PyExpression qualifier = anchorExpr.getQualifier(); + if (qualifier == null) return null; + PyClassType enumType = as(context.getType(qualifier), PyClassType.class); + if (enumType == null || !isCustomEnum(enumType.getPyClass(), context)) return null; + String memberName = enumType instanceof PyLiteralType literalType ? literalType.getEnumMemberName() : null; + return getEnumValueType(enumType.getPyClass(), memberName, context); + } + // Returns the type of enum attribute value transformed by 'EnumType' metaclass or null, if the attribute value is not transformed private static @Nullable Ref getTransformedEnumAttributeType(@NotNull PsiElement element, @NotNull TypeEvalContext context) { return RecursionManager.doPreventingRecursion(element, false, () -> getTransformedEnumAttributeTypeImpl(element, context)); @@ -255,6 +263,16 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase { PyExpression value = targetExpression.findAssignedValue(); if (value == null) return null; + // Framework enums (e.g. Django Choices) transform the declaration before constructing the member: their + // metaclass/constructor consumes extra tuple elements such as a trailing label. Let a provider model that so + // the resulting member value type is consistent across the inspection, the `value` attribute, and completion. + for (var provider : PyEnumMemberDeclarationProvider.EP_NAME.getExtensionList()) { + Ref transformed = provider.getMemberValueType(enumClass, value, context); + if (transformed != null) { + return getEnumAttributeInfo(enumClass, transformed.get(), context); + } + } + var type = context.getType(value); return getEnumAttributeInfo(enumClass, type, context); } @@ -407,34 +425,12 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase { } /** - * Whether {@code enumClass} accepts tuple member declarations whose first element is the member value. - *

- * A custom {@code __new__}/{@code __init__} receives the tuple elements as arguments. Framework-specific enum - * implementations can support other declaration transformations through {@link PyEnumMemberDeclarationProvider}. - * Plain stdlib enums do neither: their tuple is passed to {@code str(...)}/{@code int(...)} and is not a valid value. + * The primitive value type contributed by an enum's data-type mixin (e.g. {@code int} for {@code IntEnum}/{@code IntFlag}, + * {@code str} for {@code StrEnum} or a {@code str}-mixed enum), or {@code null} if the enum has no such mixin. Such a + * mixin constructs the member via {@code DataType(*value)}, so a (possibly transformed) tuple value collapses to a scalar. */ @ApiStatus.Internal - public static boolean allowsTupleEnumMemberDeclaration(@NotNull PyClass enumClass, @NotNull TypeEvalContext context) { - if (ContainerUtil.exists(PyEnumMemberDeclarationProvider.EP_NAME.getExtensionList(), - provider -> provider.allowsTupleMemberDeclaration(enumClass, context))) return true; - - PyFunction constructor = enumClass.findInitOrNew(true, context); - if (constructor == null) return false; - PyClass owner = constructor.getContainingClass(); - if (owner == null) return false; - // Mixed-in primitive data type (str/int/bytes/float/object): not a custom enum constructor. - if (PyBuiltinCache.getInstance(enumClass).isBuiltin(owner)) return false; - // Builtin `enum` module classes (Enum/IntEnum/StrEnum/Flag/...): not custom. - String qualifiedName = owner.getQualifiedName(); - return qualifiedName == null || !qualifiedName.startsWith("enum."); - } - - /** - * returns the type of the {@code value} attribute of an enum member, or the union of all members' value types - */ - private static @Nullable PyType getEnumValueType(@NotNull PyClass enumClass, - @Nullable String memberName, - @NotNull TypeEvalContext context) { + public static @Nullable PyType getEnumMixinValueType(@NotNull PyClass enumClass, @NotNull TypeEvalContext context) { PyBuiltinCache cache = PyBuiltinCache.getInstance(enumClass); if (enumClass.isSubclass("enum.IntEnum", context) || @@ -445,7 +441,6 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase { if (enumClass.isSubclass("enum.StrEnum", context)) { return cache.getStrType(); } - if (enumClass.isSubclass(PyNames.FQN.STR, context)) { return cache.getStrType(); } @@ -458,18 +453,35 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase { if (enumClass.isSubclass(PyNames.FQN.FLOAT, context)) { return cache.getFloatType(); } + return null; + } + /** + * returns the type of the {@code value} attribute of an enum member, or the union of all members' value types + */ + private static @Nullable PyType getEnumValueType(@NotNull PyClass enumClass, + @Nullable String memberName, + @NotNull TypeEvalContext context) { // Infer from the MEMBERS' assigned values (not non-members like helpers/descriptors). List memberValueTypes = new ArrayList<>(); for (PyTargetExpression targetExpr : enumClass.getClassAttributes()) { EnumAttributeInfo attributeInfo = getEnumAttributeInfo(enumClass, targetExpr, context); if (attributeInfo != null && attributeInfo.attributeKind == EnumAttributeKind.MEMBER) { + // For a specific member, prefer its own (possibly transformed) value type so literals are preserved, + // e.g. `MyIntChoices.OK.value` is `Literal[1]`, not the widened `int` from the data-type mixin. if (Objects.equals(memberName, targetExpr.getName())) { return attributeInfo.assignedValueType; } memberValueTypes.add(attributeInfo.assignedValueType); } } + + // The enum-wide value type (no specific member, or the member is inherited) is the data-type mixin's type when + // present, otherwise the union of all members' value types. + PyType mixinValueType = getEnumMixinValueType(enumClass, context); + if (mixinValueType != null) { + return mixinValueType; + } if (memberValueTypes.isEmpty()) { return PyAnyType.getUnknown(); } diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.kt index 2dc95ff94b3c..7c91b22455d1 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.kt @@ -18,7 +18,6 @@ import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil.getScopeOwner import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider.getEnumAttributeInfo import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider.getEnumValueType -import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider.allowsTupleEnumMemberDeclaration import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider.isCustomEnum import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.Companion.coroutineOrGeneratorElementType @@ -494,34 +493,15 @@ open class PyTypeCheckerInspection : PyInspection() { if (info == null || info.attributeKind != PyStdlibTypeProvider.EnumAttributeKind.MEMBER) return val expected = getEnumValueType(scopeOwner, myTypeEvalContext) + // `assignedValueType` is the member value type produced by the enum's metaclass/constructor. For enums that + // transform the declaration (e.g. Django's `ChoicesType.__new__` drops a trailing label), the type provider + // already stripped the extra elements, so a plain match here is correct for both stdlib and framework enums. val actual = info.assignedValueType - // The whole assigned value matches the member value type: fine (also covers plain enums whose value type - // is itself inferred as a tuple, e.g. `class C(Enum): A = 1, 2`). - if (match(expected, actual, myTypeEvalContext)) return - - // Enums with a custom __new__/__init__ forward a tuple member `(value, label, ...)` to that constructor. - // Framework-specific enum implementations may transform the declaration before constructing the member. - // By convention the first element is the actual member value, so validate only it against the scalar value - // type; the remaining elements are extra metadata or constructor arguments (label, encoding, ...). - // Plain stdlib enums (StrEnum/IntEnum/...) have no such constructor: their tuple is not a valid value and - // falls through to the whole-value mismatch report below. - if (allowsTupleEnumMemberDeclaration(scopeOwner, myTypeEvalContext)) { - val tuple = flattenParens(assignedValue) as? PyTupleExpression - if (tuple != null) { - val valueExpression = tuple.elements.firstOrNull() ?: return - val valueType = myTypeEvalContext.getType(valueExpression) - if (!match(expected, valueType, myTypeEvalContext)) { - registerTypeMismatch(PyTypeCheckerSuppressionCode.BAD_ASSIGNMENT, valueExpression, expected, valueType, - typeMismatchMessage(expected, valueType, valueExpression), - ProblemHighlightType.GENERIC_ERROR_OR_WARNING) - } - return - } + if (!match(expected, actual, myTypeEvalContext)) { + registerTypeMismatch(PyTypeCheckerSuppressionCode.BAD_ASSIGNMENT, assignedValue, expected, actual, + typeMismatchMessage(expected, actual, assignedValue), + ProblemHighlightType.GENERIC_ERROR_OR_WARNING) } - - registerTypeMismatch(PyTypeCheckerSuppressionCode.BAD_ASSIGNMENT, assignedValue, expected, actual, - typeMismatchMessage(expected, actual, assignedValue), - ProblemHighlightType.GENERIC_ERROR_OR_WARNING) return } diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index 53df75baa994..3c9487196843 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -234,7 +234,7 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { class StrMixin(str, Enum): C = "c" s2: str = StrMixin.C.value - s3: int = StrMixin.C.value + s3: int = StrMixin.C.value # Empty str mixin should also infer str class EmptyStrMixin(str, Enum): diff --git a/python/testSrc/com/jetbrains/python/types/PyEnumTypeTest.kt b/python/testSrc/com/jetbrains/python/types/PyEnumTypeTest.kt index 6b233d493090..45f426c4779c 100644 --- a/python/testSrc/com/jetbrains/python/types/PyEnumTypeTest.kt +++ b/python/testSrc/com/jetbrains/python/types/PyEnumTypeTest.kt @@ -672,39 +672,21 @@ class PyEnumTypeTest : PyCodeInsightTestCase() { @Test @TestFor(issues = ["PY-88892"]) - fun `enum member declared as a value-label tuple validates only the value element when supported`() = test(""" - from enum import Enum, EnumType, StrEnum, IntEnum - from choices import Choices + fun `plain enum member declared as a tuple is validated as the whole value`() = test(""" + from enum import Enum, StrEnum, IntEnum + # Plain stdlib enums have no metaclass/constructor that consumes extra elements: the whole tuple is the value, + # so a str/int-based enum rejects it (runtime TypeError). Only framework enums (e.g. django Choices) relax this + # via PyEnumMemberDeclarationProvider; see DjangoEnumTypeTest. class PlainStr(StrEnum): - # plain StrEnum passes the tuple to str(...): not a valid value (runtime TypeError) A = "A", "A" # WARNING Expected type 'str', got 'tuple[Literal["A"], Literal["A"]]' instead class PlainInt(IntEnum): - # same for the plain int-based enum (runtime TypeError) B = 1, "label" # WARNING Expected type 'int', got 'tuple[Literal[1], Literal["label"]]' instead class PlainEnum(Enum): C = 1, 2 # OK: value type is inferred as the tuple itself - - # Choices has a custom __init__ that consumes the extra element (django TextChoices form), - # so the tuple is forwarded to it and only the first element is the member value. - class MyChoices(Choices): - OK = "x", "label" # OK: the value element 'x' is a str - BAD = 1, "label" # WARNING Expected type 'str', got 'Literal[1]' instead - - class PassthroughEnumType(EnumType): ... - - class CustomMetaclassStrEnum(StrEnum, metaclass=PassthroughEnumType): - VALUE = "x", "label" # WARNING Expected type 'str', got 'tuple[Literal["x"], Literal["label"]]' instead - """, - "choices.py" to """ - from enum import Enum - - class Choices(str, Enum): - def __init__(self, value, label=""): ... - """, - ) + """) @Test @TestFor(issues = ["PY-87344"])