diff --git a/python/python-psi-impl/resources/messages/PyPsiBundle.properties b/python/python-psi-impl/resources/messages/PyPsiBundle.properties
index f6d8776db16d..4e542f343a52 100644
--- a/python/python-psi-impl/resources/messages/PyPsiBundle.properties
+++ b/python/python-psi-impl/resources/messages/PyPsiBundle.properties
@@ -910,6 +910,7 @@ INSP.argument.equals.to.default=Argument equals to the default parameter value
#PyAsyncCallInspection
INSP.NAME.coroutine.is.not.awaited=Coroutine ''{0}'' is not awaited
INSP.async.call=Missing `await` syntax in coroutine calls
+INSP.await.call.on.imported.untyped.function=Function ''{0}'' neither declared as ''async'' nor with ''Awaitable'' as return type
# PyAttributeOutsideInitInspection
INSP.NAME.attribute.outside.init=An instance attribute is defined outside `__init__`
diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java b/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java
index 7c97579ffaed..f4c5e3960e3a 100644
--- a/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java
+++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/unresolvedReference/PyUnresolvedReferencesVisitor.java
@@ -134,9 +134,15 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor
if (unresolved) {
boolean ignoreUnresolved = ignoreUnresolved(node, reference) || !evaluateVersionsForElement(node).contains(myVersion);
if (!ignoreUnresolved) {
- final HighlightSeverity severity = reference instanceof PsiReferenceEx
+ HighlightSeverity severity = reference instanceof PsiReferenceEx
? ((PsiReferenceEx)reference).getUnresolvedHighlightSeverity(myTypeEvalContext)
: HighlightSeverity.ERROR;
+ if (severity == null) {
+ if (isAwaitCallToImportedNonAsyncFunction(reference)) {
+ // special case: type of prefixExpression.getQualifier() is null but we want to check whether the called function is async
+ severity = HighlightSeverity.WEAK_WARNING;
+ }
+ }
if (severity == null) return;
registerUnresolvedReferenceProblem(node, reference, severity);
}
@@ -148,6 +154,27 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor
}
}
+ private boolean isAwaitCallToImportedNonAsyncFunction(@NotNull PsiReference reference) {
+ if (reference.getElement() instanceof PyPrefixExpression prefixExpression
+ && PyNames.DUNDER_AWAIT.equals(prefixExpression.getOperator().getSpecialMethodName())
+ && getReferenceQualifier(reference) instanceof PyCallExpression callExpression) {
+
+ @NotNull List<@NotNull PyCallable> callees =
+ callExpression.multiResolveCalleeFunction(PyResolveContext.defaultContext(myTypeEvalContext));
+
+ if (callees.isEmpty()) {
+ return false;
+ }
+ for (PyCallable callee : callees) {
+ if (callee instanceof PyFunction pyFunction && pyFunction.isAsync()) {
+ return false;
+ }
+ }
+ return true; // no signature is declared async -> warning
+ }
+ return false;
+ }
+
private void registerUnresolvedReferenceProblem(@NotNull PyElement node, final @NotNull PsiReference reference,
@NotNull HighlightSeverity severity) {
if (reference instanceof DocStringTypeReference) {
@@ -263,6 +290,14 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor
}
markedQualified = true;
}
+ else {
+ if (isAwaitCallToImportedNonAsyncFunction(reference)) {
+ description = PyPsiBundle.message("INSP.await.call.on.imported.untyped.function", qualifier.getText());
+ node = qualifier; // show warning on the function call
+ rangeInElement = TextRange.create(0, qualifier.getTextRange().getLength());
+ markedQualified = true;
+ }
+ }
}
if (!markedQualified) {
description = PyPsiBundle.message("INSP.unresolved.refs.unresolved.reference", refText);
diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFun/a.py b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFun/a.py
new file mode 100644
index 000000000000..7135b2811861
--- /dev/null
+++ b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFun/a.py
@@ -0,0 +1,17 @@
+from b import fun_async, fun_non_async
+
+
+async def expect_no_warning():
+ await fun_async()
+
+
+async def expect_new_warning():
+ await fun_non_async()
+
+
+def local_fun_non_async():
+ pass
+
+
+async def expect_warning():
+ await local_fun_non_async()
diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFun/b.py b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFun/b.py
new file mode 100644
index 000000000000..44495165e1f3
--- /dev/null
+++ b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFun/b.py
@@ -0,0 +1,7 @@
+
+async def fun_async():
+ return 3
+
+
+def fun_non_async():
+ return 3
\ No newline at end of file
diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunOverloaded/a.py b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunOverloaded/a.py
new file mode 100644
index 000000000000..299f41db069f
--- /dev/null
+++ b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunOverloaded/a.py
@@ -0,0 +1,5 @@
+from b import overloaded_fun_async_with_implicit_return_type
+
+
+async def expect_no_warning():
+ await overloaded_fun_async_with_implicit_return_type(1)
diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunOverloaded/b.py b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunOverloaded/b.py
new file mode 100644
index 000000000000..0deca23878f4
--- /dev/null
+++ b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunOverloaded/b.py
@@ -0,0 +1,15 @@
+from typing import overload
+
+
+@overload
+async def overloaded_fun_async_with_implicit_return_type(arg0: str):
+ ...
+
+
+@overload
+def overloaded_fun_async_with_implicit_return_type(arg0: int):
+ ...
+
+
+def overloaded_fun_async_with_implicit_return_type(arg0):
+ return 3
diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunReturnAwaitable/a.py b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunReturnAwaitable/a.py
new file mode 100644
index 000000000000..403734c64dd0
--- /dev/null
+++ b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunReturnAwaitable/a.py
@@ -0,0 +1,18 @@
+from b import fun_awaitable_imported, MyAwaitable
+
+
+async def expect_false_positive_warning():
+ await fun_awaitable_imported()
+
+
+async def expect_pass_1():
+ await MyAwaitable()
+
+
+def fun_awaitable_local():
+ return MyAwaitable()
+
+
+async def expect_pass_2():
+ await fun_awaitable_local()
+
diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunReturnAwaitable/b.py b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunReturnAwaitable/b.py
new file mode 100644
index 000000000000..926dd5a4b5ea
--- /dev/null
+++ b/python/testData/inspections/PyUnresolvedReferencesInspection/AsyncAwaitWarningOnImportedFunReturnAwaitable/b.py
@@ -0,0 +1,9 @@
+
+class MyAwaitable:
+ def __await__(self):
+ yield from []
+ return "done"
+
+
+def fun_awaitable_imported():
+ return MyAwaitable()
diff --git a/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java
index a9090920da9a..6b816079ce64 100644
--- a/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java
+++ b/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java
@@ -815,14 +815,14 @@ public class PyUnresolvedReferencesInspectionTest extends PyInspectionTestCase {
runWithLanguageLevel(
LanguageLevel.getLatest(),
() -> doTestByText("""
- def foo(cls):
- return cls
-
-
- @foo
- class Bar2(object):
- def __init__(self):
- print(self.hello)
+ def foo(cls):
+ return cls
+
+
+ @foo
+ class Bar2(object):
+ def __init__(self):
+ print(self.hello)
""")
);
}
@@ -891,6 +891,27 @@ public class PyUnresolvedReferencesInspectionTest extends PyInspectionTestCase {
});
}
+ // PY-78413
+ public void testAsyncAwaitWarningOnImportedFun() {
+ runWithLanguageLevel(LanguageLevel.getLatest(), () -> {
+ doMultiFileTest();
+ });
+ }
+
+ // PY-78413
+ public void testAsyncAwaitWarningOnImportedFunReturnAwaitable() {
+ runWithLanguageLevel(LanguageLevel.getLatest(), () -> {
+ doMultiFileTest();
+ });
+ }
+
+ // PY-78413
+ public void testAsyncAwaitWarningOnImportedFunOverloaded() {
+ runWithLanguageLevel(LanguageLevel.getLatest(), () -> {
+ doMultiFileTest();
+ });
+ }
+
@NotNull
@Override
protected Class extends PyInspection> getInspectionClass() {