diff --git a/CHANGELOG.md b/CHANGELOG.md index 1ca5c7f..eb2719a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -35,6 +35,8 @@ - Simplify `backoff.constant` wait generator implementation [#181](https://github.com/python-backoff/backoff/pull/181) +- Move retry loop logic into a dedicated object [#189](https://github.com/python-backoff/backoff/pull/189) + ## [v2.3.1] - 2025-12-18 ### Fixed diff --git a/backoff/_async.py b/backoff/_async.py index 7f905eb..4d35a0a 100644 --- a/backoff/_async.py +++ b/backoff/_async.py @@ -3,10 +3,9 @@ import asyncio import functools import inspect -import time from typing import TYPE_CHECKING, Any, Callable, TypeVar -from backoff._common import _init_wait_gen, _maybe_call, _next_wait +from backoff._common import _RetryState if TYPE_CHECKING: import sys @@ -107,37 +106,30 @@ def retry_predicate( @functools.wraps(target) async def retry(*args: P.args, **kwargs: P.kwargs) -> T: - # update variables from outer function args - max_tries_value: int | None = _maybe_call(max_tries) - max_time_value: float | None = _maybe_call(max_time) - - tries = 0 - start = time.monotonic() - wait = _init_wait_gen(wait_gen, wait_gen_kwargs) + state = _RetryState( + wait_gen, + wait_gen_kwargs, + max_tries=max_tries, + max_time=max_time, + ) while True: - tries += 1 + state.start_attempt() ret = await target(*args, **kwargs) - elapsed = time.monotonic() - start details: _BaseDetails = { "target": target, "args": args, "kwargs": kwargs, - "tries": tries, - "elapsed": elapsed, + "tries": state.tries, + "elapsed": state.record_elapsed(), } if predicate(ret): - max_tries_exceeded = tries == max_tries_value - max_time_exceeded = ( - max_time_value is not None and elapsed >= max_time_value - ) - - if max_tries_exceeded or max_time_exceeded: + if state.exhausted(): await _call_handlers(on_giveup, **details, value=ret) break try: - seconds = _next_wait(wait, ret, jitter, elapsed, max_time_value) + seconds = state.next_wait(ret, jitter) except StopIteration: await _call_handlers(on_giveup, **details, value=ret) break @@ -192,41 +184,36 @@ async def retry( *args: P.args, **kwargs: P.kwargs, ) -> T: - max_tries_value: int | None = _maybe_call(max_tries) - max_time_value: float | None = _maybe_call(max_time) - - tries = 0 - start = time.monotonic() - wait = _init_wait_gen(wait_gen, wait_gen_kwargs) + state = _RetryState( + wait_gen, + wait_gen_kwargs, + max_tries=max_tries, + max_time=max_time, + ) while True: - tries += 1 + state.start_attempt() details: _BaseDetails = { "target": target, "args": args, "kwargs": kwargs, - "tries": tries, + "tries": state.tries, "elapsed": 0, } try: ret = await target(*args, **kwargs) # type: ignore[misc] # ty:ignore[invalid-await] except exception as e: # type: ignore[misc] # ty:ignore[invalid-exception-caught] - elapsed = time.monotonic() - start - details["elapsed"] = elapsed + details["elapsed"] = state.record_elapsed() giveup_result = await giveup(e) - max_tries_exceeded = tries == max_tries_value - max_time_exceeded = ( - max_time_value is not None and elapsed >= max_time_value - ) - if giveup_result or max_tries_exceeded or max_time_exceeded: + if giveup_result or state.exhausted(): await _call_handlers(on_giveup, **details, exception=e) if raise_on_giveup: raise return None # type: ignore[return-value] # ty:ignore[invalid-return-type] try: - seconds = _next_wait(wait, e, jitter, elapsed, max_time_value) + seconds = state.next_wait(e, jitter) except StopIteration: await _call_handlers(on_giveup, **details, exception=e) raise e from None @@ -244,7 +231,7 @@ async def retry( # await asyncio.sleep(seconds) else: - details["elapsed"] = time.monotonic() - start + details["elapsed"] = state.record_elapsed() await _call_handlers(on_success, **details) return ret diff --git a/backoff/_common.py b/backoff/_common.py index b344184..6c1f387 100644 --- a/backoff/_common.py +++ b/backoff/_common.py @@ -3,6 +3,7 @@ import functools import logging import sys +import time import traceback import warnings from typing import TYPE_CHECKING, Any, TypeVar @@ -51,7 +52,7 @@ def _maybe_call(f: _MaybeCallable[T] | None, *args: Any, **kwargs: Any) -> T | N def _init_wait_gen( wait_gen: _WaitGenerator, wait_gen_kwargs: dict[str, Any], -) -> Generator[Any, Any, None]: +) -> Generator[float, Any, None]: kwargs = {k: _maybe_call(v) for k, v in wait_gen_kwargs.items()} initialized = wait_gen(**kwargs) initialized.send(None) # Initialize with an empty send @@ -59,7 +60,7 @@ def _init_wait_gen( def _next_wait( - wait: Generator[float, int | None, None], + wait: Generator[float, None, None], send_value: Any, jitter: _Jitterer | None, elapsed: float, @@ -86,6 +87,54 @@ def _next_wait( return seconds +class _RetryState: + """Bookkeeping shared by the sync and async retry loops. + + Tracks try count and elapsed time, and owns the initialized wait + generator, so `retry_predicate`/`retry_exception` only need to drive + the call/check/sleep sequence around it. + """ + + __slots__ = ( + "elapsed", + "max_time", + "max_tries", + "start", + "tries", + "wait", + ) + + def __init__( + self, + wait_gen: _WaitGenerator, + wait_gen_kwargs: dict[str, Any], + *, + max_tries: _MaybeCallable[int] | None = None, + max_time: _MaybeCallable[float] | None = None, + ) -> None: + self.tries = 0 + self.elapsed: float = 0 + self.start = time.monotonic() + self.max_tries = _maybe_call(max_tries) + self.max_time = _maybe_call(max_time) + self.wait = _init_wait_gen(wait_gen, wait_gen_kwargs) + + def start_attempt(self) -> None: + self.tries += 1 + + def record_elapsed(self) -> float: + self.elapsed = time.monotonic() - self.start + return self.elapsed + + def exhausted(self) -> bool: + max_tries_exceeded = self.tries == self.max_tries + max_time_exceeded = self.max_time is not None and self.elapsed >= self.max_time + return max_tries_exceeded or max_time_exceeded + + def next_wait(self, send_value: Any, jitter: _Jitterer | None) -> float: + return _next_wait(self.wait, send_value, jitter, self.elapsed, self.max_time) + + def _prepare_logger( logger: str | logging.Logger | logging.LoggerAdapter | None, ) -> logging.Logger | logging.LoggerAdapter | None: diff --git a/backoff/_sync.py b/backoff/_sync.py index 8296cce..86fd63e 100644 --- a/backoff/_sync.py +++ b/backoff/_sync.py @@ -4,7 +4,7 @@ import time from typing import TYPE_CHECKING, Any, Callable, TypeVar -from backoff._common import _init_wait_gen, _maybe_call, _next_wait +from backoff._common import _RetryState if TYPE_CHECKING: import sys @@ -73,36 +73,30 @@ def retry_predicate( ) -> Callable[P, T]: @functools.wraps(target) def retry(*args: P.args, **kwargs: P.kwargs) -> T: - max_tries_value: int | None = _maybe_call(max_tries) - max_time_value: float | None = _maybe_call(max_time) - - tries = 0 - start = time.monotonic() - wait = _init_wait_gen(wait_gen, wait_gen_kwargs) + state = _RetryState( + wait_gen, + wait_gen_kwargs, + max_tries=max_tries, + max_time=max_time, + ) while True: - tries += 1 + state.start_attempt() ret = target(*args, **kwargs) - elapsed = time.monotonic() - start details: _BaseDetails = { "target": target, "args": args, "kwargs": kwargs, - "tries": tries, - "elapsed": elapsed, + "tries": state.tries, + "elapsed": state.record_elapsed(), } if predicate(ret): - max_tries_exceeded = tries == max_tries_value - max_time_exceeded = ( - max_time_value is not None and elapsed >= max_time_value - ) - - if max_tries_exceeded or max_time_exceeded: + if state.exhausted(): _call_handlers(on_giveup, **details, value=ret) break try: - seconds = _next_wait(wait, ret, jitter, elapsed, max_time_value) + seconds = state.next_wait(ret, jitter) except StopIteration: _call_handlers(on_giveup, **details) break @@ -136,40 +130,35 @@ def retry_exception( ) -> Callable[P, T]: @functools.wraps(target) def retry(*args: P.args, **kwargs: P.kwargs) -> T: # type: ignore[return] # ty:ignore[invalid-return-type] - max_tries_value: int | None = _maybe_call(max_tries) - max_time_value: float | None = _maybe_call(max_time) - - tries = 0 - start = time.monotonic() - wait = _init_wait_gen(wait_gen, wait_gen_kwargs) + state = _RetryState( + wait_gen, + wait_gen_kwargs, + max_tries=max_tries, + max_time=max_time, + ) while True: - tries += 1 + state.start_attempt() details: _BaseDetails = { "target": target, "args": args, "kwargs": kwargs, - "tries": tries, + "tries": state.tries, "elapsed": 0, } try: ret = target(*args, **kwargs) except exception as e: # type: ignore[misc] # ty:ignore[invalid-exception-caught] - elapsed = time.monotonic() - start - details["elapsed"] = elapsed - max_tries_exceeded = tries == max_tries_value - max_time_exceeded = ( - max_time_value is not None and elapsed >= max_time_value - ) - - if giveup(e) or max_tries_exceeded or max_time_exceeded: + details["elapsed"] = state.record_elapsed() + + if giveup(e) or state.exhausted(): _call_handlers(on_giveup, **details, exception=e) if raise_on_giveup: raise break try: - seconds = _next_wait(wait, e, jitter, elapsed, max_time_value) + seconds = state.next_wait(e, jitter) except StopIteration: _call_handlers(on_giveup, **details, exception=e) raise e from None @@ -178,7 +167,7 @@ def retry(*args: P.args, **kwargs: P.kwargs) -> T: # type: ignore[return] # ty time.sleep(seconds) else: - details["elapsed"] = time.monotonic() - start + details["elapsed"] = state.record_elapsed() _call_handlers(on_success, **details) return ret diff --git a/backoff/_typing.py b/backoff/_typing.py index ec59ef0..c084115 100644 --- a/backoff/_typing.py +++ b/backoff/_typing.py @@ -44,4 +44,4 @@ class Details(_BaseDetails, _CallDetails, total=False): Callable[[T], bool], Callable[[T], Coroutine[Any, Any, bool]], ] -_WaitGenerator = Callable[..., Generator[Union[float, None], None, None]] +_WaitGenerator = Callable[..., Generator[float, Any, None]] diff --git a/tests/test_backoff.py b/tests/test_backoff.py index c94ac00..f2f7309 100644 --- a/tests/test_backoff.py +++ b/tests/test_backoff.py @@ -16,6 +16,8 @@ from tests.common import _save_target if TYPE_CHECKING: + from collections.abc import Generator + from backoff._typing import Details @@ -693,7 +695,7 @@ def test_on_exception_callable_gen_kwargs(): def lookup_foo(): return "foo" - def wait_gen(foo=None, bar=None): + def wait_gen(foo=None, bar=None) -> Generator[float, None, None]: assert foo == "foo" assert bar == "bar" diff --git a/tests/test_backoff_async.py b/tests/test_backoff_async.py index b0809c4..160b1c4 100644 --- a/tests/test_backoff_async.py +++ b/tests/test_backoff_async.py @@ -13,6 +13,8 @@ from tests.common import _log_hdlrs, _save_target if TYPE_CHECKING: + from collections.abc import Generator + from backoff._typing import Details @@ -681,7 +683,7 @@ async def test_on_exception_callable_gen_kwargs() -> None: def lookup_foo(): return "foo" - def wait_gen(foo=None, bar=None): + def wait_gen(foo=None, bar=None) -> Generator[float, None, None]: assert foo == "foo" assert bar == "bar"