diff --git a/pyproject.toml b/pyproject.toml index 6c01067c..dba836ff 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -117,4 +117,19 @@ strict = true files = ["tenacity", "tests"] show_error_codes = true exclude = ["tenacity/_version\\.py"] +# `strict` is not "everything" -- these are checks it leaves off. +warn_unreachable = true +disallow_any_unimported = true +extra_checks = true +enable_error_code = [ + "deprecated", + "exhaustive-match", + "ignore-without-code", + "mutable-override", + "possibly-undefined", + "redundant-expr", + "truthy-bool", + "truthy-iterable", + "unused-awaitable", +] diff --git a/tenacity/__init__.py b/tenacity/__init__.py index 21d7f507..e993396c 100644 --- a/tenacity/__init__.py +++ b/tenacity/__init__.py @@ -88,6 +88,11 @@ except ImportError: tornado = None # type: ignore[assignment] +# mypy resolves `tornado` to the module (it only ever sees the `try` branch), +# so testing the module object for truthiness reads as an always-true check. +# Keep the availability answer in a plain bool instead. +_HAS_TORNADO = tornado is not None + if t.TYPE_CHECKING: if sys.version_info >= (3, 11): from typing import Self @@ -148,8 +153,8 @@ class BaseAction: - NAME: for identification in retry object methods and callbacks """ - REPR_FIELDS: t.Sequence[str] = () - NAME: str | None = None + REPR_FIELDS: t.ClassVar[t.Sequence[str]] = () + NAME: t.ClassVar[str | None] = None def __repr__(self) -> str: state_str = ", ".join( @@ -162,8 +167,8 @@ def __str__(self) -> str: class RetryAction(BaseAction): - REPR_FIELDS = ("sleep",) - NAME = "retry" + 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) @@ -403,12 +408,7 @@ 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: - if self.wait: - sleep = self.wait(retry_state) - else: - sleep = 0.0 - - retry_state.upcoming_sleep = sleep + 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 @@ -770,7 +770,7 @@ def wrap(f: t.Callable[P, R]) -> _RetryDecorated[P, R]: ): r = AsyncRetrying(*dargs, **dkw) elif ( - tornado + _HAS_TORNADO and hasattr(tornado.gen, "is_coroutine_function") and tornado.gen.is_coroutine_function(f) ): @@ -785,7 +785,7 @@ def wrap(f: t.Callable[P, R]) -> _RetryDecorated[P, R]: from tenacity.asyncio import AsyncRetrying # noqa: E402 -if tornado: +if _HAS_TORNADO: from tenacity.tornadoweb import TornadoRetrying diff --git a/tenacity/asyncio/__init__.py b/tenacity/asyncio/__init__.py index 6291b02f..a91ca577 100644 --- a/tenacity/asyncio/__init__.py +++ b/tenacity/asyncio/__init__.py @@ -142,12 +142,9 @@ async def _run_retry(self, retry_state: "RetryCallState") -> None: # type: igno ) 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 + ) async def _run_stop(self, retry_state: "RetryCallState") -> None: # type: ignore[override] self.statistics["delay_since_first_attempt"] = retry_state.seconds_since_start diff --git a/tests/test_asyncio.py b/tests/test_asyncio.py index 96560b07..c9d58027 100644 --- a/tests/test_asyncio.py +++ b/tests/test_asyncio.py @@ -85,7 +85,6 @@ async def test_retry(self) -> None: @asynctest async def test_iscoroutinefunction(self) -> None: - assert asyncio.iscoroutinefunction(_retryable_coroutine) assert inspect.iscoroutinefunction(_retryable_coroutine) @asynctest diff --git a/tests/test_tenacity.py b/tests/test_tenacity.py index 8cfd34e5..d5dd8f17 100644 --- a/tests/test_tenacity.py +++ b/tests/test_tenacity.py @@ -692,10 +692,9 @@ def waitfunc(retry_state: RetryCallState) -> float: def returnval() -> int: return 123 - try: + with self.assertRaises(ExtractCallState) as caught: retrying(returnval) - except ExtractCallState as err: - retry_state = err.args[0] + retry_state = caught.exception.args[0] self.assertIs(retry_state.fn, returnval) self.assertEqual(retry_state.args, ()) self.assertEqual(retry_state.kwargs, {}) @@ -706,10 +705,9 @@ def returnval() -> int: def dying() -> None: raise Exception("Broken") - try: + with self.assertRaises(ExtractCallState) as caught: retrying(dying) - except ExtractCallState as err: - retry_state = err.args[0] + retry_state = caught.exception.args[0] self.assertIs(retry_state.fn, dying) self.assertEqual(retry_state.args, ()) self.assertEqual(retry_state.kwargs, {})