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
This commit is contained in:
Mikhail Golubev
2025-12-22 19:38:56 +00:00
committed by intellij-monorepo-bot
parent 453bfc627f
commit 405dafb8fe
3 changed files with 56 additions and 33 deletions
@@ -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<PyTypeVarType, Ref<PyType>> 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<PyParamSpecType, PyCallableParameterVariadicType> e : this.paramSpecs.entrySet()) {
PyCallableParameterVariadicType mapped = e.getValue();
if (mapped != null && !hasGenerics(mapped, context)) {
concreteExpectedSubs.paramSpecs.put(e.getKey(), mapped);
}
}
for (Map.Entry<PyTypeVarTupleType, PyPositionalVariadicType> 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{" +
@@ -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());
@@ -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(<warning descr="Expected type 'P2[str]', got 'Impl' instead">Impl()</warning>)
""");
}
// PY-76822
public void testExplicitAnyInConcreteType() {
doTestByText("""