PY-85836 inlay-hints: show enum.auto() value hints

(cherry picked from commit 0143c413d5c2f5147f6af397ccc161bf7cf17da9)

IJ-MR-208400

GitOrigin-RevId: 496c792f126d362ee1499bd14b500f7e52f6bff9
This commit is contained in:
Morgan Bartholomew
2026-08-06 01:56:43 +00:00
committed by intellij-monorepo-bot
parent 4b858866f6
commit 7135abd19f
5 changed files with 359 additions and 0 deletions
@@ -131,6 +131,15 @@
nameKey="INLAY.numeric.promotion"
providerId="python.numeric.promotion.inlays"/>
<codeInsight.declarativeInlayProvider group="VALUES_GROUP"
bundle="messages.PyPsiBundle"
implementationClass="com.jetbrains.python.inlayHints.PyEnumAutoValueInlayHintsProvider"
isEnabledByDefault="true"
language="Python"
nameKey="INLAY.enum.auto.value"
descriptionKey="INLAY.enum.auto.value.description"
providerId="python.enum.auto.value.inlays"/>
<codeInsight.declarativeInlayProvider group="PARAMETERS_GROUP"
bundle="messages.PyPsiBundle"
implementationClass="com.jetbrains.python.inlayHints.PyTestParametrizeInlayHintsProvider"
@@ -89,6 +89,9 @@ INLAY.solved.type.parameters.function.description=Shows the type arguments infer
INLAY.pytest.parametrize.tuple=Pytest parametrize names
INLAY.positional.args.patterns=Positional arguments in class patterns
INLAY.numeric.promotion=Numeric type promotions
INLAY.enum.auto.value=Enum auto values
INLAY.enum.auto.value.description=Shows the value that <code>enum.auto()</code> generates for each enum member: consecutive integers for a plain Enum, successive powers of two for a Flag, and the lower-cased member name for a StrEnum.
INLAY.enum.auto.value.tooltip=Value generated by <code>enum.auto()</code>
### Annotators ###
ANN.deleting.none=Deleting None
@@ -0,0 +1,22 @@
# enum.auto() value hints show the generated member value
from enum import Enum, Flag, StrEnum, auto
# Consecutive integers for a plain Enum
class Priority(Enum):
LOW = auto()/*<# = 1#>*/
MEDIUM = 5
HIGH = auto()/*<# = 6#>*/
# Successive powers of two for a Flag
class Permission(Flag):
READ = auto()/*<# = 1#>*/
WRITE = auto()/*<# = 2#>*/
EXECUTE = auto()/*<# = 4#>*/
# Lower-cased member name for a StrEnum
class Color(StrEnum):
RED = auto()/*<# = 'red'#>*/
GREEN = auto()/*<# = 'green'#>*/
@@ -0,0 +1,212 @@
// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.inlayHints
import com.intellij.codeInsight.hints.declarative.HintFormat
import com.intellij.codeInsight.hints.declarative.InlayHintsCollector
import com.intellij.codeInsight.hints.declarative.InlayHintsProvider
import com.intellij.codeInsight.hints.declarative.InlayTreeSink
import com.intellij.codeInsight.hints.declarative.InlineInlayPosition
import com.intellij.codeInsight.hints.declarative.SharedBypassCollector
import com.intellij.openapi.editor.Editor
import com.intellij.psi.PsiElement
import com.intellij.psi.PsiFile
import com.intellij.psi.util.endOffset
import com.jetbrains.python.PyNames
import com.jetbrains.python.PyPsiBundle
import com.jetbrains.python.PyTokenTypes
import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil
import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider
import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider.EnumAttributeKind
import com.jetbrains.python.psi.PyCallExpression
import com.jetbrains.python.psi.PyClass
import com.jetbrains.python.psi.PyExpression
import com.jetbrains.python.psi.PyNumericLiteralExpression
import com.jetbrains.python.psi.PyPrefixExpression
import com.jetbrains.python.psi.PyReferenceExpression
import com.jetbrains.python.psi.PyStringLiteralExpression
import com.jetbrains.python.psi.PyTargetExpression
import com.jetbrains.python.psi.PyTupleExpression
import com.jetbrains.python.psi.impl.PyPsiUtils
import com.jetbrains.python.psi.resolve.PyResolveContext
import com.jetbrains.python.psi.types.TypeEvalContext
/**
* Shows the value generated by `enum.auto()` as an inlay hint next to the member declaration.
*
* The value is computed statically following the semantics of `Enum._generate_next_value_`:
* - for a plain `Enum`/`IntEnum`: the previous integer value plus one (starting from `1`);
* - for a `Flag`/`IntFlag`: the next power of two;
* - for a `StrEnum`: the lower-cased member name.
*
* No hint is shown when the value cannot be determined statically (a custom `_generate_next_value_`,
* a preceding member whose value is not a literal, etc.).
*
* For example:
* ```
* class E(Enum):
* A = auto() # = 1
* B = auto() # = 2
* C = auto() # = 3
* ```
*/
class PyEnumAutoValueInlayHintsProvider : InlayHintsProvider {
override fun createCollector(file: PsiFile, editor: Editor): InlayHintsCollector = Collector()
private class Collector : SharedBypassCollector {
/** Per-file cache: auto member -> rendered hint text, computed lazily for the whole enum class at once. */
private val cache = HashMap<PyClass, Map<PyTargetExpression, String>>()
override fun collectFromElement(element: PsiElement, sink: InlayTreeSink) {
if (element !is PyTargetExpression) return
// Cheap pre-filter: only members assigned an `auto()` call are interesting.
val assignedValue = element.findAssignedValue() as? PyCallExpression ?: return
val enumClass = ScopeUtil.getScopeOwner(element) as? PyClass ?: return
val hint = cache.getOrPut(enumClass) { computeAutoValues(enumClass) }[element] ?: return
sink.addPresentation(
position = InlineInlayPosition(assignedValue.endOffset, true),
hintFormat = HintFormat.default,
tooltip = PyPsiBundle.message("INLAY.enum.auto.value.tooltip"),
) {
text(hint)
}
}
private fun computeAutoValues(enumClass: PyClass): Map<PyTargetExpression, String> {
val context = TypeEvalContext.codeAnalysis(enumClass.project, enumClass.containingFile)
if (!PyStdlibTypeProvider.isCustomEnum(enumClass, context)) return emptyMap()
val kind = autoValueKind(enumClass, context) ?: return emptyMap()
val result = HashMap<PyTargetExpression, String>()
val lastValues = ArrayList<MemberValue>()
for (attribute in enumClass.classAttributes) {
val info = PyStdlibTypeProvider.getEnumAttributeInfo(enumClass, attribute, context)
if (info == null || info.attributeKind == EnumAttributeKind.NONMEMBER) continue
val name = attribute.name
val value = attribute.findAssignedValue()
if (name != null && isEnumAutoCall(value, context)) {
when (kind) {
AutoKind.STR -> {
result[attribute] = " = '${name.lowercase()}'"
lastValues.add(MemberValue.NonInt)
}
AutoKind.INT -> {
val next = nextIntValue(lastValues)
if (next != null) {
result[attribute] = " = $next"
lastValues.add(MemberValue.Int(next))
}
else {
lastValues.add(MemberValue.Unknown)
}
}
AutoKind.FLAG -> {
val next = nextFlagValue(lastValues)
if (next != null) {
result[attribute] = " = $next"
lastValues.add(MemberValue.Int(next))
}
else {
lastValues.add(MemberValue.Unknown)
}
}
}
}
else {
lastValues.add(evaluate(value))
}
}
return result
}
private fun autoValueKind(enumClass: PyClass, context: TypeEvalContext): AutoKind? {
// `Flag`/`IntFlag` generate powers of two; they don't override `_generate_next_value_` in typeshed.
if (enumClass.isSubclass(PyNames.TYPE_ENUM_FLAG, context)) return AutoKind.FLAG
val generateNextValue = enumClass.findMethodByName(GENERATE_NEXT_VALUE, true, context) ?: return null
return when (generateNextValue.containingClass?.qualifiedName) {
PyNames.TYPE_ENUM -> AutoKind.INT
STR_ENUM -> AutoKind.STR
else -> null // a custom `_generate_next_value_`: cannot compute statically
}
}
private fun isEnumAutoCall(value: PyExpression?, context: TypeEvalContext): Boolean {
val call = value as? PyCallExpression ?: return false
val callee = call.callee as? PyReferenceExpression ?: return false
val resolveContext = PyResolveContext.defaultContext(context)
return callee.getReference(resolveContext).multiResolve(false).any {
val resolved = it.element
resolved is PyClass && PyNames.TYPE_ENUM_AUTO == resolved.qualifiedName
}
}
/** Mirrors the default `Enum._generate_next_value_`: most recent integer value plus one, or `1` to start. */
private fun nextIntValue(lastValues: List<MemberValue>): Long? {
for (value in lastValues.asReversed()) {
when (value) {
is MemberValue.Int -> return value.value + 1
MemberValue.NonInt -> continue // non-integer values raise `TypeError` at runtime and are skipped
MemberValue.Unknown -> return null // an unevaluable value: bail to avoid showing a wrong hint
}
}
return 1
}
/** Mirrors `Flag._generate_next_value_`: the next power of two above the highest value seen so far. */
private fun nextFlagValue(lastValues: List<MemberValue>): Long? {
if (lastValues.isEmpty()) return 1
var max = -1L
for (value in lastValues) {
if (value !is MemberValue.Int) return null
if (value.value > max) max = value.value
}
if (max < 0) return null
val bitLength = java.lang.Long.SIZE - java.lang.Long.numberOfLeadingZeros(max)
if (bitLength >= java.lang.Long.SIZE - 1) return null // guard against overflow for unrealistically large flags
return 1L shl bitLength
}
private fun evaluate(value: PyExpression?): MemberValue {
return when (val expression = PyPsiUtils.flattenParens(value)) {
is PyNumericLiteralExpression ->
if (expression.isIntegerLiteral) expression.longValue?.let { MemberValue.Int(it) } ?: MemberValue.Unknown
else MemberValue.NonInt
is PyPrefixExpression -> {
if (expression.operator == PyTokenTypes.MINUS) {
val operand = PyPsiUtils.flattenParens(expression.operand)
if (operand is PyNumericLiteralExpression && operand.isIntegerLiteral) {
operand.longValue?.let { MemberValue.Int(-it) } ?: MemberValue.Unknown
}
else MemberValue.Unknown
}
else MemberValue.Unknown
}
is PyStringLiteralExpression, is PyTupleExpression -> MemberValue.NonInt
else -> MemberValue.Unknown
}
}
}
private enum class AutoKind { INT, FLAG, STR }
private sealed interface MemberValue {
/** A statically known integer value. */
data class Int(val value: Long) : MemberValue
/** A statically known non-integer value (string, tuple, ...). */
object NonInt : MemberValue
/** A value that could not be evaluated statically. */
object Unknown : MemberValue
}
private companion object {
private const val GENERATE_NEXT_VALUE = "_generate_next_value_"
private const val STR_ENUM = "enum.StrEnum"
}
}
@@ -0,0 +1,113 @@
// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.inlayHints
import com.intellij.idea.TestFor
import com.intellij.testFramework.LightProjectDescriptor
import com.intellij.testFramework.utils.inlays.declarative.DeclarativeInlayHintsProviderTestCase
import com.jetbrains.python.fixtures.PyLightProjectDescriptor
import com.jetbrains.python.psi.LanguageLevel
@TestFor(issues = ["PY-85836"])
class PyEnumAutoValueInlayHintsProviderTest : DeclarativeInlayHintsProviderTestCase() {
fun `test consecutive auto values`() = doTest("""
from enum import Enum, auto
class E(Enum):
A = auto()/*<# = 1 #>*/
B = auto()/*<# = 2 #>*/
C = auto()/*<# = 3 #>*/
""")
fun `test qualified auto`() = doTest("""
import enum
class E(enum.Enum):
A = enum.auto()/*<# = 1 #>*/
B = enum.auto()/*<# = 2 #>*/
""")
fun `test int enum`() = doTest("""
from enum import IntEnum, auto
class E(IntEnum):
A = auto()/*<# = 1 #>*/
B = auto()/*<# = 2 #>*/
""")
fun `test auto continues from explicit int value`() = doTest("""
from enum import Enum, auto
class E(Enum):
A = auto()/*<# = 1 #>*/
B = 5
C = auto()/*<# = 6 #>*/
D = auto()/*<# = 7 #>*/
""")
fun `test auto skips non-int explicit value`() = doTest("""
from enum import Enum, auto
class E(Enum):
A = auto()/*<# = 1 #>*/
B = "x"
C = auto()/*<# = 2 #>*/
""")
fun `test flag uses powers of two`() = doTest("""
from enum import Flag, auto
class E(Flag):
A = auto()/*<# = 1 #>*/
B = auto()/*<# = 2 #>*/
C = auto()/*<# = 4 #>*/
D = auto()/*<# = 8 #>*/
""")
fun `test int flag uses powers of two`() = doTest("""
from enum import IntFlag, auto
class E(IntFlag):
READ = auto()/*<# = 1 #>*/
WRITE = auto()/*<# = 2 #>*/
EXECUTE = auto()/*<# = 4 #>*/
""")
fun `test str enum uses lower-cased name`() = doTest("""
from enum import StrEnum, auto
class Color(StrEnum):
RED = auto()/*<# = 'red' #>*/
Green = auto()/*<# = 'green' #>*/
""")
fun `test no hint for custom generate next value`() = doTest("""
from enum import Enum, auto
class E(Enum):
@staticmethod
def _generate_next_value_(name, start, count, last_values):
return name
A = auto()
B = auto()
""", verifyHintsPresence = false)
fun `test no hint after unevaluable value`() = doTest("""
from enum import Enum, auto
def factory(): ...
class E(Enum):
A = factory()
B = auto()
""", verifyHintsPresence = false)
fun `test no hint outside enum`() = doTest("""
from enum import auto
x = auto()
""", verifyHintsPresence = false)
private fun doTest(text: String, verifyHintsPresence: Boolean = true) {
doTestProvider(
"test.py",
text.trimIndent(),
PyEnumAutoValueInlayHintsProvider(),
emptyMap(),
verifyHintsPresence = verifyHintsPresence,
testMode = ProviderTestMode.SIMPLE
)
}
override fun getProjectDescriptor(): LightProjectDescriptor {
return PyLightProjectDescriptor(LanguageLevel.getLatest())
}
}