PY-52839 Overload types: initial

- is not yet applied at call-sites

(cherry picked from commit ce966e6b5d9abe33d7e51f95e4e7083fcd76ecc7)

GitOrigin-RevId: 69bb2e09a9ab4fa447c8ec4b23d48cd4b807e419
This commit is contained in:
Morgan Bartholomew
2026-03-25 07:22:50 +00:00
committed by intellij-monorepo-bot
parent b3b8047d22
commit b3ea8e8736
10 changed files with 364 additions and 5 deletions
@@ -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 = "<unnamed>"
@@ -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;
@@ -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));
@@ -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<PyType> {
);
}
@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;
@@ -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<PyCallableType?>, val impl: Ref<PyType?>?) : PyType {
override fun resolveMember(
name: String,
location: PyExpression?,
direction: AccessDirection,
resolveContext: PyResolveContext,
): List<RatedResolveResult> = emptyList()
override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array<out Any> =
emptyArray()
override val name: String = PyNames.OVERLOAD_TYPE
override val isBuiltin: Boolean = false
override fun assertValid(message: String?) {
}
override fun <T> acceptTypeVisitor(visitor: PyTypeVisitor<T>): T? {
if (visitor is PyTypeVisitorExt) {
return visitor.visitPyOverloadType(this);
}
return visitor.visitPyType(this)
}
}
@@ -144,6 +144,11 @@ public final class PyRecursiveTypeVisitor extends PyTypeVisitorExt<PyRecursiveTy
return result;
}
@Override
public @NotNull List<@Nullable PyType> visitPyOverloadType(@NotNull PyOverloadType overloadType) {
return Collections.unmodifiableList(new ArrayList<>(overloadType.getItems()));
}
@Override
public @NotNull List<@Nullable PyType> visitPyGenericType(@NotNull PyCollectionType genericType) {
return genericType.getElementTypes();
@@ -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)
@@ -62,4 +62,8 @@ public abstract class PyTypeVisitorExt<T> extends PyTypeVisitor<T> {
public T visitPyConcatenateType(@NotNull PyConcatenateType concatenateType) {
return visitPyType(concatenateType);
}
public T visitPyOverloadType(@NotNull PyOverloadType overloadType) {
return visitPyType(overloadType);
}
}
@@ -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);
@@ -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] = <warning descr="Expected type '(int) -> str', got 'Overload[(x: int) -> int, (x: str) -> str]' instead">foo</warning>
""");
}
@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(<warning descr="Expected type 'Overload[(x: int) -> int, (x: str) -> str]' (matched generic type '_T'), got 'Overload[(x: int) -> int, (x: str) -> int]' instead">bar</warning>)
l.append(<warning descr="Expected type 'Overload[(x: int) -> int, (x: str) -> str]' (matched generic type '_T'), got '(x: int) -> int' instead">baz</warning>)
""");
}
@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 = <warning descr="Expected type 'ConverterProtocol', got 'IncompatibleCallable' instead">IncompatibleCallable()</warning>
c3: ConverterProtocol = converter_func # ok
c3: ConverterProtocol = <warning descr="Expected type 'ConverterProtocol', got 'Overload[(x: str) -> str, (x: int) -> int]' instead">bad_converter_func</warning>
def t(c: ConverterProtocol):
l3 = [converter_func]
l3.append(c)
l4 = [bad_converter_func]
l4.append(<warning descr="Expected type 'Overload[(x: str) -> str, (x: int) -> int]' (matched generic type '_T'), got 'ConverterProtocol' instead">c</warning>)
""");
}
@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(<warning descr="Expected type 'Overload[(x: int) -> int, (x: str) -> str, (x: float) -> float]' (matched generic type '_T'), got 'Overload[(x: str) -> str, (x: int) -> int]' instead">few_overloads</warning>)
""");
}
}