Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
93 changes: 92 additions & 1 deletion flake8_async/visitors/visitor91x.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,72 @@ def func_empty_body(node: cst.FunctionDef) -> bool:
)


def func_has_await(node: cst.FunctionDef) -> bool:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

you're adding a lot of code that duplicates the existing logic of async124, which is very bad in case they were ever to diverge. There's way better ways of tackling this

"""Check if function body contains any await, async with, or async for.

This matches Visitor124's logic for determining if ASYNC124 should fire.
Nested functions are not checked - they're handled separately.
"""
return _func_has_await_impl(node.body)


def _func_has_await_impl(body: cst.BaseSuite) -> bool:
"""Check for await/async with/async for in function body, excluding nested functions."""
visitor = _AwaitFinderVisitor()
if isinstance(body, cst.SimpleStatementSuite) or isinstance(
body, cst.IndentedBlock
):
for stmt in body.body:
stmt.visit(visitor)
if visitor.found:
return True
return False
return False


class _AwaitFinderVisitor(cst.CSTVisitor):
"""Visitor that finds await/async with/async for, but not inside nested functions."""

def __init__(self):
super().__init__()
self.found = False
self.nesting_depth = 0

def visit_FunctionDef(self, node: cst.FunctionDef) -> bool:
self.nesting_depth += 1
return True

def leave_FunctionDef(self, original_node: cst.FunctionDef) -> None:
self.nesting_depth -= 1

def visit_Lambda(self, node: cst.Lambda) -> bool:
self.nesting_depth += 1
return True

def leave_Lambda(self, original_node: cst.Lambda) -> None:
self.nesting_depth -= 1

def visit_Await(self, node: cst.Await) -> bool:
if self.nesting_depth == 0:
self.found = True
return False # don't need to visit children

def visit_With(self, node: cst.With) -> bool:
if self.nesting_depth == 0 and getattr(node, "asynchronous", None):
self.found = True
return True

def visit_For(self, node: cst.For) -> bool:
if self.nesting_depth == 0 and getattr(node, "asynchronous", None):
self.found = True
return True

def visit_CompFor(self, node: cst.CompFor) -> bool:
if self.nesting_depth == 0 and node.asynchronous:
self.found = True
return True


# this could've been implemented as part of visitor91x, but /shrug
@error_class_cst
class Visitor124(Flake8AsyncVisitor_cst):
Expand Down Expand Up @@ -442,6 +508,9 @@ def __init__(self, *args: Any, **kwargs: Any):
self.async_function = False
self.uncheckpointed_statements: set[Statement] = set()
self.comp_unknown = False
self.has_await = False
self.function_has_await = False
self.in_class = False

