From 50ddce24c75495167fed97c17c541a55bd5ac32b Mon Sep 17 00:00:00 2001 From: Marcus Mews Date: Fri, 4 Jul 2025 07:17:36 +0000 Subject: [PATCH] PY-79622: `@typing.overload` does not work correctly with decorator application syntax - check all overloads of each decorator to retrieve type annotations GitOrigin-RevId: 916e4f796e3295c413f93e3adebf64f4692c8c11 --- .../PyDecoratedFunctionTypeProvider.java | 23 ++-- .../python/psi/impl/PyCallExpressionHelper.kt | 9 +- .../python/psi/resolve/PyResolveUtil.java | 4 +- .../a.py | 9 ++ .../ext.py | 9 ++ .../a.py | 9 ++ .../ext.py | 12 ++ .../a.py | 9 ++ .../ext.py | 12 ++ .../PyDecoratedFunctionTypeProviderTest.java | 119 +++++++++++++++--- .../PyTypeCheckerInspectionTest.java | 8 +- 11 files changed, 190 insertions(+), 33 deletions(-) create mode 100644 python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/a.py create mode 100644 python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/ext.py create mode 100644 python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/a.py create mode 100644 python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/ext.py create mode 100644 python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/a.py create mode 100644 python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/ext.py diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/decorator/PyDecoratedFunctionTypeProvider.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/decorator/PyDecoratedFunctionTypeProvider.java index 995f8c72b7eb..cefdb5bd7f92 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/decorator/PyDecoratedFunctionTypeProvider.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/decorator/PyDecoratedFunctionTypeProvider.java @@ -13,6 +13,8 @@ import com.jetbrains.python.psi.resolve.PyResolveUtil; import com.jetbrains.python.psi.types.PyType; import com.jetbrains.python.psi.types.PyTypeProviderBase; import com.jetbrains.python.psi.types.TypeEvalContext; +import com.jetbrains.python.pyi.PyiUtil; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -35,20 +37,21 @@ public final class PyDecoratedFunctionTypeProvider extends PyTypeProviderBase { return null; } - List decorators = ContainerUtil.filter(decoratorList.getDecorators(), d -> !isTransparentDecorator(d, context)); - if (decorators.isEmpty()) { + PyDecorator[] decorators = decoratorList.getDecorators(); + List explicitlyTypedDecorators = ContainerUtil.filter(decorators, d -> !isTransparentDecorator(d, context)); + if (explicitlyTypedDecorators.isEmpty()) { return null; } - /* Our goal is to infer the type of reference of decorated object. - * For that we going to infer a type of expression decorator(reference). + /* Our goal is to infer the type of reference of a decorated object. + * For that we are going to infer a type of expression decorator(reference). * The expression contains a new reference for the same object, and our * type inferring engine will ask this provider about the type of the object again. * To prevent an infinite loop, we need to add a recursion guard here. */ return RecursionManager.doPreventingRecursion( Pair.create(referenceTarget, context), false, - () -> evaluateType(pyDecoratable, context, decorators) + () -> evaluateType(pyDecoratable, context, explicitlyTypedDecorators) ); } @@ -60,11 +63,11 @@ public final class PyDecoratedFunctionTypeProvider extends PyTypeProviderBase { QualifiedName qualifiedName = decorator.getQualifiedName(); if (qualifiedName == null) return false; // Decorator is the only expression persisted in PSI stubs. - // Calling getReference().resolve() will cause un-stubbing of the containing file + // Calling decorator.getCallee() to retrieve a reference will cause un-stubbing of the containing file List resolved = PyResolveUtil.resolveQualifiedNameInScope(qualifiedName, (PyFile)decorator.getContainingFile(), context); for (PsiElement res : resolved) { if (res instanceof PyFunction function) { - if (function.getTypeCommentAnnotation() != null || function.getAnnotation() != null) { + if (hasReturnTypeHint(function, context)) { return false; } } @@ -75,6 +78,12 @@ public final class PyDecoratedFunctionTypeProvider extends PyTypeProviderBase { return true; } + private static boolean hasReturnTypeHint(@NotNull PyFunction function, @NotNull TypeEvalContext context) { + return StreamEx.of(function) + .append(PyiUtil.getOverloads(function, context)) + .anyMatch(f -> f.getTypeCommentAnnotation() != null || f.getAnnotation() != null); + } + private static @Nullable Ref evaluateType(@NotNull PyDecoratable referenceTarget, @NotNull TypeEvalContext context, @NotNull List decorators) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt index 70b2a7291df9..db2fa82ccd13 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.kt @@ -120,13 +120,12 @@ fun multiResolveCallee(x: PyCallExpression, resolveContext: PyResolveContext): L * It is not the same as [getCalleeType] since * this method returns callable types that would be actually called, the mentioned method returns type of underlying callee. * Compare: - *
- * `class A:
- * pass
+ * ```
+ * class A:
+ *     pass
  * a = A()
  * b = a()  # callee type is A, resolved callee is A.__call__
-` *
-
* + * ``` */ fun PyCallExpression.multipleResolveCallee(resolveContext: PyResolveContext): List { return PyUtil.getParameterizedCachedValue( diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/resolve/PyResolveUtil.java b/python/python-psi-impl/src/com/jetbrains/python/psi/resolve/PyResolveUtil.java index b530caf4cc39..7af943123ac2 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/resolve/PyResolveUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/resolve/PyResolveUtil.java @@ -241,8 +241,8 @@ public final class PyResolveUtil { * @return all possible candidates that can be found by the given qualified name */ public static @NotNull List resolveQualifiedNameInScope(@NotNull QualifiedName qualifiedName, - @NotNull ScopeOwner scopeOwner, - @NotNull TypeEvalContext context) { + @NotNull ScopeOwner scopeOwner, + @NotNull TypeEvalContext context) { return PyUtil.getParameterizedCachedValue(scopeOwner, Pair.create(qualifiedName, context), (param) -> { return doResolveQualifiedNameInScope(param.getFirst(), scopeOwner, param.getSecond()); }); diff --git a/python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/a.py b/python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/a.py new file mode 100644 index 000000000000..645a815286b3 --- /dev/null +++ b/python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/a.py @@ -0,0 +1,9 @@ +from ext import my_decorator + +@my_decorator +def f() -> int: + pass + + +value = f() +dec_func = f \ No newline at end of file diff --git a/python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/ext.py b/python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/ext.py new file mode 100644 index 000000000000..74e657c241c6 --- /dev/null +++ b/python/testData/types/DecoratedTransformsReturnTypeOneImportedOverload/ext.py @@ -0,0 +1,9 @@ +from typing import Callable, overload + + +@overload +def my_decorator(fn: Callable[[], int]) -> Callable[[], str]: + ... + +def my_decorator(fn: Callable[[], int]) : + pass \ No newline at end of file diff --git a/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/a.py b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/a.py new file mode 100644 index 000000000000..8e90293494c8 --- /dev/null +++ b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/a.py @@ -0,0 +1,9 @@ +from ext import my_decorator + +@my_decorator +def f(a: Any) -> int: + return 3 + + +value = f() +dec_func = f \ No newline at end of file diff --git a/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/ext.py b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/ext.py new file mode 100644 index 000000000000..e90b0135c182 --- /dev/null +++ b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst/ext.py @@ -0,0 +1,12 @@ +from typing import Any, Callable, overload + +@overload +def my_decorator(fn: Callable[[Any], int]) -> Callable[[str], str]: + ... + +@overload +def my_decorator(fn: Callable[[Any], int]) -> Callable[[bool], bool]: + ... + +def my_decorator(fn: Callable[[Any], int]) : + pass \ No newline at end of file diff --git a/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/a.py b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/a.py new file mode 100644 index 000000000000..f68f6b71b55d --- /dev/null +++ b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/a.py @@ -0,0 +1,9 @@ +from ext import my_decorator + +@my_decorator +def f(a: str) -> int: + return 3 + + +value = f() +dec_func = f \ No newline at end of file diff --git a/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/ext.py b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/ext.py new file mode 100644 index 000000000000..d743ad7caa9f --- /dev/null +++ b/python/testData/types/DecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond/ext.py @@ -0,0 +1,12 @@ +from typing import Any, Callable, overload + +@overload +def my_decorator(fn: Callable[[int], int]) -> Callable[[str], str]: + ... + +@overload +def my_decorator(fn: Callable[[str], int]) -> Callable[[bool], bool]: + ... + +def my_decorator(fn: Callable[[Any], int]) : + pass \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/PyDecoratedFunctionTypeProviderTest.java b/python/testSrc/com/jetbrains/python/PyDecoratedFunctionTypeProviderTest.java index cd87c1e8cf70..e56419b07644 100644 --- a/python/testSrc/com/jetbrains/python/PyDecoratedFunctionTypeProviderTest.java +++ b/python/testSrc/com/jetbrains/python/PyDecoratedFunctionTypeProviderTest.java @@ -245,21 +245,21 @@ public class PyDecoratedFunctionTypeProviderTest extends PyTestCase { // PY-60104 public void testNotAnnotatedClassDecoratorIsIgnored() { doTest("C", "type[C]", """ - from typing import reveal_type - - def changing_interface(cls): - cls.new_attr = 42 - return cls - - - @changing_interface - class C: - def __init__(self, p: int) -> None: - self.attr = p - - - value = C() - dec_func = C + from typing import reveal_type + + def changing_interface(cls): + cls.new_attr = 42 + return cls + + + @changing_interface + class C: + def __init__(self, p: int) -> None: + self.attr = p + + + value = C() + dec_func = C """); } @@ -376,6 +376,95 @@ public class PyDecoratedFunctionTypeProviderTest extends PyTestCase { """); } + // PY-79622 + public void testDecoratedTransformsReturnTypeOneOverload() { + doTest("str", "() -> str", """ + from typing import Callable, overload + + @overload + def my_decorator(fn: Callable[[], int]) -> Callable[[], str]: + ... + + def my_decorator(fn: Callable[[], int]) : + pass + + @my_decorator + def f() -> int: + pass + + + value = f() + dec_func = f + """); + } + + // PY-79622 + public void testDecoratedTransformsReturnTypeTwoMatchingOverloadsChoseFirst() { + doTest("str", "(str) -> str", """ + from typing import Any, Callable, overload + + @overload + def my_decorator(fn: Callable[[Any], int]) -> Callable[[str], str]: + ... + + @overload + def my_decorator(fn: Callable[[Any], int]) -> Callable[[bool], bool]: + ... + + def my_decorator(fn: Callable[[Any], int]) : + pass + + @my_decorator + def f(a: Any) -> int: + return 3 + + # first overload is chosen + value = f() + dec_func = f + """); + } + + // PY-79622 + public void testDecoratedTransformsReturnTypeTwoMatchingOverloadsChoseSecond() { + doTest("bool", "(bool) -> bool", """ + from typing import Any, Callable, overload + + @overload + def my_decorator(fn: Callable[[int], int]) -> Callable[[str], str]: + ... + + @overload + def my_decorator(fn: Callable[[str], int]) -> Callable[[bool], bool]: + ... + + def my_decorator(fn: Callable[[Any], int]) : + pass + + @my_decorator + def f(a: str) -> int: + return 3 + + # second overload is chosen + value = f() + dec_func = f + """); + } + + // PY-79622 + public void testDecoratedTransformsReturnTypeOneImportedOverload() { + doMultiFileTest("str", "() -> str"); + } + + // PY-79622 + public void testDecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseFirst() { + doMultiFileTest("str", "(str) -> str"); + } + + // PY-79622 + public void testDecoratedTransformsReturnTypeTwoMatchingImportedOverloadsChoseSecond() { + doMultiFileTest("bool", "(bool) -> bool"); + } + private void doTest(@NotNull String expectedValueType, @NotNull String expectedFuncType, @NotNull String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); checkTypes(expectedValueType, expectedFuncType, allContexts()); diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java index aca164cbe5c6..528ea695dfa0 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java @@ -1566,11 +1566,11 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { LanguageLevel.getLatest(), () -> doTestByText(""" from typing import NoReturn - + class Test: def __init__(self) -> NoReturn: raise Exception() - + """) ); } @@ -1582,7 +1582,7 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { () -> doTestByText(""" def foo[T: str](p: T): return p - + expr = foo(42) """) ); @@ -1595,7 +1595,7 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { () -> doTestByText(""" def foo[T: (str, bool)](p: T): return p - + expr = foo(42) """) );