diff --git a/tenacity/__init__.py b/tenacity/__init__.py index e4660039..132da4ba 100644 --- a/tenacity/__init__.py +++ b/tenacity/__init__.py @@ -1,870 +1,884 @@ -# Copyright 2016-2018 Julien Danjou -# Copyright 2017 Elisey Zanko -# Copyright 2016 Étienne Bersac -# Copyright 2016 Joshua Harlow -# Copyright 2013-2014 Ray Holder -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import dataclasses -import functools -import sys -import threading -import time -import typing as t -import warnings -from abc import ABC, abstractmethod -from concurrent import futures - -from . import _utils -from ._utils import override - -# Import all built-in after strategies for easier usage. -from .after import after_log, after_nothing - -# Import all built-in before strategies for easier usage. -from .before import before_log, before_nothing - -# Import all built-in before sleep strategies for easier usage. -from .before_sleep import before_sleep_log, before_sleep_nothing - -# Import all nap strategies for easier usage. -from .nap import sleep, sleep_using_event - -# Import all built-in retry strategies for easier usage. -from .retry import ( - retry_all, - retry_always, - retry_any, - retry_base, - retry_if_exception, - retry_if_exception_cause_type, - retry_if_exception_message, - retry_if_exception_type, - retry_if_not_exception_message, - retry_if_not_exception_type, - retry_if_not_result, - retry_if_result, - retry_never, - retry_unless_exception_type, -) - -# Import all built-in stop strategies for easier usage. -from .stop import ( - stop_after_attempt, - stop_after_delay, - stop_all, - stop_any, - stop_before_delay, - stop_never, - stop_when_event_set, -) - -# Import all built-in wait strategies for easier usage. -from .wait import ( - wait_chain, - wait_combine, - wait_exception, - wait_exponential, - wait_exponential_jitter, - wait_fixed, - wait_incrementing, - wait_none, - wait_random, - wait_random_exponential, -) -from .wait import wait_random_exponential as wait_full_jitter - -try: - import tornado -except ImportError: - tornado = None # type: ignore[assignment] - - -def _has_tornado() -> bool: - # A function, not a module-level constant: test suites force the - # non-tornado path by setting `tenacity.tornado = None`, so the answer has - # to be computed from the live global every time it is asked for. - return tornado is not None - - -if t.TYPE_CHECKING: - if sys.version_info >= (3, 11): - from typing import Self - else: - from typing_extensions import Self - - import types - - from . import asyncio as tasyncio - from .retry import RetryBaseT - from .stop import StopBaseT - from .wait import WaitBaseT - - -WrappedFnReturnT = t.TypeVar("WrappedFnReturnT") -WrappedFn = t.TypeVar("WrappedFn", bound=t.Callable[..., t.Any]) -P = t.ParamSpec("P") -R = t.TypeVar("R") - - -@dataclasses.dataclass(slots=True) -class IterState: - actions: list[t.Callable[["RetryCallState"], t.Any]] = dataclasses.field( - default_factory=list - ) - retry_run_result: bool = False - stop_run_result: bool = False - is_explicit_retry: bool = False - - def reset(self) -> None: - self.actions = [] - self.retry_run_result = False - self.stop_run_result = False - self.is_explicit_retry = False - - -class TryAgain(Exception): - """Always retry the executed function when raised.""" - - -NO_RESULT = object() - - -class DoAttempt: - pass - - -class DoSleep(float): - pass - - -class BaseAction: - """Base class for representing actions to take by retry object. - - Concrete implementations must define: - - __init__: to initialize all necessary fields - - REPR_FIELDS: class variable specifying attributes to include in repr(self) - - NAME: for identification in retry object methods and callbacks - """ - - REPR_FIELDS: t.ClassVar[t.Sequence[str]] = () - NAME: t.ClassVar[str | None] = None - - @override - def __repr__(self) -> str: - state_str = ", ".join( - f"{field}={getattr(self, field)!r}" for field in self.REPR_FIELDS - ) - return f"{self.__class__.__name__}({state_str})" - - @override - def __str__(self) -> str: - return repr(self) - - -class RetryAction(BaseAction): - REPR_FIELDS: t.ClassVar[t.Sequence[str]] = ("sleep",) - NAME: t.ClassVar[str | None] = "retry" - - def __init__(self, sleep: t.SupportsFloat) -> None: - self.sleep = float(sleep) - - -_unset = object() - - -def _first_set(first: t.Any | object, second: t.Any) -> t.Any: - return second if first is _unset else first - - -class RetryError(Exception): - """Encapsulates the last attempt instance right before giving up.""" - - def __init__(self, last_attempt: "Future") -> None: - self.last_attempt = last_attempt - super().__init__(last_attempt) - - def reraise(self) -> t.NoReturn: - if self.last_attempt.failed: - exc = self.last_attempt.exception() - # When the user explicitly raises TryAgain (typically from within - # an "except" block), surface the underlying exception that caused - # the retry rather than the opaque TryAgain sentinel. - if isinstance(exc, TryAgain): - cause = exc.__cause__ or exc.__context__ - if cause is not None: - raise cause.with_traceback(cause.__traceback__) from None - raise self.last_attempt.result() - raise self - - @override - def __str__(self) -> str: - return f"{self.__class__.__name__}[{self.last_attempt}]" - - -class AttemptManager: - """Manage attempt context.""" - - def __init__(self, retry_state: "RetryCallState") -> None: - self.retry_state = retry_state - - def __enter__(self) -> None: - pass - - def __exit__( - self, - exc_type: type[BaseException] | None, - exc_value: BaseException | None, - traceback: "types.TracebackType | None", - ) -> bool | None: - if exc_type is not None and exc_value is not None: - self.retry_state.set_exception((exc_type, exc_value, traceback)) - return True # Swallow exception. - # We don't have the result, actually. - self.retry_state.set_result(None) - return None - - async def __aenter__(self) -> None: - pass - - async def __aexit__( - self, - exc_type: type[BaseException] | None, - exc_value: BaseException | None, - traceback: "types.TracebackType | None", - ) -> bool | None: - return self.__exit__(exc_type, exc_value, traceback) - - -class BaseRetrying(ABC): - def __init__( - self, - sleep: t.Callable[[int | float], None] = sleep, - stop: "StopBaseT" = stop_never, - wait: "WaitBaseT" = wait_none(), - retry: "RetryBaseT" = retry_if_exception_type(), - before: t.Callable[["RetryCallState"], None] = before_nothing, - after: t.Callable[["RetryCallState"], None] = after_nothing, - before_sleep: t.Callable[["RetryCallState"], None] | None = None, - reraise: bool = False, - retry_error_cls: type[RetryError] = RetryError, - retry_error_callback: t.Callable[["RetryCallState"], t.Any] | None = None, - name: str | None = None, - enabled: bool = True, - ) -> None: - self.sleep = sleep - self.stop = stop - self.wait = wait - self.retry = retry - self.before = before - self.after = after - self.before_sleep = before_sleep - self.reraise = reraise - self._local = threading.local() - self.retry_error_cls = retry_error_cls - self.retry_error_callback = retry_error_callback - self._name = name - self.enabled = enabled - - def copy( - self, - sleep: t.Callable[[int | float], None] | object = _unset, - stop: "StopBaseT | object" = _unset, - wait: "WaitBaseT | object" = _unset, - retry: retry_base | object = _unset, - before: t.Callable[["RetryCallState"], None] | object = _unset, - after: t.Callable[["RetryCallState"], None] | object = _unset, - before_sleep: t.Callable[["RetryCallState"], None] | object | None = _unset, - reraise: bool | object = _unset, - retry_error_cls: type[RetryError] | object = _unset, - retry_error_callback: t.Callable[["RetryCallState"], t.Any] - | object - | None = _unset, - name: str | object | None = _unset, - enabled: bool | object = _unset, - ) -> "Self": - """Copy this object with some parameters changed if needed.""" - return self.__class__( - sleep=_first_set(sleep, self.sleep), - stop=_first_set(stop, self.stop), - wait=_first_set(wait, self.wait), - retry=_first_set(retry, self.retry), - before=_first_set(before, self.before), - after=_first_set(after, self.after), - before_sleep=_first_set(before_sleep, self.before_sleep), - reraise=_first_set(reraise, self.reraise), - retry_error_cls=_first_set(retry_error_cls, self.retry_error_cls), - retry_error_callback=_first_set( - retry_error_callback, self.retry_error_callback - ), - name=_first_set(name, self._name), - enabled=_first_set(enabled, self.enabled), - ) - - # No @override: object.__getstate__ only exists from Python 3.11 on. - def __getstate__(self) -> dict[str, t.Any]: - # Exclude threading.local which cannot be pickled - return {k: v for k, v in self.__dict__.items() if k != "_local"} - - def __setstate__(self, state: dict[str, t.Any]) -> None: - self.__dict__.update(state) - self._local = threading.local() - - @override - def __str__(self) -> str: - return self._name if self._name is not None else "" - - @override - def __repr__(self) -> str: - return ( - f"<{self.__class__.__name__} object at 0x{id(self):x} (" - f"stop={self.stop}, " - f"wait={self.wait}, " - f"sleep={self.sleep}, " - f"retry={self.retry}, " - f"before={self.before}, " - f"after={self.after}, " - f"name={self._name!r})>" - ) - - @property - def statistics(self) -> dict[str, t.Any]: - """Return a dictionary of runtime statistics. - - This dictionary will be empty when the controller has never been - ran. When it is running or has ran previously it should have (but - may not) have useful and/or informational keys and values when - running is underway and/or completed. - - .. warning:: The keys in this dictionary **should** be somewhat - stable (not changing), but their existence **may** - change between major releases as new statistics are - gathered or removed so before accessing keys ensure that - they actually exist and handle when they do not. - - .. note:: The values in this dictionary are local to the thread - running call (so if multiple threads share the same retrying - object - either directly or indirectly) they will each have - their own view of statistics they have collected (in the - future we may provide a way to aggregate the various - statistics from each thread). - """ - if not hasattr(self._local, "statistics"): - self._local.statistics = t.cast("dict[str, t.Any]", {}) - return self._local.statistics # type: ignore[no-any-return] - - @property - def iter_state(self) -> IterState: - if not hasattr(self._local, "iter_state"): - self._local.iter_state = IterState() - return self._local.iter_state # type: ignore[no-any-return] - - def wraps(self, f: t.Callable[P, R]) -> "_RetryDecorated[P, R]": - """Wrap a function for retrying. - - :param f: A function to wrap for retrying. - """ - - @functools.wraps( - f, functools.WRAPPER_ASSIGNMENTS + ("__defaults__", "__kwdefaults__") - ) - def wrapped_f(*args: t.Any, **kw: t.Any) -> t.Any: - if not self.enabled: - return f(*args, **kw) - # Always create a copy to prevent overwriting the local contexts when - # calling the same wrapped functions multiple times in the same stack - copy = self.copy() - # Reuse the same statistics dict rather than rebinding the attribute - # so that the stats stay visible through additional decorators that - # copy attributes via functools.wraps (which copies the reference to - # this dict into the outer wrapper's __dict__). See issue #519. - stats = wrapped_f.statistics # type: ignore[attr-defined] - stats.clear() - copy._local.statistics = stats # noqa: SLF001 - self._local.statistics = stats - return copy(f, *args, **kw) - - def retry_with(*args: t.Any, **kwargs: t.Any) -> "_RetryDecorated[P, R]": - return self.copy(*args, **kwargs).wraps(f) - - # Preserve attributes - wrapped_f.retry = self # type: ignore[attr-defined] - wrapped_f.retry_with = retry_with # type: ignore[attr-defined] - wrapped_f.statistics = {} # type: ignore[attr-defined] - - return t.cast("_RetryDecorated[P, R]", wrapped_f) - - def begin(self) -> None: - self.statistics.clear() - self.statistics["start_time"] = time.monotonic() - self.statistics["attempt_number"] = 1 - self.statistics["idle_for"] = 0 - self.statistics["delay_since_first_attempt"] = 0 - - def _add_action_func(self, fn: t.Callable[..., t.Any]) -> None: - self.iter_state.actions.append(fn) - - def _run_retry(self, retry_state: "RetryCallState") -> None: - self.iter_state.retry_run_result = self.retry(retry_state) - - def _run_wait(self, retry_state: "RetryCallState") -> None: - # `wait` is annotated as always set, so a type checker sees this guard - # as always true -- but untyped callers legitimately pass `None` or `0` - # to mean "no wait", and `sum([])` over an empty list of strategies - # yields the int 0. Keep honouring those. - if not self.wait: # type: ignore[truthy-bool] - retry_state.upcoming_sleep = 0.0 - else: - retry_state.upcoming_sleep = self.wait(retry_state) - - def _run_stop(self, retry_state: "RetryCallState") -> None: - self.statistics["delay_since_first_attempt"] = retry_state.seconds_since_start - self.iter_state.stop_run_result = self.stop(retry_state) - - def iter(self, retry_state: "RetryCallState") -> DoAttempt | DoSleep | t.Any: - self._begin_iter(retry_state) - result = None - for action in self.iter_state.actions: - result = action(retry_state) - return result - - def _begin_iter(self, retry_state: "RetryCallState") -> None: - self.iter_state.reset() - - fut = retry_state.outcome - if fut is None: - if self.before is not None: - self._add_action_func(self.before) - self._add_action_func(lambda rs: DoAttempt()) - return - - self.iter_state.is_explicit_retry = fut.failed and isinstance( - fut.exception(), TryAgain - ) - if not self.iter_state.is_explicit_retry: - self._add_action_func(self._run_retry) - self._add_action_func(self._post_retry_check_actions) - - def _post_retry_check_actions(self, retry_state: "RetryCallState") -> None: - if not (self.iter_state.is_explicit_retry or self.iter_state.retry_run_result): - self._add_action_func(lambda rs: rs.outcome.result()) - return - - if self.after is not None: - self._add_action_func(self.after) - - self._add_action_func(self._run_wait) - self._add_action_func(self._run_stop) - self._add_action_func(self._post_stop_check_actions) - - def _post_stop_check_actions(self, retry_state: "RetryCallState") -> None: - if self.iter_state.stop_run_result: - if self.retry_error_callback: - self._add_action_func(self.retry_error_callback) - return - - def exc_check(rs: "RetryCallState") -> None: - fut = t.cast("Future", rs.outcome) - retry_exc = self.retry_error_cls(fut) - if self.reraise: - retry_exc.reraise() - raise retry_exc from fut.exception() - - self._add_action_func(exc_check) - return - - def next_action(rs: "RetryCallState") -> None: - sleep = rs.upcoming_sleep - rs.next_action = RetryAction(sleep) - rs.idle_for += sleep - self.statistics["idle_for"] += sleep - self.statistics["attempt_number"] += 1 - - self._add_action_func(next_action) - - if self.before_sleep is not None: - self._add_action_func(self.before_sleep) - - self._add_action_func(lambda rs: DoSleep(rs.upcoming_sleep)) - - def __iter__(self) -> t.Generator[AttemptManager, None, None]: - if not self.enabled: - retry_state = RetryCallState(self, fn=None, args=(), kwargs={}) - yield AttemptManager(retry_state=retry_state) - if retry_state.outcome is not None and retry_state.outcome.failed: - raise retry_state.outcome.exception() # type: ignore[misc] - return - - self.begin() - - retry_state = RetryCallState(self, fn=None, args=(), kwargs={}) - while True: - do = self.iter(retry_state=retry_state) - if isinstance(do, DoAttempt): - yield AttemptManager(retry_state=retry_state) - elif isinstance(do, DoSleep): - retry_state.prepare_for_next_attempt() - self.sleep(do) - else: - break - - @abstractmethod - def __call__( - self, - fn: t.Callable[..., WrappedFnReturnT], - *args: t.Any, - **kwargs: t.Any, - ) -> WrappedFnReturnT: - pass - - -class Retrying(BaseRetrying): - """Retrying controller.""" - - @override - def __call__( - self, - fn: t.Callable[..., WrappedFnReturnT], - *args: t.Any, - **kwargs: t.Any, - ) -> WrappedFnReturnT: - if not self.enabled: - return fn(*args, **kwargs) - - self.begin() - - retry_state = RetryCallState(retry_object=self, fn=fn, args=args, kwargs=kwargs) - while True: - do = self.iter(retry_state=retry_state) - if isinstance(do, DoAttempt): - try: - result = fn(*args, **kwargs) - except BaseException: - retry_state.set_exception(sys.exc_info()) # type: ignore[arg-type] - else: - retry_state.set_result(result) - elif isinstance(do, DoSleep): - retry_state.prepare_for_next_attempt() - self.sleep(do) - else: - return do # type: ignore[no-any-return] - - -class Future(futures.Future[t.Any]): - """Encapsulates a (future or past) attempted call to a target function.""" - - def __init__(self, attempt_number: int) -> None: - super().__init__() - self.attempt_number = attempt_number - - @property - def failed(self) -> bool: - """Return whether a exception is being held in this future.""" - return self.exception() is not None - - @classmethod - def construct( - cls, attempt_number: int, value: t.Any, has_exception: bool - ) -> "Future": - """Construct a new Future object.""" - fut = cls(attempt_number) - if has_exception: - fut.set_exception(value) - else: - fut.set_result(value) - return fut - - -class RetryCallState: - """State related to a single call wrapped with Retrying.""" - - def __init__( - self, - retry_object: BaseRetrying, - fn: WrappedFn | None, - args: t.Any, - kwargs: t.Any, - ) -> None: - #: Retry call start timestamp - self.start_time = time.monotonic() - #: Retry manager object - self.retry_object = retry_object - #: Function wrapped by this retry call - self.fn = fn - #: Arguments of the function wrapped by this retry call - self.args = args - #: Keyword arguments of the function wrapped by this retry call - self.kwargs = kwargs - - #: The number of the current attempt - self.attempt_number: int = 1 - #: Last outcome (result or exception) produced by the function - self.outcome: Future | None = None - #: Timestamp of the last outcome - self.outcome_timestamp: float | None = None - #: Time spent sleeping in retries - self.idle_for: float = 0.0 - #: Next action as decided by the retry manager - self.next_action: RetryAction | None = None - #: Next sleep time as decided by the retry manager. - self.upcoming_sleep: float = 0.0 - - def get_fn_name(self) -> str: - """Get the name of the function being retried. - - Returns the fully-qualified name of the wrapped function when used as a - decorator, the ``name`` passed to the retrying object when used as a - context manager, or ``""`` if neither is available. - """ - if self.fn is not None: - return _utils.get_callback_name(self.fn) - return str(self.retry_object) - - @property - def seconds_since_start(self) -> float | None: - if self.outcome_timestamp is None: - return None - return self.outcome_timestamp - self.start_time - - def prepare_for_next_attempt(self) -> None: - self.outcome = None - self.outcome_timestamp = None - self.attempt_number += 1 - self.next_action = None - - def set_result(self, val: t.Any) -> None: - ts = time.monotonic() - fut = Future(self.attempt_number) - fut.set_result(val) - self.outcome, self.outcome_timestamp = fut, ts - - def set_exception( - self, - exc_info: tuple[ - type[BaseException], BaseException, "types.TracebackType | None" - ], - ) -> None: - ts = time.monotonic() - fut = Future(self.attempt_number) - fut.set_exception(exc_info[1]) - self.outcome, self.outcome_timestamp = fut, ts - - @override - def __repr__(self) -> str: - if self.outcome is None: - result = "none yet" - elif self.outcome.failed: - exception = self.outcome.exception() - result = f"failed ({exception.__class__.__name__} {exception})" - else: - result = f"returned {self.outcome.result()}" - - slept = float(round(self.idle_for, 2)) - clsname = self.__class__.__name__ - return f"<{clsname} {id(self)}: attempt #{self.attempt_number}; slept for {slept}; last result: {result}>" - - -class _RetryDecorated(t.Protocol[P, R]): - """Protocol for functions decorated with @retry. - - Provides the original callable signature plus retry control attributes. - """ - - retry: "BaseRetrying" - statistics: dict[str, t.Any] - # Set by functools.wraps on the retry wrapper. Declared so that the - # statistics stay accessible in a type-safe way even when the decorated - # function is further wrapped by another functools.wraps-based decorator - # (which copies these attributes onto the outer wrapper). See issue #519. - __wrapped__: "_RetryDecorated[P, R]" - - def retry_with(self, *args: t.Any, **kwargs: t.Any) -> "_RetryDecorated[P, R]": ... - - def __call__(self, *args: P.args, **kwargs: P.kwargs) -> R: ... - - # Without an explicit `__get__`, type checkers treat `_RetryDecorated` as a - # plain callable attribute rather than a descriptor, so accessing a - # decorated method through an instance keeps demanding the original - # unbound signature (including `self`). Declaring `__get__` here fixes - # attribute access on both the class and an instance; the runtime object - # is a real function (see `wraps` above), which already binds correctly, - # so this only affects static analysis. See issue #532. - @t.overload - def __get__( - self, instance: None, owner: type | None = None - ) -> "_RetryDecorated[P, R]": ... - @t.overload - def __get__( - self, instance: object, owner: type | None = None - ) -> "_RetryDecorated[..., R]": ... - - -class _AsyncRetryDecorator(t.Protocol): - @t.overload - def __call__( - self, fn: "t.Callable[P, types.CoroutineType[t.Any, t.Any, R]]" - ) -> "_RetryDecorated[P, types.CoroutineType[t.Any, t.Any, R]]": ... - @t.overload - def __call__( - self, fn: t.Callable[P, t.Coroutine[t.Any, t.Any, R]] - ) -> "_RetryDecorated[P, t.Coroutine[t.Any, t.Any, R]]": ... - @t.overload - def __call__( - self, fn: t.Callable[P, t.Awaitable[R]] - ) -> "_RetryDecorated[P, t.Awaitable[R]]": ... - @t.overload - def __call__( - self, fn: t.Callable[P, R] - ) -> "_RetryDecorated[P, t.Awaitable[R]]": ... - - -@t.overload -def retry(func: t.Callable[P, R]) -> _RetryDecorated[P, R]: ... - - -@t.overload -def retry( - *, - sleep: t.Callable[[int | float], t.Awaitable[None]], - stop: "StopBaseT" = ..., - wait: "WaitBaseT" = ..., - retry: "RetryBaseT | tasyncio.retry.RetryBaseT" = ..., - before: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = ..., - after: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = ..., - before_sleep: t.Callable[["RetryCallState"], t.Awaitable[None] | None] | None = ..., - reraise: bool = ..., - retry_error_cls: type["RetryError"] = ..., - retry_error_callback: t.Callable[["RetryCallState"], t.Any | t.Awaitable[t.Any]] - | None = ..., - enabled: bool = ..., -) -> _AsyncRetryDecorator: ... - - -@t.overload -def retry( - sleep: t.Callable[[int | float], None] = sleep, - stop: "StopBaseT" = stop_never, - wait: "WaitBaseT" = wait_none(), - retry: "RetryBaseT | tasyncio.retry.RetryBaseT" = retry_if_exception_type(), - before: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = before_nothing, - after: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = after_nothing, - before_sleep: t.Callable[["RetryCallState"], t.Awaitable[None] | None] - | None = None, - reraise: bool = False, - retry_error_cls: type["RetryError"] = RetryError, - retry_error_callback: t.Callable[["RetryCallState"], t.Any | t.Awaitable[t.Any]] - | None = None, - enabled: bool = True, -) -> t.Callable[[t.Callable[P, R]], _RetryDecorated[P, R]]: ... - - -def retry(*dargs: t.Any, **dkw: t.Any) -> t.Any: - """Wrap a function with a new `Retrying` object. - - :param dargs: positional arguments passed to Retrying object - :param dkw: keyword arguments passed to the Retrying object - """ - # support both @retry and @retry() as valid syntax - if len(dargs) == 1 and callable(dargs[0]): - return retry()(dargs[0]) - - def wrap(f: t.Callable[P, R]) -> _RetryDecorated[P, R]: - if isinstance(f, retry_base): - warnings.warn( - f"Got retry_base instance ({f.__class__.__name__}) as callable argument, " - f"this will probably hang indefinitely (did you mean retry={f.__class__.__name__}(...)?)", - stacklevel=2, - ) - r: BaseRetrying - sleep = dkw.get("sleep") - if _utils.is_coroutine_callable(f) or ( - sleep is not None and _utils.is_coroutine_callable(sleep) - ): - r = AsyncRetrying(*dargs, **dkw) - elif ( - _has_tornado() - and hasattr(tornado.gen, "is_coroutine_function") - and tornado.gen.is_coroutine_function(f) - ): - r = TornadoRetrying(*dargs, **dkw) - else: - r = Retrying(*dargs, **dkw) - - return r.wraps(f) - - return wrap - - -from tenacity.asyncio import AsyncRetrying # noqa: E402 - -if _has_tornado(): - from tenacity.tornadoweb import TornadoRetrying - - -__all__ = [ - "NO_RESULT", - "AsyncRetrying", - "AttemptManager", - "BaseAction", - "BaseRetrying", - "DoAttempt", - "DoSleep", - "Future", - "RetryAction", - "RetryCallState", - "RetryError", - "Retrying", - "TryAgain", - "WrappedFn", - "after_log", - "after_nothing", - "before_log", - "before_nothing", - "before_sleep_log", - "before_sleep_nothing", - "retry", - "retry_all", - "retry_always", - "retry_any", - "retry_base", - "retry_if_exception", - "retry_if_exception_cause_type", - "retry_if_exception_message", - "retry_if_exception_type", - "retry_if_not_exception_message", - "retry_if_not_exception_type", - "retry_if_not_result", - "retry_if_result", - "retry_never", - "retry_unless_exception_type", - "sleep", - "sleep_using_event", - "stop_after_attempt", - "stop_after_delay", - "stop_all", - "stop_any", - "stop_before_delay", - "stop_never", - "stop_when_event_set", - "wait_chain", - "wait_combine", - "wait_exception", - "wait_exponential", - "wait_exponential_jitter", - "wait_fixed", - "wait_full_jitter", - "wait_incrementing", - "wait_none", - "wait_random", - "wait_random_exponential", -] +# Copyright 2016-2018 Julien Danjou +# Copyright 2017 Elisey Zanko +# Copyright 2016 Étienne Bersac +# Copyright 2016 Joshua Harlow +# Copyright 2013-2014 Ray Holder +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import dataclasses +import functools +import sys +import threading +import time +import typing as t +import warnings +from abc import ABC, abstractmethod +from concurrent import futures + +from . import _utils +from ._utils import override + +# Import all built-in after strategies for easier usage. +from .after import after_log, after_nothing + +# Import all built-in before strategies for easier usage. +from .before import before_log, before_nothing + +# Import all built-in before sleep strategies for easier usage. +from .before_sleep import before_sleep_log, before_sleep_nothing + +# Import all nap strategies for easier usage. +from .nap import sleep, sleep_using_event + +# Import all built-in retry strategies for easier usage. +from .retry import ( + retry_all, + retry_always, + retry_any, + retry_base, + retry_if_exception, + retry_if_exception_cause_type, + retry_if_exception_message, + retry_if_exception_type, + retry_if_not_exception_message, + retry_if_not_exception_type, + retry_if_not_result, + retry_if_result, + retry_never, + retry_unless_exception_type, +) + +# Import all built-in stop strategies for easier usage. +from .stop import ( + stop_after_attempt, + stop_after_delay, + stop_all, + stop_any, + stop_before_delay, + stop_never, + stop_when_event_set, +) + +# Import all built-in wait strategies for easier usage. +from .wait import ( + wait_chain, + wait_combine, + wait_exception, + wait_exponential, + wait_exponential_jitter, + wait_fixed, + wait_incrementing, + wait_none, + wait_random, + wait_random_exponential, +) +from .wait import wait_random_exponential as wait_full_jitter + +try: + import tornado +except ImportError: + tornado = None # type: ignore[assignment] + + +def _has_tornado() -> bool: + # A function, not a module-level constant: test suites force the + # non-tornado path by setting `tenacity.tornado = None`, so the answer has + # to be computed from the live global every time it is asked for. + return tornado is not None + + +if t.TYPE_CHECKING: + if sys.version_info >= (3, 11): + from typing import Self + else: + from typing_extensions import Self + + import types + + from . import asyncio as tasyncio + from .retry import RetryBaseT + from .stop import StopBaseT + from .wait import WaitBaseT + + +WrappedFnReturnT = t.TypeVar("WrappedFnReturnT") +WrappedFn = t.TypeVar("WrappedFn", bound=t.Callable[..., t.Any]) +P = t.ParamSpec("P") +R = t.TypeVar("R") + + +@dataclasses.dataclass(slots=True) +class IterState: + actions: list[t.Callable[["RetryCallState"], t.Any]] = dataclasses.field( + default_factory=list + ) + retry_run_result: bool = False + stop_run_result: bool = False + is_explicit_retry: bool = False + + def reset(self) -> None: + self.actions = [] + self.retry_run_result = False + self.stop_run_result = False + self.is_explicit_retry = False + + +class TryAgain(Exception): + """Always retry the executed function when raised.""" + + +NO_RESULT = object() + + +class DoAttempt: + pass + + +class DoSleep(float): + pass + + +class BaseAction: + """Base class for representing actions to take by retry object. + + Concrete implementations must define: + - __init__: to initialize all necessary fields + - REPR_FIELDS: class variable specifying attributes to include in repr(self) + - NAME: for identification in retry object methods and callbacks + """ + + REPR_FIELDS: t.ClassVar[t.Sequence[str]] = () + NAME: t.ClassVar[str | None] = None + + @override + def __repr__(self) -> str: + state_str = ", ".join( + f"{field}={getattr(self, field)!r}" for field in self.REPR_FIELDS + ) + return f"{self.__class__.__name__}({state_str})" + + @override + def __str__(self) -> str: + return repr(self) + + +class RetryAction(BaseAction): + REPR_FIELDS: t.ClassVar[t.Sequence[str]] = ("sleep",) + NAME: t.ClassVar[str | None] = "retry" + + def __init__(self, sleep: t.SupportsFloat) -> None: + self.sleep = float(sleep) + + +_unset = object() + + +def _first_set(first: t.Any | object, second: t.Any) -> t.Any: + return second if first is _unset else first + + +class RetryError(Exception): + """Encapsulates the last attempt instance right before giving up.""" + + def __init__(self, last_attempt: "Future") -> None: + self.last_attempt = last_attempt + super().__init__(last_attempt) + + def reraise(self) -> t.NoReturn: + if self.last_attempt.failed: + exc = self.last_attempt.exception() + # When the user explicitly raises TryAgain (typically from within + # an "except" block), surface the underlying exception that caused + # the retry rather than the opaque TryAgain sentinel. + if isinstance(exc, TryAgain): + cause = exc.__cause__ or exc.__context__ + if cause is not None: + raise cause.with_traceback(cause.__traceback__) from None + raise self.last_attempt.result() + raise self + + @override + def __str__(self) -> str: + return f"{self.__class__.__name__}[{self.last_attempt}]" + + +class AttemptManager: + """Manage attempt context.""" + + def __init__(self, retry_state: "RetryCallState") -> None: + self.retry_state = retry_state + + def __enter__(self) -> None: + pass + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: "types.TracebackType | None", + ) -> bool | None: + if exc_type is not None and exc_value is not None: + self.retry_state.set_exception((exc_type, exc_value, traceback)) + return True # Swallow exception. + # We don't have the result, actually. + self.retry_state.set_result(None) + return None + + async def __aenter__(self) -> None: + pass + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: "types.TracebackType | None", + ) -> bool | None: + return self.__exit__(exc_type, exc_value, traceback) + + +class BaseRetrying(ABC): + def __init__( + self, + sleep: t.Callable[[int | float], None] = sleep, + stop: "StopBaseT" = stop_never, + wait: "WaitBaseT" = wait_none(), + retry: "RetryBaseT" = retry_if_exception_type(), + before: t.Callable[["RetryCallState"], None] = before_nothing, + after: t.Callable[["RetryCallState"], None] = after_nothing, + before_sleep: t.Callable[["RetryCallState"], None] | None = None, + reraise: bool = False, + retry_error_cls: type[RetryError] = RetryError, + retry_error_callback: t.Callable[["RetryCallState"], t.Any] | None = None, + name: str | None = None, + enabled: bool = True, + ) -> None: + self.sleep = sleep + self.stop = stop + self.wait = wait + self.retry = retry + self.before = before + self.after = after + self.before_sleep = before_sleep + self.reraise = reraise + self._local = threading.local() + self.retry_error_cls = retry_error_cls + self.retry_error_callback = retry_error_callback + self._name = name + self.enabled = enabled + + def copy( + self, + sleep: t.Callable[[int | float], None] | object = _unset, + stop: "StopBaseT | object" = _unset, + wait: "WaitBaseT | object" = _unset, + retry: retry_base | object = _unset, + before: t.Callable[["RetryCallState"], None] | object = _unset, + after: t.Callable[["RetryCallState"], None] | object = _unset, + before_sleep: t.Callable[["RetryCallState"], None] | object | None = _unset, + reraise: bool | object = _unset, + retry_error_cls: type[RetryError] | object = _unset, + retry_error_callback: t.Callable[["RetryCallState"], t.Any] + | object + | None = _unset, + name: str | object | None = _unset, + enabled: bool | object = _unset, + ) -> "Self": + """Copy this object with some parameters changed if needed.""" + return self.__class__( + sleep=_first_set(sleep, self.sleep), + stop=_first_set(stop, self.stop), + wait=_first_set(wait, self.wait), + retry=_first_set(retry, self.retry), + before=_first_set(before, self.before), + after=_first_set(after, self.after), + before_sleep=_first_set(before_sleep, self.before_sleep), + reraise=_first_set(reraise, self.reraise), + retry_error_cls=_first_set(retry_error_cls, self.retry_error_cls), + retry_error_callback=_first_set( + retry_error_callback, self.retry_error_callback + ), + name=_first_set(name, self._name), + enabled=_first_set(enabled, self.enabled), + ) + + # No @override: object.__getstate__ only exists from Python 3.11 on. + def __getstate__(self) -> dict[str, t.Any]: + # Exclude threading.local which cannot be pickled + return {k: v for k, v in self.__dict__.items() if k != "_local"} + + def __setstate__(self, state: dict[str, t.Any]) -> None: + self.__dict__.update(state) + self._local = threading.local() + + @override + def __str__(self) -> str: + return self._name if self._name is not None else "" + + @override + def __repr__(self) -> str: + return ( + f"<{self.__class__.__name__} object at 0x{id(self):x} (" + f"stop={self.stop}, " + f"wait={self.wait}, " + f"sleep={self.sleep}, " + f"retry={self.retry}, " + f"before={self.before}, " + f"after={self.after}, " + f"name={self._name!r})>" + ) + + @property + def statistics(self) -> dict[str, t.Any]: + """Return a dictionary of runtime statistics. + + This dictionary will be empty when the controller has never been + ran. When it is running or has ran previously it should have (but + may not) have useful and/or informational keys and values when + running is underway and/or completed. + + .. warning:: The keys in this dictionary **should** be somewhat + stable (not changing), but their existence **may** + change between major releases as new statistics are + gathered or removed so before accessing keys ensure that + they actually exist and handle when they do not. + + .. note:: The values in this dictionary are local to the thread + running call (so if multiple threads share the same retrying + object - either directly or indirectly) they will each have + their own view of statistics they have collected (in the + future we may provide a way to aggregate the various + statistics from each thread). + """ + if not hasattr(self._local, "statistics"): + self._local.statistics = t.cast("dict[str, t.Any]", {}) + return self._local.statistics # type: ignore[no-any-return] + + @property + def iter_state(self) -> IterState: + if not hasattr(self._local, "iter_state"): + self._local.iter_state = IterState() + return self._local.iter_state # type: ignore[no-any-return] + + def wraps(self, f: t.Callable[P, R]) -> "_RetryDecorated[P, R]": + """Wrap a function for retrying. + + :param f: A function to wrap for retrying. + """ + + @functools.wraps( + f, functools.WRAPPER_ASSIGNMENTS + ("__defaults__", "__kwdefaults__") + ) + def wrapped_f(*args: t.Any, **kw: t.Any) -> t.Any: + if not self.enabled: + return f(*args, **kw) + # Always create a copy to prevent overwriting the local contexts when + # calling the same wrapped functions multiple times in the same stack + copy = self.copy() + # Reuse the same statistics dict rather than rebinding the attribute + # so that the stats stay visible through additional decorators that + # copy attributes via functools.wraps (which copies the reference to + # this dict into the outer wrapper's __dict__). See issue #519. + stats = wrapped_f.statistics # type: ignore[attr-defined] + stats.clear() + copy._local.statistics = stats # noqa: SLF001 + self._local.statistics = stats + return copy(f, *args, **kw) + + def retry_with(*args: t.Any, **kwargs: t.Any) -> "_RetryDecorated[P, R]": + return self.copy(*args, **kwargs).wraps(f) + + # Preserve attributes + wrapped_f.retry = self # type: ignore[attr-defined] + wrapped_f.retry_with = retry_with # type: ignore[attr-defined] + wrapped_f.statistics = {} # type: ignore[attr-defined] + + return t.cast("_RetryDecorated[P, R]", wrapped_f) + + def begin(self) -> None: + self.statistics.clear() + self.statistics["start_time"] = time.monotonic() + self.statistics["attempt_number"] = 1 + self.statistics["idle_for"] = 0 + self.statistics["delay_since_first_attempt"] = 0 + + def _add_action_func(self, fn: t.Callable[..., t.Any]) -> None: + self.iter_state.actions.append(fn) + + def _run_retry(self, retry_state: "RetryCallState") -> None: + self.iter_state.retry_run_result = self.retry(retry_state) + + def _run_wait(self, retry_state: "RetryCallState") -> None: + # `wait` is annotated as always set, so a type checker sees this guard + # as always true -- but untyped callers legitimately pass `None` or `0` + # to mean "no wait", and `sum([])` over an empty list of strategies + # yields the int 0. Keep honouring those. + if not self.wait: # type: ignore[truthy-bool] + retry_state.upcoming_sleep = 0.0 + else: + retry_state.upcoming_sleep = self.wait(retry_state) + + def _run_stop(self, retry_state: "RetryCallState") -> None: + self.statistics["delay_since_first_attempt"] = retry_state.seconds_since_start + self.iter_state.stop_run_result = self.stop(retry_state) + + def iter(self, retry_state: "RetryCallState") -> DoAttempt | DoSleep | t.Any: + self._begin_iter(retry_state) + result = None + for action in self.iter_state.actions: + result = action(retry_state) + return result + + def _begin_iter(self, retry_state: "RetryCallState") -> None: + self.iter_state.reset() + + fut = retry_state.outcome + if fut is None: + if self.before is not None: + self._add_action_func(self.before) + self._add_action_func(lambda rs: DoAttempt()) + return + + self.iter_state.is_explicit_retry = fut.failed and isinstance( + fut.exception(), TryAgain + ) + if not self.iter_state.is_explicit_retry: + self._add_action_func(self._run_retry) + self._add_action_func(self._post_retry_check_actions) + + def _post_retry_check_actions(self, retry_state: "RetryCallState") -> None: + if not (self.iter_state.is_explicit_retry or self.iter_state.retry_run_result): + self._add_action_func(lambda rs: rs.outcome.result()) + return + + if self.after is not None: + self._add_action_func(self.after) + + self._add_action_func(self._run_wait) + self._add_action_func(self._run_stop) + self._add_action_func(self._post_stop_check_actions) + + def _post_stop_check_actions(self, retry_state: "RetryCallState") -> None: + if self.iter_state.stop_run_result: + if self.retry_error_callback: + self._add_action_func(self.retry_error_callback) + return + + def exc_check(rs: "RetryCallState") -> None: + fut = t.cast("Future", rs.outcome) + retry_exc = self.retry_error_cls(fut) + if self.reraise: + retry_exc.reraise() + raise retry_exc from fut.exception() + + self._add_action_func(exc_check) + return + + def next_action(rs: "RetryCallState") -> None: + sleep = rs.upcoming_sleep + rs.next_action = RetryAction(sleep) + rs.idle_for += sleep + self.statistics["idle_for"] += sleep + self.statistics["attempt_number"] += 1 + + self._add_action_func(next_action) + + if self.before_sleep is not None: + self._add_action_func(self.before_sleep) + + self._add_action_func(lambda rs: DoSleep(rs.upcoming_sleep)) + + def __iter__(self) -> t.Generator[AttemptManager, None, None]: + if not self.enabled: + retry_state = RetryCallState(self, fn=None, args=(), kwargs={}) + yield AttemptManager(retry_state=retry_state) + if retry_state.outcome is not None and retry_state.outcome.failed: + raise retry_state.outcome.exception() # type: ignore[misc] + return + + self.begin() + + retry_state = RetryCallState(self, fn=None, args=(), kwargs={}) + # When used as a context manager without an explicit name=, infer the + # caller's qualified name from the frame that called next() on this + # generator (i.e. the for-loop body in user code). + if self._name is None: + _frame = sys._getframe(1) # noqa: SLF001 + retry_state._inferred_name = ( # noqa: SLF001 + getattr(_frame.f_code, "co_qualname", None) or _frame.f_code.co_name + ) + while True: + do = self.iter(retry_state=retry_state) + if isinstance(do, DoAttempt): + yield AttemptManager(retry_state=retry_state) + elif isinstance(do, DoSleep): + retry_state.prepare_for_next_attempt() + self.sleep(do) + else: + break + + @abstractmethod + def __call__( + self, + fn: t.Callable[..., WrappedFnReturnT], + *args: t.Any, + **kwargs: t.Any, + ) -> WrappedFnReturnT: + pass + + +class Retrying(BaseRetrying): + """Retrying controller.""" + + @override + def __call__( + self, + fn: t.Callable[..., WrappedFnReturnT], + *args: t.Any, + **kwargs: t.Any, + ) -> WrappedFnReturnT: + if not self.enabled: + return fn(*args, **kwargs) + + self.begin() + + retry_state = RetryCallState(retry_object=self, fn=fn, args=args, kwargs=kwargs) + while True: + do = self.iter(retry_state=retry_state) + if isinstance(do, DoAttempt): + try: + result = fn(*args, **kwargs) + except BaseException: + retry_state.set_exception(sys.exc_info()) # type: ignore[arg-type] + else: + retry_state.set_result(result) + elif isinstance(do, DoSleep): + retry_state.prepare_for_next_attempt() + self.sleep(do) + else: + return do # type: ignore[no-any-return] + + +class Future(futures.Future[t.Any]): + """Encapsulates a (future or past) attempted call to a target function.""" + + def __init__(self, attempt_number: int) -> None: + super().__init__() + self.attempt_number = attempt_number + + @property + def failed(self) -> bool: + """Return whether a exception is being held in this future.""" + return self.exception() is not None + + @classmethod + def construct( + cls, attempt_number: int, value: t.Any, has_exception: bool + ) -> "Future": + """Construct a new Future object.""" + fut = cls(attempt_number) + if has_exception: + fut.set_exception(value) + else: + fut.set_result(value) + return fut + + +class RetryCallState: + """State related to a single call wrapped with Retrying.""" + + def __init__( + self, + retry_object: BaseRetrying, + fn: WrappedFn | None, + args: t.Any, + kwargs: t.Any, + ) -> None: + #: Retry call start timestamp + self.start_time = time.monotonic() + #: Retry manager object + self.retry_object = retry_object + #: Function wrapped by this retry call + self.fn = fn + #: Arguments of the function wrapped by this retry call + self.args = args + #: Keyword arguments of the function wrapped by this retry call + self.kwargs = kwargs + + #: The number of the current attempt + self.attempt_number: int = 1 + #: Last outcome (result or exception) produced by the function + self.outcome: Future | None = None + #: Timestamp of the last outcome + self.outcome_timestamp: float | None = None + #: Time spent sleeping in retries + self.idle_for: float = 0.0 + #: Next action as decided by the retry manager + self.next_action: RetryAction | None = None + #: Next sleep time as decided by the retry manager. + self.upcoming_sleep: float = 0.0 + #: Inferred caller name for context-manager usage without explicit name= + self._inferred_name: str | None = None + + def get_fn_name(self) -> str: + """Get the name of the function being retried. + + Returns the fully-qualified name of the wrapped function when used as a + decorator, the ``name`` passed to the retrying object when used as a + context manager, or ``""`` if neither is available. + """ + if self.fn is not None: + return _utils.get_callback_name(self.fn) + # Inferred from the caller's frame in __iter__ (context-manager usage) + inferred: str | None = getattr(self, "_inferred_name", None) + if inferred is not None: + return inferred + return str(self.retry_object) + + @property + def seconds_since_start(self) -> float | None: + if self.outcome_timestamp is None: + return None + return self.outcome_timestamp - self.start_time + + def prepare_for_next_attempt(self) -> None: + self.outcome = None + self.outcome_timestamp = None + self.attempt_number += 1 + self.next_action = None + + def set_result(self, val: t.Any) -> None: + ts = time.monotonic() + fut = Future(self.attempt_number) + fut.set_result(val) + self.outcome, self.outcome_timestamp = fut, ts + + def set_exception( + self, + exc_info: tuple[ + type[BaseException], BaseException, "types.TracebackType | None" + ], + ) -> None: + ts = time.monotonic() + fut = Future(self.attempt_number) + fut.set_exception(exc_info[1]) + self.outcome, self.outcome_timestamp = fut, ts + + @override + def __repr__(self) -> str: + if self.outcome is None: + result = "none yet" + elif self.outcome.failed: + exception = self.outcome.exception() + result = f"failed ({exception.__class__.__name__} {exception})" + else: + result = f"returned {self.outcome.result()}" + + slept = float(round(self.idle_for, 2)) + clsname = self.__class__.__name__ + return f"<{clsname} {id(self)}: attempt #{self.attempt_number}; slept for {slept}; last result: {result}>" + + +class _RetryDecorated(t.Protocol[P, R]): + """Protocol for functions decorated with @retry. + + Provides the original callable signature plus retry control attributes. + """ + + retry: "BaseRetrying" + statistics: dict[str, t.Any] + # Set by functools.wraps on the retry wrapper. Declared so that the + # statistics stay accessible in a type-safe way even when the decorated + # function is further wrapped by another functools.wraps-based decorator + # (which copies these attributes onto the outer wrapper). See issue #519. + __wrapped__: "_RetryDecorated[P, R]" + + def retry_with(self, *args: t.Any, **kwargs: t.Any) -> "_RetryDecorated[P, R]": ... + + def __call__(self, *args: P.args, **kwargs: P.kwargs) -> R: ... + + # Without an explicit `__get__`, type checkers treat `_RetryDecorated` as a + # plain callable attribute rather than a descriptor, so accessing a + # decorated method through an instance keeps demanding the original + # unbound signature (including `self`). Declaring `__get__` here fixes + # attribute access on both the class and an instance; the runtime object + # is a real function (see `wraps` above), which already binds correctly, + # so this only affects static analysis. See issue #532. + @t.overload + def __get__( + self, instance: None, owner: type | None = None + ) -> "_RetryDecorated[P, R]": ... + @t.overload + def __get__( + self, instance: object, owner: type | None = None + ) -> "_RetryDecorated[..., R]": ... + + +class _AsyncRetryDecorator(t.Protocol): + @t.overload + def __call__( + self, fn: "t.Callable[P, types.CoroutineType[t.Any, t.Any, R]]" + ) -> "_RetryDecorated[P, types.CoroutineType[t.Any, t.Any, R]]": ... + @t.overload + def __call__( + self, fn: t.Callable[P, t.Coroutine[t.Any, t.Any, R]] + ) -> "_RetryDecorated[P, t.Coroutine[t.Any, t.Any, R]]": ... + @t.overload + def __call__( + self, fn: t.Callable[P, t.Awaitable[R]] + ) -> "_RetryDecorated[P, t.Awaitable[R]]": ... + @t.overload + def __call__( + self, fn: t.Callable[P, R] + ) -> "_RetryDecorated[P, t.Awaitable[R]]": ... + + +@t.overload +def retry(func: t.Callable[P, R]) -> _RetryDecorated[P, R]: ... + + +@t.overload +def retry( + *, + sleep: t.Callable[[int | float], t.Awaitable[None]], + stop: "StopBaseT" = ..., + wait: "WaitBaseT" = ..., + retry: "RetryBaseT | tasyncio.retry.RetryBaseT" = ..., + before: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = ..., + after: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = ..., + before_sleep: t.Callable[["RetryCallState"], t.Awaitable[None] | None] | None = ..., + reraise: bool = ..., + retry_error_cls: type["RetryError"] = ..., + retry_error_callback: t.Callable[["RetryCallState"], t.Any | t.Awaitable[t.Any]] + | None = ..., + enabled: bool = ..., +) -> _AsyncRetryDecorator: ... + + +@t.overload +def retry( + sleep: t.Callable[[int | float], None] = sleep, + stop: "StopBaseT" = stop_never, + wait: "WaitBaseT" = wait_none(), + retry: "RetryBaseT | tasyncio.retry.RetryBaseT" = retry_if_exception_type(), + before: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = before_nothing, + after: t.Callable[["RetryCallState"], t.Awaitable[None] | None] = after_nothing, + before_sleep: t.Callable[["RetryCallState"], t.Awaitable[None] | None] + | None = None, + reraise: bool = False, + retry_error_cls: type["RetryError"] = RetryError, + retry_error_callback: t.Callable[["RetryCallState"], t.Any | t.Awaitable[t.Any]] + | None = None, + enabled: bool = True, +) -> t.Callable[[t.Callable[P, R]], _RetryDecorated[P, R]]: ... + + +def retry(*dargs: t.Any, **dkw: t.Any) -> t.Any: + """Wrap a function with a new `Retrying` object. + + :param dargs: positional arguments passed to Retrying object + :param dkw: keyword arguments passed to the Retrying object + """ + # support both @retry and @retry() as valid syntax + if len(dargs) == 1 and callable(dargs[0]): + return retry()(dargs[0]) + + def wrap(f: t.Callable[P, R]) -> _RetryDecorated[P, R]: + if isinstance(f, retry_base): + warnings.warn( + f"Got retry_base instance ({f.__class__.__name__}) as callable argument, " + f"this will probably hang indefinitely (did you mean retry={f.__class__.__name__}(...)?)", + stacklevel=2, + ) + r: BaseRetrying + sleep = dkw.get("sleep") + if _utils.is_coroutine_callable(f) or ( + sleep is not None and _utils.is_coroutine_callable(sleep) + ): + r = AsyncRetrying(*dargs, **dkw) + elif ( + _has_tornado() + and hasattr(tornado.gen, "is_coroutine_function") + and tornado.gen.is_coroutine_function(f) + ): + r = TornadoRetrying(*dargs, **dkw) + else: + r = Retrying(*dargs, **dkw) + + return r.wraps(f) + + return wrap + + +from tenacity.asyncio import AsyncRetrying # noqa: E402 + +if _has_tornado(): + from tenacity.tornadoweb import TornadoRetrying + + +__all__ = [ + "NO_RESULT", + "AsyncRetrying", + "AttemptManager", + "BaseAction", + "BaseRetrying", + "DoAttempt", + "DoSleep", + "Future", + "RetryAction", + "RetryCallState", + "RetryError", + "Retrying", + "TryAgain", + "WrappedFn", + "after_log", + "after_nothing", + "before_log", + "before_nothing", + "before_sleep_log", + "before_sleep_nothing", + "retry", + "retry_all", + "retry_always", + "retry_any", + "retry_base", + "retry_if_exception", + "retry_if_exception_cause_type", + "retry_if_exception_message", + "retry_if_exception_type", + "retry_if_not_exception_message", + "retry_if_not_exception_type", + "retry_if_not_result", + "retry_if_result", + "retry_never", + "retry_unless_exception_type", + "sleep", + "sleep_using_event", + "stop_after_attempt", + "stop_after_delay", + "stop_all", + "stop_any", + "stop_before_delay", + "stop_never", + "stop_when_event_set", + "wait_chain", + "wait_combine", + "wait_exception", + "wait_exponential", + "wait_exponential_jitter", + "wait_fixed", + "wait_full_jitter", + "wait_incrementing", + "wait_none", + "wait_random", + "wait_random_exponential", +] diff --git a/tests/test_tenacity.py b/tests/test_tenacity.py index 95ab8117..5f4a5a68 100644 --- a/tests/test_tenacity.py +++ b/tests/test_tenacity.py @@ -1,2410 +1,2446 @@ -# Copyright 2016–2021 Julien Danjou -# Copyright 2016 Joshua Harlow -# Copyright 2013 Ray Holder -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import contextlib -import datetime -import logging -import pickle -import re -import time -import typing -import unittest -from fractions import Fraction -from unittest import mock - -import pytest - -import tenacity -from tenacity import RetryCallState, RetryError, Retrying, retry -from tenacity._utils import override -from tenacity.retry import retry_all, retry_any - -_unset = object() - - -def _make_unset_exception(func_name: str, **kwargs: typing.Any) -> TypeError: - missing = [] - for k, v in kwargs.items(): - if v is _unset: - missing.append(k) - missing_str = ", ".join(repr(s) for s in missing) - return TypeError(func_name + " func missing parameters: " + missing_str) - - -def _set_delay_since_start(retry_state: RetryCallState, delay: typing.Any) -> None: - # Ensure outcome_timestamp - start_time is *exactly* equal to the delay to - # avoid complexity in test code. - retry_state.start_time = Fraction(retry_state.start_time) # type: ignore[assignment] - retry_state.outcome_timestamp = retry_state.start_time + Fraction(delay) - assert retry_state.seconds_since_start == delay - - -def make_retry_state( - previous_attempt_number: typing.Any, - delay_since_first_attempt: typing.Any, - last_result: typing.Any = None, - upcoming_sleep: typing.Any = 0, -) -> RetryCallState: - """Construct RetryCallState for given attempt number & delay. - - Only used in testing and thus is extra careful about timestamp arithmetics. - """ - required_parameter_unset = ( - previous_attempt_number is _unset or delay_since_first_attempt is _unset - ) - if required_parameter_unset: - raise _make_unset_exception( - "wait/stop", - previous_attempt_number=previous_attempt_number, - delay_since_first_attempt=delay_since_first_attempt, - ) - - retry_state = RetryCallState(None, None, (), {}) # type: ignore[arg-type] - retry_state.attempt_number = previous_attempt_number - if last_result is not None: - retry_state.outcome = last_result - else: - retry_state.set_result(None) - - retry_state.upcoming_sleep = upcoming_sleep - - _set_delay_since_start(retry_state, delay_since_first_attempt) - return retry_state - - -class TestBase(unittest.TestCase): - def test_retrying_repr(self) -> None: - class ConcreteRetrying(tenacity.BaseRetrying): - @override - def __call__( - self, fn: typing.Any, *args: typing.Any, **kwargs: typing.Any - ) -> typing.Any: - pass - - repr(ConcreteRetrying()) - - def test_callstate_repr(self) -> None: - rs = RetryCallState(None, None, (), {}) # type: ignore[arg-type] - rs.idle_for = 1.1111111 - assert repr(rs).endswith("attempt #1; slept for 1.11; last result: none yet>") - rs = make_retry_state(2, 5) - assert repr(rs).endswith( - "attempt #2; slept for 0.0; last result: returned None>" - ) - rs = make_retry_state( - 0, 0, last_result=tenacity.Future.construct(1, ValueError("aaa"), True) - ) - assert repr(rs).endswith( - "attempt #0; slept for 0.0; last result: failed (ValueError aaa)>" - ) - - -class TestRetryingName(unittest.TestCase): - def test_str_default(self) -> None: - """Without a name, str() returns ''.""" - assert str(Retrying()) == "" - - def test_str_with_name(self) -> None: - """With a name, str() returns the given name.""" - assert str(Retrying(name="my_block")) == "my_block" - - def test_str_preserved_by_copy(self) -> None: - """copy() preserves the name.""" - r = Retrying(name="my_block") - assert str(r.copy()) == "my_block" - - def test_str_overridden_by_copy(self) -> None: - """copy() allows overriding the name.""" - r = Retrying(name="original") - assert str(r.copy(name="overridden")) == "overridden" - - def test_get_fn_name_decorator(self) -> None: - """get_fn_name() returns the function's qualified name when used as decorator.""" - captured: list[RetryCallState] = [] - - @tenacity.retry( - stop=tenacity.stop_after_attempt(1), - after=lambda rs: captured.append(rs), - ) - def my_func() -> None: - raise ValueError - - with contextlib.suppress(Exception): - my_func() - assert captured - assert "my_func" in captured[0].get_fn_name() - - def test_get_fn_name_context_manager_no_name(self) -> None: - """get_fn_name() returns '' in context manager mode without a name.""" - r = Retrying(stop=tenacity.stop_after_attempt(1)) - rs = RetryCallState(r, None, (), {}) - assert rs.get_fn_name() == "" - - def test_get_fn_name_context_manager_with_name(self) -> None: - """get_fn_name() returns the given name in context manager mode.""" - r = Retrying(name="ws_listener", stop=tenacity.stop_after_attempt(1)) - rs = RetryCallState(r, None, (), {}) - assert rs.get_fn_name() == "ws_listener" - - def test_logging_uses_name(self) -> None: - """before_log uses the name parameter in context manager mode.""" - import unittest.mock - - log = unittest.mock.MagicMock() - logger = unittest.mock.MagicMock(log=log) - - with contextlib.suppress(Exception): - for attempt in Retrying( - name="my_block", - before=tenacity.before_log(logger, logging.INFO), - stop=tenacity.stop_after_attempt(1), - ): - with attempt: - raise ValueError - - args = log.call_args[0] - assert "my_block" in args[1] - - -class TestStopConditions(unittest.TestCase): - def test_never_stop(self) -> None: - r = Retrying() - self.assertFalse(r.stop(make_retry_state(3, 6546))) - - def test_stop_any(self) -> None: - stop = tenacity.stop_any( - tenacity.stop_after_delay(1), tenacity.stop_after_attempt(4) - ) - - def s(*args: typing.Any) -> bool: - return stop(make_retry_state(*args)) - - self.assertFalse(s(1, 0.1)) - self.assertFalse(s(2, 0.2)) - self.assertFalse(s(2, 0.8)) - self.assertTrue(s(4, 0.8)) - self.assertTrue(s(3, 1.8)) - self.assertTrue(s(4, 1.8)) - - def test_stop_all(self) -> None: - stop = tenacity.stop_all( - tenacity.stop_after_delay(1), tenacity.stop_after_attempt(4) - ) - - def s(*args: typing.Any) -> bool: - return stop(make_retry_state(*args)) - - self.assertFalse(s(1, 0.1)) - self.assertFalse(s(2, 0.2)) - self.assertFalse(s(2, 0.8)) - self.assertFalse(s(4, 0.8)) - self.assertFalse(s(3, 1.8)) - self.assertTrue(s(4, 1.8)) - - def test_stop_or(self) -> None: - stop = tenacity.stop_after_delay(1) | tenacity.stop_after_attempt(4) - - def s(*args: typing.Any) -> bool: - return stop(make_retry_state(*args)) - - self.assertFalse(s(1, 0.1)) - self.assertFalse(s(2, 0.2)) - self.assertFalse(s(2, 0.8)) - self.assertTrue(s(4, 0.8)) - self.assertTrue(s(3, 1.8)) - self.assertTrue(s(4, 1.8)) - - def test_stop_and(self) -> None: - stop = tenacity.stop_after_delay(1) & tenacity.stop_after_attempt(4) - - def s(*args: typing.Any) -> bool: - return stop(make_retry_state(*args)) - - self.assertFalse(s(1, 0.1)) - self.assertFalse(s(2, 0.2)) - self.assertFalse(s(2, 0.8)) - self.assertFalse(s(4, 0.8)) - self.assertFalse(s(3, 1.8)) - self.assertTrue(s(4, 1.8)) - - def test_stop_after_attempt(self) -> None: - r = Retrying(stop=tenacity.stop_after_attempt(3)) - self.assertFalse(r.stop(make_retry_state(2, 6546))) - self.assertTrue(r.stop(make_retry_state(3, 6546))) - self.assertTrue(r.stop(make_retry_state(4, 6546))) - - def test_stop_after_delay(self) -> None: - for delay in (1, datetime.timedelta(seconds=1)): - with self.subTest(): - r = Retrying(stop=tenacity.stop_after_delay(delay)) - self.assertFalse(r.stop(make_retry_state(2, 0.999))) - self.assertTrue(r.stop(make_retry_state(2, 1))) - self.assertTrue(r.stop(make_retry_state(2, 1.001))) - - def test_stop_before_delay(self) -> None: - for delay in (1, datetime.timedelta(seconds=1)): - with self.subTest(): - r = Retrying(stop=tenacity.stop_before_delay(delay)) - self.assertFalse( - r.stop(make_retry_state(2, 0.999, upcoming_sleep=0.0001)) - ) - self.assertTrue(r.stop(make_retry_state(2, 1, upcoming_sleep=0.001))) - self.assertTrue(r.stop(make_retry_state(2, 1, upcoming_sleep=1))) - - # It should act the same as stop_after_delay if upcoming sleep is 0 - self.assertFalse(r.stop(make_retry_state(2, 0.999, upcoming_sleep=0))) - self.assertTrue(r.stop(make_retry_state(2, 1, upcoming_sleep=0))) - self.assertTrue(r.stop(make_retry_state(2, 1.001, upcoming_sleep=0))) - - def test_legacy_explicit_stop_type(self) -> None: - Retrying(stop="stop_after_attempt") # type: ignore[arg-type] - - def test_stop_func_with_retry_state(self) -> None: - def stop_func(retry_state: RetryCallState) -> bool: - rs = retry_state - return rs.attempt_number == rs.seconds_since_start - - r = Retrying(stop=stop_func) - self.assertFalse(r.stop(make_retry_state(1, 3))) - self.assertFalse(r.stop(make_retry_state(100, 99))) - self.assertTrue(r.stop(make_retry_state(101, 101))) - - -class TestWaitConditions(unittest.TestCase): - def test_no_sleep(self) -> None: - r = Retrying() - self.assertEqual(0, r.wait(make_retry_state(18, 9879))) - - def test_fixed_sleep(self) -> None: - for wait in (1, datetime.timedelta(seconds=1)): - with self.subTest(): - r = Retrying(wait=tenacity.wait_fixed(wait)) - self.assertEqual(1, r.wait(make_retry_state(12, 6546))) - - def test_incrementing_sleep(self) -> None: - for start, increment in ( - (500, 100), - (datetime.timedelta(seconds=500), datetime.timedelta(seconds=100)), - ): - with self.subTest(): - r = Retrying( - wait=tenacity.wait_incrementing(start=start, increment=increment) - ) - self.assertEqual(500, r.wait(make_retry_state(1, 6546))) - self.assertEqual(600, r.wait(make_retry_state(2, 6546))) - self.assertEqual(700, r.wait(make_retry_state(3, 6546))) - - def test_random_sleep(self) -> None: - for min_, max_ in ( - (1, 20), - (datetime.timedelta(seconds=1), datetime.timedelta(seconds=20)), - ): - with self.subTest(): - r = Retrying(wait=tenacity.wait_random(min=min_, max=max_)) - times = set() - for _ in range(1000): - times.add(r.wait(make_retry_state(1, 6546))) - - # this is kind of non-deterministic... - self.assertTrue(len(times) > 1) - for t in times: - self.assertTrue(t >= 1) - self.assertTrue(t < 20) - - def test_random_sleep_withoutmin_(self) -> None: - r = Retrying(wait=tenacity.wait_random(max=2)) - times = set() - times.add(r.wait(make_retry_state(1, 6546))) - times.add(r.wait(make_retry_state(1, 6546))) - times.add(r.wait(make_retry_state(1, 6546))) - times.add(r.wait(make_retry_state(1, 6546))) - - # this is kind of non-deterministic... - self.assertTrue(len(times) > 1) - for t in times: - self.assertTrue(t >= 0) - self.assertTrue(t <= 2) - - def test_exponential(self) -> None: - r = Retrying(wait=tenacity.wait_exponential()) - self.assertEqual(r.wait(make_retry_state(1, 0)), 1) - self.assertEqual(r.wait(make_retry_state(2, 0)), 2) - self.assertEqual(r.wait(make_retry_state(3, 0)), 4) - self.assertEqual(r.wait(make_retry_state(4, 0)), 8) - self.assertEqual(r.wait(make_retry_state(5, 0)), 16) - self.assertEqual(r.wait(make_retry_state(6, 0)), 32) - self.assertEqual(r.wait(make_retry_state(7, 0)), 64) - self.assertEqual(r.wait(make_retry_state(8, 0)), 128) - - def test_exponential_with_max_wait(self) -> None: - r = Retrying(wait=tenacity.wait_exponential(max=40)) - self.assertEqual(r.wait(make_retry_state(1, 0)), 1) - self.assertEqual(r.wait(make_retry_state(2, 0)), 2) - self.assertEqual(r.wait(make_retry_state(3, 0)), 4) - self.assertEqual(r.wait(make_retry_state(4, 0)), 8) - self.assertEqual(r.wait(make_retry_state(5, 0)), 16) - self.assertEqual(r.wait(make_retry_state(6, 0)), 32) - self.assertEqual(r.wait(make_retry_state(7, 0)), 40) - self.assertEqual(r.wait(make_retry_state(8, 0)), 40) - self.assertEqual(r.wait(make_retry_state(50, 0)), 40) - - def test_exponential_skips_power_after_reaching_max(self) -> None: - class ExplodingPower(float): - @override - def __pow__(self, exponent: float, modulo: int | None = None) -> float: - raise AssertionError("power should not be calculated above the maximum") - - r = Retrying(wait=tenacity.wait_exponential(max=40, exp_base=ExplodingPower(2))) - self.assertEqual(r.wait(make_retry_state(50, 0)), 40) - - def test_exponential_caps_when_max_multiplier_ratio_underflows(self) -> None: - r = Retrying(wait=tenacity.wait_exponential(multiplier=1e300, max=1e-100)) - self.assertEqual(r.wait(make_retry_state(1, 0)), 1e-100) - - def test_exponential_with_min_wait(self) -> None: - r = Retrying(wait=tenacity.wait_exponential(min=20)) - self.assertEqual(r.wait(make_retry_state(1, 0)), 20) - self.assertEqual(r.wait(make_retry_state(2, 0)), 20) - self.assertEqual(r.wait(make_retry_state(3, 0)), 20) - self.assertEqual(r.wait(make_retry_state(4, 0)), 20) - self.assertEqual(r.wait(make_retry_state(5, 0)), 20) - self.assertEqual(r.wait(make_retry_state(6, 0)), 32) - self.assertEqual(r.wait(make_retry_state(7, 0)), 64) - self.assertEqual(r.wait(make_retry_state(8, 0)), 128) - self.assertEqual(r.wait(make_retry_state(20, 0)), 524288) - - def test_exponential_with_max_wait_and_multiplier(self) -> None: - r = Retrying(wait=tenacity.wait_exponential(max=50, multiplier=1)) - self.assertEqual(r.wait(make_retry_state(1, 0)), 1) - self.assertEqual(r.wait(make_retry_state(2, 0)), 2) - self.assertEqual(r.wait(make_retry_state(3, 0)), 4) - self.assertEqual(r.wait(make_retry_state(4, 0)), 8) - self.assertEqual(r.wait(make_retry_state(5, 0)), 16) - self.assertEqual(r.wait(make_retry_state(6, 0)), 32) - self.assertEqual(r.wait(make_retry_state(7, 0)), 50) - self.assertEqual(r.wait(make_retry_state(8, 0)), 50) - self.assertEqual(r.wait(make_retry_state(50, 0)), 50) - - def test_exponential_with_min_wait_and_multiplier(self) -> None: - r = Retrying(wait=tenacity.wait_exponential(min=20, multiplier=2)) - self.assertEqual(r.wait(make_retry_state(1, 0)), 20) - self.assertEqual(r.wait(make_retry_state(2, 0)), 20) - self.assertEqual(r.wait(make_retry_state(3, 0)), 20) - self.assertEqual(r.wait(make_retry_state(4, 0)), 20) - self.assertEqual(r.wait(make_retry_state(5, 0)), 32) - self.assertEqual(r.wait(make_retry_state(6, 0)), 64) - self.assertEqual(r.wait(make_retry_state(7, 0)), 128) - self.assertEqual(r.wait(make_retry_state(8, 0)), 256) - self.assertEqual(r.wait(make_retry_state(20, 0)), 1048576) - - def test_exponential_with_min_wait_andmax__wait(self) -> None: - for min_, max_ in ( - (10, 100), - (datetime.timedelta(seconds=10), datetime.timedelta(seconds=100)), - ): - with self.subTest(): - r = Retrying(wait=tenacity.wait_exponential(min=min_, max=max_)) - self.assertEqual(r.wait(make_retry_state(1, 0)), 10) - self.assertEqual(r.wait(make_retry_state(2, 0)), 10) - self.assertEqual(r.wait(make_retry_state(3, 0)), 10) - self.assertEqual(r.wait(make_retry_state(4, 0)), 10) - self.assertEqual(r.wait(make_retry_state(5, 0)), 16) - self.assertEqual(r.wait(make_retry_state(6, 0)), 32) - self.assertEqual(r.wait(make_retry_state(7, 0)), 64) - self.assertEqual(r.wait(make_retry_state(8, 0)), 100) - self.assertEqual(r.wait(make_retry_state(9, 0)), 100) - self.assertEqual(r.wait(make_retry_state(20, 0)), 100) - - def test_legacy_explicit_wait_type(self) -> None: - Retrying(wait="exponential_sleep") # type: ignore[arg-type] - - def test_wait_func(self) -> None: - def wait_func(retry_state: RetryCallState) -> typing.Any: - return retry_state.attempt_number * retry_state.seconds_since_start # type: ignore[operator] - - r = Retrying(wait=wait_func) - self.assertEqual(r.wait(make_retry_state(1, 5)), 5) - self.assertEqual(r.wait(make_retry_state(2, 11)), 22) - self.assertEqual(r.wait(make_retry_state(10, 100)), 1000) - - def test_wait_combine(self) -> None: - r = Retrying( - wait=tenacity.wait_combine( - tenacity.wait_random(0, 3), tenacity.wait_fixed(5) - ) - ) - # Test it a few time since it's random - for _i in range(1000): - w = r.wait(make_retry_state(1, 5)) - self.assertLess(w, 8) - self.assertGreaterEqual(w, 5) - - def test_wait_exception(self) -> None: - def predicate(exc: BaseException) -> float: - if isinstance(exc, ValueError): - return 3.5 - return 10.0 - - r = Retrying(wait=tenacity.wait_exception(predicate)) - - fut1 = tenacity.Future.construct(1, ValueError(), True) - self.assertEqual(r.wait(make_retry_state(1, 0, last_result=fut1)), 3.5) - - fut2 = tenacity.Future.construct(1, KeyError(), True) - self.assertEqual(r.wait(make_retry_state(1, 0, last_result=fut2)), 10.0) - - fut3 = tenacity.Future.construct(1, None, False) - with self.assertRaises(RuntimeError): - r.wait(make_retry_state(1, 0, last_result=fut3)) - - def test_wait_double_sum(self) -> None: - r = Retrying(wait=tenacity.wait_random(0, 3) + tenacity.wait_fixed(5)) - # Test it a few time since it's random - for _i in range(1000): - w = r.wait(make_retry_state(1, 5)) - self.assertLess(w, 8) - self.assertGreaterEqual(w, 5) - - def test_wait_triple_sum(self) -> None: - r = Retrying( - wait=tenacity.wait_fixed(1) - + tenacity.wait_random(0, 3) - + tenacity.wait_fixed(5) - ) - # Test it a few time since it's random - for _i in range(1000): - w = r.wait(make_retry_state(1, 5)) - self.assertLess(w, 9) - self.assertGreaterEqual(w, 6) - - def test_wait_arbitrary_sum(self) -> None: - r = Retrying( - wait=sum( # type: ignore[arg-type] - [ - tenacity.wait_fixed(1), - tenacity.wait_random(0, 3), - tenacity.wait_fixed(5), - tenacity.wait_none(), - ] - ) - ) - # Test it a few time since it's random - for _ in range(1000): - w = r.wait(make_retry_state(1, 5)) - self.assertLess(w, 9) - self.assertGreaterEqual(w, 6) - - def test_wait_falsy_values_mean_no_wait(self) -> None: - # Untyped callers pass None or 0 to mean "no wait", and `sum([])` - # over an empty list of strategies yields the int 0. All must reach - # the retry path without blowing up inside iter(). - def make_flaky() -> typing.Callable[[], str]: - attempts = [] - - def flaky() -> str: - attempts.append(1) - if len(attempts) < 2: - raise ValueError("boom") - return "ok" - - return flaky - - for wait in (None, 0, sum([])): - with self.subTest(wait=wait): - flaky = make_flaky() - r = Retrying( - wait=wait, # type: ignore[arg-type] - stop=tenacity.stop_after_attempt(3), - ) - self.assertEqual(r(flaky), "ok") - - def test_wait_radd_plain_callable(self) -> None: - # A plain callable is a valid WaitBaseT, and functions have no - # __add__, so `callable + strategy` goes through wait_base.__radd__. - def cb(retry_state: RetryCallState) -> float: - return 2.0 - - combined = cb + tenacity.wait_fixed(1) - self.assertIsInstance(combined, tenacity.wait_combine) - self.assertEqual(combined(make_retry_state(1, 5)), 3.0) - - def test_wait_combine_passes_state_positionally(self) -> None: - # A WaitBaseT callable only promises to take the state positionally; - # its parameter name is its own business. - combined = tenacity.wait_combine( - tenacity.wait_fixed(1), - lambda rs: 2.0, - ) - self.assertEqual(combined(make_retry_state(1, 5)), 3.0) - - def test_wait_radd_rejects_non_zero_number(self) -> None: - with self.assertRaises(TypeError): - # Statically accepted -- see the comment on wait_base.__radd__ -- - # so the runtime rejection is what has to be tested. - 5 + tenacity.wait_fixed(1) - - def _assert_range(self, wait: float, min_: float, max_: float) -> None: - self.assertLess(wait, max_) - self.assertGreaterEqual(wait, min_) - - def _assert_inclusive_range(self, wait: float, low: float, high: float) -> None: - self.assertLessEqual(wait, high) - self.assertGreaterEqual(wait, low) - - def _assert_inclusive_epsilon( - self, wait: float, target: float, epsilon: float - ) -> None: - self.assertLessEqual(wait, target + epsilon) - self.assertGreaterEqual(wait, target - epsilon) - - def test_wait_chain(self) -> None: - r = Retrying( - wait=tenacity.wait_chain( - *[tenacity.wait_fixed(1) for i in range(2)] - + [tenacity.wait_fixed(4) for i in range(2)] - + [tenacity.wait_fixed(8) for i in range(1)] - ) - ) - - for i in range(10): - w = r.wait(make_retry_state(i + 1, 1)) - if i < 2: - self._assert_range(w, 1, 2) - elif i < 4: - self._assert_range(w, 4, 5) - else: - self._assert_range(w, 8, 9) - - def test_wait_chain_multiple_invocations(self) -> None: - sleep_intervals: list[float] = [] - r = Retrying( - sleep=sleep_intervals.append, - wait=tenacity.wait_chain(*[tenacity.wait_fixed(i + 1) for i in range(3)]), - stop=tenacity.stop_after_attempt(5), - retry=tenacity.retry_if_result(lambda x: x == 1), - ) - - @r.wraps - def always_return_1() -> int: - return 1 - - self.assertRaises(tenacity.RetryError, always_return_1) - self.assertEqual(sleep_intervals, [1.0, 2.0, 3.0, 3.0]) - sleep_intervals[:] = [] - - def test_wait_chain_requires_at_least_one_strategy(self) -> None: - with self.assertRaises(ValueError): - tenacity.wait_chain() - # Confirm the wrapped Retrying path surfaces the same ValueError - # instead of an opaque IndexError from self.strategies[-1]. - with self.assertRaises(ValueError): - Retrying(wait=tenacity.wait_chain()) - - def test_wait_random_exponential(self) -> None: - fn = tenacity.wait_random_exponential(0.5, 60.0) - - for _ in range(1000): - self._assert_inclusive_range(fn(make_retry_state(1, 0)), 0, 0.5) - self._assert_inclusive_range(fn(make_retry_state(2, 0)), 0, 1.0) - self._assert_inclusive_range(fn(make_retry_state(3, 0)), 0, 2.0) - self._assert_inclusive_range(fn(make_retry_state(4, 0)), 0, 4.0) - self._assert_inclusive_range(fn(make_retry_state(5, 0)), 0, 8.0) - self._assert_inclusive_range(fn(make_retry_state(6, 0)), 0, 16.0) - self._assert_inclusive_range(fn(make_retry_state(7, 0)), 0, 32.0) - self._assert_inclusive_range(fn(make_retry_state(8, 0)), 0, 60.0) - self._assert_inclusive_range(fn(make_retry_state(9, 0)), 0, 60.0) - - # max wait - max_wait = 5 - fn = tenacity.wait_random_exponential(10, max_wait) - for _ in range(1000): - self._assert_inclusive_range(fn(make_retry_state(1, 0)), 0.00, max_wait) - - # min wait - min_wait = 5 - fn = tenacity.wait_random_exponential(min=min_wait) - for _ in range(1000): - self._assert_inclusive_range(fn(make_retry_state(1, 0)), min_wait, 5) - - # Default arguments exist - fn = tenacity.wait_random_exponential() - fn(make_retry_state(0, 0)) - - def test_wait_random_exponential_statistically(self) -> None: - fn = tenacity.wait_random_exponential(0.5, 60.0) - - attempt = [[fn(make_retry_state(i, 0)) for _ in range(4000)] for i in range(10)] - - def mean(lst: list[float]) -> float: - return float(sum(lst)) / float(len(lst)) - - # skipping attempt 0 - self._assert_inclusive_epsilon(mean(attempt[1]), 0.25, 0.02) - self._assert_inclusive_epsilon(mean(attempt[2]), 0.50, 0.04) - self._assert_inclusive_epsilon(mean(attempt[3]), 1, 0.08) - self._assert_inclusive_epsilon(mean(attempt[4]), 2, 0.16) - self._assert_inclusive_epsilon(mean(attempt[5]), 4, 0.32) - self._assert_inclusive_epsilon(mean(attempt[6]), 8, 0.64) - self._assert_inclusive_epsilon(mean(attempt[7]), 16, 1.28) - self._assert_inclusive_epsilon(mean(attempt[8]), 30, 2.56) - self._assert_inclusive_epsilon(mean(attempt[9]), 30, 2.56) - - def test_wait_exponential_jitter(self) -> None: - fn = tenacity.wait_exponential_jitter(max=60) - - for _ in range(1000): - self._assert_inclusive_range(fn(make_retry_state(1, 0)), 1, 2) - self._assert_inclusive_range(fn(make_retry_state(2, 0)), 2, 3) - self._assert_inclusive_range(fn(make_retry_state(3, 0)), 4, 5) - self._assert_inclusive_range(fn(make_retry_state(4, 0)), 8, 9) - self._assert_inclusive_range(fn(make_retry_state(5, 0)), 16, 17) - self._assert_inclusive_range(fn(make_retry_state(6, 0)), 32, 33) - self.assertEqual(fn(make_retry_state(7, 0)), 60) - self.assertEqual(fn(make_retry_state(8, 0)), 60) - self.assertEqual(fn(make_retry_state(9, 0)), 60) - - with self.assertWarns(DeprecationWarning): - fn = tenacity.wait_exponential_jitter(10, 5) - for _ in range(1000): - self.assertEqual(fn(make_retry_state(1, 0)), 5) - - # Default arguments exist - fn = tenacity.wait_exponential_jitter() - fn(make_retry_state(0, 0)) - - def test_wait_exponential_jitter_min(self) -> None: - fn = tenacity.wait_exponential_jitter(initial=1, max=60, jitter=1, min=5) - for _ in range(1000): - # Even for attempt 1 (base wait=1 + jitter 0..1 = 1..2), min=5 applies - self._assert_inclusive_range(fn(make_retry_state(1, 0)), 5, 5) - self._assert_inclusive_range(fn(make_retry_state(2, 0)), 5, 5) - self._assert_inclusive_range(fn(make_retry_state(3, 0)), 5, 5) - # For attempt 4, base wait=8 + jitter 0..1 = 8..9, above min - self._assert_inclusive_range(fn(make_retry_state(4, 0)), 8, 9) - - def test_wait_exponential_jitter_timedelta(self) -> None: - from datetime import timedelta - - fn = tenacity.wait_exponential_jitter( - max=timedelta(seconds=60), - jitter=timedelta(seconds=1), - min=timedelta(seconds=5), - ) - for _ in range(1000): - self._assert_inclusive_range(fn(make_retry_state(1, 0)), 5, 5) - self._assert_inclusive_range(fn(make_retry_state(5, 0)), 16, 17) - self.assertEqual(fn(make_retry_state(7, 0)), 60) - - def test_wait_exponential_jitter_multiplier(self) -> None: - fn = tenacity.wait_exponential_jitter(multiplier=10, max=60, jitter=0) - self.assertEqual(fn(make_retry_state(1, 0)), 10) - self.assertEqual(fn(make_retry_state(2, 0)), 20) - self.assertEqual(fn(make_retry_state(3, 0)), 40) - self.assertEqual(fn(make_retry_state(4, 0)), 60) - - def test_wait_exponential_jitter_initial_deprecated(self) -> None: - with self.assertWarns(DeprecationWarning): - fn = tenacity.wait_exponential_jitter(initial=10, max=60, jitter=0) - self.assertEqual(fn(make_retry_state(1, 0)), 10) - self.assertEqual(fn(make_retry_state(2, 0)), 20) - - def test_wait_exponential_jitter_initial_and_multiplier_raises(self) -> None: - with self.assertRaises(ValueError): - tenacity.wait_exponential_jitter(initial=5, multiplier=10) - - def test_wait_retry_state_attributes(self) -> None: - class ExtractCallState(Exception): - pass - - # retry_state is mutable, so return it as an exception to extract the - # exact values it has when wait is called and bypass any other logic. - def waitfunc(retry_state: RetryCallState) -> float: - raise ExtractCallState(retry_state) - - retrying = Retrying( - wait=waitfunc, - retry=( - tenacity.retry_if_exception_type() - | tenacity.retry_if_result(lambda result: result == 123) - ), - ) - - def returnval() -> int: - return 123 - - with self.assertRaises(ExtractCallState) as caught: - retrying(returnval) - retry_state = caught.exception.args[0] - self.assertIs(retry_state.fn, returnval) - self.assertEqual(retry_state.args, ()) - self.assertEqual(retry_state.kwargs, {}) - self.assertEqual(retry_state.outcome.result(), 123) - self.assertEqual(retry_state.attempt_number, 1) - self.assertGreaterEqual(retry_state.outcome_timestamp, retry_state.start_time) - - def dying() -> None: - raise Exception("Broken") - - with self.assertRaises(ExtractCallState) as caught: - retrying(dying) - retry_state = caught.exception.args[0] - self.assertIs(retry_state.fn, dying) - self.assertEqual(retry_state.args, ()) - self.assertEqual(retry_state.kwargs, {}) - self.assertEqual(str(retry_state.outcome.exception()), "Broken") - self.assertEqual(retry_state.attempt_number, 1) - self.assertGreaterEqual(retry_state.outcome_timestamp, retry_state.start_time) - - -class TestRetryConditions(unittest.TestCase): - def test_retry_if_result(self) -> None: - retry = tenacity.retry_if_result(lambda x: x == 1) - - def r(fut: tenacity.Future) -> bool: - retry_state = make_retry_state(1, 1.0, last_result=fut) - return retry(retry_state) - - self.assertTrue(r(tenacity.Future.construct(1, 1, False))) - self.assertFalse(r(tenacity.Future.construct(1, 2, False))) - - def test_retry_if_not_result(self) -> None: - retry = tenacity.retry_if_not_result(lambda x: x == 1) - - def r(fut: tenacity.Future) -> bool: - retry_state = make_retry_state(1, 1.0, last_result=fut) - return retry(retry_state) - - self.assertTrue(r(tenacity.Future.construct(1, 2, False))) - self.assertFalse(r(tenacity.Future.construct(1, 1, False))) - - def test_retry_any(self) -> None: - retry = tenacity.retry_any( - tenacity.retry_if_result(lambda x: x == 1), - tenacity.retry_if_result(lambda x: x == 2), - ) - - def r(fut: tenacity.Future) -> bool: - retry_state = make_retry_state(1, 1.0, last_result=fut) - return retry(retry_state) - - self.assertTrue(r(tenacity.Future.construct(1, 1, False))) - self.assertTrue(r(tenacity.Future.construct(1, 2, False))) - self.assertFalse(r(tenacity.Future.construct(1, 3, False))) - self.assertFalse(r(tenacity.Future.construct(1, 1, True))) - - def test_retry_all(self) -> None: - retry = tenacity.retry_all( - tenacity.retry_if_result(lambda x: x == 1), - tenacity.retry_if_result(lambda x: isinstance(x, int)), - ) - - def r(fut: tenacity.Future) -> bool: - retry_state = make_retry_state(1, 1.0, last_result=fut) - return retry(retry_state) - - self.assertTrue(r(tenacity.Future.construct(1, 1, False))) - self.assertFalse(r(tenacity.Future.construct(1, 2, False))) - self.assertFalse(r(tenacity.Future.construct(1, 3, False))) - self.assertFalse(r(tenacity.Future.construct(1, 1, True))) - - def test_retry_and(self) -> None: - retry = tenacity.retry_if_result(lambda x: x == 1) & tenacity.retry_if_result( - lambda x: isinstance(x, int) - ) - - def r(fut: tenacity.Future) -> bool: - retry_state = make_retry_state(1, 1.0, last_result=fut) - return retry(retry_state) - - self.assertTrue(r(tenacity.Future.construct(1, 1, False))) - self.assertFalse(r(tenacity.Future.construct(1, 2, False))) - self.assertFalse(r(tenacity.Future.construct(1, 3, False))) - self.assertFalse(r(tenacity.Future.construct(1, 1, True))) - - def test_retry_or(self) -> None: - retry = tenacity.retry_if_result( - lambda x: x == "foo" - ) | tenacity.retry_if_result(lambda x: isinstance(x, int)) - - def r(fut: tenacity.Future) -> bool: - retry_state = make_retry_state(1, 1.0, last_result=fut) - return retry(retry_state) - - self.assertTrue(r(tenacity.Future.construct(1, "foo", False))) - self.assertFalse(r(tenacity.Future.construct(1, "foobar", False))) - self.assertFalse(r(tenacity.Future.construct(1, 2.2, False))) - self.assertFalse(r(tenacity.Future.construct(1, 42, True))) - - def test_retry_or_with_plain_function(self) -> None: - """Plain callables can be composed with retry_base via |.""" - - def my_retry(retry_state: tenacity.RetryCallState) -> bool: - return retry_state.outcome is not None and not retry_state.outcome.failed - - # retry_base | plain_callable (exercises __or__ fallback) - retry = tenacity.retry_if_exception_type(Exception) | my_retry - retry_state = make_retry_state( - 1, 1.0, last_result=tenacity.Future.construct(1, "ok", False) - ) - self.assertTrue(retry(retry_state)) - - # plain_callable | retry_base (exercises __ror__ via reflection) - retry2 = my_retry | tenacity.retry_if_exception_type(Exception) - self.assertTrue(retry2(retry_state)) - - def test_retry_and_with_plain_function(self) -> None: - """Plain callables can be composed with retry_base via &.""" - - def my_retry(retry_state: tenacity.RetryCallState) -> bool: - return True - - # retry_base & plain_callable (exercises __and__ fallback) - retry = tenacity.retry_if_result(lambda x: x == 1) & my_retry - retry_state = make_retry_state( - 1, 1.0, last_result=tenacity.Future.construct(1, 1, False) - ) - self.assertTrue(retry(retry_state)) - - # plain_callable & retry_base (exercises __rand__ via reflection) - retry2 = my_retry & tenacity.retry_if_result(lambda x: x == 1) - self.assertTrue(retry2(retry_state)) - - def test_retry_or_coalesces(self) -> None: - """Multiple | operations flatten into a single retry_any.""" - a = tenacity.retry_if_exception_type(IOError) - b = tenacity.retry_if_exception_type(OSError) - c = tenacity.retry_if_exception_type(ValueError) - - combined = a | b | c - self.assertIsInstance(combined, retry_any) - self.assertEqual(len(combined.retries), 3) - - def test_retry_and_coalesces(self) -> None: - """Multiple & operations flatten into a single retry_all.""" - a = tenacity.retry_if_result(lambda x: x == 1) - b = tenacity.retry_if_result(lambda x: x > 0) - c = tenacity.retry_if_result(lambda x: x < 10) - - combined = a & b & c - self.assertIsInstance(combined, retry_all) - self.assertEqual(len(combined.retries), 3) - - def _raise_try_again(self) -> None: - self._attempts += 1 - if self._attempts < 3: - raise tenacity.TryAgain - - def test_retry_try_again(self) -> None: - self._attempts = 0 - Retrying(stop=tenacity.stop_after_attempt(5), retry=tenacity.retry_never)( - self._raise_try_again - ) - self.assertEqual(3, self._attempts) - - def test_retry_try_again_forever(self) -> None: - def _r() -> None: - raise tenacity.TryAgain - - r = Retrying(stop=tenacity.stop_after_attempt(5), retry=tenacity.retry_never) - self.assertRaises(tenacity.RetryError, r, _r) - self.assertEqual(5, r.statistics["attempt_number"]) - - def test_retry_try_again_with_cause_reraise(self) -> None: - # When TryAgain is raised from within an "except" block, reraise=True - # should surface the underlying exception rather than TryAgain itself. - class UnderlyingError(Exception): - pass - - def _r() -> None: - try: - raise UnderlyingError("boom") - except UnderlyingError: - # Implicit chaining via __context__ is exactly what we test. - raise tenacity.TryAgain # noqa: B904 - - r = Retrying( - stop=tenacity.stop_after_attempt(5), - retry=tenacity.retry_never, - reraise=True, - ) - self.assertRaises(UnderlyingError, r, _r) - self.assertEqual(5, r.statistics["attempt_number"]) - - def test_retry_try_again_from_cause_reraise(self) -> None: - # An explicit "raise TryAgain from exc" should also be unwrapped. - class UnderlyingError(Exception): - pass - - def _r() -> None: - raise tenacity.TryAgain from UnderlyingError("boom") - - r = Retrying( - stop=tenacity.stop_after_attempt(5), - retry=tenacity.retry_never, - reraise=True, - ) - self.assertRaises(UnderlyingError, r, _r) - self.assertEqual(5, r.statistics["attempt_number"]) - - def test_retry_try_again_forever_reraise(self) -> None: - def _r() -> None: - raise tenacity.TryAgain - - r = Retrying( - stop=tenacity.stop_after_attempt(5), - retry=tenacity.retry_never, - reraise=True, - ) - self.assertRaises(tenacity.TryAgain, r, _r) - self.assertEqual(5, r.statistics["attempt_number"]) - - def test_retry_if_exception_message_negative_no_inputs(self) -> None: - with self.assertRaises(TypeError): - tenacity.retry_if_exception_message() - - def test_retry_if_exception_message_negative_too_many_inputs(self) -> None: - with self.assertRaises(TypeError): - tenacity.retry_if_exception_message(message="negative", match="negative") - - def test_retry_if_exception_message_empty_message(self) -> None: - # An exception whose str() is empty (e.g. a bare ``RuntimeError()``) is - # a valid target. ``message=""`` must be accepted and match it, rather - # than being treated as "no message given" by a truthiness check. - r = tenacity.retry_if_exception_message(message="") - self.assertEqual(r.message, "") - empty = make_retry_state( - 1, 0, last_result=tenacity.Future.construct(1, RuntimeError(), True) - ) - nonempty = make_retry_state( - 1, 0, last_result=tenacity.Future.construct(1, RuntimeError("boom"), True) - ) - self.assertTrue(r(empty)) - self.assertFalse(r(nonempty)) - - def test_retry_if_not_exception_message_empty_message(self) -> None: - r = tenacity.retry_if_not_exception_message(message="") - empty = make_retry_state( - 1, 0, last_result=tenacity.Future.construct(1, RuntimeError(), True) - ) - nonempty = make_retry_state( - 1, 0, last_result=tenacity.Future.construct(1, RuntimeError("boom"), True) - ) - self.assertFalse(r(empty)) - self.assertTrue(r(nonempty)) - - -class NoneReturnUntilAfterCount: - """Holds counter state for invoking a method several times in a row.""" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go(self) -> typing.Any: - """Return None until after count threshold has been crossed. - - Then return True. - """ - if self.counter < self.count: - self.counter += 1 - return None - return True - - -class NoIOErrorAfterCount: - """Holds counter state for invoking a method several times in a row.""" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go(self) -> typing.Any: - """Raise an IOError until after count threshold has been crossed. - - Then return True. - """ - if self.counter < self.count: - self.counter += 1 - raise OSError("Hi there, I'm an IOError") - return True - - -class NoNameErrorAfterCount: - """Holds counter state for invoking a method several times in a row.""" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go(self) -> typing.Any: - """Raise a NameError until after count threshold has been crossed. - - Then return True. - """ - if self.counter < self.count: - self.counter += 1 - raise NameError("Hi there, I'm a NameError") - return True - - -class NoNameErrorCauseAfterCount: - """Holds counter state for invoking a method several times in a row.""" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go2(self) -> typing.Any: - raise NameError("Hi there, I'm a NameError") - - def go(self) -> typing.Any: - """Raise an IOError with a NameError as cause until after count threshold has been crossed. - - Then return True. - """ - if self.counter < self.count: - self.counter += 1 - try: - self.go2() - except NameError as e: - raise OSError from e - - return True - - -class NoIOErrorCauseAfterCount: - """Holds counter state for invoking a method several times in a row.""" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go2(self) -> typing.Any: - raise OSError("Hi there, I'm an IOError") - - def go(self) -> typing.Any: - """Raise a NameError with an IOError as cause until after count threshold has been crossed. - - Then return True. - """ - if self.counter < self.count: - self.counter += 1 - try: - self.go2() - except OSError as e: - raise NameError from e - - return True - - -class NameErrorUntilCount: - """Holds counter state for invoking a method several times in a row.""" - - derived_message = "Hi there, I'm a NameError" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go(self) -> typing.Any: - """Return True until after count threshold has been crossed. - - Then raise a NameError. - """ - if self.counter < self.count: - self.counter += 1 - return True - raise NameError(self.derived_message) - - -class IOErrorUntilCount: - """Holds counter state for invoking a method several times in a row.""" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go(self) -> typing.Any: - """Return True until after count threshold has been crossed. - - Then raise an IOError. - """ - if self.counter < self.count: - self.counter += 1 - return True - raise OSError("Hi there, I'm an IOError") - - -class CustomError(Exception): - """This is a custom exception class. - - Note that For Python 2.x, we don't strictly need to extend BaseException, - however, Python 3.x will complain. While this test suite won't run - correctly under Python 3.x without extending from the Python exception - hierarchy, the actual module code is backwards compatible Python 2.x and - will allow for cases where exception classes don't extend from the - hierarchy. - """ - - def __init__(self, value: str) -> None: - self.value = value - - @override - def __str__(self) -> str: - return self.value - - -class NoCustomErrorAfterCount: - """Holds counter state for invoking a method several times in a row.""" - - derived_message = "This is a Custom exception class" - - def __init__(self, count: int) -> None: - self.counter = 0 - self.count = count - - def go(self) -> typing.Any: - """Raise a CustomError until after count threshold has been crossed. - - Then return True. - """ - if self.counter < self.count: - self.counter += 1 - raise CustomError(self.derived_message) - return True - - -class CapturingHandler(logging.Handler): - """Captures log records for inspection.""" - - def __init__(self, *args: typing.Any, **kwargs: typing.Any) -> None: - super().__init__(*args, **kwargs) - self.records: list[logging.LogRecord] = [] - - @override - def emit(self, record: logging.LogRecord) -> None: - self.records.append(record) - - -def current_time_ms() -> int: - return round(time.time() * 1000) - - -@retry( - wait=tenacity.wait_fixed(0.05), - retry=tenacity.retry_if_result(lambda result: result is None), -) -def _retryable_test_with_wait(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - stop=tenacity.stop_after_attempt(3), - retry=tenacity.retry_if_result(lambda result: result is None), -) -def _retryable_test_with_stop(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry(retry=tenacity.retry_if_exception_cause_type(NameError)) -def _retryable_test_with_exception_cause_type(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry(retry=tenacity.retry_if_exception_type(IOError)) -def _retryable_test_with_exception_type_io(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry(retry=tenacity.retry_if_not_exception_type(IOError)) -def _retryable_test_if_not_exception_type_io(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - stop=tenacity.stop_after_attempt(3), retry=tenacity.retry_if_exception_type(IOError) -) -def _retryable_test_with_exception_type_io_attempt_limit( - thing: typing.Any, -) -> typing.Any: - return thing.go() - - -@retry(retry=tenacity.retry_unless_exception_type(NameError)) -def _retryable_test_with_unless_exception_type_name(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - stop=tenacity.stop_after_attempt(3), - retry=tenacity.retry_unless_exception_type(NameError), -) -def _retryable_test_with_unless_exception_type_name_attempt_limit( - thing: typing.Any, -) -> typing.Any: - return thing.go() - - -@retry(retry=tenacity.retry_unless_exception_type()) -def _retryable_test_with_unless_exception_type_no_input( - thing: typing.Any, -) -> typing.Any: - return thing.go() - - -@retry( - stop=tenacity.stop_after_attempt(5), - retry=tenacity.retry_if_exception_message( - message=NoCustomErrorAfterCount.derived_message - ), -) -def _retryable_test_if_exception_message_message(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - retry=tenacity.retry_if_not_exception_message( - message=NoCustomErrorAfterCount.derived_message - ) -) -def _retryable_test_if_not_exception_message_message(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - retry=tenacity.retry_if_exception_message( - match=NoCustomErrorAfterCount.derived_message[:3] + ".*" - ) -) -def _retryable_test_if_exception_message_match(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - retry=tenacity.retry_if_not_exception_message( - match=NoCustomErrorAfterCount.derived_message[:3] + ".*" - ) -) -def _retryable_test_if_not_exception_message_match(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - retry=tenacity.retry_if_not_exception_message( - message=NameErrorUntilCount.derived_message - ) -) -def _retryable_test_not_exception_message_delay(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry -def _retryable_default(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry() -def _retryable_default_f(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry(retry=tenacity.retry_if_exception_type(CustomError)) -def _retryable_test_with_exception_type_custom(thing: typing.Any) -> typing.Any: - return thing.go() - - -@retry( - stop=tenacity.stop_after_attempt(3), - retry=tenacity.retry_if_exception_type(CustomError), -) -def _retryable_test_with_exception_type_custom_attempt_limit( - thing: typing.Any, -) -> typing.Any: - return thing.go() - - -class TestDecoratorWrapper(unittest.TestCase): - def test_with_wait(self) -> None: - start = current_time_ms() - result = _retryable_test_with_wait(NoneReturnUntilAfterCount(5)) - t = current_time_ms() - start - self.assertGreaterEqual(t, 250) - self.assertTrue(result) - - def test_with_stop_on_return_value(self) -> None: - try: - _retryable_test_with_stop(NoneReturnUntilAfterCount(5)) - self.fail("Expected RetryError after 3 attempts") - except RetryError as re: - self.assertFalse(re.last_attempt.failed) - self.assertEqual(3, re.last_attempt.attempt_number) - self.assertTrue(re.last_attempt.result() is None) - print(re) - - def test_with_stop_on_exception(self) -> None: - try: - _retryable_test_with_stop(NoIOErrorAfterCount(5)) - self.fail("Expected IOError") - except OSError as re: - self.assertTrue(isinstance(re, IOError)) - print(re) - - def test_retry_if_exception_of_type(self) -> None: - self.assertTrue(_retryable_test_with_exception_type_io(NoIOErrorAfterCount(5))) - - try: - _retryable_test_with_exception_type_io(NoNameErrorAfterCount(5)) - self.fail("Expected NameError") - except NameError as n: - self.assertTrue(isinstance(n, NameError)) - print(n) - - self.assertTrue( - _retryable_test_with_exception_type_custom(NoCustomErrorAfterCount(5)) - ) - - try: - _retryable_test_with_exception_type_custom(NoNameErrorAfterCount(5)) - self.fail("Expected NameError") - except NameError as n: - self.assertTrue(isinstance(n, NameError)) - print(n) - - def test_retry_except_exception_of_type(self) -> None: - self.assertTrue( - _retryable_test_if_not_exception_type_io(NoNameErrorAfterCount(5)) - ) - - try: - _retryable_test_if_not_exception_type_io(NoIOErrorAfterCount(5)) - self.fail("Expected IOError") - except OSError as err: - self.assertTrue(isinstance(err, IOError)) - print(err) - - def test_retry_until_exception_of_type_attempt_number(self) -> None: - try: - self.assertTrue( - _retryable_test_with_unless_exception_type_name(NameErrorUntilCount(5)) - ) - except NameError as e: - s = _retryable_test_with_unless_exception_type_name.statistics - self.assertTrue(s["attempt_number"] == 6) - print(e) - else: - self.fail("Expected NameError") - - def test_retry_until_exception_of_type_no_type(self) -> None: - try: - # no input should catch all subclasses of Exception - self.assertTrue( - _retryable_test_with_unless_exception_type_no_input( - NameErrorUntilCount(5) - ) - ) - except NameError as e: - s = _retryable_test_with_unless_exception_type_no_input.statistics - self.assertTrue(s["attempt_number"] == 6) - print(e) - else: - self.fail("Expected NameError") - - def test_retry_until_exception_of_type_wrong_exception(self) -> None: - try: - # two iterations with IOError, one that returns True - _retryable_test_with_unless_exception_type_name_attempt_limit( - IOErrorUntilCount(2) - ) - self.fail("Expected RetryError") - except RetryError as e: - self.assertTrue(isinstance(e, RetryError)) - print(e) - - def test_retry_if_exception_message(self) -> None: - try: - self.assertTrue( - _retryable_test_if_exception_message_message(NoCustomErrorAfterCount(3)) - ) - except CustomError: - print(_retryable_test_if_exception_message_message.statistics) - self.fail("CustomError should've been retried from errormessage") - - def test_retry_if_not_exception_message(self) -> None: - try: - self.assertTrue( - _retryable_test_if_not_exception_message_message( - NoCustomErrorAfterCount(2) - ) - ) - except CustomError: - s = _retryable_test_if_not_exception_message_message.statistics - self.assertTrue(s["attempt_number"] == 1) - - def test_retry_if_not_exception_message_delay(self) -> None: - try: - self.assertTrue( - _retryable_test_not_exception_message_delay(NameErrorUntilCount(3)) - ) - except NameError: - s = _retryable_test_not_exception_message_delay.statistics - print(s["attempt_number"]) - self.assertTrue(s["attempt_number"] == 4) - - def test_retry_if_exception_message_match(self) -> None: - try: - self.assertTrue( - _retryable_test_if_exception_message_match(NoCustomErrorAfterCount(3)) - ) - except CustomError: - self.fail("CustomError should've been retried from errormessage") - - def test_retry_if_not_exception_message_match(self) -> None: - try: - self.assertTrue( - _retryable_test_if_not_exception_message_message( - NoCustomErrorAfterCount(2) - ) - ) - except CustomError: - s = _retryable_test_if_not_exception_message_message.statistics - self.assertTrue(s["attempt_number"] == 1) - - def test_retry_if_exception_cause_type(self) -> None: - self.assertTrue( - _retryable_test_with_exception_cause_type(NoNameErrorCauseAfterCount(5)) - ) - - try: - _retryable_test_with_exception_cause_type(NoIOErrorCauseAfterCount(5)) - self.fail("Expected exception without NameError as cause") - except NameError: - pass - - def test_retry_if_exception_cause_type_handles_cause_cycles(self) -> None: - """Cyclic __cause__ chains must not hang the retry predicate (#658).""" - - def boom_self_cause() -> None: - try: - raise ValueError("inner") - except ValueError as e: - raise e from e - - def boom_two_node_cycle() -> None: - a = ValueError("a") - b = RuntimeError("b") - a.__cause__ = b - b.__cause__ = a - raise a - - for boom in (boom_self_cause, boom_two_node_cycle): - r = tenacity.Retrying( - retry=tenacity.retry_if_exception_cause_type(KeyError), - stop=tenacity.stop_after_attempt(2), - reraise=True, - ) - with self.assertRaises((ValueError, RuntimeError)): - r(boom) - - def test_retry_preserves_argument_defaults(self) -> None: - def function_with_defaults(a: int = 1) -> int: - return a - - def function_with_kwdefaults(*, a: int = 1) -> int: - return a - - retrying = Retrying( - wait=tenacity.wait_fixed(0.01), stop=tenacity.stop_after_attempt(3) - ) - wrapped_defaults_function = retrying.wraps(function_with_defaults) - wrapped_kwdefaults_function = retrying.wraps(function_with_kwdefaults) - - self.assertEqual( - function_with_defaults.__defaults__, - wrapped_defaults_function.__defaults__, # type: ignore[attr-defined] - ) - self.assertEqual( - function_with_kwdefaults.__kwdefaults__, - wrapped_kwdefaults_function.__kwdefaults__, # type: ignore[attr-defined] - ) - - def test_defaults(self) -> None: - self.assertTrue(_retryable_default(NoNameErrorAfterCount(5))) - self.assertTrue(_retryable_default_f(NoNameErrorAfterCount(5))) - self.assertTrue(_retryable_default(NoCustomErrorAfterCount(5))) - self.assertTrue(_retryable_default_f(NoCustomErrorAfterCount(5))) - - def test_retry_function_object(self) -> None: - """Test that functools.wraps doesn't cause problems with callable objects. - - It raises an error upon trying to wrap it in Py2, because __name__ - attribute is missing. It's fixed in Py3 but was never backported. - """ - - class Hello: - def __call__(self) -> str: - return "Hello" - - retrying = Retrying( - wait=tenacity.wait_fixed(0.01), stop=tenacity.stop_after_attempt(3) - ) - h = retrying.wraps(Hello()) - self.assertEqual(h(), "Hello") - - def test_retry_function_attributes(self) -> None: - """Test that the wrapped function attributes are exposed as intended. - - - statistics contains the value for the latest function run - - retry object can be modified to change its behaviour (useful to patch in tests) - - retry object statistics are synced with function statistics - """ - - self.assertTrue(_retryable_test_with_stop(NoneReturnUntilAfterCount(2))) - - expected_stats = { - "attempt_number": 3, - "delay_since_first_attempt": mock.ANY, - "idle_for": mock.ANY, - "start_time": mock.ANY, - } - self.assertEqual(_retryable_test_with_stop.statistics, expected_stats) - self.assertEqual(_retryable_test_with_stop.retry.statistics, expected_stats) - - with mock.patch.object( - _retryable_test_with_stop.retry, - "stop", - tenacity.stop_after_attempt(1), - ): - try: - self.assertTrue(_retryable_test_with_stop(NoneReturnUntilAfterCount(2))) - except RetryError as exc: - expected_stats = { - "attempt_number": 1, - "delay_since_first_attempt": mock.ANY, - "idle_for": mock.ANY, - "start_time": mock.ANY, - } - self.assertEqual(_retryable_test_with_stop.statistics, expected_stats) - self.assertEqual(exc.last_attempt.attempt_number, 1) - self.assertEqual( - _retryable_test_with_stop.retry.statistics, expected_stats - ) - else: - self.fail("RetryError should have been raised after 1 attempt") - - -class TestStatisticsKeys: - def test_delay_since_first_attempt_available_on_first_attempt(self) -> None: - """delay_since_first_attempt should be in statistics from the start.""" - - @retry( - stop=tenacity.stop_after_attempt(3), - retry=tenacity.retry_if_result(lambda x: x is None), - ) - def succeeds_first_try() -> bool: - assert "delay_since_first_attempt" in succeeds_first_try.statistics - assert succeeds_first_try.statistics["delay_since_first_attempt"] == 0 - return True - - succeeds_first_try() - assert succeeds_first_try.statistics["delay_since_first_attempt"] == 0 - - def test_statistics_visible_through_outer_decorator(self) -> None: - """Statistics must resolve when @retry is wrapped by another decorator. - - A well-behaved outer decorator uses functools.wraps, which copies the - inner wrapper's ``__dict__`` (including ``statistics``). Rebinding the - attribute on each call left the outer wrapper pointing at a stale empty - dict. The statistics must instead stay visible through the wrapper - chain. See issue #519. - """ - import functools - - _F = typing.TypeVar("_F", bound=typing.Callable[..., typing.Any]) - - def outer(fn: _F) -> _F: - @functools.wraps(fn) - def wrapper(*args: typing.Any, **kwargs: typing.Any) -> typing.Any: - return fn(*args, **kwargs) - - return typing.cast("_F", wrapper) - - @outer - @retry(stop=tenacity.stop_after_attempt(3)) - def my_call() -> str: - return "ok" - - assert my_call() == "ok" - assert my_call.statistics["attempt_number"] == 1 - assert my_call.statistics is my_call.__wrapped__.statistics - - -class TestEnabled: - def test_enabled_false_skips_retry(self) -> None: - """When enabled=False, the function is called directly without retrying.""" - call_count = 0 - - @retry(enabled=False, stop=tenacity.stop_after_attempt(3)) - def always_fails() -> None: - nonlocal call_count - call_count += 1 - raise ValueError("fail") - - with pytest.raises(ValueError, match="fail"): - always_fails() - assert call_count == 1 - - def test_enabled_false_preserves_attributes(self) -> None: - """When enabled=False, .retry, .retry_with, .statistics are still available.""" - - @retry(enabled=False, stop=tenacity.stop_after_attempt(3)) - def my_func() -> str: - return "ok" - - assert hasattr(my_func, "retry") - assert hasattr(my_func, "retry_with") - assert hasattr(my_func, "statistics") - assert my_func() == "ok" - - def test_enabled_false_via_retry_with(self) -> None: - """retry_with(enabled=False) disables retrying.""" - call_count = 0 - - @retry(stop=tenacity.stop_after_attempt(3)) - def always_fails() -> None: - nonlocal call_count - call_count += 1 - raise ValueError("fail") - - disabled = always_fails.retry_with(enabled=False) - with pytest.raises(ValueError, match="fail"): - disabled() - assert call_count == 1 - - def test_enabled_true_retries_normally(self) -> None: - """When enabled=True (default), retrying works as usual.""" - call_count = 0 - - @retry(enabled=True, stop=tenacity.stop_after_attempt(3), reraise=True) - def fails_twice() -> bool: - nonlocal call_count - call_count += 1 - if call_count < 3: - raise ValueError("fail") - return True - - assert fails_twice() is True - assert call_count == 3 - - def test_enabled_false_iter_raises_original_exception(self) -> None: - """When enabled=False, the iterator protocol raises the original exception, - not a RetryError, and the body executes exactly once.""" - call_count = 0 - retrying = Retrying( - enabled=False, - stop=tenacity.stop_after_attempt(5), - wait=tenacity.wait_none(), - ) - with pytest.raises(ValueError, match="fail"): - for attempt in retrying: - with attempt: - call_count += 1 - raise ValueError("fail") - assert call_count == 1 - - def test_enabled_false_iter_succeeds_on_first_attempt(self) -> None: - """When enabled=False, the iterator protocol runs the body once and stops.""" - call_count = 0 - retrying = Retrying( - enabled=False, - stop=tenacity.stop_after_attempt(5), - wait=tenacity.wait_none(), - ) - for attempt in retrying: - with attempt: - call_count += 1 - assert call_count == 1 - - def test_enabled_false_call_raises_original_exception(self) -> None: - """When enabled=False, calling the controller directly raises the original - exception, not a RetryError, and the function executes exactly once.""" - call_count = 0 - - def fails() -> None: - nonlocal call_count - call_count += 1 - raise ValueError("fail") - - retrying = Retrying( - enabled=False, - stop=tenacity.stop_after_attempt(5), - wait=tenacity.wait_none(), - ) - with pytest.raises(ValueError, match="fail"): - retrying(fails) - assert call_count == 1 - - def test_enabled_false_call_succeeds_on_first_attempt(self) -> None: - """When enabled=False, calling the controller directly runs the function once - and returns its result.""" - call_count = 0 - - def succeeds() -> str: - nonlocal call_count - call_count += 1 - return "ok" - - retrying = Retrying( - enabled=False, - stop=tenacity.stop_after_attempt(5), - wait=tenacity.wait_none(), - ) - assert retrying(succeeds) == "ok" - assert call_count == 1 - - -class TestRetryWith: - def test_redefine_wait(self) -> None: - start = current_time_ms() - result = _retryable_test_with_wait.retry_with(wait=tenacity.wait_fixed(0.1))( - NoneReturnUntilAfterCount(5) - ) - t = current_time_ms() - start - assert t >= 500 - assert result is True - - def test_redefine_stop(self) -> None: - result = _retryable_test_with_stop.retry_with( - stop=tenacity.stop_after_attempt(5) - )(NoneReturnUntilAfterCount(4)) - assert result is True - - def test_retry_error_cls_should_be_preserved(self) -> None: - @retry(stop=tenacity.stop_after_attempt(10), retry_error_cls=ValueError) # type: ignore[arg-type] - def _retryable() -> None: - raise Exception("raised for test purposes") - - with pytest.raises(Exception) as exc_ctx: - _retryable.retry_with(stop=tenacity.stop_after_attempt(2))() - - assert exc_ctx.type is ValueError, "Should remap to specific exception type" - - def test_retry_error_callback_should_be_preserved(self) -> None: - def return_text(retry_state: RetryCallState) -> str: - return f"Calling {retry_state.fn.__name__} keeps raising errors after {retry_state.attempt_number} attempts" # type: ignore[union-attr] - - @retry(stop=tenacity.stop_after_attempt(10), retry_error_callback=return_text) - def _retryable() -> None: - raise Exception("raised for test purposes") - - result = _retryable.retry_with(stop=tenacity.stop_after_attempt(5))() - assert result == "Calling _retryable keeps raising errors after 5 attempts" - - -class TestBeforeAfterAttempts(unittest.TestCase): - _attempt_number = 0 - - def test_before_attempts(self) -> None: - TestBeforeAfterAttempts._attempt_number = 0 - - def _before(retry_state: RetryCallState) -> None: - TestBeforeAfterAttempts._attempt_number = retry_state.attempt_number - - @retry( - wait=tenacity.wait_fixed(1), - stop=tenacity.stop_after_attempt(1), - before=_before, - ) - def _test_before() -> None: - pass - - _test_before() - - self.assertTrue(TestBeforeAfterAttempts._attempt_number == 1) - - def test_after_attempts(self) -> None: - TestBeforeAfterAttempts._attempt_number = 0 - - def _after(retry_state: RetryCallState) -> None: - TestBeforeAfterAttempts._attempt_number = retry_state.attempt_number - - @retry( - wait=tenacity.wait_fixed(0.1), - stop=tenacity.stop_after_attempt(3), - after=_after, - ) - def _test_after() -> None: - if TestBeforeAfterAttempts._attempt_number < 2: - raise Exception("testing after_attempts handler") - - _test_after() - - self.assertTrue(TestBeforeAfterAttempts._attempt_number == 2) - - def test_before_sleep(self) -> None: - def _before_sleep(retry_state: RetryCallState) -> None: - self.assertGreater(retry_state.next_action.sleep, 0) # type: ignore[union-attr] - _before_sleep.attempt_number = retry_state.attempt_number # type: ignore[attr-defined] - - @retry( - wait=tenacity.wait_fixed(0.01), - stop=tenacity.stop_after_attempt(3), - before_sleep=_before_sleep, - ) - def _test_before_sleep() -> None: - if _before_sleep.attempt_number < 2: # type: ignore[attr-defined] - raise Exception("testing before_sleep_attempts handler") - - _test_before_sleep() - self.assertEqual(_before_sleep.attempt_number, 2) # type: ignore[attr-defined] - - def _before_sleep_log_raises( - self, get_call_fn: typing.Callable[..., typing.Any] - ) -> None: - thing = NoIOErrorAfterCount(2) - logger = logging.getLogger(self.id()) - logger.propagate = False - logger.setLevel(logging.INFO) - handler = CapturingHandler() - logger.addHandler(handler) - try: - _before_sleep = tenacity.before_sleep_log(logger, logging.INFO) - retrying = Retrying( - wait=tenacity.wait_fixed(0.01), - stop=tenacity.stop_after_attempt(3), - before_sleep=_before_sleep, - ) - get_call_fn(retrying)(thing.go) - finally: - logger.removeHandler(handler) - - etalon_re = ( - r"^Retrying .* in 0\.01 seconds as it raised " - r"(IO|OS)Error: Hi there, I'm an IOError\.$" - ) - self.assertEqual(len(handler.records), 2) - fmt = logging.Formatter().format - self.assertRegex(fmt(handler.records[0]), etalon_re) - self.assertRegex(fmt(handler.records[1]), etalon_re) - - def test_before_sleep_log_raises(self) -> None: - self._before_sleep_log_raises(lambda x: x) - - def test_before_sleep_log_raises_with_exc_info(self) -> None: - thing = NoIOErrorAfterCount(2) - logger = logging.getLogger(self.id()) - logger.propagate = False - logger.setLevel(logging.INFO) - handler = CapturingHandler() - logger.addHandler(handler) - try: - _before_sleep = tenacity.before_sleep_log( - logger, logging.INFO, exc_info=True - ) - retrying = Retrying( - wait=tenacity.wait_fixed(0.01), - stop=tenacity.stop_after_attempt(3), - before_sleep=_before_sleep, - ) - retrying(thing.go) - finally: - logger.removeHandler(handler) - - etalon_re = re.compile( - r"^Retrying .* in 0\.01 seconds as it raised " - r"(IO|OS)Error: Hi there, I'm an IOError\.{0}" - r"Traceback \(most recent call last\):{0}" - r".*$".format("\n"), - flags=re.MULTILINE, - ) - self.assertEqual(len(handler.records), 2) - fmt = logging.Formatter().format - self.assertRegex(fmt(handler.records[0]), etalon_re) - self.assertRegex(fmt(handler.records[1]), etalon_re) - - def test_before_sleep_log_returns(self, exc_info: bool = False) -> None: - thing = NoneReturnUntilAfterCount(2) - logger = logging.getLogger(self.id()) - logger.propagate = False - logger.setLevel(logging.INFO) - handler = CapturingHandler() - logger.addHandler(handler) - try: - _before_sleep = tenacity.before_sleep_log( - logger, logging.INFO, exc_info=exc_info - ) - _retry = tenacity.retry_if_result(lambda result: result is None) - retrying = Retrying( - wait=tenacity.wait_fixed(0.01), - stop=tenacity.stop_after_attempt(3), - retry=_retry, - before_sleep=_before_sleep, - ) - retrying(thing.go) - finally: - logger.removeHandler(handler) - - etalon_re = r"^Retrying .* in 0\.01 seconds as it returned None\.$" - self.assertEqual(len(handler.records), 2) - fmt = logging.Formatter().format - self.assertRegex(fmt(handler.records[0]), etalon_re) - self.assertRegex(fmt(handler.records[1]), etalon_re) - - def test_before_sleep_log_returns_with_exc_info(self) -> None: - self.test_before_sleep_log_returns(exc_info=True) - - -class TestReraiseExceptions(unittest.TestCase): - def test_reraise_by_default(self) -> None: - calls = [] - - @retry( - wait=tenacity.wait_fixed(0.1), - stop=tenacity.stop_after_attempt(2), - reraise=True, - ) - def _reraised_by_default() -> None: - calls.append("x") - raise KeyError("Bad key") - - self.assertRaises(KeyError, _reraised_by_default) - self.assertEqual(2, len(calls)) - - def test_reraise_from_retry_error(self) -> None: - calls = [] - - @retry(wait=tenacity.wait_fixed(0.1), stop=tenacity.stop_after_attempt(2)) - def _raise_key_error() -> None: - calls.append("x") - raise KeyError("Bad key") - - def _reraised_key_error() -> None: - try: - _raise_key_error() - except tenacity.RetryError as retry_err: - retry_err.reraise() - - self.assertRaises(KeyError, _reraised_key_error) - self.assertEqual(2, len(calls)) - - def test_reraise_timeout_from_retry_error(self) -> None: - calls = [] - - @retry( - wait=tenacity.wait_fixed(0.1), - stop=tenacity.stop_after_attempt(2), - retry=lambda retry_state: True, - ) - def _mock_fn() -> None: - calls.append("x") - - def _reraised_mock_fn() -> None: - try: - _mock_fn() - except tenacity.RetryError as retry_err: - retry_err.reraise() - - self.assertRaises(tenacity.RetryError, _reraised_mock_fn) - self.assertEqual(2, len(calls)) - - def test_reraise_no_exception(self) -> None: - calls = [] - - @retry( - wait=tenacity.wait_fixed(0.1), - stop=tenacity.stop_after_attempt(2), - retry=lambda retry_state: True, - reraise=True, - ) - def _mock_fn() -> None: - calls.append("x") - - self.assertRaises(tenacity.RetryError, _mock_fn) - self.assertEqual(2, len(calls)) - - -class TestStatistics(unittest.TestCase): - def test_stats(self) -> None: - @retry() - def _foobar() -> int: - return 42 - - self.assertEqual({}, _foobar.statistics) - _foobar() - self.assertEqual(1, _foobar.statistics["attempt_number"]) - - def test_stats_failing(self) -> None: - @retry(stop=tenacity.stop_after_attempt(2)) - def _foobar() -> None: - raise ValueError(42) - - self.assertEqual({}, _foobar.statistics) - with contextlib.suppress(Exception): - _foobar() - self.assertEqual(2, _foobar.statistics["attempt_number"]) - - def test_retry_object_statistics_synced(self) -> None: - """Test that func.retry.statistics is synced with func.statistics.""" - - @retry(stop=tenacity.stop_after_attempt(3)) - def _foobar() -> int: - return 42 - - _foobar() - self.assertEqual( - _foobar.retry.statistics["attempt_number"], - _foobar.statistics["attempt_number"], - ) - - def test_retry_object_statistics_during_execution(self) -> None: - """Test that func.retry.statistics is accessible during execution.""" - attempts: list[int] = [] - - @retry( - stop=tenacity.stop_after_attempt(3), - retry=tenacity.retry_if_exception_type(ValueError), - reraise=True, - ) - def _foobar() -> int: - attempts.append(_foobar.retry.statistics["attempt_number"]) - if len(attempts) < 3: - raise ValueError("retry") - return 42 - - _foobar() - self.assertEqual(attempts, [1, 2, 3]) - - -class TestRetryErrorCallback(unittest.TestCase): - @override - def setUp(self) -> None: - self._attempt_number = 0 - self._callback_called = False - - def _callback(self, fut: tenacity.Future) -> tenacity.Future: - self._callback_called = True - return fut - - def test_retry_error_callback(self) -> None: - num_attempts = 3 - - def retry_error_callback(retry_state: RetryCallState) -> typing.Any: - retry_error_callback.called_times += 1 # type: ignore[attr-defined] - return retry_state.outcome - - retry_error_callback.called_times = 0 # type: ignore[attr-defined] - - @retry( - stop=tenacity.stop_after_attempt(num_attempts), - retry_error_callback=retry_error_callback, - ) - def _foobar() -> None: - self._attempt_number += 1 - raise Exception("This exception should not be raised") - - result = _foobar() - - self.assertEqual(retry_error_callback.called_times, 1) # type: ignore[attr-defined] - self.assertEqual(num_attempts, self._attempt_number) - self.assertIsInstance(result, tenacity.Future) - - -class TestContextManager(unittest.TestCase): - def test_context_manager_retry_one(self) -> None: - from tenacity import Retrying - - raise_ = True - - for attempt in Retrying(): - with attempt: - if raise_: - raise_ = False - raise Exception("Retry it!") - - def test_context_manager_on_error(self) -> None: - from tenacity import Retrying - - class CustomError(Exception): - pass - - retry = Retrying(retry=tenacity.retry_if_exception_type(IOError)) - - def test() -> None: - for attempt in retry: - with attempt: - raise CustomError("Don't retry!") - - self.assertRaises(CustomError, test) - - def test_context_manager_retry_error(self) -> None: - from tenacity import Retrying - - retry = Retrying(stop=tenacity.stop_after_attempt(2)) - - def test() -> None: - for attempt in retry: - with attempt: - raise Exception("Retry it!") - - self.assertRaises(RetryError, test) - - def test_context_manager_reraise(self) -> None: - from tenacity import Retrying - - class CustomError(Exception): - pass - - retry = Retrying(reraise=True, stop=tenacity.stop_after_attempt(2)) - - def test() -> None: - for attempt in retry: - with attempt: - raise CustomError("Don't retry!") - - self.assertRaises(CustomError, test) - - -class TestInvokeAsCallable: - """Test direct invocation of Retrying as a callable.""" - - @staticmethod - def invoke(retry: Retrying, f: typing.Callable[..., typing.Any]) -> typing.Any: - """ - Invoke Retrying logic. - - Wrapper allows testing different call mechanisms in test sub-classes. - """ - return retry(f) - - def test_retry_one(self) -> None: - def f() -> typing.Any: - f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] - if len(f.calls) <= 1: # type: ignore[attr-defined] - raise Exception("Retry it!") - return 42 - - f.calls = [] # type: ignore[attr-defined] - - retry = Retrying() - assert self.invoke(retry, f) == 42 - assert f.calls == [1, 2] # type: ignore[attr-defined] - - def test_on_error(self) -> None: - class CustomError(Exception): - pass - - def f() -> typing.Any: - f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] - if len(f.calls) <= 1: # type: ignore[attr-defined] - raise CustomError("Don't retry!") - return 42 - - f.calls = [] # type: ignore[attr-defined] - - retry = Retrying(retry=tenacity.retry_if_exception_type(IOError)) - with pytest.raises(CustomError): - self.invoke(retry, f) - assert f.calls == [1] # type: ignore[attr-defined] - - def test_retry_error(self) -> None: - def f() -> typing.Any: - f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] - raise Exception("Retry it!") - - f.calls = [] # type: ignore[attr-defined] - - retry = Retrying(stop=tenacity.stop_after_attempt(2)) - with pytest.raises(RetryError): - self.invoke(retry, f) - assert f.calls == [1, 2] # type: ignore[attr-defined] - - def test_reraise(self) -> None: - class CustomError(Exception): - pass - - def f() -> typing.Any: - f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] - raise CustomError("Retry it!") - - f.calls = [] # type: ignore[attr-defined] - - retry = Retrying(reraise=True, stop=tenacity.stop_after_attempt(2)) - with pytest.raises(CustomError): - self.invoke(retry, f) - assert f.calls == [1, 2] # type: ignore[attr-defined] - - -class TestRetryException(unittest.TestCase): - def test_retry_error_is_pickleable(self) -> None: - import pickle - - expected = RetryError(last_attempt=123) # type: ignore[arg-type] - pickled = pickle.dumps(expected) - actual = pickle.loads(pickled) - self.assertEqual(expected.last_attempt, actual.last_attempt) - - -class TestRetryTyping(unittest.TestCase): - def test_retry_type_annotations(self) -> None: - """The decorator should maintain types of decorated functions. - - The annotations below are the assertions; mypy checks them when it runs - over this file. The negative case leans on warn_unused_ignores: should - @retry ever decay to returning Any, that assignment would stop being an - error and the now-dead ignore would fail the type check. - """ - - def num_to_str(number: int) -> str: - return str(number) - - # equivalent to a raw @retry decoration - with_raw = retry(num_to_str) - with_raw_result = with_raw(1) - - # equivalent to a @retry(...) decoration - with_constructor = retry()(num_to_str) - with_constructor_result = with_constructor(1) - - # The wrapper stays usable wherever the undecorated function was. - _raw_signature: typing.Callable[[int], str] = with_raw - _constructor_signature: typing.Callable[[int], str] = with_constructor - - # ...and an incompatible signature is still rejected. - _mismatch: typing.Callable[[str], int] = with_raw # type: ignore[assignment] - - self.assertEqual(with_raw_result, "1") - self.assertEqual(with_constructor_result, "1") - - def test_retry_decorated_method_keeps_bound_signature(self) -> None: - """A decorated instance method must type-check like a bound method. - - Without a descriptor (``__get__``) on ``_RetryDecorated``, static type - checkers treat ``instance.method`` the same as the unbound - ``Class.method``, so a normal call with only the non-``self`` keyword - arguments looks like a type error and the return type resolves to - ``Any``/``Unknown``. This does not fail at runtime (functools.wraps - returns a real function, which Python always binds correctly), but it - does fail under `mypy --strict`, which also type-checks this file. - See issue #532. - """ - - class Doubler: - @retry(stop=tenacity.stop_after_attempt(3)) - def double(self, value: int) -> int: - return value * 2 - - doubler = Doubler() - result: int = doubler.double(value=21) - self.assertEqual(result, 42) - - -class TestMockingSleep: - RETRY_ARGS = { - "wait": tenacity.wait_fixed(0.1), - "stop": tenacity.stop_after_attempt(5), - } - - def _fail(self) -> None: - raise NotImplementedError - - @retry(**RETRY_ARGS) # type: ignore[call-overload, untyped-decorator] - def _decorated_fail(self) -> None: - self._fail() - - @pytest.fixture() - def mock_sleep( - self, monkeypatch: typing.Any - ) -> typing.Generator[typing.Any, None, None]: - class MockSleep: - call_count = 0 - - def __call__(self, seconds: float) -> None: - self.call_count += 1 - - sleep = MockSleep() - monkeypatch.setattr(tenacity.nap.time, "sleep", sleep) # type: ignore[attr-defined] - yield sleep - - def test_decorated(self, mock_sleep: typing.Any) -> None: - with pytest.raises(RetryError): - self._decorated_fail() - assert mock_sleep.call_count == 4 - - def test_decorated_retry_with(self, mock_sleep: typing.Any) -> None: - fail_faster = self._decorated_fail.retry_with( - stop=tenacity.stop_after_attempt(2), - ) - with pytest.raises(RetryError): - fail_faster() - assert mock_sleep.call_count == 1 - - -class TestPickle(unittest.TestCase): - def test_retrying_picklable(self) -> None: - """Retrying objects can be pickled for multiprocessing support.""" - retrying = Retrying(stop=tenacity.stop_after_attempt(3)) - pickled = pickle.dumps(retrying) - restored = pickle.loads(pickled) - assert isinstance(restored, Retrying) - assert isinstance(restored.stop, tenacity.stop_after_attempt) - - def test_retrying_picklable_after_run(self) -> None: - """Retrying objects can be pickled even after being used.""" - retrying = Retrying(stop=tenacity.stop_after_attempt(3)) - # Access statistics to populate _local - _ = retrying.statistics - pickled = pickle.dumps(retrying) - restored = pickle.loads(pickled) - assert isinstance(restored, Retrying) - # Statistics should be reset on the restored object - assert restored.statistics == {} - - def test_retry_strategies_picklable(self) -> None: - """All built-in retry strategies can be pickled.""" - strategies = [ - tenacity.retry_if_exception_type(ValueError), - tenacity.retry_if_not_exception_type(ValueError), - tenacity.retry_if_exception_message(message="fail"), - tenacity.retry_if_exception_message(match="fail.*"), - tenacity.retry_if_not_exception_message(message="fail"), - ] - for strategy in strategies: - restored = pickle.loads(pickle.dumps(strategy)) - assert type(restored) is type(strategy) - - def test_retrying_pickle_round_trip_works(self) -> None: - """A pickled-then-restored Retrying object retries correctly.""" - retrying = Retrying( - stop=tenacity.stop_after_attempt(3), - retry=tenacity.retry_if_exception_type(ValueError), - reraise=True, - ) - restored = pickle.loads(pickle.dumps(retrying)) - - calls = 0 - - def succeed_on_third() -> str: - nonlocal calls - calls += 1 - if calls < 3: - raise ValueError("not yet") - return "ok" - - result = restored(succeed_on_third) - assert result == "ok" - assert calls == 3 - - -if __name__ == "__main__": - unittest.main() +# Copyright 2016–2021 Julien Danjou +# Copyright 2016 Joshua Harlow +# Copyright 2013 Ray Holder +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import contextlib +import datetime +import logging +import pickle +import re +import time +import typing +import unittest +from fractions import Fraction +from unittest import mock + +import pytest + +import tenacity +from tenacity import RetryCallState, RetryError, Retrying, retry +from tenacity._utils import override +from tenacity.retry import retry_all, retry_any + +_unset = object() + + +def _make_unset_exception(func_name: str, **kwargs: typing.Any) -> TypeError: + missing = [] + for k, v in kwargs.items(): + if v is _unset: + missing.append(k) + missing_str = ", ".join(repr(s) for s in missing) + return TypeError(func_name + " func missing parameters: " + missing_str) + + +def _set_delay_since_start(retry_state: RetryCallState, delay: typing.Any) -> None: + # Ensure outcome_timestamp - start_time is *exactly* equal to the delay to + # avoid complexity in test code. + retry_state.start_time = Fraction(retry_state.start_time) # type: ignore[assignment] + retry_state.outcome_timestamp = retry_state.start_time + Fraction(delay) + assert retry_state.seconds_since_start == delay + + +def make_retry_state( + previous_attempt_number: typing.Any, + delay_since_first_attempt: typing.Any, + last_result: typing.Any = None, + upcoming_sleep: typing.Any = 0, +) -> RetryCallState: + """Construct RetryCallState for given attempt number & delay. + + Only used in testing and thus is extra careful about timestamp arithmetics. + """ + required_parameter_unset = ( + previous_attempt_number is _unset or delay_since_first_attempt is _unset + ) + if required_parameter_unset: + raise _make_unset_exception( + "wait/stop", + previous_attempt_number=previous_attempt_number, + delay_since_first_attempt=delay_since_first_attempt, + ) + + retry_state = RetryCallState(None, None, (), {}) # type: ignore[arg-type] + retry_state.attempt_number = previous_attempt_number + if last_result is not None: + retry_state.outcome = last_result + else: + retry_state.set_result(None) + + retry_state.upcoming_sleep = upcoming_sleep + + _set_delay_since_start(retry_state, delay_since_first_attempt) + return retry_state + + +class TestBase(unittest.TestCase): + def test_retrying_repr(self) -> None: + class ConcreteRetrying(tenacity.BaseRetrying): + @override + def __call__( + self, fn: typing.Any, *args: typing.Any, **kwargs: typing.Any + ) -> typing.Any: + pass + + repr(ConcreteRetrying()) + + def test_callstate_repr(self) -> None: + rs = RetryCallState(None, None, (), {}) # type: ignore[arg-type] + rs.idle_for = 1.1111111 + assert repr(rs).endswith("attempt #1; slept for 1.11; last result: none yet>") + rs = make_retry_state(2, 5) + assert repr(rs).endswith( + "attempt #2; slept for 0.0; last result: returned None>" + ) + rs = make_retry_state( + 0, 0, last_result=tenacity.Future.construct(1, ValueError("aaa"), True) + ) + assert repr(rs).endswith( + "attempt #0; slept for 0.0; last result: failed (ValueError aaa)>" + ) + + +class TestRetryingName(unittest.TestCase): + def test_str_default(self) -> None: + """Without a name, str() returns ''.""" + assert str(Retrying()) == "" + + def test_str_with_name(self) -> None: + """With a name, str() returns the given name.""" + assert str(Retrying(name="my_block")) == "my_block" + + def test_str_preserved_by_copy(self) -> None: + """copy() preserves the name.""" + r = Retrying(name="my_block") + assert str(r.copy()) == "my_block" + + def test_str_overridden_by_copy(self) -> None: + """copy() allows overriding the name.""" + r = Retrying(name="original") + assert str(r.copy(name="overridden")) == "overridden" + + def test_get_fn_name_decorator(self) -> None: + """get_fn_name() returns the function's qualified name when used as decorator.""" + captured: list[RetryCallState] = [] + + @tenacity.retry( + stop=tenacity.stop_after_attempt(1), + after=lambda rs: captured.append(rs), + ) + def my_func() -> None: + raise ValueError + + with contextlib.suppress(Exception): + my_func() + assert captured + assert "my_func" in captured[0].get_fn_name() + + def test_get_fn_name_context_manager_no_name(self) -> None: + """get_fn_name() returns '' in context manager mode without a name.""" + r = Retrying(stop=tenacity.stop_after_attempt(1)) + rs = RetryCallState(r, None, (), {}) + assert rs.get_fn_name() == "" + + def test_get_fn_name_context_manager_with_name(self) -> None: + """get_fn_name() returns the given name in context manager mode.""" + r = Retrying(name="ws_listener", stop=tenacity.stop_after_attempt(1)) + rs = RetryCallState(r, None, (), {}) + assert rs.get_fn_name() == "ws_listener" + + def test_logging_uses_name(self) -> None: + """before_log uses the name parameter in context manager mode.""" + import unittest.mock + + log = unittest.mock.MagicMock() + logger = unittest.mock.MagicMock(log=log) + + with contextlib.suppress(Exception): + for attempt in Retrying( + name="my_block", + before=tenacity.before_log(logger, logging.INFO), + stop=tenacity.stop_after_attempt(1), + ): + with attempt: + raise ValueError + + args = log.call_args[0] + assert "my_block" in args[1] + + def test_logging_infers_caller_name(self) -> None: + """before_log infers the enclosing function name when no name= is given (#511). + + When Retrying is used as a context manager without an explicit name= + parameter, retry_state.fn is None. Before the fix, get_fn_name() would + fall back to str(retry_object) == "". Now __iter__ captures + sys._getframe(1) -- the for-loop frame -- and stores co_qualname / co_name + as retry_state._inferred_name so that before_log() can log something + meaningful instead of ''. + """ + import unittest.mock + + log = unittest.mock.MagicMock() + logger = unittest.mock.MagicMock(log=log) + + def my_retry_function() -> None: + with contextlib.suppress(Exception): + for attempt in Retrying( + before=tenacity.before_log(logger, logging.INFO), + stop=tenacity.stop_after_attempt(1), + ): + with attempt: + raise ValueError("boom") + + my_retry_function() + + # before_log must have been called at least once + assert log.call_args is not None, "before_log was never called" + msg = log.call_args[0][1] + assert "" not in msg, ( + f"Expected an inferred caller name in the log message, got: {msg!r}" + ) + assert "my_retry_function" in msg, ( + f"Expected 'my_retry_function' in the log message, got: {msg!r}" + ) + + +class TestStopConditions(unittest.TestCase): + def test_never_stop(self) -> None: + r = Retrying() + self.assertFalse(r.stop(make_retry_state(3, 6546))) + + def test_stop_any(self) -> None: + stop = tenacity.stop_any( + tenacity.stop_after_delay(1), tenacity.stop_after_attempt(4) + ) + + def s(*args: typing.Any) -> bool: + return stop(make_retry_state(*args)) + + self.assertFalse(s(1, 0.1)) + self.assertFalse(s(2, 0.2)) + self.assertFalse(s(2, 0.8)) + self.assertTrue(s(4, 0.8)) + self.assertTrue(s(3, 1.8)) + self.assertTrue(s(4, 1.8)) + + def test_stop_all(self) -> None: + stop = tenacity.stop_all( + tenacity.stop_after_delay(1), tenacity.stop_after_attempt(4) + ) + + def s(*args: typing.Any) -> bool: + return stop(make_retry_state(*args)) + + self.assertFalse(s(1, 0.1)) + self.assertFalse(s(2, 0.2)) + self.assertFalse(s(2, 0.8)) + self.assertFalse(s(4, 0.8)) + self.assertFalse(s(3, 1.8)) + self.assertTrue(s(4, 1.8)) + + def test_stop_or(self) -> None: + stop = tenacity.stop_after_delay(1) | tenacity.stop_after_attempt(4) + + def s(*args: typing.Any) -> bool: + return stop(make_retry_state(*args)) + + self.assertFalse(s(1, 0.1)) + self.assertFalse(s(2, 0.2)) + self.assertFalse(s(2, 0.8)) + self.assertTrue(s(4, 0.8)) + self.assertTrue(s(3, 1.8)) + self.assertTrue(s(4, 1.8)) + + def test_stop_and(self) -> None: + stop = tenacity.stop_after_delay(1) & tenacity.stop_after_attempt(4) + + def s(*args: typing.Any) -> bool: + return stop(make_retry_state(*args)) + + self.assertFalse(s(1, 0.1)) + self.assertFalse(s(2, 0.2)) + self.assertFalse(s(2, 0.8)) + self.assertFalse(s(4, 0.8)) + self.assertFalse(s(3, 1.8)) + self.assertTrue(s(4, 1.8)) + + def test_stop_after_attempt(self) -> None: + r = Retrying(stop=tenacity.stop_after_attempt(3)) + self.assertFalse(r.stop(make_retry_state(2, 6546))) + self.assertTrue(r.stop(make_retry_state(3, 6546))) + self.assertTrue(r.stop(make_retry_state(4, 6546))) + + def test_stop_after_delay(self) -> None: + for delay in (1, datetime.timedelta(seconds=1)): + with self.subTest(): + r = Retrying(stop=tenacity.stop_after_delay(delay)) + self.assertFalse(r.stop(make_retry_state(2, 0.999))) + self.assertTrue(r.stop(make_retry_state(2, 1))) + self.assertTrue(r.stop(make_retry_state(2, 1.001))) + + def test_stop_before_delay(self) -> None: + for delay in (1, datetime.timedelta(seconds=1)): + with self.subTest(): + r = Retrying(stop=tenacity.stop_before_delay(delay)) + self.assertFalse( + r.stop(make_retry_state(2, 0.999, upcoming_sleep=0.0001)) + ) + self.assertTrue(r.stop(make_retry_state(2, 1, upcoming_sleep=0.001))) + self.assertTrue(r.stop(make_retry_state(2, 1, upcoming_sleep=1))) + + # It should act the same as stop_after_delay if upcoming sleep is 0 + self.assertFalse(r.stop(make_retry_state(2, 0.999, upcoming_sleep=0))) + self.assertTrue(r.stop(make_retry_state(2, 1, upcoming_sleep=0))) + self.assertTrue(r.stop(make_retry_state(2, 1.001, upcoming_sleep=0))) + + def test_legacy_explicit_stop_type(self) -> None: + Retrying(stop="stop_after_attempt") # type: ignore[arg-type] + + def test_stop_func_with_retry_state(self) -> None: + def stop_func(retry_state: RetryCallState) -> bool: + rs = retry_state + return rs.attempt_number == rs.seconds_since_start + + r = Retrying(stop=stop_func) + self.assertFalse(r.stop(make_retry_state(1, 3))) + self.assertFalse(r.stop(make_retry_state(100, 99))) + self.assertTrue(r.stop(make_retry_state(101, 101))) + + +class TestWaitConditions(unittest.TestCase): + def test_no_sleep(self) -> None: + r = Retrying() + self.assertEqual(0, r.wait(make_retry_state(18, 9879))) + + def test_fixed_sleep(self) -> None: + for wait in (1, datetime.timedelta(seconds=1)): + with self.subTest(): + r = Retrying(wait=tenacity.wait_fixed(wait)) + self.assertEqual(1, r.wait(make_retry_state(12, 6546))) + + def test_incrementing_sleep(self) -> None: + for start, increment in ( + (500, 100), + (datetime.timedelta(seconds=500), datetime.timedelta(seconds=100)), + ): + with self.subTest(): + r = Retrying( + wait=tenacity.wait_incrementing(start=start, increment=increment) + ) + self.assertEqual(500, r.wait(make_retry_state(1, 6546))) + self.assertEqual(600, r.wait(make_retry_state(2, 6546))) + self.assertEqual(700, r.wait(make_retry_state(3, 6546))) + + def test_random_sleep(self) -> None: + for min_, max_ in ( + (1, 20), + (datetime.timedelta(seconds=1), datetime.timedelta(seconds=20)), + ): + with self.subTest(): + r = Retrying(wait=tenacity.wait_random(min=min_, max=max_)) + times = set() + for _ in range(1000): + times.add(r.wait(make_retry_state(1, 6546))) + + # this is kind of non-deterministic... + self.assertTrue(len(times) > 1) + for t in times: + self.assertTrue(t >= 1) + self.assertTrue(t < 20) + + def test_random_sleep_withoutmin_(self) -> None: + r = Retrying(wait=tenacity.wait_random(max=2)) + times = set() + times.add(r.wait(make_retry_state(1, 6546))) + times.add(r.wait(make_retry_state(1, 6546))) + times.add(r.wait(make_retry_state(1, 6546))) + times.add(r.wait(make_retry_state(1, 6546))) + + # this is kind of non-deterministic... + self.assertTrue(len(times) > 1) + for t in times: + self.assertTrue(t >= 0) + self.assertTrue(t <= 2) + + def test_exponential(self) -> None: + r = Retrying(wait=tenacity.wait_exponential()) + self.assertEqual(r.wait(make_retry_state(1, 0)), 1) + self.assertEqual(r.wait(make_retry_state(2, 0)), 2) + self.assertEqual(r.wait(make_retry_state(3, 0)), 4) + self.assertEqual(r.wait(make_retry_state(4, 0)), 8) + self.assertEqual(r.wait(make_retry_state(5, 0)), 16) + self.assertEqual(r.wait(make_retry_state(6, 0)), 32) + self.assertEqual(r.wait(make_retry_state(7, 0)), 64) + self.assertEqual(r.wait(make_retry_state(8, 0)), 128) + + def test_exponential_with_max_wait(self) -> None: + r = Retrying(wait=tenacity.wait_exponential(max=40)) + self.assertEqual(r.wait(make_retry_state(1, 0)), 1) + self.assertEqual(r.wait(make_retry_state(2, 0)), 2) + self.assertEqual(r.wait(make_retry_state(3, 0)), 4) + self.assertEqual(r.wait(make_retry_state(4, 0)), 8) + self.assertEqual(r.wait(make_retry_state(5, 0)), 16) + self.assertEqual(r.wait(make_retry_state(6, 0)), 32) + self.assertEqual(r.wait(make_retry_state(7, 0)), 40) + self.assertEqual(r.wait(make_retry_state(8, 0)), 40) + self.assertEqual(r.wait(make_retry_state(50, 0)), 40) + + def test_exponential_skips_power_after_reaching_max(self) -> None: + class ExplodingPower(float): + @override + def __pow__(self, exponent: float, modulo: int | None = None) -> float: + raise AssertionError("power should not be calculated above the maximum") + + r = Retrying(wait=tenacity.wait_exponential(max=40, exp_base=ExplodingPower(2))) + self.assertEqual(r.wait(make_retry_state(50, 0)), 40) + + def test_exponential_caps_when_max_multiplier_ratio_underflows(self) -> None: + r = Retrying(wait=tenacity.wait_exponential(multiplier=1e300, max=1e-100)) + self.assertEqual(r.wait(make_retry_state(1, 0)), 1e-100) + + def test_exponential_with_min_wait(self) -> None: + r = Retrying(wait=tenacity.wait_exponential(min=20)) + self.assertEqual(r.wait(make_retry_state(1, 0)), 20) + self.assertEqual(r.wait(make_retry_state(2, 0)), 20) + self.assertEqual(r.wait(make_retry_state(3, 0)), 20) + self.assertEqual(r.wait(make_retry_state(4, 0)), 20) + self.assertEqual(r.wait(make_retry_state(5, 0)), 20) + self.assertEqual(r.wait(make_retry_state(6, 0)), 32) + self.assertEqual(r.wait(make_retry_state(7, 0)), 64) + self.assertEqual(r.wait(make_retry_state(8, 0)), 128) + self.assertEqual(r.wait(make_retry_state(20, 0)), 524288) + + def test_exponential_with_max_wait_and_multiplier(self) -> None: + r = Retrying(wait=tenacity.wait_exponential(max=50, multiplier=1)) + self.assertEqual(r.wait(make_retry_state(1, 0)), 1) + self.assertEqual(r.wait(make_retry_state(2, 0)), 2) + self.assertEqual(r.wait(make_retry_state(3, 0)), 4) + self.assertEqual(r.wait(make_retry_state(4, 0)), 8) + self.assertEqual(r.wait(make_retry_state(5, 0)), 16) + self.assertEqual(r.wait(make_retry_state(6, 0)), 32) + self.assertEqual(r.wait(make_retry_state(7, 0)), 50) + self.assertEqual(r.wait(make_retry_state(8, 0)), 50) + self.assertEqual(r.wait(make_retry_state(50, 0)), 50) + + def test_exponential_with_min_wait_and_multiplier(self) -> None: + r = Retrying(wait=tenacity.wait_exponential(min=20, multiplier=2)) + self.assertEqual(r.wait(make_retry_state(1, 0)), 20) + self.assertEqual(r.wait(make_retry_state(2, 0)), 20) + self.assertEqual(r.wait(make_retry_state(3, 0)), 20) + self.assertEqual(r.wait(make_retry_state(4, 0)), 20) + self.assertEqual(r.wait(make_retry_state(5, 0)), 32) + self.assertEqual(r.wait(make_retry_state(6, 0)), 64) + self.assertEqual(r.wait(make_retry_state(7, 0)), 128) + self.assertEqual(r.wait(make_retry_state(8, 0)), 256) + self.assertEqual(r.wait(make_retry_state(20, 0)), 1048576) + + def test_exponential_with_min_wait_andmax__wait(self) -> None: + for min_, max_ in ( + (10, 100), + (datetime.timedelta(seconds=10), datetime.timedelta(seconds=100)), + ): + with self.subTest(): + r = Retrying(wait=tenacity.wait_exponential(min=min_, max=max_)) + self.assertEqual(r.wait(make_retry_state(1, 0)), 10) + self.assertEqual(r.wait(make_retry_state(2, 0)), 10) + self.assertEqual(r.wait(make_retry_state(3, 0)), 10) + self.assertEqual(r.wait(make_retry_state(4, 0)), 10) + self.assertEqual(r.wait(make_retry_state(5, 0)), 16) + self.assertEqual(r.wait(make_retry_state(6, 0)), 32) + self.assertEqual(r.wait(make_retry_state(7, 0)), 64) + self.assertEqual(r.wait(make_retry_state(8, 0)), 100) + self.assertEqual(r.wait(make_retry_state(9, 0)), 100) + self.assertEqual(r.wait(make_retry_state(20, 0)), 100) + + def test_legacy_explicit_wait_type(self) -> None: + Retrying(wait="exponential_sleep") # type: ignore[arg-type] + + def test_wait_func(self) -> None: + def wait_func(retry_state: RetryCallState) -> typing.Any: + return retry_state.attempt_number * retry_state.seconds_since_start # type: ignore[operator] + + r = Retrying(wait=wait_func) + self.assertEqual(r.wait(make_retry_state(1, 5)), 5) + self.assertEqual(r.wait(make_retry_state(2, 11)), 22) + self.assertEqual(r.wait(make_retry_state(10, 100)), 1000) + + def test_wait_combine(self) -> None: + r = Retrying( + wait=tenacity.wait_combine( + tenacity.wait_random(0, 3), tenacity.wait_fixed(5) + ) + ) + # Test it a few time since it's random + for _i in range(1000): + w = r.wait(make_retry_state(1, 5)) + self.assertLess(w, 8) + self.assertGreaterEqual(w, 5) + + def test_wait_exception(self) -> None: + def predicate(exc: BaseException) -> float: + if isinstance(exc, ValueError): + return 3.5 + return 10.0 + + r = Retrying(wait=tenacity.wait_exception(predicate)) + + fut1 = tenacity.Future.construct(1, ValueError(), True) + self.assertEqual(r.wait(make_retry_state(1, 0, last_result=fut1)), 3.5) + + fut2 = tenacity.Future.construct(1, KeyError(), True) + self.assertEqual(r.wait(make_retry_state(1, 0, last_result=fut2)), 10.0) + + fut3 = tenacity.Future.construct(1, None, False) + with self.assertRaises(RuntimeError): + r.wait(make_retry_state(1, 0, last_result=fut3)) + + def test_wait_double_sum(self) -> None: + r = Retrying(wait=tenacity.wait_random(0, 3) + tenacity.wait_fixed(5)) + # Test it a few time since it's random + for _i in range(1000): + w = r.wait(make_retry_state(1, 5)) + self.assertLess(w, 8) + self.assertGreaterEqual(w, 5) + + def test_wait_triple_sum(self) -> None: + r = Retrying( + wait=tenacity.wait_fixed(1) + + tenacity.wait_random(0, 3) + + tenacity.wait_fixed(5) + ) + # Test it a few time since it's random + for _i in range(1000): + w = r.wait(make_retry_state(1, 5)) + self.assertLess(w, 9) + self.assertGreaterEqual(w, 6) + + def test_wait_arbitrary_sum(self) -> None: + r = Retrying( + wait=sum( # type: ignore[arg-type] + [ + tenacity.wait_fixed(1), + tenacity.wait_random(0, 3), + tenacity.wait_fixed(5), + tenacity.wait_none(), + ] + ) + ) + # Test it a few time since it's random + for _ in range(1000): + w = r.wait(make_retry_state(1, 5)) + self.assertLess(w, 9) + self.assertGreaterEqual(w, 6) + + def test_wait_falsy_values_mean_no_wait(self) -> None: + # Untyped callers pass None or 0 to mean "no wait", and `sum([])` + # over an empty list of strategies yields the int 0. All must reach + # the retry path without blowing up inside iter(). + def make_flaky() -> typing.Callable[[], str]: + attempts = [] + + def flaky() -> str: + attempts.append(1) + if len(attempts) < 2: + raise ValueError("boom") + return "ok" + + return flaky + + for wait in (None, 0, sum([])): + with self.subTest(wait=wait): + flaky = make_flaky() + r = Retrying( + wait=wait, # type: ignore[arg-type] + stop=tenacity.stop_after_attempt(3), + ) + self.assertEqual(r(flaky), "ok") + + def test_wait_radd_plain_callable(self) -> None: + # A plain callable is a valid WaitBaseT, and functions have no + # __add__, so `callable + strategy` goes through wait_base.__radd__. + def cb(retry_state: RetryCallState) -> float: + return 2.0 + + combined = cb + tenacity.wait_fixed(1) + self.assertIsInstance(combined, tenacity.wait_combine) + self.assertEqual(combined(make_retry_state(1, 5)), 3.0) + + def test_wait_combine_passes_state_positionally(self) -> None: + # A WaitBaseT callable only promises to take the state positionally; + # its parameter name is its own business. + combined = tenacity.wait_combine( + tenacity.wait_fixed(1), + lambda rs: 2.0, + ) + self.assertEqual(combined(make_retry_state(1, 5)), 3.0) + + def test_wait_radd_rejects_non_zero_number(self) -> None: + with self.assertRaises(TypeError): + # Statically accepted -- see the comment on wait_base.__radd__ -- + # so the runtime rejection is what has to be tested. + 5 + tenacity.wait_fixed(1) + + def _assert_range(self, wait: float, min_: float, max_: float) -> None: + self.assertLess(wait, max_) + self.assertGreaterEqual(wait, min_) + + def _assert_inclusive_range(self, wait: float, low: float, high: float) -> None: + self.assertLessEqual(wait, high) + self.assertGreaterEqual(wait, low) + + def _assert_inclusive_epsilon( + self, wait: float, target: float, epsilon: float + ) -> None: + self.assertLessEqual(wait, target + epsilon) + self.assertGreaterEqual(wait, target - epsilon) + + def test_wait_chain(self) -> None: + r = Retrying( + wait=tenacity.wait_chain( + *[tenacity.wait_fixed(1) for i in range(2)] + + [tenacity.wait_fixed(4) for i in range(2)] + + [tenacity.wait_fixed(8) for i in range(1)] + ) + ) + + for i in range(10): + w = r.wait(make_retry_state(i + 1, 1)) + if i < 2: + self._assert_range(w, 1, 2) + elif i < 4: + self._assert_range(w, 4, 5) + else: + self._assert_range(w, 8, 9) + + def test_wait_chain_multiple_invocations(self) -> None: + sleep_intervals: list[float] = [] + r = Retrying( + sleep=sleep_intervals.append, + wait=tenacity.wait_chain(*[tenacity.wait_fixed(i + 1) for i in range(3)]), + stop=tenacity.stop_after_attempt(5), + retry=tenacity.retry_if_result(lambda x: x == 1), + ) + + @r.wraps + def always_return_1() -> int: + return 1 + + self.assertRaises(tenacity.RetryError, always_return_1) + self.assertEqual(sleep_intervals, [1.0, 2.0, 3.0, 3.0]) + sleep_intervals[:] = [] + + def test_wait_chain_requires_at_least_one_strategy(self) -> None: + with self.assertRaises(ValueError): + tenacity.wait_chain() + # Confirm the wrapped Retrying path surfaces the same ValueError + # instead of an opaque IndexError from self.strategies[-1]. + with self.assertRaises(ValueError): + Retrying(wait=tenacity.wait_chain()) + + def test_wait_random_exponential(self) -> None: + fn = tenacity.wait_random_exponential(0.5, 60.0) + + for _ in range(1000): + self._assert_inclusive_range(fn(make_retry_state(1, 0)), 0, 0.5) + self._assert_inclusive_range(fn(make_retry_state(2, 0)), 0, 1.0) + self._assert_inclusive_range(fn(make_retry_state(3, 0)), 0, 2.0) + self._assert_inclusive_range(fn(make_retry_state(4, 0)), 0, 4.0) + self._assert_inclusive_range(fn(make_retry_state(5, 0)), 0, 8.0) + self._assert_inclusive_range(fn(make_retry_state(6, 0)), 0, 16.0) + self._assert_inclusive_range(fn(make_retry_state(7, 0)), 0, 32.0) + self._assert_inclusive_range(fn(make_retry_state(8, 0)), 0, 60.0) + self._assert_inclusive_range(fn(make_retry_state(9, 0)), 0, 60.0) + + # max wait + max_wait = 5 + fn = tenacity.wait_random_exponential(10, max_wait) + for _ in range(1000): + self._assert_inclusive_range(fn(make_retry_state(1, 0)), 0.00, max_wait) + + # min wait + min_wait = 5 + fn = tenacity.wait_random_exponential(min=min_wait) + for _ in range(1000): + self._assert_inclusive_range(fn(make_retry_state(1, 0)), min_wait, 5) + + # Default arguments exist + fn = tenacity.wait_random_exponential() + fn(make_retry_state(0, 0)) + + def test_wait_random_exponential_statistically(self) -> None: + fn = tenacity.wait_random_exponential(0.5, 60.0) + + attempt = [[fn(make_retry_state(i, 0)) for _ in range(4000)] for i in range(10)] + + def mean(lst: list[float]) -> float: + return float(sum(lst)) / float(len(lst)) + + # skipping attempt 0 + self._assert_inclusive_epsilon(mean(attempt[1]), 0.25, 0.02) + self._assert_inclusive_epsilon(mean(attempt[2]), 0.50, 0.04) + self._assert_inclusive_epsilon(mean(attempt[3]), 1, 0.08) + self._assert_inclusive_epsilon(mean(attempt[4]), 2, 0.16) + self._assert_inclusive_epsilon(mean(attempt[5]), 4, 0.32) + self._assert_inclusive_epsilon(mean(attempt[6]), 8, 0.64) + self._assert_inclusive_epsilon(mean(attempt[7]), 16, 1.28) + self._assert_inclusive_epsilon(mean(attempt[8]), 30, 2.56) + self._assert_inclusive_epsilon(mean(attempt[9]), 30, 2.56) + + def test_wait_exponential_jitter(self) -> None: + fn = tenacity.wait_exponential_jitter(max=60) + + for _ in range(1000): + self._assert_inclusive_range(fn(make_retry_state(1, 0)), 1, 2) + self._assert_inclusive_range(fn(make_retry_state(2, 0)), 2, 3) + self._assert_inclusive_range(fn(make_retry_state(3, 0)), 4, 5) + self._assert_inclusive_range(fn(make_retry_state(4, 0)), 8, 9) + self._assert_inclusive_range(fn(make_retry_state(5, 0)), 16, 17) + self._assert_inclusive_range(fn(make_retry_state(6, 0)), 32, 33) + self.assertEqual(fn(make_retry_state(7, 0)), 60) + self.assertEqual(fn(make_retry_state(8, 0)), 60) + self.assertEqual(fn(make_retry_state(9, 0)), 60) + + with self.assertWarns(DeprecationWarning): + fn = tenacity.wait_exponential_jitter(10, 5) + for _ in range(1000): + self.assertEqual(fn(make_retry_state(1, 0)), 5) + + # Default arguments exist + fn = tenacity.wait_exponential_jitter() + fn(make_retry_state(0, 0)) + + def test_wait_exponential_jitter_min(self) -> None: + fn = tenacity.wait_exponential_jitter(initial=1, max=60, jitter=1, min=5) + for _ in range(1000): + # Even for attempt 1 (base wait=1 + jitter 0..1 = 1..2), min=5 applies + self._assert_inclusive_range(fn(make_retry_state(1, 0)), 5, 5) + self._assert_inclusive_range(fn(make_retry_state(2, 0)), 5, 5) + self._assert_inclusive_range(fn(make_retry_state(3, 0)), 5, 5) + # For attempt 4, base wait=8 + jitter 0..1 = 8..9, above min + self._assert_inclusive_range(fn(make_retry_state(4, 0)), 8, 9) + + def test_wait_exponential_jitter_timedelta(self) -> None: + from datetime import timedelta + + fn = tenacity.wait_exponential_jitter( + max=timedelta(seconds=60), + jitter=timedelta(seconds=1), + min=timedelta(seconds=5), + ) + for _ in range(1000): + self._assert_inclusive_range(fn(make_retry_state(1, 0)), 5, 5) + self._assert_inclusive_range(fn(make_retry_state(5, 0)), 16, 17) + self.assertEqual(fn(make_retry_state(7, 0)), 60) + + def test_wait_exponential_jitter_multiplier(self) -> None: + fn = tenacity.wait_exponential_jitter(multiplier=10, max=60, jitter=0) + self.assertEqual(fn(make_retry_state(1, 0)), 10) + self.assertEqual(fn(make_retry_state(2, 0)), 20) + self.assertEqual(fn(make_retry_state(3, 0)), 40) + self.assertEqual(fn(make_retry_state(4, 0)), 60) + + def test_wait_exponential_jitter_initial_deprecated(self) -> None: + with self.assertWarns(DeprecationWarning): + fn = tenacity.wait_exponential_jitter(initial=10, max=60, jitter=0) + self.assertEqual(fn(make_retry_state(1, 0)), 10) + self.assertEqual(fn(make_retry_state(2, 0)), 20) + + def test_wait_exponential_jitter_initial_and_multiplier_raises(self) -> None: + with self.assertRaises(ValueError): + tenacity.wait_exponential_jitter(initial=5, multiplier=10) + + def test_wait_retry_state_attributes(self) -> None: + class ExtractCallState(Exception): + pass + + # retry_state is mutable, so return it as an exception to extract the + # exact values it has when wait is called and bypass any other logic. + def waitfunc(retry_state: RetryCallState) -> float: + raise ExtractCallState(retry_state) + + retrying = Retrying( + wait=waitfunc, + retry=( + tenacity.retry_if_exception_type() + | tenacity.retry_if_result(lambda result: result == 123) + ), + ) + + def returnval() -> int: + return 123 + + with self.assertRaises(ExtractCallState) as caught: + retrying(returnval) + retry_state = caught.exception.args[0] + self.assertIs(retry_state.fn, returnval) + self.assertEqual(retry_state.args, ()) + self.assertEqual(retry_state.kwargs, {}) + self.assertEqual(retry_state.outcome.result(), 123) + self.assertEqual(retry_state.attempt_number, 1) + self.assertGreaterEqual(retry_state.outcome_timestamp, retry_state.start_time) + + def dying() -> None: + raise Exception("Broken") + + with self.assertRaises(ExtractCallState) as caught: + retrying(dying) + retry_state = caught.exception.args[0] + self.assertIs(retry_state.fn, dying) + self.assertEqual(retry_state.args, ()) + self.assertEqual(retry_state.kwargs, {}) + self.assertEqual(str(retry_state.outcome.exception()), "Broken") + self.assertEqual(retry_state.attempt_number, 1) + self.assertGreaterEqual(retry_state.outcome_timestamp, retry_state.start_time) + + +class TestRetryConditions(unittest.TestCase): + def test_retry_if_result(self) -> None: + retry = tenacity.retry_if_result(lambda x: x == 1) + + def r(fut: tenacity.Future) -> bool: + retry_state = make_retry_state(1, 1.0, last_result=fut) + return retry(retry_state) + + self.assertTrue(r(tenacity.Future.construct(1, 1, False))) + self.assertFalse(r(tenacity.Future.construct(1, 2, False))) + + def test_retry_if_not_result(self) -> None: + retry = tenacity.retry_if_not_result(lambda x: x == 1) + + def r(fut: tenacity.Future) -> bool: + retry_state = make_retry_state(1, 1.0, last_result=fut) + return retry(retry_state) + + self.assertTrue(r(tenacity.Future.construct(1, 2, False))) + self.assertFalse(r(tenacity.Future.construct(1, 1, False))) + + def test_retry_any(self) -> None: + retry = tenacity.retry_any( + tenacity.retry_if_result(lambda x: x == 1), + tenacity.retry_if_result(lambda x: x == 2), + ) + + def r(fut: tenacity.Future) -> bool: + retry_state = make_retry_state(1, 1.0, last_result=fut) + return retry(retry_state) + + self.assertTrue(r(tenacity.Future.construct(1, 1, False))) + self.assertTrue(r(tenacity.Future.construct(1, 2, False))) + self.assertFalse(r(tenacity.Future.construct(1, 3, False))) + self.assertFalse(r(tenacity.Future.construct(1, 1, True))) + + def test_retry_all(self) -> None: + retry = tenacity.retry_all( + tenacity.retry_if_result(lambda x: x == 1), + tenacity.retry_if_result(lambda x: isinstance(x, int)), + ) + + def r(fut: tenacity.Future) -> bool: + retry_state = make_retry_state(1, 1.0, last_result=fut) + return retry(retry_state) + + self.assertTrue(r(tenacity.Future.construct(1, 1, False))) + self.assertFalse(r(tenacity.Future.construct(1, 2, False))) + self.assertFalse(r(tenacity.Future.construct(1, 3, False))) + self.assertFalse(r(tenacity.Future.construct(1, 1, True))) + + def test_retry_and(self) -> None: + retry = tenacity.retry_if_result(lambda x: x == 1) & tenacity.retry_if_result( + lambda x: isinstance(x, int) + ) + + def r(fut: tenacity.Future) -> bool: + retry_state = make_retry_state(1, 1.0, last_result=fut) + return retry(retry_state) + + self.assertTrue(r(tenacity.Future.construct(1, 1, False))) + self.assertFalse(r(tenacity.Future.construct(1, 2, False))) + self.assertFalse(r(tenacity.Future.construct(1, 3, False))) + self.assertFalse(r(tenacity.Future.construct(1, 1, True))) + + def test_retry_or(self) -> None: + retry = tenacity.retry_if_result( + lambda x: x == "foo" + ) | tenacity.retry_if_result(lambda x: isinstance(x, int)) + + def r(fut: tenacity.Future) -> bool: + retry_state = make_retry_state(1, 1.0, last_result=fut) + return retry(retry_state) + + self.assertTrue(r(tenacity.Future.construct(1, "foo", False))) + self.assertFalse(r(tenacity.Future.construct(1, "foobar", False))) + self.assertFalse(r(tenacity.Future.construct(1, 2.2, False))) + self.assertFalse(r(tenacity.Future.construct(1, 42, True))) + + def test_retry_or_with_plain_function(self) -> None: + """Plain callables can be composed with retry_base via |.""" + + def my_retry(retry_state: tenacity.RetryCallState) -> bool: + return retry_state.outcome is not None and not retry_state.outcome.failed + + # retry_base | plain_callable (exercises __or__ fallback) + retry = tenacity.retry_if_exception_type(Exception) | my_retry + retry_state = make_retry_state( + 1, 1.0, last_result=tenacity.Future.construct(1, "ok", False) + ) + self.assertTrue(retry(retry_state)) + + # plain_callable | retry_base (exercises __ror__ via reflection) + retry2 = my_retry | tenacity.retry_if_exception_type(Exception) + self.assertTrue(retry2(retry_state)) + + def test_retry_and_with_plain_function(self) -> None: + """Plain callables can be composed with retry_base via &.""" + + def my_retry(retry_state: tenacity.RetryCallState) -> bool: + return True + + # retry_base & plain_callable (exercises __and__ fallback) + retry = tenacity.retry_if_result(lambda x: x == 1) & my_retry + retry_state = make_retry_state( + 1, 1.0, last_result=tenacity.Future.construct(1, 1, False) + ) + self.assertTrue(retry(retry_state)) + + # plain_callable & retry_base (exercises __rand__ via reflection) + retry2 = my_retry & tenacity.retry_if_result(lambda x: x == 1) + self.assertTrue(retry2(retry_state)) + + def test_retry_or_coalesces(self) -> None: + """Multiple | operations flatten into a single retry_any.""" + a = tenacity.retry_if_exception_type(IOError) + b = tenacity.retry_if_exception_type(OSError) + c = tenacity.retry_if_exception_type(ValueError) + + combined = a | b | c + self.assertIsInstance(combined, retry_any) + self.assertEqual(len(combined.retries), 3) + + def test_retry_and_coalesces(self) -> None: + """Multiple & operations flatten into a single retry_all.""" + a = tenacity.retry_if_result(lambda x: x == 1) + b = tenacity.retry_if_result(lambda x: x > 0) + c = tenacity.retry_if_result(lambda x: x < 10) + + combined = a & b & c + self.assertIsInstance(combined, retry_all) + self.assertEqual(len(combined.retries), 3) + + def _raise_try_again(self) -> None: + self._attempts += 1 + if self._attempts < 3: + raise tenacity.TryAgain + + def test_retry_try_again(self) -> None: + self._attempts = 0 + Retrying(stop=tenacity.stop_after_attempt(5), retry=tenacity.retry_never)( + self._raise_try_again + ) + self.assertEqual(3, self._attempts) + + def test_retry_try_again_forever(self) -> None: + def _r() -> None: + raise tenacity.TryAgain + + r = Retrying(stop=tenacity.stop_after_attempt(5), retry=tenacity.retry_never) + self.assertRaises(tenacity.RetryError, r, _r) + self.assertEqual(5, r.statistics["attempt_number"]) + + def test_retry_try_again_with_cause_reraise(self) -> None: + # When TryAgain is raised from within an "except" block, reraise=True + # should surface the underlying exception rather than TryAgain itself. + class UnderlyingError(Exception): + pass + + def _r() -> None: + try: + raise UnderlyingError("boom") + except UnderlyingError: + # Implicit chaining via __context__ is exactly what we test. + raise tenacity.TryAgain # noqa: B904 + + r = Retrying( + stop=tenacity.stop_after_attempt(5), + retry=tenacity.retry_never, + reraise=True, + ) + self.assertRaises(UnderlyingError, r, _r) + self.assertEqual(5, r.statistics["attempt_number"]) + + def test_retry_try_again_from_cause_reraise(self) -> None: + # An explicit "raise TryAgain from exc" should also be unwrapped. + class UnderlyingError(Exception): + pass + + def _r() -> None: + raise tenacity.TryAgain from UnderlyingError("boom") + + r = Retrying( + stop=tenacity.stop_after_attempt(5), + retry=tenacity.retry_never, + reraise=True, + ) + self.assertRaises(UnderlyingError, r, _r) + self.assertEqual(5, r.statistics["attempt_number"]) + + def test_retry_try_again_forever_reraise(self) -> None: + def _r() -> None: + raise tenacity.TryAgain + + r = Retrying( + stop=tenacity.stop_after_attempt(5), + retry=tenacity.retry_never, + reraise=True, + ) + self.assertRaises(tenacity.TryAgain, r, _r) + self.assertEqual(5, r.statistics["attempt_number"]) + + def test_retry_if_exception_message_negative_no_inputs(self) -> None: + with self.assertRaises(TypeError): + tenacity.retry_if_exception_message() + + def test_retry_if_exception_message_negative_too_many_inputs(self) -> None: + with self.assertRaises(TypeError): + tenacity.retry_if_exception_message(message="negative", match="negative") + + def test_retry_if_exception_message_empty_message(self) -> None: + # An exception whose str() is empty (e.g. a bare ``RuntimeError()``) is + # a valid target. ``message=""`` must be accepted and match it, rather + # than being treated as "no message given" by a truthiness check. + r = tenacity.retry_if_exception_message(message="") + self.assertEqual(r.message, "") + empty = make_retry_state( + 1, 0, last_result=tenacity.Future.construct(1, RuntimeError(), True) + ) + nonempty = make_retry_state( + 1, 0, last_result=tenacity.Future.construct(1, RuntimeError("boom"), True) + ) + self.assertTrue(r(empty)) + self.assertFalse(r(nonempty)) + + def test_retry_if_not_exception_message_empty_message(self) -> None: + r = tenacity.retry_if_not_exception_message(message="") + empty = make_retry_state( + 1, 0, last_result=tenacity.Future.construct(1, RuntimeError(), True) + ) + nonempty = make_retry_state( + 1, 0, last_result=tenacity.Future.construct(1, RuntimeError("boom"), True) + ) + self.assertFalse(r(empty)) + self.assertTrue(r(nonempty)) + + +class NoneReturnUntilAfterCount: + """Holds counter state for invoking a method several times in a row.""" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go(self) -> typing.Any: + """Return None until after count threshold has been crossed. + + Then return True. + """ + if self.counter < self.count: + self.counter += 1 + return None + return True + + +class NoIOErrorAfterCount: + """Holds counter state for invoking a method several times in a row.""" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go(self) -> typing.Any: + """Raise an IOError until after count threshold has been crossed. + + Then return True. + """ + if self.counter < self.count: + self.counter += 1 + raise OSError("Hi there, I'm an IOError") + return True + + +class NoNameErrorAfterCount: + """Holds counter state for invoking a method several times in a row.""" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go(self) -> typing.Any: + """Raise a NameError until after count threshold has been crossed. + + Then return True. + """ + if self.counter < self.count: + self.counter += 1 + raise NameError("Hi there, I'm a NameError") + return True + + +class NoNameErrorCauseAfterCount: + """Holds counter state for invoking a method several times in a row.""" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go2(self) -> typing.Any: + raise NameError("Hi there, I'm a NameError") + + def go(self) -> typing.Any: + """Raise an IOError with a NameError as cause until after count threshold has been crossed. + + Then return True. + """ + if self.counter < self.count: + self.counter += 1 + try: + self.go2() + except NameError as e: + raise OSError from e + + return True + + +class NoIOErrorCauseAfterCount: + """Holds counter state for invoking a method several times in a row.""" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go2(self) -> typing.Any: + raise OSError("Hi there, I'm an IOError") + + def go(self) -> typing.Any: + """Raise a NameError with an IOError as cause until after count threshold has been crossed. + + Then return True. + """ + if self.counter < self.count: + self.counter += 1 + try: + self.go2() + except OSError as e: + raise NameError from e + + return True + + +class NameErrorUntilCount: + """Holds counter state for invoking a method several times in a row.""" + + derived_message = "Hi there, I'm a NameError" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go(self) -> typing.Any: + """Return True until after count threshold has been crossed. + + Then raise a NameError. + """ + if self.counter < self.count: + self.counter += 1 + return True + raise NameError(self.derived_message) + + +class IOErrorUntilCount: + """Holds counter state for invoking a method several times in a row.""" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go(self) -> typing.Any: + """Return True until after count threshold has been crossed. + + Then raise an IOError. + """ + if self.counter < self.count: + self.counter += 1 + return True + raise OSError("Hi there, I'm an IOError") + + +class CustomError(Exception): + """This is a custom exception class. + + Note that For Python 2.x, we don't strictly need to extend BaseException, + however, Python 3.x will complain. While this test suite won't run + correctly under Python 3.x without extending from the Python exception + hierarchy, the actual module code is backwards compatible Python 2.x and + will allow for cases where exception classes don't extend from the + hierarchy. + """ + + def __init__(self, value: str) -> None: + self.value = value + + @override + def __str__(self) -> str: + return self.value + + +class NoCustomErrorAfterCount: + """Holds counter state for invoking a method several times in a row.""" + + derived_message = "This is a Custom exception class" + + def __init__(self, count: int) -> None: + self.counter = 0 + self.count = count + + def go(self) -> typing.Any: + """Raise a CustomError until after count threshold has been crossed. + + Then return True. + """ + if self.counter < self.count: + self.counter += 1 + raise CustomError(self.derived_message) + return True + + +class CapturingHandler(logging.Handler): + """Captures log records for inspection.""" + + def __init__(self, *args: typing.Any, **kwargs: typing.Any) -> None: + super().__init__(*args, **kwargs) + self.records: list[logging.LogRecord] = [] + + @override + def emit(self, record: logging.LogRecord) -> None: + self.records.append(record) + + +def current_time_ms() -> int: + return round(time.time() * 1000) + + +@retry( + wait=tenacity.wait_fixed(0.05), + retry=tenacity.retry_if_result(lambda result: result is None), +) +def _retryable_test_with_wait(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + stop=tenacity.stop_after_attempt(3), + retry=tenacity.retry_if_result(lambda result: result is None), +) +def _retryable_test_with_stop(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry(retry=tenacity.retry_if_exception_cause_type(NameError)) +def _retryable_test_with_exception_cause_type(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry(retry=tenacity.retry_if_exception_type(IOError)) +def _retryable_test_with_exception_type_io(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry(retry=tenacity.retry_if_not_exception_type(IOError)) +def _retryable_test_if_not_exception_type_io(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + stop=tenacity.stop_after_attempt(3), retry=tenacity.retry_if_exception_type(IOError) +) +def _retryable_test_with_exception_type_io_attempt_limit( + thing: typing.Any, +) -> typing.Any: + return thing.go() + + +@retry(retry=tenacity.retry_unless_exception_type(NameError)) +def _retryable_test_with_unless_exception_type_name(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + stop=tenacity.stop_after_attempt(3), + retry=tenacity.retry_unless_exception_type(NameError), +) +def _retryable_test_with_unless_exception_type_name_attempt_limit( + thing: typing.Any, +) -> typing.Any: + return thing.go() + + +@retry(retry=tenacity.retry_unless_exception_type()) +def _retryable_test_with_unless_exception_type_no_input( + thing: typing.Any, +) -> typing.Any: + return thing.go() + + +@retry( + stop=tenacity.stop_after_attempt(5), + retry=tenacity.retry_if_exception_message( + message=NoCustomErrorAfterCount.derived_message + ), +) +def _retryable_test_if_exception_message_message(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + retry=tenacity.retry_if_not_exception_message( + message=NoCustomErrorAfterCount.derived_message + ) +) +def _retryable_test_if_not_exception_message_message(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + retry=tenacity.retry_if_exception_message( + match=NoCustomErrorAfterCount.derived_message[:3] + ".*" + ) +) +def _retryable_test_if_exception_message_match(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + retry=tenacity.retry_if_not_exception_message( + match=NoCustomErrorAfterCount.derived_message[:3] + ".*" + ) +) +def _retryable_test_if_not_exception_message_match(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + retry=tenacity.retry_if_not_exception_message( + message=NameErrorUntilCount.derived_message + ) +) +def _retryable_test_not_exception_message_delay(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry +def _retryable_default(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry() +def _retryable_default_f(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry(retry=tenacity.retry_if_exception_type(CustomError)) +def _retryable_test_with_exception_type_custom(thing: typing.Any) -> typing.Any: + return thing.go() + + +@retry( + stop=tenacity.stop_after_attempt(3), + retry=tenacity.retry_if_exception_type(CustomError), +) +def _retryable_test_with_exception_type_custom_attempt_limit( + thing: typing.Any, +) -> typing.Any: + return thing.go() + + +class TestDecoratorWrapper(unittest.TestCase): + def test_with_wait(self) -> None: + start = current_time_ms() + result = _retryable_test_with_wait(NoneReturnUntilAfterCount(5)) + t = current_time_ms() - start + self.assertGreaterEqual(t, 250) + self.assertTrue(result) + + def test_with_stop_on_return_value(self) -> None: + try: + _retryable_test_with_stop(NoneReturnUntilAfterCount(5)) + self.fail("Expected RetryError after 3 attempts") + except RetryError as re: + self.assertFalse(re.last_attempt.failed) + self.assertEqual(3, re.last_attempt.attempt_number) + self.assertTrue(re.last_attempt.result() is None) + print(re) + + def test_with_stop_on_exception(self) -> None: + try: + _retryable_test_with_stop(NoIOErrorAfterCount(5)) + self.fail("Expected IOError") + except OSError as re: + self.assertTrue(isinstance(re, IOError)) + print(re) + + def test_retry_if_exception_of_type(self) -> None: + self.assertTrue(_retryable_test_with_exception_type_io(NoIOErrorAfterCount(5))) + + try: + _retryable_test_with_exception_type_io(NoNameErrorAfterCount(5)) + self.fail("Expected NameError") + except NameError as n: + self.assertTrue(isinstance(n, NameError)) + print(n) + + self.assertTrue( + _retryable_test_with_exception_type_custom(NoCustomErrorAfterCount(5)) + ) + + try: + _retryable_test_with_exception_type_custom(NoNameErrorAfterCount(5)) + self.fail("Expected NameError") + except NameError as n: + self.assertTrue(isinstance(n, NameError)) + print(n) + + def test_retry_except_exception_of_type(self) -> None: + self.assertTrue( + _retryable_test_if_not_exception_type_io(NoNameErrorAfterCount(5)) + ) + + try: + _retryable_test_if_not_exception_type_io(NoIOErrorAfterCount(5)) + self.fail("Expected IOError") + except OSError as err: + self.assertTrue(isinstance(err, IOError)) + print(err) + + def test_retry_until_exception_of_type_attempt_number(self) -> None: + try: + self.assertTrue( + _retryable_test_with_unless_exception_type_name(NameErrorUntilCount(5)) + ) + except NameError as e: + s = _retryable_test_with_unless_exception_type_name.statistics + self.assertTrue(s["attempt_number"] == 6) + print(e) + else: + self.fail("Expected NameError") + + def test_retry_until_exception_of_type_no_type(self) -> None: + try: + # no input should catch all subclasses of Exception + self.assertTrue( + _retryable_test_with_unless_exception_type_no_input( + NameErrorUntilCount(5) + ) + ) + except NameError as e: + s = _retryable_test_with_unless_exception_type_no_input.statistics + self.assertTrue(s["attempt_number"] == 6) + print(e) + else: + self.fail("Expected NameError") + + def test_retry_until_exception_of_type_wrong_exception(self) -> None: + try: + # two iterations with IOError, one that returns True + _retryable_test_with_unless_exception_type_name_attempt_limit( + IOErrorUntilCount(2) + ) + self.fail("Expected RetryError") + except RetryError as e: + self.assertTrue(isinstance(e, RetryError)) + print(e) + + def test_retry_if_exception_message(self) -> None: + try: + self.assertTrue( + _retryable_test_if_exception_message_message(NoCustomErrorAfterCount(3)) + ) + except CustomError: + print(_retryable_test_if_exception_message_message.statistics) + self.fail("CustomError should've been retried from errormessage") + + def test_retry_if_not_exception_message(self) -> None: + try: + self.assertTrue( + _retryable_test_if_not_exception_message_message( + NoCustomErrorAfterCount(2) + ) + ) + except CustomError: + s = _retryable_test_if_not_exception_message_message.statistics + self.assertTrue(s["attempt_number"] == 1) + + def test_retry_if_not_exception_message_delay(self) -> None: + try: + self.assertTrue( + _retryable_test_not_exception_message_delay(NameErrorUntilCount(3)) + ) + except NameError: + s = _retryable_test_not_exception_message_delay.statistics + print(s["attempt_number"]) + self.assertTrue(s["attempt_number"] == 4) + + def test_retry_if_exception_message_match(self) -> None: + try: + self.assertTrue( + _retryable_test_if_exception_message_match(NoCustomErrorAfterCount(3)) + ) + except CustomError: + self.fail("CustomError should've been retried from errormessage") + + def test_retry_if_not_exception_message_match(self) -> None: + try: + self.assertTrue( + _retryable_test_if_not_exception_message_message( + NoCustomErrorAfterCount(2) + ) + ) + except CustomError: + s = _retryable_test_if_not_exception_message_message.statistics + self.assertTrue(s["attempt_number"] == 1) + + def test_retry_if_exception_cause_type(self) -> None: + self.assertTrue( + _retryable_test_with_exception_cause_type(NoNameErrorCauseAfterCount(5)) + ) + + try: + _retryable_test_with_exception_cause_type(NoIOErrorCauseAfterCount(5)) + self.fail("Expected exception without NameError as cause") + except NameError: + pass + + def test_retry_if_exception_cause_type_handles_cause_cycles(self) -> None: + """Cyclic __cause__ chains must not hang the retry predicate (#658).""" + + def boom_self_cause() -> None: + try: + raise ValueError("inner") + except ValueError as e: + raise e from e + + def boom_two_node_cycle() -> None: + a = ValueError("a") + b = RuntimeError("b") + a.__cause__ = b + b.__cause__ = a + raise a + + for boom in (boom_self_cause, boom_two_node_cycle): + r = tenacity.Retrying( + retry=tenacity.retry_if_exception_cause_type(KeyError), + stop=tenacity.stop_after_attempt(2), + reraise=True, + ) + with self.assertRaises((ValueError, RuntimeError)): + r(boom) + + def test_retry_preserves_argument_defaults(self) -> None: + def function_with_defaults(a: int = 1) -> int: + return a + + def function_with_kwdefaults(*, a: int = 1) -> int: + return a + + retrying = Retrying( + wait=tenacity.wait_fixed(0.01), stop=tenacity.stop_after_attempt(3) + ) + wrapped_defaults_function = retrying.wraps(function_with_defaults) + wrapped_kwdefaults_function = retrying.wraps(function_with_kwdefaults) + + self.assertEqual( + function_with_defaults.__defaults__, + wrapped_defaults_function.__defaults__, # type: ignore[attr-defined] + ) + self.assertEqual( + function_with_kwdefaults.__kwdefaults__, + wrapped_kwdefaults_function.__kwdefaults__, # type: ignore[attr-defined] + ) + + def test_defaults(self) -> None: + self.assertTrue(_retryable_default(NoNameErrorAfterCount(5))) + self.assertTrue(_retryable_default_f(NoNameErrorAfterCount(5))) + self.assertTrue(_retryable_default(NoCustomErrorAfterCount(5))) + self.assertTrue(_retryable_default_f(NoCustomErrorAfterCount(5))) + + def test_retry_function_object(self) -> None: + """Test that functools.wraps doesn't cause problems with callable objects. + + It raises an error upon trying to wrap it in Py2, because __name__ + attribute is missing. It's fixed in Py3 but was never backported. + """ + + class Hello: + def __call__(self) -> str: + return "Hello" + + retrying = Retrying( + wait=tenacity.wait_fixed(0.01), stop=tenacity.stop_after_attempt(3) + ) + h = retrying.wraps(Hello()) + self.assertEqual(h(), "Hello") + + def test_retry_function_attributes(self) -> None: + """Test that the wrapped function attributes are exposed as intended. + + - statistics contains the value for the latest function run + - retry object can be modified to change its behaviour (useful to patch in tests) + - retry object statistics are synced with function statistics + """ + + self.assertTrue(_retryable_test_with_stop(NoneReturnUntilAfterCount(2))) + + expected_stats = { + "attempt_number": 3, + "delay_since_first_attempt": mock.ANY, + "idle_for": mock.ANY, + "start_time": mock.ANY, + } + self.assertEqual(_retryable_test_with_stop.statistics, expected_stats) + self.assertEqual(_retryable_test_with_stop.retry.statistics, expected_stats) + + with mock.patch.object( + _retryable_test_with_stop.retry, + "stop", + tenacity.stop_after_attempt(1), + ): + try: + self.assertTrue(_retryable_test_with_stop(NoneReturnUntilAfterCount(2))) + except RetryError as exc: + expected_stats = { + "attempt_number": 1, + "delay_since_first_attempt": mock.ANY, + "idle_for": mock.ANY, + "start_time": mock.ANY, + } + self.assertEqual(_retryable_test_with_stop.statistics, expected_stats) + self.assertEqual(exc.last_attempt.attempt_number, 1) + self.assertEqual( + _retryable_test_with_stop.retry.statistics, expected_stats + ) + else: + self.fail("RetryError should have been raised after 1 attempt") + + +class TestStatisticsKeys: + def test_delay_since_first_attempt_available_on_first_attempt(self) -> None: + """delay_since_first_attempt should be in statistics from the start.""" + + @retry( + stop=tenacity.stop_after_attempt(3), + retry=tenacity.retry_if_result(lambda x: x is None), + ) + def succeeds_first_try() -> bool: + assert "delay_since_first_attempt" in succeeds_first_try.statistics + assert succeeds_first_try.statistics["delay_since_first_attempt"] == 0 + return True + + succeeds_first_try() + assert succeeds_first_try.statistics["delay_since_first_attempt"] == 0 + + def test_statistics_visible_through_outer_decorator(self) -> None: + """Statistics must resolve when @retry is wrapped by another decorator. + + A well-behaved outer decorator uses functools.wraps, which copies the + inner wrapper's ``__dict__`` (including ``statistics``). Rebinding the + attribute on each call left the outer wrapper pointing at a stale empty + dict. The statistics must instead stay visible through the wrapper + chain. See issue #519. + """ + import functools + + _F = typing.TypeVar("_F", bound=typing.Callable[..., typing.Any]) + + def outer(fn: _F) -> _F: + @functools.wraps(fn) + def wrapper(*args: typing.Any, **kwargs: typing.Any) -> typing.Any: + return fn(*args, **kwargs) + + return typing.cast("_F", wrapper) + + @outer + @retry(stop=tenacity.stop_after_attempt(3)) + def my_call() -> str: + return "ok" + + assert my_call() == "ok" + assert my_call.statistics["attempt_number"] == 1 + assert my_call.statistics is my_call.__wrapped__.statistics + + +class TestEnabled: + def test_enabled_false_skips_retry(self) -> None: + """When enabled=False, the function is called directly without retrying.""" + call_count = 0 + + @retry(enabled=False, stop=tenacity.stop_after_attempt(3)) + def always_fails() -> None: + nonlocal call_count + call_count += 1 + raise ValueError("fail") + + with pytest.raises(ValueError, match="fail"): + always_fails() + assert call_count == 1 + + def test_enabled_false_preserves_attributes(self) -> None: + """When enabled=False, .retry, .retry_with, .statistics are still available.""" + + @retry(enabled=False, stop=tenacity.stop_after_attempt(3)) + def my_func() -> str: + return "ok" + + assert hasattr(my_func, "retry") + assert hasattr(my_func, "retry_with") + assert hasattr(my_func, "statistics") + assert my_func() == "ok" + + def test_enabled_false_via_retry_with(self) -> None: + """retry_with(enabled=False) disables retrying.""" + call_count = 0 + + @retry(stop=tenacity.stop_after_attempt(3)) + def always_fails() -> None: + nonlocal call_count + call_count += 1 + raise ValueError("fail") + + disabled = always_fails.retry_with(enabled=False) + with pytest.raises(ValueError, match="fail"): + disabled() + assert call_count == 1 + + def test_enabled_true_retries_normally(self) -> None: + """When enabled=True (default), retrying works as usual.""" + call_count = 0 + + @retry(enabled=True, stop=tenacity.stop_after_attempt(3), reraise=True) + def fails_twice() -> bool: + nonlocal call_count + call_count += 1 + if call_count < 3: + raise ValueError("fail") + return True + + assert fails_twice() is True + assert call_count == 3 + + def test_enabled_false_iter_raises_original_exception(self) -> None: + """When enabled=False, the iterator protocol raises the original exception, + not a RetryError, and the body executes exactly once.""" + call_count = 0 + retrying = Retrying( + enabled=False, + stop=tenacity.stop_after_attempt(5), + wait=tenacity.wait_none(), + ) + with pytest.raises(ValueError, match="fail"): + for attempt in retrying: + with attempt: + call_count += 1 + raise ValueError("fail") + assert call_count == 1 + + def test_enabled_false_iter_succeeds_on_first_attempt(self) -> None: + """When enabled=False, the iterator protocol runs the body once and stops.""" + call_count = 0 + retrying = Retrying( + enabled=False, + stop=tenacity.stop_after_attempt(5), + wait=tenacity.wait_none(), + ) + for attempt in retrying: + with attempt: + call_count += 1 + assert call_count == 1 + + def test_enabled_false_call_raises_original_exception(self) -> None: + """When enabled=False, calling the controller directly raises the original + exception, not a RetryError, and the function executes exactly once.""" + call_count = 0 + + def fails() -> None: + nonlocal call_count + call_count += 1 + raise ValueError("fail") + + retrying = Retrying( + enabled=False, + stop=tenacity.stop_after_attempt(5), + wait=tenacity.wait_none(), + ) + with pytest.raises(ValueError, match="fail"): + retrying(fails) + assert call_count == 1 + + def test_enabled_false_call_succeeds_on_first_attempt(self) -> None: + """When enabled=False, calling the controller directly runs the function once + and returns its result.""" + call_count = 0 + + def succeeds() -> str: + nonlocal call_count + call_count += 1 + return "ok" + + retrying = Retrying( + enabled=False, + stop=tenacity.stop_after_attempt(5), + wait=tenacity.wait_none(), + ) + assert retrying(succeeds) == "ok" + assert call_count == 1 + + +class TestRetryWith: + def test_redefine_wait(self) -> None: + start = current_time_ms() + result = _retryable_test_with_wait.retry_with(wait=tenacity.wait_fixed(0.1))( + NoneReturnUntilAfterCount(5) + ) + t = current_time_ms() - start + assert t >= 500 + assert result is True + + def test_redefine_stop(self) -> None: + result = _retryable_test_with_stop.retry_with( + stop=tenacity.stop_after_attempt(5) + )(NoneReturnUntilAfterCount(4)) + assert result is True + + def test_retry_error_cls_should_be_preserved(self) -> None: + @retry(stop=tenacity.stop_after_attempt(10), retry_error_cls=ValueError) # type: ignore[arg-type] + def _retryable() -> None: + raise Exception("raised for test purposes") + + with pytest.raises(Exception) as exc_ctx: + _retryable.retry_with(stop=tenacity.stop_after_attempt(2))() + + assert exc_ctx.type is ValueError, "Should remap to specific exception type" + + def test_retry_error_callback_should_be_preserved(self) -> None: + def return_text(retry_state: RetryCallState) -> str: + return f"Calling {retry_state.fn.__name__} keeps raising errors after {retry_state.attempt_number} attempts" # type: ignore[union-attr] + + @retry(stop=tenacity.stop_after_attempt(10), retry_error_callback=return_text) + def _retryable() -> None: + raise Exception("raised for test purposes") + + result = _retryable.retry_with(stop=tenacity.stop_after_attempt(5))() + assert result == "Calling _retryable keeps raising errors after 5 attempts" + + +class TestBeforeAfterAttempts(unittest.TestCase): + _attempt_number = 0 + + def test_before_attempts(self) -> None: + TestBeforeAfterAttempts._attempt_number = 0 + + def _before(retry_state: RetryCallState) -> None: + TestBeforeAfterAttempts._attempt_number = retry_state.attempt_number + + @retry( + wait=tenacity.wait_fixed(1), + stop=tenacity.stop_after_attempt(1), + before=_before, + ) + def _test_before() -> None: + pass + + _test_before() + + self.assertTrue(TestBeforeAfterAttempts._attempt_number == 1) + + def test_after_attempts(self) -> None: + TestBeforeAfterAttempts._attempt_number = 0 + + def _after(retry_state: RetryCallState) -> None: + TestBeforeAfterAttempts._attempt_number = retry_state.attempt_number + + @retry( + wait=tenacity.wait_fixed(0.1), + stop=tenacity.stop_after_attempt(3), + after=_after, + ) + def _test_after() -> None: + if TestBeforeAfterAttempts._attempt_number < 2: + raise Exception("testing after_attempts handler") + + _test_after() + + self.assertTrue(TestBeforeAfterAttempts._attempt_number == 2) + + def test_before_sleep(self) -> None: + def _before_sleep(retry_state: RetryCallState) -> None: + self.assertGreater(retry_state.next_action.sleep, 0) # type: ignore[union-attr] + _before_sleep.attempt_number = retry_state.attempt_number # type: ignore[attr-defined] + + @retry( + wait=tenacity.wait_fixed(0.01), + stop=tenacity.stop_after_attempt(3), + before_sleep=_before_sleep, + ) + def _test_before_sleep() -> None: + if _before_sleep.attempt_number < 2: # type: ignore[attr-defined] + raise Exception("testing before_sleep_attempts handler") + + _test_before_sleep() + self.assertEqual(_before_sleep.attempt_number, 2) # type: ignore[attr-defined] + + def _before_sleep_log_raises( + self, get_call_fn: typing.Callable[..., typing.Any] + ) -> None: + thing = NoIOErrorAfterCount(2) + logger = logging.getLogger(self.id()) + logger.propagate = False + logger.setLevel(logging.INFO) + handler = CapturingHandler() + logger.addHandler(handler) + try: + _before_sleep = tenacity.before_sleep_log(logger, logging.INFO) + retrying = Retrying( + wait=tenacity.wait_fixed(0.01), + stop=tenacity.stop_after_attempt(3), + before_sleep=_before_sleep, + ) + get_call_fn(retrying)(thing.go) + finally: + logger.removeHandler(handler) + + etalon_re = ( + r"^Retrying .* in 0\.01 seconds as it raised " + r"(IO|OS)Error: Hi there, I'm an IOError\.$" + ) + self.assertEqual(len(handler.records), 2) + fmt = logging.Formatter().format + self.assertRegex(fmt(handler.records[0]), etalon_re) + self.assertRegex(fmt(handler.records[1]), etalon_re) + + def test_before_sleep_log_raises(self) -> None: + self._before_sleep_log_raises(lambda x: x) + + def test_before_sleep_log_raises_with_exc_info(self) -> None: + thing = NoIOErrorAfterCount(2) + logger = logging.getLogger(self.id()) + logger.propagate = False + logger.setLevel(logging.INFO) + handler = CapturingHandler() + logger.addHandler(handler) + try: + _before_sleep = tenacity.before_sleep_log( + logger, logging.INFO, exc_info=True + ) + retrying = Retrying( + wait=tenacity.wait_fixed(0.01), + stop=tenacity.stop_after_attempt(3), + before_sleep=_before_sleep, + ) + retrying(thing.go) + finally: + logger.removeHandler(handler) + + etalon_re = re.compile( + r"^Retrying .* in 0\.01 seconds as it raised " + r"(IO|OS)Error: Hi there, I'm an IOError\.{0}" + r"Traceback \(most recent call last\):{0}" + r".*$".format("\n"), + flags=re.MULTILINE, + ) + self.assertEqual(len(handler.records), 2) + fmt = logging.Formatter().format + self.assertRegex(fmt(handler.records[0]), etalon_re) + self.assertRegex(fmt(handler.records[1]), etalon_re) + + def test_before_sleep_log_returns(self, exc_info: bool = False) -> None: + thing = NoneReturnUntilAfterCount(2) + logger = logging.getLogger(self.id()) + logger.propagate = False + logger.setLevel(logging.INFO) + handler = CapturingHandler() + logger.addHandler(handler) + try: + _before_sleep = tenacity.before_sleep_log( + logger, logging.INFO, exc_info=exc_info + ) + _retry = tenacity.retry_if_result(lambda result: result is None) + retrying = Retrying( + wait=tenacity.wait_fixed(0.01), + stop=tenacity.stop_after_attempt(3), + retry=_retry, + before_sleep=_before_sleep, + ) + retrying(thing.go) + finally: + logger.removeHandler(handler) + + etalon_re = r"^Retrying .* in 0\.01 seconds as it returned None\.$" + self.assertEqual(len(handler.records), 2) + fmt = logging.Formatter().format + self.assertRegex(fmt(handler.records[0]), etalon_re) + self.assertRegex(fmt(handler.records[1]), etalon_re) + + def test_before_sleep_log_returns_with_exc_info(self) -> None: + self.test_before_sleep_log_returns(exc_info=True) + + +class TestReraiseExceptions(unittest.TestCase): + def test_reraise_by_default(self) -> None: + calls = [] + + @retry( + wait=tenacity.wait_fixed(0.1), + stop=tenacity.stop_after_attempt(2), + reraise=True, + ) + def _reraised_by_default() -> None: + calls.append("x") + raise KeyError("Bad key") + + self.assertRaises(KeyError, _reraised_by_default) + self.assertEqual(2, len(calls)) + + def test_reraise_from_retry_error(self) -> None: + calls = [] + + @retry(wait=tenacity.wait_fixed(0.1), stop=tenacity.stop_after_attempt(2)) + def _raise_key_error() -> None: + calls.append("x") + raise KeyError("Bad key") + + def _reraised_key_error() -> None: + try: + _raise_key_error() + except tenacity.RetryError as retry_err: + retry_err.reraise() + + self.assertRaises(KeyError, _reraised_key_error) + self.assertEqual(2, len(calls)) + + def test_reraise_timeout_from_retry_error(self) -> None: + calls = [] + + @retry( + wait=tenacity.wait_fixed(0.1), + stop=tenacity.stop_after_attempt(2), + retry=lambda retry_state: True, + ) + def _mock_fn() -> None: + calls.append("x") + + def _reraised_mock_fn() -> None: + try: + _mock_fn() + except tenacity.RetryError as retry_err: + retry_err.reraise() + + self.assertRaises(tenacity.RetryError, _reraised_mock_fn) + self.assertEqual(2, len(calls)) + + def test_reraise_no_exception(self) -> None: + calls = [] + + @retry( + wait=tenacity.wait_fixed(0.1), + stop=tenacity.stop_after_attempt(2), + retry=lambda retry_state: True, + reraise=True, + ) + def _mock_fn() -> None: + calls.append("x") + + self.assertRaises(tenacity.RetryError, _mock_fn) + self.assertEqual(2, len(calls)) + + +class TestStatistics(unittest.TestCase): + def test_stats(self) -> None: + @retry() + def _foobar() -> int: + return 42 + + self.assertEqual({}, _foobar.statistics) + _foobar() + self.assertEqual(1, _foobar.statistics["attempt_number"]) + + def test_stats_failing(self) -> None: + @retry(stop=tenacity.stop_after_attempt(2)) + def _foobar() -> None: + raise ValueError(42) + + self.assertEqual({}, _foobar.statistics) + with contextlib.suppress(Exception): + _foobar() + self.assertEqual(2, _foobar.statistics["attempt_number"]) + + def test_retry_object_statistics_synced(self) -> None: + """Test that func.retry.statistics is synced with func.statistics.""" + + @retry(stop=tenacity.stop_after_attempt(3)) + def _foobar() -> int: + return 42 + + _foobar() + self.assertEqual( + _foobar.retry.statistics["attempt_number"], + _foobar.statistics["attempt_number"], + ) + + def test_retry_object_statistics_during_execution(self) -> None: + """Test that func.retry.statistics is accessible during execution.""" + attempts: list[int] = [] + + @retry( + stop=tenacity.stop_after_attempt(3), + retry=tenacity.retry_if_exception_type(ValueError), + reraise=True, + ) + def _foobar() -> int: + attempts.append(_foobar.retry.statistics["attempt_number"]) + if len(attempts) < 3: + raise ValueError("retry") + return 42 + + _foobar() + self.assertEqual(attempts, [1, 2, 3]) + + +class TestRetryErrorCallback(unittest.TestCase): + @override + def setUp(self) -> None: + self._attempt_number = 0 + self._callback_called = False + + def _callback(self, fut: tenacity.Future) -> tenacity.Future: + self._callback_called = True + return fut + + def test_retry_error_callback(self) -> None: + num_attempts = 3 + + def retry_error_callback(retry_state: RetryCallState) -> typing.Any: + retry_error_callback.called_times += 1 # type: ignore[attr-defined] + return retry_state.outcome + + retry_error_callback.called_times = 0 # type: ignore[attr-defined] + + @retry( + stop=tenacity.stop_after_attempt(num_attempts), + retry_error_callback=retry_error_callback, + ) + def _foobar() -> None: + self._attempt_number += 1 + raise Exception("This exception should not be raised") + + result = _foobar() + + self.assertEqual(retry_error_callback.called_times, 1) # type: ignore[attr-defined] + self.assertEqual(num_attempts, self._attempt_number) + self.assertIsInstance(result, tenacity.Future) + + +class TestContextManager(unittest.TestCase): + def test_context_manager_retry_one(self) -> None: + from tenacity import Retrying + + raise_ = True + + for attempt in Retrying(): + with attempt: + if raise_: + raise_ = False + raise Exception("Retry it!") + + def test_context_manager_on_error(self) -> None: + from tenacity import Retrying + + class CustomError(Exception): + pass + + retry = Retrying(retry=tenacity.retry_if_exception_type(IOError)) + + def test() -> None: + for attempt in retry: + with attempt: + raise CustomError("Don't retry!") + + self.assertRaises(CustomError, test) + + def test_context_manager_retry_error(self) -> None: + from tenacity import Retrying + + retry = Retrying(stop=tenacity.stop_after_attempt(2)) + + def test() -> None: + for attempt in retry: + with attempt: + raise Exception("Retry it!") + + self.assertRaises(RetryError, test) + + def test_context_manager_reraise(self) -> None: + from tenacity import Retrying + + class CustomError(Exception): + pass + + retry = Retrying(reraise=True, stop=tenacity.stop_after_attempt(2)) + + def test() -> None: + for attempt in retry: + with attempt: + raise CustomError("Don't retry!") + + self.assertRaises(CustomError, test) + + +class TestInvokeAsCallable: + """Test direct invocation of Retrying as a callable.""" + + @staticmethod + def invoke(retry: Retrying, f: typing.Callable[..., typing.Any]) -> typing.Any: + """ + Invoke Retrying logic. + + Wrapper allows testing different call mechanisms in test sub-classes. + """ + return retry(f) + + def test_retry_one(self) -> None: + def f() -> typing.Any: + f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] + if len(f.calls) <= 1: # type: ignore[attr-defined] + raise Exception("Retry it!") + return 42 + + f.calls = [] # type: ignore[attr-defined] + + retry = Retrying() + assert self.invoke(retry, f) == 42 + assert f.calls == [1, 2] # type: ignore[attr-defined] + + def test_on_error(self) -> None: + class CustomError(Exception): + pass + + def f() -> typing.Any: + f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] + if len(f.calls) <= 1: # type: ignore[attr-defined] + raise CustomError("Don't retry!") + return 42 + + f.calls = [] # type: ignore[attr-defined] + + retry = Retrying(retry=tenacity.retry_if_exception_type(IOError)) + with pytest.raises(CustomError): + self.invoke(retry, f) + assert f.calls == [1] # type: ignore[attr-defined] + + def test_retry_error(self) -> None: + def f() -> typing.Any: + f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] + raise Exception("Retry it!") + + f.calls = [] # type: ignore[attr-defined] + + retry = Retrying(stop=tenacity.stop_after_attempt(2)) + with pytest.raises(RetryError): + self.invoke(retry, f) + assert f.calls == [1, 2] # type: ignore[attr-defined] + + def test_reraise(self) -> None: + class CustomError(Exception): + pass + + def f() -> typing.Any: + f.calls.append(len(f.calls) + 1) # type: ignore[attr-defined] + raise CustomError("Retry it!") + + f.calls = [] # type: ignore[attr-defined] + + retry = Retrying(reraise=True, stop=tenacity.stop_after_attempt(2)) + with pytest.raises(CustomError): + self.invoke(retry, f) + assert f.calls == [1, 2] # type: ignore[attr-defined] + + +class TestRetryException(unittest.TestCase): + def test_retry_error_is_pickleable(self) -> None: + import pickle + + expected = RetryError(last_attempt=123) # type: ignore[arg-type] + pickled = pickle.dumps(expected) + actual = pickle.loads(pickled) + self.assertEqual(expected.last_attempt, actual.last_attempt) + + +class TestRetryTyping(unittest.TestCase): + def test_retry_type_annotations(self) -> None: + """The decorator should maintain types of decorated functions. + + The annotations below are the assertions; mypy checks them when it runs + over this file. The negative case leans on warn_unused_ignores: should + @retry ever decay to returning Any, that assignment would stop being an + error and the now-dead ignore would fail the type check. + """ + + def num_to_str(number: int) -> str: + return str(number) + + # equivalent to a raw @retry decoration + with_raw = retry(num_to_str) + with_raw_result = with_raw(1) + + # equivalent to a @retry(...) decoration + with_constructor = retry()(num_to_str) + with_constructor_result = with_constructor(1) + + # The wrapper stays usable wherever the undecorated function was. + _raw_signature: typing.Callable[[int], str] = with_raw + _constructor_signature: typing.Callable[[int], str] = with_constructor + + # ...and an incompatible signature is still rejected. + _mismatch: typing.Callable[[str], int] = with_raw # type: ignore[assignment] + + self.assertEqual(with_raw_result, "1") + self.assertEqual(with_constructor_result, "1") + + def test_retry_decorated_method_keeps_bound_signature(self) -> None: + """A decorated instance method must type-check like a bound method. + + Without a descriptor (``__get__``) on ``_RetryDecorated``, static type + checkers treat ``instance.method`` the same as the unbound + ``Class.method``, so a normal call with only the non-``self`` keyword + arguments looks like a type error and the return type resolves to + ``Any``/``Unknown``. This does not fail at runtime (functools.wraps + returns a real function, which Python always binds correctly), but it + does fail under `mypy --strict`, which also type-checks this file. + See issue #532. + """ + + class Doubler: + @retry(stop=tenacity.stop_after_attempt(3)) + def double(self, value: int) -> int: + return value * 2 + + doubler = Doubler() + result: int = doubler.double(value=21) + self.assertEqual(result, 42) + + +class TestMockingSleep: + RETRY_ARGS = { + "wait": tenacity.wait_fixed(0.1), + "stop": tenacity.stop_after_attempt(5), + } + + def _fail(self) -> None: + raise NotImplementedError + + @retry(**RETRY_ARGS) # type: ignore[call-overload, untyped-decorator] + def _decorated_fail(self) -> None: + self._fail() + + @pytest.fixture() + def mock_sleep( + self, monkeypatch: typing.Any + ) -> typing.Generator[typing.Any, None, None]: + class MockSleep: + call_count = 0 + + def __call__(self, seconds: float) -> None: + self.call_count += 1 + + sleep = MockSleep() + monkeypatch.setattr(tenacity.nap.time, "sleep", sleep) # type: ignore[attr-defined] + yield sleep + + def test_decorated(self, mock_sleep: typing.Any) -> None: + with pytest.raises(RetryError): + self._decorated_fail() + assert mock_sleep.call_count == 4 + + def test_decorated_retry_with(self, mock_sleep: typing.Any) -> None: + fail_faster = self._decorated_fail.retry_with( + stop=tenacity.stop_after_attempt(2), + ) + with pytest.raises(RetryError): + fail_faster() + assert mock_sleep.call_count == 1 + + +class TestPickle(unittest.TestCase): + def test_retrying_picklable(self) -> None: + """Retrying objects can be pickled for multiprocessing support.""" + retrying = Retrying(stop=tenacity.stop_after_attempt(3)) + pickled = pickle.dumps(retrying) + restored = pickle.loads(pickled) + assert isinstance(restored, Retrying) + assert isinstance(restored.stop, tenacity.stop_after_attempt) + + def test_retrying_picklable_after_run(self) -> None: + """Retrying objects can be pickled even after being used.""" + retrying = Retrying(stop=tenacity.stop_after_attempt(3)) + # Access statistics to populate _local + _ = retrying.statistics + pickled = pickle.dumps(retrying) + restored = pickle.loads(pickled) + assert isinstance(restored, Retrying) + # Statistics should be reset on the restored object + assert restored.statistics == {} + + def test_retry_strategies_picklable(self) -> None: + """All built-in retry strategies can be pickled.""" + strategies = [ + tenacity.retry_if_exception_type(ValueError), + tenacity.retry_if_not_exception_type(ValueError), + tenacity.retry_if_exception_message(message="fail"), + tenacity.retry_if_exception_message(match="fail.*"), + tenacity.retry_if_not_exception_message(message="fail"), + ] + for strategy in strategies: + restored = pickle.loads(pickle.dumps(strategy)) + assert type(restored) is type(strategy) + + def test_retrying_pickle_round_trip_works(self) -> None: + """A pickled-then-restored Retrying object retries correctly.""" + retrying = Retrying( + stop=tenacity.stop_after_attempt(3), + retry=tenacity.retry_if_exception_type(ValueError), + reraise=True, + ) + restored = pickle.loads(pickle.dumps(retrying)) + + calls = 0 + + def succeed_on_third() -> str: + nonlocal calls + calls += 1 + if calls < 3: + raise ValueError("not yet") + return "ok" + + result = restored(succeed_on_third) + assert result == "ok" + assert calls == 3 + + +if __name__ == "__main__": + unittest.main()