From 128fbd3bb021fb88099f198e18b654e2a03d7abf Mon Sep 17 00:00:00 2001 From: Daniil Kalinin Date: Fri, 6 Mar 2026 17:18:24 +0100 Subject: [PATCH] PY-87344 Bind actual class types to Self when matching protocols The issue was highlighted by the following semantics: `Enum`'s metaclass is `EnumMeta` and implements iterable protocol as ``` def __iter__(self: type[_EnumMemberT]) -> Iterator[_EnumMemberT]: ... ``` Thus, we need to correctly bind the actual type of `Enum` implementation we're working with (say, `MyEnum`) to `_EnumMemberT` to be able to infer `Iterator[MyEnum]` and then get the correct type of e.g., `set(MyEnum)` (cherry picked from commit 53b26f963bdf8746c1f8fde5ee41aa37c415d920) (cherry picked from commit b0ce2298ae6a3e4706b55a56a6ba357ae9f6470d) IJ-CR-196674 GitOrigin-RevId: 49990c2629af7a8674a8db52eee1cb6751b9343a --- .../python/psi/types/PyTypeChecker.kt | 65 ++++++++++++------- .../com/jetbrains/python/Py3TypeTest.java | 32 +++++++++ .../Py3TypeCheckerInspectionTest.java | 39 +++++++++++ 3 files changed, 112 insertions(+), 24 deletions(-) diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt index 81ac564a6129..e1a31332195f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt @@ -667,9 +667,9 @@ object PyTypeChecker { } private fun matchProtocols(expected: PyClassType, actual: PyClassType, matchContext: MatchContext): Boolean { - val expectedSubstitutions = collectTypeSubstitutions(expected, matchContext.context) - val actualSubstitutions = collectTypeSubstitutions(actual, matchContext.context) - + val context = matchContext.context + val expectedSubstitutions = collectTypeSubstitutions(expected, context) + val actualSubstitutions = collectTypeSubstitutions(actual, context) // See https://typing.python.org/en/latest/spec/generics.html#use-in-protocols // > If a protocol uses Self in methods or attribute annotations, then a class Foo is assignable to the protocol @@ -678,20 +678,17 @@ object PyTypeChecker { // It should be equivalent to replacing Self in the protocol with the Foo class we're matching it with. val protocolSubstitutions = GenericSubstitutions() protocolSubstitutions.qualifierType = actual.toInstance() - val protocolContext = - MatchContext(matchContext.context, protocolSubstitutions, matchContext.reversedSubstitutions) - for (pair in inspectProtocolSubclass( - expected, actual, - matchContext.context - )) { + val protocolContext = MatchContext(context, protocolSubstitutions, matchContext.reversedSubstitutions) + + for (pair in inspectProtocolSubclass(expected, actual, context)) { val protocolMember = pair.first val subclassElementMembers = pair.second if (ContainerUtil.isEmpty(subclassElementMembers)) { return false } - val rawProtocolElementType = dropSelfInProtocolMember(protocolMember.type, matchContext.context) + val rawProtocolElementType = dropSelfInProtocolMember(expected, protocolMember.type, context) - val protocolElementType = substitute(rawProtocolElementType, expectedSubstitutions, matchContext.context) + val protocolElementType = substitute(rawProtocolElementType, expectedSubstitutions, context) val elementResult: Boolean = subclassElementMembers.any { subclassElementMember: PyTypeMember? -> if (protocolMember.isWritable && !subclassElementMember!!.isWritable) { @@ -706,8 +703,10 @@ object PyTypeChecker { return@any false } - var subclassElementType = dropSelfInProtocolMember(subclassElementMember.type, matchContext.context) - subclassElementType = substitute(subclassElementType, actualSubstitutions, matchContext.context) + var subclassElementType = substituteSelfInProtocolMember(actual, subclassElementMember.type, context) + subclassElementType = dropSelfInProtocolMember(expected, subclassElementType, context) + subclassElementType = substitute(subclassElementType, actualSubstitutions, context) + match(protocolElementType, subclassElementType, protocolContext).orElse(true)!! } @@ -732,6 +731,7 @@ object PyTypeChecker { } private fun match(expectedProtocol: PyClassType, actualModule: PyModuleType, matchContext: MatchContext): Boolean { + val context = matchContext.context val module = actualModule.module val moduleElements = @@ -744,11 +744,11 @@ object PyTypeChecker { .associateBy { it.name } val protocolElements = - inspectProtocolSubclass(expectedProtocol, expectedProtocol, matchContext.context) + inspectProtocolSubclass(expectedProtocol, expectedProtocol, context) if (protocolElements.size != moduleElements.size) return false - val substitutions = collectTypeSubstitutions(expectedProtocol, matchContext.context) + val substitutions = collectTypeSubstitutions(expectedProtocol, context) for (pair in protocolElements) { val pm = pair.first.element if (pm !is PsiNamedElement) { @@ -758,12 +758,9 @@ object PyTypeChecker { val moduleElement = moduleElements[name] if (moduleElement != null) { val expectedProtocolMemberType = - substitute( - dropSelfInProtocolMember( pair.first.type, matchContext.context), substitutions, - matchContext.context - ) - val actualModuleElementType = matchContext.context.getType(moduleElement) - if (!match(expectedProtocolMemberType, actualModuleElementType, matchContext.context)) { + substitute(dropSelfInProtocolMember(expectedProtocol, pair.first.type, context), substitutions, context) + val actualModuleElementType = context.getType(moduleElement) + if (!match(expectedProtocolMemberType, actualModuleElementType, context)) { return false } continue @@ -773,13 +770,30 @@ object PyTypeChecker { return true } - private fun dropSelfInProtocolMember(elementType: PyType?, context: TypeEvalContext): PyType? { + private fun dropSelfInProtocolMember(classType: PyClassType, elementType: PyType?, context: TypeEvalContext): PyType? { if (elementType is PyCallableType) { - return elementType.dropSelf(context) + if (PyUtil.isInitOrNewMethod(elementType.callable) || !classType.isDefinition) { + return elementType.dropSelf(context) + } } return elementType } + /** + * Binds TypeVars from the self parameter annotation of protocol member to [classType]. + */ + private fun substituteSelfInProtocolMember(classType: PyClassType, elementType: PyType?, context: TypeEvalContext): PyType? { + if (elementType !is PyCallableType) return elementType + val parameters = elementType.getParameters(context) + if (parameters.isNullOrEmpty() || !parameters.first().isSelf) return elementType + val selfParamType = parameters.first().getType(context) + if (selfParamType !is PyTypeVarType) return elementType + val selfSubstitutions = GenericSubstitutions() + match(selfParamType, classType, MatchContext(context, selfSubstitutions, false)) + if (selfSubstitutions.typeVars.isEmpty()) return elementType + return substitute(elementType, selfSubstitutions, context) as? PyCallableType ?: elementType + } + // https://typing.python.org/en/latest/spec/tuples.html#type-compatibility-rules private fun match( expected: PyTupleType, @@ -1433,7 +1447,7 @@ object PyTypeChecker { // A.a_method # type: Callable[[A], A] // B wasn't considered as the receiver type in the first place, instead of filtering it out during substitution // (see PyTypingTest.testMatchSelfUnionType) - return qualifierType.toStream() + val result = qualifierType.toStream() .map { qType: PyType? -> if (qType is PyInstantiableType<*>) { return@map if (selfScopeClassType.isDefinition) qType.toClass() else qType.toInstance() @@ -1442,6 +1456,9 @@ object PyTypeChecker { } .filter { normalizedQType: PyType? -> match(selfScopeClassType, normalizedQType, context) } .collect(PyTypeUtil.toUnion(qualifierType)) + // If no qualifier type matched Self's scope class, Self was probably inferred from a different context + // (e.g. protocol matching for a parameter type) and should be preserved as-is. + return result ?: selfType } override fun visitPyGenericType(genericType: PyCollectionType): PyType { diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index e2c7f45ffea2..7b0ef342015b 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -5182,6 +5182,38 @@ public class Py3TypeTest extends PyTestCase { """); } + // PY-87344 + public void testIteratorTypeCorrectlyInferredFromStrEnum() { + doTest("set[Variant]", """ + from enum import StrEnum + from typing import Self + + class Variant(StrEnum): + CREATED = "created" + + @classmethod + def values(cls) -> set[Self]: + return set(cls) + + expr = set(Variant) + """); + } + + // PY-87344 + public void testTypeOfSetOfStrEnumViaCls() { + doTest("set[Self@Variant]", """ + from enum import StrEnum + from typing import Self + + class Variant(StrEnum): + CREATED = "created" + + @classmethod + def values(cls): + expr = set(cls) + """); + } + private void doTest(final String expectedType, final String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class); diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index 9ac3a6948e09..31cc1f54c8d8 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -4316,5 +4316,44 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { assert_type(iter(MyFromStr()), Iterator[str]) """); } + + // PY-87344 + public void testTypeOfIteratorOfEnumTypeAndInstance() { + doTestByText(""" + from enum import Enum + from typing import Self + + class Color(Enum): + RED = "red" + + @classmethod + def all(cls) -> set[Self]: + return set(cls) + + def foo(self): + # __iter__ is defined in EnumMeta, thus, for definitions only + return set(self) + """); + } + + // PY-87344 + public void testTypeOfIteratorOfStrEnumTypeAndInstance() { + doTestByText(""" + from enum import StrEnum + from typing import Self + + + class Variant(StrEnum): + CREATED = "created" + + @classmethod + def values(cls) -> set[Self]: + return set(cls) + + def foo(self): + # StrEnum inherits str which inherits Iterable[str], thus, iterable for both instance and definition + return set(self) # OK + """); + } }