diff --git a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java index a77a323acb75..02f3d761f659 100644 --- a/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java +++ b/python/python-psi-api/src/com/jetbrains/python/psi/types/PyCallableType.java @@ -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 acceptTypeVisitor(@NotNull PyTypeVisitor visitor) { return visitor.visitPyCallableType(this); } - - @ApiStatus.Experimental - @NotNull - default PyCallableType dropSelf(@NotNull TypeEvalContext context) { - return this; - } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/ParamHelper.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/ParamHelper.java index 849ced8c77e6..20f936ca4511 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/ParamHelper.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/ParamHelper.java @@ -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 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 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. diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java index 69131f24fc43..5128a283eaf7 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableTypeImpl.java @@ -147,6 +147,21 @@ public class PyCallableTypeImpl implements PyCallableType { return myCallable; } + @Override + public @NotNull PyCallableType dropSelf(@NotNull TypeEvalContext context) { + final List parameters = getParameters(context); + if (parameters != null && myCallable instanceof PyFunction function) { + final List functionParameters = function.getParameters(context); + + if (!ContainerUtil.isEmpty(functionParameters) && functionParameters.get(0).isSelf()) { + List 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 parameters = getParameters(context); - if (!ContainerUtil.isEmpty(parameters) && myCallable instanceof PyFunction function) { - final List functionParameters = function.getParameters(context); - - if (!ContainerUtil.isEmpty(functionParameters) && functionParameters.get(0).isSelf()) { - List newParameters = ContainerUtil.subList(parameters, 1); - return new PyCallableTypeImpl(newParameters, myReturnType, myCallable, myModifier, myImplicitOffset); - } - } - - return this; - } } 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 4b450d2110ad..1cd0e68fab1c 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 @@ -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> 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 expected, @NotNull List 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 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 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> 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/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index 1173f35abd47..326bd308cb20 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -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); diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index a2f80307b7a2..3e8912aefbea 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -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, object()) - - + + def bar[T: int, str](v1: T, v2: T) -> T: if (bool(input())): return v1 return v2 - - + + _ = bar(1, "a") """); } @@ -3639,3 +3705,4 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { """); } } +