From 405dafb8fe810a01cbe9e4876cdef352b9a4118a Mon Sep 17 00:00:00 2001 From: Mikhail Golubev Date: Fri, 19 Dec 2025 16:41:51 +0200 Subject: [PATCH] PY-86463 Handle generic protocol inheritance during matching When we have a protocol hierarchy, such as the following ```python class P[T](Protocol): def method(self, x: T) -> T: pass class P2[T](P[T], Protocol): pass def expects_generic(x: P2[str]): ... ``` substitutions extracted from the expected type `P2[str]` of `x` are: ``` T@P2 -> str T@P -> T@P2 ``` Doing `expectedSubstitutions.withOnlySolvedSubstitutions(context)) then throws away `T@P -> T@P2` mapping, so we lose that `T@P` is supposed to be str. It means that during matching of `P.method` we match its unbound T@P with any actual type, causing false positives. On the other hand, if `expects_generic` was supposed to bind its own parameter, e.g. as if in ```python def expects_generic[T](x: P2[T]) -> T: ... ``` we cannot restore the actual type of `T@expects_generic` in the ```java if (expected instanceof PyCollectionType) ... ``` branch at the end, because we try to match `P2[T@expects_generic]` (expected type) with `P2[T@P2]` with substitution `T@P -> str` (for example). Again, because the mapping `T@P -> T@P2` is lost, we end up only binding `T@expects_generic -> T@P2`. Now we substitute everything into the "raw" protocol member type before matching it with the implementation, so we map `T@expects_generic` right away when matching `P.method` with implementation, and then we only need to insert this type into the expected `P2[T@expects_generic]`. GitOrigin-RevId: 88c91add1cca181db14f3276f944a84dffcd6517 --- .../python/psi/types/PyTypeChecker.java | 43 +++++-------------- .../com/jetbrains/python/PyTypingTest.java | 23 ++++++++++ .../Py3TypeCheckerInspectionTest.java | 23 ++++++++++ 3 files changed, 56 insertions(+), 33 deletions(-) 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 1cd0e68fab1c..09a60b0b8f50 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 @@ -636,9 +636,8 @@ public final class PyTypeChecker { final PyType rawProtocolElementType = dropSelfIfNeeded(expected, protocolMember.getType(), matchContext.context); - final GenericSubstitutions concreteExpectedSubstitutions = expectedSubstitutions.withOnlySolvedSubstitutions(matchContext.context); - final PyType protocolElementType = substitute(rawProtocolElementType, concreteExpectedSubstitutions, matchContext.context); + final PyType protocolElementType = substitute(rawProtocolElementType, expectedSubstitutions, matchContext.context); final boolean elementResult = ContainerUtil.exists(subclassElementMembers, subclassElementMember -> { if (protocolMember.isWritable() && !subclassElementMember.isWritable()) { return false; @@ -664,14 +663,15 @@ public final class PyTypeChecker { } - if (expected instanceof PyCollectionType) { - PyCollectionType genericSuperClass = findGenericDefinitionType(expected.getPyClass(), matchContext.context); - if (genericSuperClass != null) { - PyCollectionType concreteSuperClass = - (PyCollectionType)substitute(genericSuperClass, protocolContext.mySubstitutions, protocolContext.context); - assert concreteSuperClass != null; - return matchGenericClassesParameterWise((PyCollectionType)expected, concreteSuperClass, matchContext); - } + if (expected instanceof PyCollectionType genericExpected && hasGenerics(expected, matchContext.context)) { + PyCollectionType concreteExpected = + (PyCollectionType)substitute(expected, protocolContext.mySubstitutions, protocolContext.context); + assert concreteExpected != null; + // This match is supposed to succeed since all protocol members were compatible. + // Effectively, we're just copying substitutions for the terminal type parameters e.g. + // extracting: T@func -> str + // from: Iterable[T@func] <- Iterable[str] + return matchGenericClassesParameterWise(genericExpected, concreteExpected, matchContext); } return true; @@ -2021,29 +2021,6 @@ public final class PyTypeChecker { return qualifierType; } - private @NotNull GenericSubstitutions withOnlySolvedSubstitutions(@NotNull TypeEvalContext context) { - GenericSubstitutions concreteExpectedSubs = new GenericSubstitutions(); - for (Map.Entry> e : this.typeVars.entrySet()) { - PyType mapped = Ref.deref(e.getValue()); - if (mapped != null && !hasGenerics(mapped, context)) { - concreteExpectedSubs.typeVars.put(e.getKey(), Ref.create(mapped)); - } - } - for (Map.Entry e : this.paramSpecs.entrySet()) { - PyCallableParameterVariadicType mapped = e.getValue(); - if (mapped != null && !hasGenerics(mapped, context)) { - concreteExpectedSubs.paramSpecs.put(e.getKey(), mapped); - } - } - for (Map.Entry e : this.typeVarTuples.entrySet()) { - PyPositionalVariadicType mapped = e.getValue(); - if (mapped != null && !hasGenerics(mapped, context)) { - concreteExpectedSubs.typeVarTuples.put(e.getKey(), mapped); - } - } - return concreteExpectedSubs; - } - @Override public String toString() { return "GenericSubstitutions{" + diff --git a/python/testSrc/com/jetbrains/python/PyTypingTest.java b/python/testSrc/com/jetbrains/python/PyTypingTest.java index 24946f5ccbb7..16d3a3717301 100644 --- a/python/testSrc/com/jetbrains/python/PyTypingTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypingTest.java @@ -6836,6 +6836,29 @@ public class PyTypingTest extends PyTestCase { """); } + // PY-86463 + public void testMatchingWithInheritedGenericProtocol() { + doTest("int", """ + from typing import Protocol + + class P[T](Protocol): + def method(self, x: T) -> T: + pass + + class P2[T](P[T], Protocol): + pass + + class Impl: + def method(self, x: int) -> int: + return 42 + + def expects_generic[T](x: P2[T]) -> T: + return x.method() + + expr = expects_generic(Impl()) + """); + } + private void doTestNoInjectedText(@NotNull String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); final InjectedLanguageManager languageManager = InjectedLanguageManager.getInstance(myFixture.getProject()); diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index 3e8912aefbea..7fef125a4c9b 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -3398,6 +3398,29 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { """); } + // PY-86463 + public void testInheritedGenericProtocol() { + doTestByText(""" + from typing import Protocol, overload + + class P[T](Protocol): + def method(self, x: T) -> T: + pass + + class P2[T](P[T], Protocol): + pass + + class Impl: + def method(self, x: int) -> int: + ... + + def expects_P2_str(x: P2[str]): + pass + + expr = expects_P2_str(Impl()) + """); + } + // PY-76822 public void testExplicitAnyInConcreteType() { doTestByText("""