mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
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:
committed by
intellij-monorepo-bot
parent
b3b8047d22
commit
b3ea8e8736
@@ -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;
|
||||
|
||||
+21
-3
@@ -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>)
|
||||
""");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user