From b3ea8e8736d13cd0be5c25c4c52cfd3abd674272 Mon Sep 17 00:00:00 2001 From: Morgan Bartholomew Date: Mon, 8 Sep 2025 11:20:00 +1000 Subject: [PATCH] PY-52839 `Overload` types: initial - is not yet applied at call-sites (cherry picked from commit ce966e6b5d9abe33d7e51f95e4e7083fcd76ecc7) GitOrigin-RevId: 69bb2e09a9ab4fa447c8ec4b23d48cd4b807e419 --- .../src/com/jetbrains/python/PyNames.kt | 3 + .../python/documentation/PyTypeRenderer.java | 16 ++ .../psi/impl/PyReferenceExpressionImpl.java | 24 ++- .../psi/types/PyCloningTypeVisitor.java | 6 +- .../python/psi/types/PyOverloadType.kt | 39 +++++ .../psi/types/PyRecursiveTypeVisitor.java | 5 + .../python/psi/types/PyTypeChecker.kt | 83 +++++++++- .../python/psi/types/PyTypeVisitorExt.java | 4 + .../com/jetbrains/python/Py3TypeTest.java | 36 +++++ .../Py3TypeCheckerInspectionTest.java | 153 ++++++++++++++++++ 10 files changed, 364 insertions(+), 5 deletions(-) create mode 100644 python/python-psi-impl/src/com/jetbrains/python/psi/types/PyOverloadType.kt diff --git a/python/python-parser/src/com/jetbrains/python/PyNames.kt b/python/python-parser/src/com/jetbrains/python/PyNames.kt index 503632d93a7c..cf9bbb5bea99 100644 --- a/python/python-parser/src/com/jetbrains/python/PyNames.kt +++ b/python/python-parser/src/com/jetbrains/python/PyNames.kt @@ -258,6 +258,9 @@ object PyNames { @NlsSafe const val UNKNOWN_TYPE: @NlsSafe String = "Unknown" + @NlsSafe + const val OVERLOAD_TYPE: @NlsSafe String = "Overload" + @NlsSafe const val UNNAMED_ELEMENT: @NlsSafe String = "" diff --git a/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java b/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java index 700c1f954e1f..c0553694ec02 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java +++ b/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java @@ -31,6 +31,7 @@ import com.jetbrains.python.psi.types.PyIntersectionType; import com.jetbrains.python.psi.types.PyLiteralType; import com.jetbrains.python.psi.types.PyNarrowedType; import com.jetbrains.python.psi.types.PyNeverType; +import com.jetbrains.python.psi.types.PyOverloadType; import com.jetbrains.python.psi.types.PyParamSpecType; import com.jetbrains.python.psi.types.PySelfType; import com.jetbrains.python.psi.types.PyTupleType; @@ -231,6 +232,11 @@ public abstract class PyTypeRenderer extends PyTypeVisitorExt<@NotNull HtmlChunk HtmlChunk selfTypeRender = className(isRenderingFqn() ? "typing.Self" : "Self"); //NON-NLS return selfType.isDefinition() ? wrapInTypingType(selfTypeRender) : selfTypeRender; } + + @Override + public @NotNull HtmlChunk visitPyOverloadType(@NotNull PyOverloadType overloadType) { + return escaped("Callable[..., object]"); //NON-NLS + } } protected boolean maxDepthExceeded() { @@ -612,6 +618,16 @@ public abstract class PyTypeRenderer extends PyTypeVisitorExt<@NotNull HtmlChunk return result.toFragment(); } + @Override + public @NotNull HtmlChunk visitPyOverloadType(@NotNull PyOverloadType overloadType) { + var result = new HtmlBuilder(); + result.append(overloadType.getName()); + result.append("["); + result.append(renderList(ContainerUtil.map(overloadType.getItems(), this::render))); + result.append("]"); + return result.toFragment(); + } + protected final @Nullable @NlsSafe String getTypeName(@NotNull PyType type) { if (isNoneType(type)) { return PyNames.NONE; diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java index e83748f73018..d6c752081247 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java @@ -65,12 +65,14 @@ import com.jetbrains.python.psi.types.PyDescriptorTypeUtil; import com.jetbrains.python.psi.types.PyImportedModuleType; import com.jetbrains.python.psi.types.PyModuleType; import com.jetbrains.python.psi.types.PyNarrowedType; +import com.jetbrains.python.psi.types.PyOverloadType; import com.jetbrains.python.psi.types.PyType; import com.jetbrains.python.psi.types.PyTypeChecker; import com.jetbrains.python.psi.types.PyTypeUtil; import com.jetbrains.python.psi.types.PyUnionType; import com.jetbrains.python.psi.types.PyUnsafeUnionType; import com.jetbrains.python.psi.types.TypeEvalContext; +import com.jetbrains.python.pyi.PyiUtil; import com.jetbrains.python.refactoring.PyDefUseUtil; import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; @@ -341,16 +343,25 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere final PsiFile realFile = FileContextUtil.getContextFile(this); if (!(getContainingFile() instanceof PyExpressionCodeFragment) || (realFile != null && context.maySwitchToAST(realFile))) { + final var overloadMembers = new ArrayList<>(); for (PsiElement target : PyUtil.multiResolveTopPriority(getReference(resolveContext))) { if (target == this) { continue; } + if (overloadMembers.contains(target)) { + continue; + } + if (!target.isValid()) { throw new PsiInvalidElementAccessException(this); } - members.add(getTypeFromTarget(target, context, this)); + var member = getTypeFromTarget(target, context, this); + if (Ref.deref(member) instanceof PyOverloadType && target instanceof PyFunction function) { + overloadMembers.addAll(PyiUtil.getOverloads(function, context)); + } + members.add(member); } } @@ -493,8 +504,8 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere } } } - if (target instanceof PyFunction) { - final PyDecoratorList decoratorList = ((PyFunction)target).getDecoratorList(); + if (target instanceof PyFunction function) { + final PyDecoratorList decoratorList = function.getDecoratorList(); if (decoratorList != null) { final PyDecorator propertyDecorator = decoratorList.findDecorator(PyNames.PROPERTY); if (propertyDecorator != null) { @@ -507,6 +518,13 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere } } } + var overloads = PyiUtil.getOverloads(function, context); + if (!overloads.isEmpty()) { + return Ref.create(new PyOverloadType( + ContainerUtil.map(overloads, overload -> (PyCallableType)context.getType(overload)), + PyiUtil.isOverload(function, context) ? null : Ref.create(context.getType(function)) + )); + } } if (target instanceof PyTypedElement) { return Ref.create(context.getType((PyTypedElement)target)); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCloningTypeVisitor.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCloningTypeVisitor.java index e241cab18cb2..087f1b9569ff 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCloningTypeVisitor.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCloningTypeVisitor.java @@ -9,7 +9,6 @@ import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; import java.util.IdentityHashMap; -import java.util.List; import java.util.Map; import java.util.Set; import java.util.stream.Collectors; @@ -208,6 +207,11 @@ public abstract class PyCloningTypeVisitor extends PyTypeVisitorExt { ); } + @Override + public PyType visitPyOverloadType(@NotNull PyOverloadType overloadType) { + return new PyOverloadType(ContainerUtil.map(overloadType.getItems(), this::clone), overloadType.getImpl()); + } + @Override public PyType visitPyTypeVarType(@NotNull PyTypeVarType typeVarType) { return typeVarType; diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyOverloadType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyOverloadType.kt new file mode 100644 index 000000000000..f74c22aada5a --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyOverloadType.kt @@ -0,0 +1,39 @@ +package com.jetbrains.python.psi.types + +import com.intellij.openapi.util.Ref +import com.intellij.psi.PsiElement +import com.intellij.util.ProcessingContext +import com.jetbrains.python.PyNames +import com.jetbrains.python.psi.AccessDirection +import com.jetbrains.python.psi.PyExpression +import com.jetbrains.python.psi.resolve.PyResolveContext +import com.jetbrains.python.psi.resolve.RatedResolveResult +import org.jetbrains.annotations.ApiStatus + +@ApiStatus.Experimental +class PyOverloadType(val items: List, val impl: Ref?) : PyType { + + override fun resolveMember( + name: String, + location: PyExpression?, + direction: AccessDirection, + resolveContext: PyResolveContext, + ): List = emptyList() + + override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array = + emptyArray() + + override val name: String = PyNames.OVERLOAD_TYPE + + override val isBuiltin: Boolean = false + + override fun assertValid(message: String?) { + } + + override fun acceptTypeVisitor(visitor: PyTypeVisitor): T? { + if (visitor is PyTypeVisitorExt) { + return visitor.visitPyOverloadType(this); + } + return visitor.visitPyType(this) + } +} diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyRecursiveTypeVisitor.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyRecursiveTypeVisitor.java index c20f8488f0ba..1e31f4653464 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyRecursiveTypeVisitor.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyRecursiveTypeVisitor.java @@ -144,6 +144,11 @@ public final class PyRecursiveTypeVisitor extends PyTypeVisitorExt visitPyOverloadType(@NotNull PyOverloadType overloadType) { + return Collections.unmodifiableList(new ArrayList<>(overloadType.getItems())); + } + @Override public @NotNull List<@Nullable PyType> visitPyGenericType(@NotNull PyCollectionType genericType) { return genericType.getElementTypes(); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt index 28a7ca74558a..1b385cc67614 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.kt @@ -256,7 +256,40 @@ object PyTypeChecker { return match(expected, actual.moduleClassType, context) } - return Optional.of(matchNumericTypes(expected, actual)) + // Handle PyOverloadType matching + if (expected is PyOverloadType) { + if (actual is PyOverloadType) { + // When both are overload types, check if all overloads in expected have a match in actual (subset matching) + return Optional.of( + expected.items.all { expectedItem -> + actual.items.any { actualItem -> + match(expectedItem, actualItem, context).orElse(false)!! + } + } + ) + } + // If expected is overload but actual is not, check if actual is a callable class/protocol + // Extract the __call__ type and compare with the overload + if (actual is PyClassLikeType && actual.isCallable) { + return Optional.of(matchOverloadWithCallable(expected, actual, context, true)) + } + return Optional.of(false) + } + + if (actual is PyOverloadType) { + // If actual is overload but expected is not, first check if expected is a callable protocol + if (expected is PyClassLikeType && expected.isCallable) { + return Optional.of(matchOverloadWithCallable(actual, expected, context, false)) + } + // Otherwise, check if any overload in actual matches expected + return Optional.of( + actual.items.any { item -> + match(expected, item, context).orElse(false)!! + } + ) + } + + return Optional.of(matchNumericTypes(expected, actual)); } private fun match( @@ -1195,6 +1228,54 @@ object PyTypeChecker { return false } + /** + * Compares an overload type against a callable class/protocol's __call__ type. + * Returns true if all expected overload items have a match in the actual overload items. + */ + private fun matchOverloadWithCallable( + overloadType: PyOverloadType, + callableType: PyClassLikeType, + context: MatchContext, + expectedIsOverload: Boolean, + ): Boolean { + val resolveContext = PyResolveContext.defaultContext(context.context) + val resolveResults = callableType.resolveMember(PyNames.CALL, null, AccessDirection.READ, resolveContext) + if (resolveResults.isNullOrEmpty()) { + return false + } + + val element = resolveResults[0].element + var callType = if (element is PyTypedElement) context.context.getType(element) else null + + if (callableType is PyClassType) { + callType = dropSelfInProtocolMember(callableType, callType, context.context) + } + + when (callType) { + is PyOverloadType -> { + // If the __call__ is overloaded, compare overload types (subset matching) + return overloadType.items.all { expectedItem -> + callType.items.any { actualItem -> + match(expectedItem, actualItem, context).orElse(false)!! + } + } + } + is PyCallableType -> { + // If __call__ is not overloaded, check if any overload matches the single callable + return overloadType.items.any { item -> + // Match with correct argument order based on which is expected + if (expectedIsOverload) { + match(item, callType, context).orElse(false)!! + } + else { + match(callType, item, context).orElse(false)!! + } + } + } + else -> return false + } + } + @JvmStatic fun isUnknown(type: PyType?, context: TypeEvalContext): Boolean { return isUnknown(type, true, context) diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeVisitorExt.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeVisitorExt.java index d8ea0d06ce3c..fc16b94eda66 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeVisitorExt.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeVisitorExt.java @@ -62,4 +62,8 @@ public abstract class PyTypeVisitorExt extends PyTypeVisitor { public T visitPyConcatenateType(@NotNull PyConcatenateType concatenateType) { return visitPyType(concatenateType); } + + public T visitPyOverloadType(@NotNull PyOverloadType overloadType) { + return visitPyType(overloadType); + } } diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index 35b9351b221b..4eb46173dff1 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -5304,6 +5304,42 @@ public class Py3TypeTest extends PyTestCase { """); } + public void testOverloadImpl() { + doTest("Overload[(x: int) -> str, (x: str) -> int]", """ + from typing import overload + + @overload + def foo(x: int) -> str: ... + + @overload + def foo(x: str) -> int: ... + + def foo(x): ... + + expr = foo + """); + } + + public void testOverloadStub() { + runWithAdditionalFileInLibDir( + "stub.pyi", """ + from typing import overload + + @overload + def foo(x: int) -> str: ... + + @overload + def foo(x: str) -> int: ... + """, (x) -> { + doTest("Overload[(x: int) -> str, (x: str) -> int]", """ + from stub import foo + + expr = foo + """); + } + ); + } + 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 e639805382e0..5d9fb758531c 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -4387,5 +4387,158 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase { _: list[tuple[int]] = [(1,)] """); } + + @TestFor(issues="PY-52839") + public void testOverloadAssignabilityToCallable() { + doTestByText(""" + from typing import Callable, overload + + @overload + def foo(x: int) -> int: ... + + @overload + def foo(x: str) -> str: ... + + def foo(x: object) -> object: ... + + _: Callable[[int], int] = foo # ok + _: Callable[[str], str] = foo # ok + _: Callable[[int], str] = foo + """); + } + + @TestFor(issues="PY-52839") + public void testAssignabilityToOverload() { + doTestByText(""" + from typing import Callable, overload + + @overload + def foo(x: int) -> int: ... + + @overload + def foo(x: str) -> str: ... + + def foo(x: object) -> object: ... + + @overload + def foo2(x: str) -> str: ... + + @overload + def foo2(x: int) -> int: ... + + def foo2(x: object) -> object: ... + + @overload + def bar(x: int) -> int: ... + + @overload + def bar(x: str) -> int: ... + + def bar(x: object) -> object: ... + + def baz(x: int) -> int: ... + + l = [foo] + l.append(foo) # ok + l.append(foo2) # ok + l.append(bar) + l.append(baz) + """); + } + + @TestFor(issues="PY-52839") + public void testOverloadWithCallableProtocol() { + doTestByText(""" + from typing import overload, Protocol + + class ConverterProtocol(Protocol): + @overload + def __call__(self, x: int) -> str: ... + + @overload + def __call__(self, x: str) -> int: ... + + class CompatibleCallable: + @overload + def __call__(self, x: str) -> int: ... + + @overload + def __call__(self, x: int) -> str: ... + + def __call__(self, x: object) -> object: ... + + class IncompatibleCallable: + @overload + def __call__(self, x: int) -> int: ... + + @overload + def __call__(self, x: str) -> str: ... + + def __call__(self, x: object) -> object: ... + + + @overload + def converter_func(x: str) -> int: ... + + @overload + def converter_func(x: int) -> str: ... + + def converter_func(x: object) -> object: ... + + @overload + def bad_converter_func(x: str) -> str: ... + + @overload + def bad_converter_func(x: int) -> int: ... + + def bad_converter_func(x: object) -> object: ... + + c1: ConverterProtocol = CompatibleCallable() # ok + c2: ConverterProtocol = IncompatibleCallable() + c3: ConverterProtocol = converter_func # ok + c3: ConverterProtocol = bad_converter_func + + def t(c: ConverterProtocol): + l3 = [converter_func] + l3.append(c) + + l4 = [bad_converter_func] + l4.append(c) + """); + } + + @TestFor(issues="PY-52839") + public void testOverloadSubsetMatching() { + doTestByText(""" + from typing import overload, Callable + + @overload + def many_overloads(x: int) -> int: ... + + @overload + def many_overloads(x: str) -> str: ... + + @overload + def many_overloads(x: float) -> float: ... + + def many_overloads(x: object) -> object: ... + + + @overload + def few_overloads(x: str) -> str: ... + + @overload + def few_overloads(x: int) -> int: ... + + def few_overloads(x: object) -> object: ... + + # Assigning to list infers the overload type + l1 = [few_overloads] + l1.append(many_overloads) # ok + + l2 = [many_overloads] + l2.append(few_overloads) + """); + } }