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("""