PY-88892 django: support Choices enum label value

(cherry picked from commit 77d144a45d221c8328b10cd9012243acd6b92f0c)

IJ-MR-213651

GitOrigin-RevId: a37ba47a791f2c05d14028652e1b40f4b836b30c
This commit is contained in:
Morgan Bartholomew
2026-07-15 13:11:00 +00:00
committed by intellij-monorepo-bot
parent 1f712ab307
commit 98c83724bf
5 changed files with 83 additions and 97 deletions
@@ -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<PyType>?
companion object {
@JvmField
@@ -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<PyType> 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<PyType> 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.
* <p>
* 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<PyType> 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();
}
@@ -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
}
@@ -234,7 +234,7 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
class StrMixin(str, Enum):
C = "c"
s2: str = StrMixin.C.value
s3: int = <warning descr="Expected type 'int', got 'str' instead">StrMixin.C.value</warning>
s3: int = <warning descr="Expected type 'int', got 'Literal[\\"c\\"]' instead">StrMixin.C.value</warning>
# Empty str mixin should also infer str
class EmptyStrMixin(str, Enum):
@@ -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"])