diff --git a/pyproject.toml b/pyproject.toml index dba836ff..6bc262d5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] @@ -124,6 +128,7 @@ extra_checks = true enable_error_code = [ "deprecated", "exhaustive-match", + "explicit-override", "ignore-without-code", "mutable-override", "possibly-undefined", diff --git a/tenacity/__init__.py b/tenacity/__init__.py index e993396c..9ceaf16d 100644 --- a/tenacity/__init__.py +++ b/tenacity/__init__.py @@ -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 @@ -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) @@ -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}]" @@ -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"} @@ -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 "" + @override def __repr__(self) -> str: return ( f"<{self.__class__.__name__} object at 0x{id(self):x} (" @@ -514,6 +521,7 @@ def __call__( class Retrying(BaseRetrying): """Retrying controller.""" + @override def __call__( self, fn: t.Callable[..., WrappedFnReturnT], @@ -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" diff --git a/tenacity/_utils.py b/tenacity/_utils.py index 195fa504..96c9e99f 100644 --- a/tenacity/_utils.py +++ b/tenacity/_utils.py @@ -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 diff --git a/tenacity/asyncio/__init__.py b/tenacity/asyncio/__init__.py index a91ca577..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 ( @@ -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,25 +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] 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 @@ -159,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") @@ -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. diff --git a/tenacity/asyncio/retry.py b/tenacity/asyncio/retry.py index e0d44e82..c2480706 100644 --- a/tenacity/asyncio/retry.py +++ b/tenacity/asyncio/retry.py @@ -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 @@ -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": @@ -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") @@ -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") @@ -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: @@ -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": @@ -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: @@ -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": diff --git a/tenacity/retry.py b/tenacity/retry.py index 61c63d2e..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") @@ -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") @@ -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) @@ -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) diff --git a/tenacity/stop.py b/tenacity/stop.py index c2251b39..c1f51ea9 100644 --- a/tenacity/stop.py +++ b/tenacity/stop.py @@ -17,6 +17,7 @@ import typing from tenacity import _utils +from tenacity._utils import override if typing.TYPE_CHECKING: import threading @@ -47,6 +48,7 @@ class stop_any(stop_base): def __init__(self, *stops: stop_base) -> None: self.stops = stops + @override def __call__(self, retry_state: "RetryCallState") -> bool: return any(x(retry_state) for x in self.stops) @@ -57,6 +59,7 @@ class stop_all(stop_base): def __init__(self, *stops: stop_base) -> None: self.stops = stops + @override def __call__(self, retry_state: "RetryCallState") -> bool: return all(x(retry_state) for x in self.stops) @@ -64,6 +67,7 @@ def __call__(self, retry_state: "RetryCallState") -> bool: class _stop_never(stop_base): """Never stop.""" + @override def __call__(self, retry_state: "RetryCallState") -> bool: return False @@ -77,6 +81,7 @@ class stop_when_event_set(stop_base): def __init__(self, event: "threading.Event") -> None: self.event = event + @override def __call__(self, retry_state: "RetryCallState") -> bool: return self.event.is_set() @@ -87,6 +92,7 @@ class stop_after_attempt(stop_base): def __init__(self, max_attempt_number: int) -> None: self.max_attempt_number = max_attempt_number + @override def __call__(self, retry_state: "RetryCallState") -> bool: return retry_state.attempt_number >= self.max_attempt_number @@ -104,6 +110,7 @@ class stop_after_delay(stop_base): def __init__(self, max_delay: _utils.time_unit_type) -> None: self.max_delay = _utils.to_seconds(max_delay) + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.seconds_since_start is None: raise RuntimeError("__call__() called but seconds_since_start is not set") @@ -121,6 +128,7 @@ class stop_before_delay(stop_base): def __init__(self, max_delay: _utils.time_unit_type) -> None: self.max_delay = _utils.to_seconds(max_delay) + @override def __call__(self, retry_state: "RetryCallState") -> bool: if retry_state.seconds_since_start is None: raise RuntimeError("__call__() called but seconds_since_start is not set") diff --git a/tenacity/tornadoweb.py b/tenacity/tornadoweb.py index 0d46ed7b..717107be 100644 --- a/tenacity/tornadoweb.py +++ b/tenacity/tornadoweb.py @@ -18,6 +18,7 @@ from tornado import gen from tenacity import BaseRetrying, DoAttempt, DoSleep, RetryCallState +from tenacity._utils import override if typing.TYPE_CHECKING: from tornado.concurrent import Future @@ -37,6 +38,7 @@ def __init__( self.sleep = sleep @gen.coroutine + @override def __call__( # type: ignore[override] self, fn: "typing.Callable[..., typing.Generator[typing.Any, typing.Any, _RetValT] | Future[_RetValT]]", diff --git a/tenacity/wait.py b/tenacity/wait.py index 3053a5c1..e137b357 100644 --- a/tenacity/wait.py +++ b/tenacity/wait.py @@ -21,6 +21,7 @@ import warnings from tenacity import _utils +from tenacity._utils import override if typing.TYPE_CHECKING: from tenacity import RetryCallState @@ -54,6 +55,7 @@ class wait_fixed(wait_base): def __init__(self, wait: _utils.time_unit_type) -> None: self.wait_fixed = _utils.to_seconds(wait) + @override def __call__(self, retry_state: "RetryCallState") -> float: return self.wait_fixed @@ -74,6 +76,7 @@ def __init__( self.wait_random_min = _utils.to_seconds(min) self.wait_random_max = _utils.to_seconds(max) + @override def __call__(self, retry_state: "RetryCallState") -> float: return self.wait_random_min + ( random.random() * (self.wait_random_max - self.wait_random_min) @@ -86,6 +89,7 @@ class wait_combine(wait_base): def __init__(self, *strategies: wait_base) -> None: self.wait_funcs = strategies + @override def __call__(self, retry_state: "RetryCallState") -> float: return sum(x(retry_state=retry_state) for x in self.wait_funcs) @@ -111,6 +115,7 @@ def __init__(self, *strategies: wait_base) -> None: raise ValueError("wait_chain() requires at least one strategy") self.strategies = strategies + @override def __call__(self, retry_state: "RetryCallState") -> float: wait_func_no = min(max(retry_state.attempt_number, 1), len(self.strategies)) wait_func = self.strategies[wait_func_no - 1] @@ -146,6 +151,7 @@ def http_get_request(url: str) -> None: def __init__(self, predicate: typing.Callable[[BaseException], float]) -> None: self.predicate = predicate + @override def __call__(self, retry_state: "RetryCallState") -> float: if retry_state.outcome is None: raise RuntimeError("__call__() called before outcome was set") @@ -173,6 +179,7 @@ def __init__( self.increment = _utils.to_seconds(increment) self.max = _utils.to_seconds(max) + @override def __call__(self, retry_state: "RetryCallState") -> float: result = self.start + (self.increment * (retry_state.attempt_number - 1)) return max(0, min(result, self.max)) @@ -203,6 +210,7 @@ def __init__( self.max = _utils.to_seconds(max) self.exp_base = exp_base + @override def __call__(self, retry_state: "RetryCallState") -> float: exponent = retry_state.attempt_number - 1 if ( @@ -249,6 +257,7 @@ class wait_random_exponential(wait_exponential): """ + @override def __call__(self, retry_state: "RetryCallState") -> float: high = super().__call__(retry_state=retry_state) return random.uniform(self.min, high) @@ -294,6 +303,7 @@ def __init__( self.jitter = _utils.to_seconds(jitter) self.min = _utils.to_seconds(min) + @override def __call__(self, retry_state: "RetryCallState") -> float: jitter = random.uniform(0, self.jitter) try: diff --git a/tests/test_after.py b/tests/test_after.py index ccf4a730..86a6e408 100644 --- a/tests/test_after.py +++ b/tests/test_after.py @@ -6,11 +6,13 @@ _utils, after_log, ) +from tenacity._utils import override from . import test_tenacity class TestAfterLogFormat(unittest.TestCase): + @override def setUp(self) -> None: self.log_level = random.choice( ( diff --git a/tests/test_tenacity.py b/tests/test_tenacity.py index d5dd8f17..1acd3f8d 100644 --- a/tests/test_tenacity.py +++ b/tests/test_tenacity.py @@ -28,6 +28,7 @@ import tenacity from tenacity import RetryCallState, RetryError, Retrying, retry +from tenacity._utils import override from tenacity.retry import retry_all, retry_any _unset = object() @@ -86,6 +87,7 @@ def make_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: @@ -360,6 +362,7 @@ def test_exponential_with_max_wait(self) -> None: 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") @@ -1109,6 +1112,7 @@ class CustomError(Exception): def __init__(self, value: str) -> None: self.value = value + @override def __str__(self) -> str: return self.value @@ -1140,6 +1144,7 @@ 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) @@ -1996,6 +2001,7 @@ def _foobar() -> int: class TestRetryErrorCallback(unittest.TestCase): + @override def setUp(self) -> None: self._attempt_number = 0 self._callback_called = False