PY-85123 Fix protocol substitution when matching it against a concrete class

GitOrigin-RevId: 46438cc4d035c75fb51217ba8eef927144d84a3b
This commit is contained in:
evgeny.bovykin
2025-12-22 19:38:56 +00:00
committed by intellij-monorepo-bot
parent e3ec62ebfc
commit 453bfc627f
6 changed files with 216 additions and 51 deletions
@@ -49,6 +49,12 @@ public interface PyCallableType extends PyType {
return null;
}
@ApiStatus.Experimental
@NotNull
default PyCallableType dropSelf(@NotNull TypeEvalContext context) {
return this;
}
default @Nullable PyFunction.Modifier getModifier() {
return null;
}
@@ -66,10 +72,4 @@ public interface PyCallableType extends PyType {
default <T> T acceptTypeVisitor(@NotNull PyTypeVisitor<T> visitor) {
return visitor.visitPyCallableType(this);
}
@ApiStatus.Experimental
@NotNull
default PyCallableType dropSelf(@NotNull TypeEvalContext context) {
return this;
}
}
@@ -23,6 +23,7 @@ import com.jetbrains.python.ast.impl.ParamHelperCore;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.types.PyCallableParameter;
import com.jetbrains.python.psi.types.PyCallableParameterImpl;
import com.jetbrains.python.psi.types.PyCallableType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import org.jetbrains.annotations.Contract;
import org.jetbrains.annotations.NotNull;
@@ -178,6 +179,23 @@ public final class ParamHelper {
return false;
}
public static boolean isSelfArgsKwargsCallable(@NotNull PyCallableType type, @NotNull TypeEvalContext context) {
final List<PyCallableParameter> parameters = type.getParameters(context);
return parameters != null
&& parameters.size() == 3 &&
parameters.get(0).isSelf() &&
parameters.get(1).isPositionalContainer() &&
parameters.get(2).isKeywordContainer();
}
public static boolean isArgsKwargsCallable(@NotNull PyCallableType type, @NotNull TypeEvalContext context) {
final List<PyCallableParameter> parameters = type.getParameters(context);
return parameters != null
&& parameters.size() == 2 &&
parameters.get(0).isPositionalContainer() &&
parameters.get(1).isKeywordContainer();
}
public interface ParamWalker {
/**
* Is called when a tuple parameter is encountered, before visiting any parameters nested in it.
@@ -147,6 +147,21 @@ public class PyCallableTypeImpl implements PyCallableType {
return myCallable;
}
@Override
public @NotNull PyCallableType dropSelf(@NotNull TypeEvalContext context) {
final List<PyCallableParameter> parameters = getParameters(context);
if (parameters != null && myCallable instanceof PyFunction function) {
final List<PyCallableParameter> functionParameters = function.getParameters(context);
if (!ContainerUtil.isEmpty(functionParameters) && functionParameters.get(0).isSelf()) {
List<PyCallableParameter> newParameters = ContainerUtil.subList(parameters, 1);
return new PyCallableTypeImpl(newParameters, myReturnType, myCallable, myModifier, myImplicitOffset);
}
}
return this;
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
@@ -163,19 +178,4 @@ public class PyCallableTypeImpl implements PyCallableType {
public int hashCode() {
return Objects.hash(myParameters, myReturnType, myCallable, myModifier, myImplicitOffset);
}
@Override
public @NotNull PyCallableType dropSelf(@NotNull TypeEvalContext context) {
final List<PyCallableParameter> parameters = getParameters(context);
if (!ContainerUtil.isEmpty(parameters) && myCallable instanceof PyFunction function) {
final List<PyCallableParameter> functionParameters = function.getParameters(context);
if (!ContainerUtil.isEmpty(functionParameters) && functionParameters.get(0).isSelf()) {
List<PyCallableParameter> newParameters = ContainerUtil.subList(parameters, 1);
return new PyCallableTypeImpl(newParameters, myReturnType, myCallable, myModifier, myImplicitOffset);
}
}
return this;
}
}
@@ -15,6 +15,7 @@ import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
import com.jetbrains.python.codeInsight.typing.PyProtocolsKt;
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.ParamHelper;
import com.jetbrains.python.psi.impl.PyBuiltinCache;
import com.jetbrains.python.psi.impl.PyPsiUtils;
import com.jetbrains.python.psi.impl.PyTypeProvider;
@@ -612,17 +613,19 @@ public final class PyTypeChecker {
}
private static boolean matchProtocols(@NotNull PyClassType expected, @NotNull PyClassType actual, @NotNull MatchContext matchContext) {
GenericSubstitutions substitutions = collectTypeSubstitutions(actual, matchContext.context);
GenericSubstitutions expectedSubstitutions = collectTypeSubstitutions(expected, matchContext.context);
GenericSubstitutions actualSubstitutions = collectTypeSubstitutions(actual, matchContext.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
// > If a protocol uses Self in methods or attribute annotations, then a class Foo is assignable to the protocol
// if its corresponding methods and attribute annotations use either Self or Foo or any of Foo’s subclasses.
//
// It should be equivalent to replacing Self in the protocol with the Foo class we're matching it with.
// It should be equivalent to replacing Self in the protocol with the Foo class we're matching it with.
GenericSubstitutions protocolSubstitutions = new GenericSubstitutions();
protocolSubstitutions.qualifierType = actual.toInstance();
MatchContext protocolContext = new MatchContext(matchContext.context, protocolSubstitutions, matchContext.reversedSubstitutions);
MatchContext protocolContext =
new MatchContext(matchContext.context, protocolSubstitutions, matchContext.reversedSubstitutions);
for (kotlin.Pair<PyTypeMember, List<PyTypeMember>> pair : PyProtocolsKt.inspectProtocolSubclass(expected, actual,
matchContext.context)) {
final PyTypeMember protocolMember = pair.getFirst();
@@ -631,7 +634,11 @@ public final class PyTypeChecker {
return false;
}
final PyType protocolElementType = dropSelfIfNeeded(expected, pair.getFirst().getType(), matchContext.context);
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 boolean elementResult = ContainerUtil.exists(subclassElementMembers, subclassElementMember -> {
if (protocolMember.isWritable() && !subclassElementMember.isWritable()) {
return false;
@@ -646,7 +653,8 @@ public final class PyTypeChecker {
}
PyType subclassElementType = dropSelfIfNeeded(actual, subclassElementMember.getType(), matchContext.context);
subclassElementType = substitute(subclassElementType, substitutions, matchContext.context);
subclassElementType = substitute(subclassElementType, actualSubstitutions, matchContext.context);
return match(protocolElementType, subclassElementType, protocolContext).orElse(true);
});
@@ -710,9 +718,9 @@ public final class PyTypeChecker {
private static @Nullable PyType dropSelfIfNeeded(@NotNull PyClassType classType,
@Nullable PyType elementType,
@NotNull TypeEvalContext context) {
if (elementType instanceof PyFunctionType functionType) {
if (PyUtil.isInitOrNewMethod(functionType.getCallable()) || !classType.isDefinition()) {
return functionType.dropSelf(context);
if (elementType instanceof PyCallableType callableType) {
if (PyUtil.isInitOrNewMethod(callableType.getCallable()) || !classType.isDefinition()) {
return callableType.dropSelf(context);
}
}
return elementType;
@@ -877,8 +885,8 @@ public final class PyTypeChecker {
})
.toList();
// TODO: substitute the param spec
final var expectedElementTypes2 = ContainerUtil.filter(expectedElementTypes, it -> !(it instanceof PyParamSpecType));
final var expectedElementTypes2 =
ContainerUtil.filter(expectedElementTypes, type -> !(type instanceof PyParamSpecType));
PyTypeParameterMapping mapping = PyTypeParameterMapping.mapWithParameterList(ContainerUtil.subList(expectedElementTypes2, startIndex),
ContainerUtil.subList(actualParameters, startIndex),
@@ -930,6 +938,17 @@ public final class PyTypeChecker {
return Optional.empty();
}
private static @Nullable PyType getActualReturnType(@NotNull PyCallableType actual, @NotNull TypeEvalContext context) {
if (actual instanceof PyCallableTypeImpl) {
return actual.getReturnType(context);
}
PyCallable callable = actual.getCallable();
if (callable instanceof PyFunction) {
return getReturnTypeToAnalyzeAsCallType((PyFunction)callable, context);
}
return actual.getReturnType(context);
}
private static boolean match(@NotNull List<PyType> expected, @NotNull List<PyType> actual, @NotNull MatchContext matchContext) {
if (expected.size() != actual.size()) return false;
for (int i = 0; i < expected.size(); ++i) {
@@ -942,13 +961,6 @@ public final class PyTypeChecker {
return PyProtocolsKt.isProtocol(expected, context) && expected.getMemberNames(false, context).contains(PyNames.CALL);
}
private static @Nullable PyType getActualReturnType(@NotNull PyCallableType actual, @NotNull TypeEvalContext context) {
if (actual instanceof PyFunctionType functionType && functionType.getCallable() instanceof PyFunction function) {
return getReturnTypeToAnalyzeAsCallType(function, context);
}
return actual.getReturnType(context);
}
private static @Nullable PyTupleType widenUnionOfTuplesToTupleOfUnions(@NotNull PyUnionType unionType, int elementCount) {
boolean consistsOfSameSizeTuples = ContainerUtil.all(unionType.getMembers(), member -> member instanceof PyTupleType tupleType &&
!tupleType.isHomogeneous() &&
@@ -1326,7 +1338,8 @@ public final class PyTypeChecker {
PyTypeVarType sameScopeSubstitution = StreamEx.of(substitutions.typeVars.keySet())
.findFirst(typeVarType2 -> {
return typeVarType2.getDeclarationElement() != null
return !typeVarType2.equals(typeVarSubstitution)
&& typeVarType2.getDeclarationElement() != null
&& typeVarType2.getDeclarationElement().equals(typeVarSubstitution.getDeclarationElement());
}).orElse(null);
@@ -1370,14 +1383,14 @@ public final class PyTypeChecker {
if (qualifierType == null) {
return selfType;
}
// TODO change unification for calls on union types
// TODO change unification for calls on union types
// so that in the following
//
// class A:
// def a_method(self) -> Self:
// ...
// x: A | B
//
//
// 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)
@@ -1420,8 +1433,31 @@ public final class PyTypeChecker {
@Override
public PyType visitPyCallableType(@NotNull PyCallableType callableType) {
@Nullable PyType parametersSubs;
List<PyCallableParameter> parameters = callableType.getParameters(context);
boolean isSelfArgsKwargs = ParamHelper.isSelfArgsKwargsCallable(callableType, context);
boolean isArgsKwargs = ParamHelper.isArgsKwargsCallable(callableType, context);
if (parameters != null && (isSelfArgsKwargs || isArgsKwargs)) {
int argsIndex;
if (isSelfArgsKwargs) {
argsIndex = 1;
}
else {
argsIndex = 0;
}
PyCallableParameter args = parameters.get(argsIndex);
PyType argsType = args.getType(context);
PyCallableParameterVariadicType substituted = substitutions.getParamSpecs().get(argsType);
if (substituted instanceof PyCallableParameterListType listType) {
List<PyCallableParameter> newParameters = listType.getParameters();
if (isSelfArgsKwargs) {
newParameters = ContainerUtil.prepend(newParameters, parameters.getFirst());
}
return new PyCallableTypeImpl(newParameters, clone(callableType.getReturnType(context)), callableType.getCallable(),
callableType.getModifier(), callableType.getImplicitOffset());
}
}
@Nullable PyType parametersSubs;
if (parameters != null) {
PyCallableParameter onlyParam = ContainerUtil.getOnlyItem(parameters);
if (onlyParam != null && onlyParam.getType(context) instanceof PyParamSpecType paramSpecType) {
@@ -1985,6 +2021,29 @@ 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{" +
@@ -4580,6 +4580,27 @@ public class Py3TypeTest extends PyTestCase {
});
}
// PY-85123
public void testGenericReturnTypeMatchInProtocol() {
doTest("int", """
from typing_extensions import reveal_type, Protocol, TypeVar
_T_co = TypeVar("_T_co", covariant=True)
class P(Protocol[_T_co]):
def f(self) -> _T_co: ...
class C:
def f(self) -> int:
return 1
def a[_T](p1: P[_T]) -> _T:
return p1.f()
expr = a(C())
""");
}
private void doTest(final String expectedType, final String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class);
@@ -3331,6 +3331,73 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
""");
}
// PY-85123
public void testOverloadedMethodInConcreteClassWithGenericProtocol() {
doTestByText("""
from typing import TypeVar, overload, Protocol
T = TypeVar("T", contravariant=True)
class SupportsWrite(Protocol[T]):
def write(self, s: T): ...
class B:
@overload
def write(self, s: int): ...
@overload
def write(self, s: str): ...
a: SupportsWrite[str] = B()
""");
}
// PY-85123
public void testProtocolPartialSpecializationFixedReturnGenericParam() {
doTestByText("""
from typing import Protocol, TypeVar, overload
T = TypeVar("T", contravariant=True)
S = TypeVar("S", covariant=True)
class P(Protocol[T, S]):
def write(self, x: T) -> S: ...
class B:
@overload
def write(self, x: int) -> str: ...
@overload
def write(self, x: str) -> str: ...
def accepts_p(arg: P[T, str]) -> None: ...
accepts_p(B())
""");
}
// PY-85123
public void testProtocolPartialSpecializationUnionConcreteAndGeneric() {
doTestByText("""
from typing import Protocol, TypeVar, overload
T = TypeVar("T", contravariant=True)
class SupportsWrite(Protocol[T]):
def write(self, s: T): ...
class B:
@overload
def write(self, s: int): ...
@overload
def write(self, s: str): ...
def accepts_union(x: SupportsWrite[str] | SupportsWrite[T]) -> None: ...
accepts_union(B())
""");
}
// PY-76822
public void testExplicitAnyInConcreteType() {
doTestByText("""
@@ -3426,25 +3493,24 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
doTestByText("""
from collections.abc import Iterable
from typing import assert_type
# PY-84544
def foo(iterable: Iterable[int] | Iterable[str]) -> None:
assert_type(next(iter(iterable)), int | str)
# PY-25989
assert_type(max(1, 2.6), float)
assert_type(max(2.6, 1), float)
max(1, <warning descr="Expected type 'int' (matched generic type 'SupportsRichComparisonT ≤: SupportsDunderLT[Any] | SupportsDunderGT[Any]'), got 'object' instead">object()</warning>)
def bar[T: int, str](v1: T, v2: T) -> T:
if (bool(input())):
return v1
return v2
_ = bar(1, <warning descr="Expected type 'int' (matched generic type 'T ≤: int'), got 'str' instead">"a"</warning>)
""");
}
@@ -3639,3 +3705,4 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
""");
}
}