Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,10 @@ ignore = [

[tool.mypy]
strict = true
# Match `requires-python`, so results do not depend on which interpreter
# happens to run mypy. Without this it defaults to the running version and
# silently skips checking that the code is valid on the oldest one we support.
python_version = "3.10"
files = ["tenacity", "tests"]
show_error_codes = true
exclude = ["tenacity/_version\\.py"]
Expand All @@ -124,6 +128,7 @@ extra_checks = true
enable_error_code = [
"deprecated",
"exhaustive-match",
"explicit-override",
"ignore-without-code",
"mutable-override",
"possibly-undefined",
Expand Down
9 changes: 9 additions & 0 deletions tenacity/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
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
Expand Down Expand Up @@ -156,12 +157,14 @@ class BaseAction:
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)

Expand Down Expand Up @@ -201,6 +204,7 @@ def reraise(self) -> t.NoReturn:
raise self.last_attempt.result()
raise self

@override
def __str__(self) -> str:
return f"{self.__class__.__name__}[{self.last_attempt}]"

Expand Down Expand Up @@ -304,6 +308,7 @@ def copy(
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"}
Expand All @@ -312,9 +317,11 @@ 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 "<unknown>"

@override
def __repr__(self) -> str:
return (
f"<{self.__class__.__name__} object at 0x{id(self):x} ("
Expand Down Expand Up @@ -514,6 +521,7 @@ def __call__(
class Retrying(BaseRetrying):
"""Retrying controller."""

@override
def __call__(
self,
fn: t.Callable[..., WrappedFnReturnT],
Expand Down Expand Up @@ -638,6 +646,7 @@ def set_exception(
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"
Expand Down
23 changes: 23 additions & 0 deletions tenacity/_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,29 @@
import typing
from datetime import timedelta

if typing.TYPE_CHECKING:
# Type checkers recognise this by name and use it to enforce
# `explicit-override`; taking it from typing_extensions means they do so
# whatever `python_version` they are run under. It is never imported at
# runtime, so it stays a type-check-only dependency.
from typing_extensions import override as override
elif sys.version_info >= (3, 12):
from typing import override
else:
_F = typing.TypeVar("_F", bound=typing.Callable[..., typing.Any])

def override(method: _F) -> _F:
"""Backport of `typing.override` for Python < 3.12.

Only the runtime half is needed: setting the PEP 698 `__override__`
marker that introspection tools look for.
"""
with contextlib.suppress(AttributeError, TypeError):
# Not every callable allows attribute assignment.
method.__override__ = True
return method


# sys.maxsize:
# An integer giving the maximum value a variable of type Py_ssize_t can take.
MAX_WAIT = sys.maxsize / 2
Expand Down
9 changes: 9 additions & 0 deletions tenacity/asyncio/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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:
Expand All @@ -133,32 +135,38 @@ 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]
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
for action in self.iter_state.actions:
result = await action(retry_state)
return result

@override
def __iter__(self) -> t.Generator[AttemptManager, None, None]:
raise TypeError("AsyncRetrying object is not iterable")

Expand Down Expand Up @@ -194,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.
Expand Down
12 changes: 12 additions & 0 deletions tenacity/asyncio/retry.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
import typing

from tenacity import _utils, retry_base
from tenacity._utils import override

if typing.TYPE_CHECKING:
from tenacity import RetryCallState
Expand All @@ -26,24 +27,29 @@ class async_retry_base(retry_base):
"""Abstract base class for async retry strategies."""

@abc.abstractmethod
@override
async def __call__(self, retry_state: "RetryCallState") -> bool: # type: ignore[override]
pass

@override
def __and__( # type: ignore[override]
self, other: "retry_base | async_retry_base"
) -> "retry_all":
return retry_all(self, other)

@override
def __rand__( # type: ignore[misc,override]
self, other: "retry_base | async_retry_base"
) -> "retry_all":
return retry_all(other, self)

@override
def __or__( # type: ignore[override]
self, other: "retry_base | async_retry_base"
) -> "retry_any":
return retry_any(self, other)

@override
def __ror__( # type: ignore[misc,override]
self, other: "retry_base | async_retry_base"
) -> "retry_any":
Expand All @@ -63,6 +69,7 @@ def __init__(
) -> None:
self.predicate = predicate

@override
async def __call__(self, retry_state: "RetryCallState") -> bool: # type: ignore[override]
if retry_state.outcome is None:
raise RuntimeError("__call__() called before outcome was set")
Expand All @@ -83,6 +90,7 @@ def __init__(
) -> None:
self.predicate = predicate

@override
async def __call__(self, retry_state: "RetryCallState") -> bool: # type: ignore[override]
if retry_state.outcome is None:
raise RuntimeError("__call__() called before outcome was set")
Expand All @@ -98,6 +106,7 @@ class retry_any(async_retry_base):
def __init__(self, *retries: retry_base | async_retry_base) -> None:
self.retries = retries

@override
async def __call__(self, retry_state: "RetryCallState") -> bool: # type: ignore[override]
result = False
for r in self.retries:
Expand All @@ -106,6 +115,7 @@ async def __call__(self, retry_state: "RetryCallState") -> bool: # type: ignore
break
return result

@override
def __ror__( # type: ignore[misc,override]
self, other: "retry_base | async_retry_base"
) -> "retry_any":
Expand All @@ -120,6 +130,7 @@ class retry_all(async_retry_base):
def __init__(self, *retries: retry_base | async_retry_base) -> None:
self.retries = retries

@override
async def __call__(self, retry_state: "RetryCallState") -> bool: # type: ignore[override]
result = True
for r in self.retries:
Expand All @@ -128,6 +139,7 @@ async def __call__(self, retry_state: "RetryCallState") -> bool: # type: ignore
break
return result

@override
def __rand__( # type: ignore[misc,override]
self, other: "retry_base | async_retry_base"
) -> "retry_all":
Expand Down
15 changes: 15 additions & 0 deletions tenacity/retry.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
import re
import typing

from tenacity._utils import override

if typing.TYPE_CHECKING:
from tenacity import RetryCallState

Expand Down Expand Up @@ -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

Expand All @@ -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

Expand All @@ -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")
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand All @@ -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")
Expand All @@ -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")
Expand Down Expand Up @@ -254,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")
Expand All @@ -276,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)
Expand All @@ -291,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)
Expand Down
Loading
Loading