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
This commit is contained in:
Marcus Mews
2025-07-04 07:17:36 +00:00
committed by intellij-monorepo-bot
parent 60dbca64fc
commit 50ddce24c7
11 changed files with 190 additions and 33 deletions
@@ -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<PyDecorator> decorators = ContainerUtil.filter(decoratorList.getDecorators(), d -> !isTransparentDecorator(d, context));
if (decorators.isEmpty()) {
PyDecorator[] decorators = decoratorList.getDecorators();
List<PyDecorator> 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 <code>decorator(reference)<code>.
/* Our goal is to infer the type of reference of a decorated object.
* For that we are going to infer a type of expression <code>decorator(reference)<code>.
* 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<PsiElement> 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<PyType> evaluateType(@NotNull PyDecoratable referenceTarget,
@NotNull TypeEvalContext context,
@NotNull List<PyDecorator> decorators) {
@@ -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:
* <pre>
* `class A:
* pass
* ```
* class A:
* pass
* a = A()
* b = a() # callee type is A, resolved callee is A.__call__
` *
</pre> *
* ```
*/
fun PyCallExpression.multipleResolveCallee(resolveContext: PyResolveContext): List<PyCallableType> {
return PyUtil.getParameterizedCachedValue(
@@ -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<PsiElement> 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());
});
@@ -0,0 +1,9 @@
from ext import my_decorator
@my_decorator
def f() -> int:
pass
value = f()
dec_func = f
@@ -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
@@ -0,0 +1,9 @@
from ext import my_decorator
@my_decorator
def f(a: Any) -> int:
return 3
value = f()
dec_func = f
@@ -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
@@ -0,0 +1,9 @@
from ext import my_decorator
@my_decorator
def f(a: str) -> int:
return 3
value = f()
dec_func = f
@@ -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
@@ -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());
@@ -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(<warning descr="Expected type 'T ≤: str', got 'int' instead">42</warning>)
""")
);
@@ -1595,7 +1595,7 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase {
() -> doTestByText("""
def foo[T: (str, bool)](p: T):
return p
expr = foo(<warning descr="Expected type 'T ≤: str | bool', got 'int' instead">42</warning>)
""")
);