self.loop_state = LoopState()
self.try_state = TryState()
Expand Down Expand Up @@ -553,7 +622,8 @@ def visit_ImportFrom(self, node: cst.ImportFrom) -> None:
# from a base class (which we charitably assume contains a checkpoint).
# See https://github.com/python-trio/flake8-async/issues/441.
def visit_ClassDef(self, node: cst.ClassDef) -> None:
self.save_state(node, "async_cm_class", "async_cm_class_has_bases")
self.save_state(node, "async_cm_class", "async_cm_class_has_bases", "in_class")
self.in_class = True
defined: dict[str, bool] = {}
checkpointy = (
m.Await()
Expand Down Expand Up @@ -615,6 +685,10 @@ def visit_FunctionDef(self, node: cst.FunctionDef) -> bool:

is_exempt_cm = self._is_exempt_async_cm_method(node)

# Pre-scan for awaits to match Visitor124's behavior for ASYNC124 suppression.
# Class methods are treated as having awaits (for ASYNC124 compatibility).
function_has_await = self.in_class or func_has_await(node)

self.save_state(
node,
"has_yield",
Expand All @@ -632,11 +706,20 @@ def visit_FunctionDef(self, node: cst.FunctionDef) -> bool:
"async_cm_class",
"async_cm_class_has_bases",
"exempt_async_cm_method",
"has_await",
"in_class",
"function_has_await",
copy=True,
)
self.uncheckpointed_statements = set()
self.has_checkpoint_stack = []
self.has_yield = False
self.has_await = self.in_class
self.function_has_await = function_has_await
# Class methods (including __aenter__/__aexit__) are treated as having
# awaits for ASYNC124 compatibility, so ASYNC910/911 suppression doesn't
# apply to them. Nested functions inside a class are not class methods.
self.in_class = False
self.loop_state = LoopState()
# try_state is reset upon entering try
self.taskgroup_has_start_soon = {}
Expand Down Expand Up @@ -834,6 +917,10 @@ def error_91x(
if self.exempt_async_cm_method:
return False

# Suppress ASYNC910/911 when ASYNC124 would fire (no awaits in function)
if not self.function_has_await:
return False

if isinstance(node, cst.FunctionDef):
msg = "exit"
else:
Expand All @@ -853,6 +940,7 @@ def leave_Await(
# so only set checkpoint after the await node

# all nodes are now checkpointed
self.has_await = True
self.checkpoint()
return updated_node

Expand Down Expand Up @@ -938,6 +1026,7 @@ def visit_With_body(self, node: cst.With):
continue

if bool(getattr(node, "asynchronous", False)):
self.has_await = True
self.checkpoint()

# not a clean function call
Expand Down Expand Up @@ -1282,6 +1371,7 @@ def visit_While_body(self, node: cst.For | cst.While):
# appropriate errors if the loop doesn't checkpoint

if getattr(node, "asynchronous", None):
self.has_await = True
self.checkpoint()
else:
self.uncheckpointed_statements = {ARTIFICIAL_STATEMENT}
Expand Down Expand Up @@ -1504,6 +1594,7 @@ def visit_CompFor(self, node: cst.CompFor):

# if async comprehension, checkpoint
if node.asynchronous:
self.has_await = True
self.checkpoint()
self.comp_unknown = False
return False
Expand Down
35 changes: 12 additions & 23 deletions tests/autofix_files/async124.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,14 @@
"""Async function with no awaits could be sync.
It currently does not care if 910/911 would also be triggered."""
ASYNC910/911 are suppressed when ASYNC124 fires (no awaits in function)."""

# ARG --enable=ASYNC124,ASYNC910,ASYNC911
# ARG --no-checkpoint-warning-decorator=custom_disabled_decorator
# NOCOMPILE: foo_nested_sync contains `await` in a sync nested function, which is
# a SyntaxError the bytecode compiler catches but ast.parse accepts. It's only here
# to make sure the plugin doesn't crash on such code.

# 910/911 will also autofix async124, in the sense of adding a checkpoint. This is perhaps
# not what the user wants though, so this would be a case in favor of making 910/911 not
# trigger when async124 does.
# AUTOFIX # all errors get "fixed" except for foo_fix_no_subfix in async124_no_autofix.py
# NOTRIO # ASYNC124 is not autofixable, skip autofix test
# NOAUTOFIX # ASYNC124 is not autofixable
# ASYNCIO_NO_AUTOFIX
from typing import Any, overload
from pytest import fixture
Expand All @@ -27,9 +25,8 @@ async def foo() -> Any:
await foo()


async def foo_print(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
async def foo_print(): # ASYNC124: 0
print("hello")
await trio.lowlevel.checkpoint()


async def conditional_wait(): # ASYNC910: 0, "exit", Statement("function definition", lineno)
Expand All @@ -38,10 +35,8 @@ async def conditional_wait(): # ASYNC910: 0, "exit", Statement("function defini
await trio.lowlevel.checkpoint()


async def foo_gen(): # ASYNC124: 0 # ASYNC911: 0, "exit", Statement("yield", lineno+1)
await trio.lowlevel.checkpoint()
yield # ASYNC911: 4, "yield", Statement("function definition", lineno-1)
await trio.lowlevel.checkpoint()
async def foo_gen(): # ASYNC124: 0
yield


async def foo_async_with():
Expand All @@ -54,16 +49,14 @@ async def foo_async_for():
...


async def foo_nested(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
async def foo_nested(): # ASYNC124: 0
async def foo_nested_2():
await foo()
await trio.lowlevel.checkpoint()


async def foo_nested_sync(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
async def foo_nested_sync(): # ASYNC124: 0
def foo_nested_sync_child():
await foo() # type: ignore[await-not-async]
await trio.lowlevel.checkpoint()


# We don't want to trigger on empty/pass functions because of inheritance.
Expand All @@ -82,17 +75,15 @@ async def foo_empty_pass():

# this was previously silenced, but pytest now gives good errors on sync test + async
# fixture; so in the rare case that it has to be async the user will be able to debug it
async def test_async_fixture( # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
async def test_async_fixture( # ASYNC124: 0
my_async_fixture,
):
assert my_async_fixture.setup_worked_correctly
await trio.lowlevel.checkpoint()


# no params -> no async fixtures
async def test_no_fixture(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
async def test_no_fixture(): # ASYNC124: 0
print("blah")
await trio.lowlevel.checkpoint()


# skip @overload. They should always be empty, but /shrug
Expand Down Expand Up @@ -140,9 +131,8 @@ class Foo:
async def bar( # ASYNC910: 4, "exit", Statement("function definition", lineno)
self,
):
async def bee(): # ASYNC124: 8 # ASYNC910: 8, "exit", Statement("function definition", lineno)
async def bee(): # ASYNC124: 8
print("blah")
await trio.lowlevel.checkpoint()
await trio.lowlevel.checkpoint()

async def later_in_class( # ASYNC910: 4, "exit", Statement("function definition", lineno)
Expand All @@ -152,9 +142,8 @@ async def later_in_class( # ASYNC910: 4, "exit", Statement("function definition
await trio.lowlevel.checkpoint()


async def after_class(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
async def after_class(): # ASYNC124: 0
print()
await trio.lowlevel.checkpoint()


@custom_disabled_decorator
Expand Down
56 changes: 5 additions & 51 deletions tests/autofix_files/async124.py.diff
Original file line number Diff line number Diff line change
Expand Up @@ -8,54 +8,14 @@

custom_disabled_decorator: Any = ...

@@ x,15 x,19 @@

async def foo_print(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
print("hello")
+ await trio.lowlevel.checkpoint()


@@ x,6 x,7 @@
async def conditional_wait(): # ASYNC910: 0, "exit", Statement("function definition", lineno)
if condition():
await foo()
+ await trio.lowlevel.checkpoint()


async def foo_gen(): # ASYNC124: 0 # ASYNC911: 0, "exit", Statement("yield", lineno+1)
+ await trio.lowlevel.checkpoint()
yield # ASYNC911: 4, "yield", Statement("function definition", lineno-1)
+ await trio.lowlevel.checkpoint()


async def foo_async_with():
@@ x,11 x,13 @@
async def foo_nested(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
async def foo_nested_2():
await foo()
+ await trio.lowlevel.checkpoint()


async def foo_nested_sync(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
def foo_nested_sync_child():
await foo() # type: ignore[await-not-async]
+ await trio.lowlevel.checkpoint()


# We don't want to trigger on empty/pass functions because of inheritance.
@@ x,11 x,13 @@
my_async_fixture,
):
assert my_async_fixture.setup_worked_correctly
+ await trio.lowlevel.checkpoint()


# no params -> no async fixtures
async def test_no_fixture(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
print("blah")
+ await trio.lowlevel.checkpoint()


# skip @overload. They should always be empty, but /shrug
async def foo_gen(): # ASYNC124: 0
@@ x,6 x,7 @@

# only the expression in genexp's get checked
Expand All @@ -64,11 +24,10 @@
return ( # ASYNC910: 4, "return", Statement("function definition", lineno-1)
await a async for a in foo_gen()
)
@@ x,15 x,19 @@
@@ x,11 x,13 @@
):
async def bee(): # ASYNC124: 8 # ASYNC910: 8, "exit", Statement("function definition", lineno)
async def bee(): # ASYNC124: 8
print("blah")
+ await trio.lowlevel.checkpoint()
+ await trio.lowlevel.checkpoint()

async def later_in_class( # ASYNC910: 4, "exit", Statement("function definition", lineno)
Expand All @@ -78,9 +37,4 @@
+ await trio.lowlevel.checkpoint()


async def after_class(): # ASYNC124: 0 # ASYNC910: 0, "exit", Statement("function definition", lineno)
print()
+ await trio.lowlevel.checkpoint()


@custom_disabled_decorator
async def after_class(): # ASYNC124: 0
Loading
Loading