mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
PY-85123 Fix protocol substitution when matching it against a concrete class
GitOrigin-RevId: 46438cc4d035c75fb51217ba8eef927144d84a3b
This commit is contained in:
committed by
intellij-monorepo-bot
parent
e3ec62ebfc
commit
453bfc627f
@@ -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 {
|
||||
""");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user