diff --git a/releasenotes/notes/await-callback-results-67c0d3be82014a95.yaml b/releasenotes/notes/await-callback-results-67c0d3be82014a95.yaml new file mode 100644 index 00000000..dfc780ab --- /dev/null +++ b/releasenotes/notes/await-callback-results-67c0d3be82014a95.yaml @@ -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. diff --git a/tenacity/asyncio/__init__.py b/tenacity/asyncio/__init__.py index 3292a6ce..d2097360 100644 --- a/tenacity/asyncio/__init__.py +++ b/tenacity/asyncio/__init__.py @@ -16,6 +16,7 @@ # limitations under the License. import functools +import inspect import sys import typing as t @@ -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] diff --git a/tests/test_asyncio.py b/tests/test_asyncio.py index 21c4e936..99020e1a 100644 --- a/tests/test_asyncio.py +++ b/tests/test_asyncio.py @@ -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 @@ -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: