diff --git a/tenacity/__init__.py b/tenacity/__init__.py index 6b591464..856c6973 100644 --- a/tenacity/__init__.py +++ b/tenacity/__init__.py @@ -491,6 +491,14 @@ def __iter__(self) -> t.Generator[AttemptManager, None, None]: 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): @@ -597,16 +605,23 @@ def __init__( 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. + context manager, the inferred caller name when used as a context manager + without an explicit ``name``, or ``""`` if none 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 @@ -701,9 +716,9 @@ def retry( stop: "StopBaseT" = ..., wait: "WaitBaseT" = ..., retry: "RetryBaseT | tasyncio.retry.RetryBaseT" = ..., - before: t.Callable[["RetryCallState"], None | t.Awaitable[None]] = ..., - after: t.Callable[["RetryCallState"], None | t.Awaitable[None]] = ..., - before_sleep: t.Callable[["RetryCallState"], None | t.Awaitable[None]] | None = ..., + 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]] @@ -718,9 +733,9 @@ def retry( stop: "StopBaseT" = stop_never, wait: "WaitBaseT" = wait_none(), retry: "RetryBaseT | tasyncio.retry.RetryBaseT" = retry_if_exception_type(), - before: t.Callable[["RetryCallState"], None | t.Awaitable[None]] = before_nothing, - after: t.Callable[["RetryCallState"], None | t.Awaitable[None]] = after_nothing, - before_sleep: t.Callable[["RetryCallState"], None | t.Awaitable[None]] + 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, diff --git a/tenacity/asyncio/__init__.py b/tenacity/asyncio/__init__.py index a20edc2f..a6c2253b 100644 --- a/tenacity/asyncio/__init__.py +++ b/tenacity/asyncio/__init__.py @@ -32,6 +32,7 @@ after_nothing, before_nothing, ) +from tenacity._utils import override # Import all built-in retry strategies for easier usage. from .retry import ( @@ -75,16 +76,16 @@ class AsyncRetrying(BaseRetrying): def __init__( self, sleep: t.Callable[ - [int | float], None | t.Awaitable[None] + [int | float], t.Awaitable[None] | None ] = _portable_async_sleep, stop: "StopBaseT" = tenacity.stop.stop_never, wait: "WaitBaseT" = tenacity.wait.wait_none(), retry: "SyncRetryBaseT | RetryBaseT" = tenacity.retry_if_exception_type(), before: t.Callable[ - ["RetryCallState"], None | t.Awaitable[None] + ["RetryCallState"], t.Awaitable[None] | None ] = before_nothing, - after: t.Callable[["RetryCallState"], None | t.Awaitable[None]] = after_nothing, - before_sleep: t.Callable[["RetryCallState"], None | t.Awaitable[None]] + 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, @@ -108,6 +109,7 @@ def __init__( enabled=enabled, ) + @override async def __call__( # type: ignore[override] self, fn: WrappedFn, *args: t.Any, **kwargs: t.Any ) -> WrappedFnReturnT: @@ -133,28 +135,30 @@ async def __call__( # type: ignore[override] else: return do # type: ignore[no-any-return] + @override def _add_action_func(self, fn: t.Callable[..., t.Any]) -> None: self.iter_state.actions.append(_utils.wrap_to_async_func(fn)) + @override async def _run_retry(self, retry_state: "RetryCallState") -> None: # type: ignore[override] self.iter_state.retry_run_result = await _utils.wrap_to_async_func(self.retry)( retry_state ) + @override async def _run_wait(self, retry_state: "RetryCallState") -> None: # type: ignore[override] - if self.wait: - sleep = await _utils.wrap_to_async_func(self.wait)(retry_state) - else: - sleep = 0.0 - - retry_state.upcoming_sleep = sleep + retry_state.upcoming_sleep = await _utils.wrap_to_async_func(self.wait)( + retry_state + ) + @override async def _run_stop(self, retry_state: "RetryCallState") -> None: # type: ignore[override] self.statistics["delay_since_first_attempt"] = retry_state.seconds_since_start self.iter_state.stop_run_result = await _utils.wrap_to_async_func(self.stop)( retry_state ) + @override async def iter(self, retry_state: "RetryCallState") -> DoAttempt | DoSleep | t.Any: self._begin_iter(retry_state) result = None @@ -162,6 +166,7 @@ async def iter(self, retry_state: "RetryCallState") -> DoAttempt | DoSleep | t.A result = await action(retry_state) return result + @override def __iter__(self) -> t.Generator[AttemptManager, None, None]: raise TypeError("AsyncRetrying object is not iterable") @@ -197,6 +202,7 @@ async def __anext__(self) -> AttemptManager: else: raise StopAsyncIteration + @override def wraps(self, fn: t.Callable[P, R]) -> _RetryDecorated[P, R]: wrapped = super().wraps(fn) # Ensure wrapper is recognized as a coroutine function. diff --git a/tenacity/retry.py b/tenacity/retry.py index b1781dae..dc7beab6 100644 --- a/tenacity/retry.py +++ b/tenacity/retry.py @@ -18,6 +18,8 @@ import re import typing +from tenacity._utils import override + if typing.TYPE_CHECKING: from tenacity import RetryCallState @@ -64,6 +66,7 @@ def __ror__(self, other: "RetryBaseT") -> "retry_any": class _retry_never(retry_base): """Retry strategy that never rejects any result.""" + @override def __call__(self, retry_state: "RetryCallState") -> bool: return False @@ -74,6 +77,7 @@ def __call__(self, retry_state: "RetryCallState") -> bool: class _retry_always(retry_base): """Retry strategy that always rejects any result.""" + @override def __call__(self, retry_state: "RetryCallState") -> bool: return True @@ -87,6 +91,7 @@ class retry_if_exception(retry_base): def __init__(self, predicate: typing.Callable[[BaseException], bool]) -> None: self.predicate = predicate + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.outcome is None: raise RuntimeError("__call__() called before outcome was set") @@ -143,6 +148,7 @@ def __init__( def _check(self, e: BaseException) -> bool: return not isinstance(e, self.exception_types) + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.outcome is None: raise RuntimeError("__call__() called before outcome was set") @@ -171,6 +177,7 @@ def __init__( ) -> None: self.exception_cause_types = exception_types + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.outcome is None: raise RuntimeError("__call__ called before outcome was set") @@ -196,6 +203,7 @@ class retry_if_result(retry_base): def __init__(self, predicate: typing.Callable[[typing.Any], bool]) -> None: self.predicate = predicate + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.outcome is None: raise RuntimeError("__call__() called before outcome was set") @@ -211,6 +219,7 @@ class retry_if_not_result(retry_base): def __init__(self, predicate: typing.Callable[[typing.Any], bool]) -> None: self.predicate = predicate + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.outcome is None: raise RuntimeError("__call__() called before outcome was set") @@ -226,7 +235,7 @@ class retry_if_exception_message(retry_if_exception): def __init__( self, message: str | None = None, - match: None | str | re.Pattern[str] = None, + match: str | re.Pattern[str] | None = None, ) -> None: if message is not None and match is not None: raise TypeError( @@ -239,7 +248,9 @@ def __init__( ) self.message = message - self.match = re.compile(match) if match is not None else None + self.match: re.Pattern[str] | None = ( + re.compile(match) if match is not None else None + ) super().__init__(self._check) def _check(self, exception: BaseException) -> bool: @@ -252,9 +263,11 @@ def _check(self, exception: BaseException) -> bool: class retry_if_not_exception_message(retry_if_exception_message): """Retries until an exception message equals or matches.""" + @override def _check(self, exception: BaseException) -> bool: return not super()._check(exception) + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.outcome is None: raise RuntimeError("__call__() called before outcome was set") @@ -274,9 +287,11 @@ class retry_any(retry_base): def __init__(self, *retries: "RetryBaseT") -> None: self.retries = retries + @override def __call__(self, retry_state: "RetryCallState") -> bool: return any(r(retry_state) for r in self.retries) + @override def __ror__(self, other: "RetryBaseT") -> "retry_any": if isinstance(other, retry_any): return retry_any(*other.retries, *self.retries) @@ -289,9 +304,11 @@ class retry_all(retry_base): def __init__(self, *retries: "RetryBaseT") -> None: self.retries = retries + @override def __call__(self, retry_state: "RetryCallState") -> bool: return all(r(retry_state) for r in self.retries) + @override def __rand__(self, other: "RetryBaseT") -> "retry_all": if isinstance(other, retry_all): return retry_all(*other.retries, *self.retries) diff --git a/tests/test_tenacity.py b/tests/test_tenacity.py index 8f74ec86..8cd2546e 100644 --- a/tests/test_tenacity.py +++ b/tests/test_tenacity.py @@ -176,6 +176,42 @@ def test_logging_uses_name(self) -> None: 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: