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
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
---
fixes:
- |
AsyncRetrying's before, after, and before_sleep callbacks now await their
returned awaitables. Callbacks returning a coroutine, Task, or Future
finish before the next retry action, and their exceptions propagate
instead of being silently ignored. The retried operation's return value
remains unchanged, including when it is itself an awaitable.
13 changes: 12 additions & 1 deletion tenacity/asyncio/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
# limitations under the License.

import functools
import inspect
import sys
import typing as t

Expand Down Expand Up @@ -142,7 +143,17 @@ async def __call__( # type: ignore[override]

@override
def _add_action_func(self, fn: t.Callable[..., t.Any]) -> None:
self.iter_state.actions.append(_utils.wrap_to_async_func(fn))
if fn is self.before or fn is self.after or fn is self.before_sleep:

async def wrapped_callback(retry_state: RetryCallState) -> t.Any:
result = fn(retry_state)
if inspect.isawaitable(result):
return await result
return result

self.iter_state.actions.append(wrapped_callback)
else:
self.iter_state.actions.append(_utils.wrap_to_async_func(fn))

@override
async def _run_retry(self, retry_state: "RetryCallState") -> None: # type: ignore[override]
Expand Down
96 changes: 95 additions & 1 deletion tests/test_asyncio.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
import asyncio
import inspect
import unittest
from collections.abc import Callable, Coroutine
from collections.abc import Awaitable, Callable, Coroutine
from functools import wraps
from typing import Any, TypeVar
from unittest import mock
Expand Down Expand Up @@ -76,6 +76,100 @@ async def _retryable_coroutine_with_2_attempts(thing: NoIOErrorAfterCount) -> An
return thing.go()


@pytest.mark.parametrize("iterate", [False, True])
@pytest.mark.parametrize("task", [False, True])
@pytest.mark.parametrize("callback", ["before", "after", "before_sleep"])
@asynctest
async def test_sync_callback_awaits_returned_awaitable(
callback: str, task: bool, iterate: bool
) -> None:
events: list[tuple[str, int]] = []
pending: list[Awaitable[None]] = []
attempts = 0

async def record(state: RetryCallState) -> None:
await asyncio.sleep(0)
events.append(("callback", state.attempt_number))

def before(state: RetryCallState) -> Awaitable[None]:
result = record(state)
awaitable = asyncio.create_task(result) if task else result
pending.append(awaitable)
return awaitable

async def operation() -> str:
nonlocal attempts
attempts += 1
events.append(("operation", attempts))
if attempts == 1:
raise ValueError("retry")
return "ok"

callbacks: dict[str, Any] = {callback: before}
retrying = AsyncRetrying(**callbacks, stop=stop_after_attempt(2))
result = None
try:
if iterate:
async for attempt in retrying:
with attempt:
result = await operation()
else:
result = await retrying(operation)
assert result == "ok"
if callback == "before":
assert events == [
("callback", 1),
("operation", 1),
("callback", 2),
("operation", 2),
]
assert len(pending) == 2
else:
assert events == [("operation", 1), ("callback", 1), ("operation", 2)]
assert len(pending) == 1
finally:
for awaitable in pending:
if inspect.iscoroutine(awaitable):
awaitable.close()
if task:
await asyncio.gather(*pending)


@asynctest
async def test_callback_future_exception_propagates() -> None:
future: asyncio.Future[None] = asyncio.get_running_loop().create_future()
future.set_exception(ValueError("callback failed"))
calls = 0

def before(state: RetryCallState) -> asyncio.Future[None]:
return future

async def operation() -> None:
nonlocal calls
calls += 1

try:
with pytest.raises(ValueError, match="callback failed"):
await AsyncRetrying(before=before)(operation)
assert calls == 0
finally:
future.exception()


@asynctest
async def test_operation_returned_future_is_not_awaited_twice() -> None:
future: asyncio.Future[str] = asyncio.get_running_loop().create_future()

async def operation() -> asyncio.Future[str]:
return future

result: asyncio.Future[str] = await asyncio.wait_for(
AsyncRetrying()(operation), timeout=1
)
assert result is future
assert not future.done()


class TestAsyncio(unittest.TestCase):
@asynctest
async def test_retry(self) -> None:
Expand Down
Loading