From 5cae04cac5e5b3eb7c8c00d2edee5445ac2d4953 Mon Sep 17 00:00:00 2001 From: arturka Date: Fri, 24 Jul 2026 11:49:23 +0200 Subject: [PATCH] feat: add per-task types_of_exceptions --- docs/available-components/middlewares.md | 26 ++++ taskiq/middlewares/simple_retry_middleware.py | 27 +++- taskiq/middlewares/smart_retry_middleware.py | 27 +++- tests/middlewares/test_simple_retry.py | 1 + tests/middlewares/test_task_retry.py | 116 ++++++++++++++++++ 5 files changed, 193 insertions(+), 4 deletions(-) diff --git a/docs/available-components/middlewares.md b/docs/available-components/middlewares.md index 9ddbdcac..adee01ac 100644 --- a/docs/available-components/middlewares.md +++ b/docs/available-components/middlewares.md @@ -34,6 +34,25 @@ async def test(): `retry_on_error` enables retries for a task. `max_retries` is the maximum number of retry attempts. +### Retrying only specific exceptions + +By default, all exceptions trigger a retry. You can limit retries to specific +exception types either broker-wide via `types_of_exceptions`, or per task via +the `types_of_exceptions` label in the task decorator. The per-task value +overrides the broker-wide setting. + +```python +broker = ZeroMQBroker().with_middlewares( + # Broker-wide default: retry only on ConnectionError. + SimpleRetryMiddleware(types_of_exceptions=(ConnectionError,)), +) + + +@broker.task(retry_on_error=True, types_of_exceptions=(ValueError, KeyError)) +async def test(): + raise ValueError("retry only on ValueError or KeyError") +``` + ## Smart retry middleware The `SmartRetryMiddleware` automatically retries tasks with flexible delay settings and retry strategies when errors occur. This is particularly useful when tasks fail due to temporary issues, such as network errors or temporary unavailability of external services. @@ -78,6 +97,13 @@ async def my_task(): * `retry_on_error`: Enables the retry mechanism for the specific task. * `max_retries`: Maximum number of retries (overrides middleware default). * `delay`: Initial delay before retrying the task, in seconds. +* `types_of_exceptions`: Exception types that trigger a retry for this task. Overrides the broker-wide `types_of_exceptions` passed to the middleware. + +```python +@broker.task(retry_on_error=True, types_of_exceptions=(ConnectionError,)) +async def my_task(): + raise ConnectionError("retrying only on ConnectionError") +``` ### Usage Recommendations diff --git a/taskiq/middlewares/simple_retry_middleware.py b/taskiq/middlewares/simple_retry_middleware.py index 9bd80f7b..4c2f9a45 100644 --- a/taskiq/middlewares/simple_retry_middleware.py +++ b/taskiq/middlewares/simple_retry_middleware.py @@ -26,6 +26,28 @@ def __init__( self.no_result_on_retry = no_result_on_retry self.types_of_exceptions = types_of_exceptions + def _get_types_of_exceptions( + self, + message: "TaskiqMessage", + ) -> Iterable[type[BaseException]] | None: + """ + Resolve retryable exception types for a task. + + Per-task ``types_of_exceptions`` set via the task decorator take + precedence over the broker-wide value. Types are read from the + registered task object, since label values are stringified when a + message is serialized and cannot carry real exception types. + + :param message: Original task message. + :return: Effective exception types or None. + """ + task = self.broker.find_task(message.task_name) + if task is not None: + task_types = task.labels.get("types_of_exceptions") + if task_types is not None: + return task_types + return self.types_of_exceptions + async def on_error( self, message: "TaskiqMessage", @@ -45,9 +67,10 @@ async def on_error( :param result: execution result. :param exception: found exception. """ - if self.types_of_exceptions is not None and not isinstance( + types_of_exceptions = self._get_types_of_exceptions(message) + if types_of_exceptions is not None and not isinstance( exception, - tuple(self.types_of_exceptions), + tuple(types_of_exceptions), ): return diff --git a/taskiq/middlewares/smart_retry_middleware.py b/taskiq/middlewares/smart_retry_middleware.py index 66874504..9752b07f 100644 --- a/taskiq/middlewares/smart_retry_middleware.py +++ b/taskiq/middlewares/smart_retry_middleware.py @@ -68,6 +68,28 @@ def __init__( "schedule_source must be an instance of ScheduleSource or None", ) + def _get_types_of_exceptions( + self, + message: TaskiqMessage, + ) -> Iterable[type[BaseException]] | None: + """ + Resolve retryable exception types for a task. + + Per-task ``types_of_exceptions`` set via the task decorator take + precedence over the broker-wide value. Types are read from the + registered task object, since label values are stringified when a + message is serialized and cannot carry real exception types. + + :param message: Original task message. + :return: Effective exception types or None. + """ + task = self.broker.find_task(message.task_name) + if task is not None: + task_types = task.labels.get("types_of_exceptions") + if task_types is not None: + return task_types + return self.types_of_exceptions + def is_retry_on_error(self, message: TaskiqMessage) -> bool: """ Check if retry is enabled for this task. @@ -142,9 +164,10 @@ async def on_error( :param result: Execution result. :param exception: Caught exception. """ - if self.types_of_exceptions is not None and not isinstance( + types_of_exceptions = self._get_types_of_exceptions(message) + if types_of_exceptions is not None and not isinstance( exception, - tuple(self.types_of_exceptions), + tuple(types_of_exceptions), ): return diff --git a/tests/middlewares/test_simple_retry.py b/tests/middlewares/test_simple_retry.py index 783d98b7..f07b7c55 100644 --- a/tests/middlewares/test_simple_retry.py +++ b/tests/middlewares/test_simple_retry.py @@ -14,6 +14,7 @@ def broker() -> AsyncMock: mocked_broker = AsyncMock() mocked_broker.id_generator = lambda: uuid.uuid4().hex mocked_broker.formatter = JSONFormatter() + mocked_broker.find_task = lambda task_name: None return mocked_broker diff --git a/tests/middlewares/test_task_retry.py b/tests/middlewares/test_task_retry.py index 59798633..62a13c46 100644 --- a/tests/middlewares/test_task_retry.py +++ b/tests/middlewares/test_task_retry.py @@ -200,6 +200,122 @@ def run_task2() -> None: assert runs == 1 +@pytest.mark.parametrize( + "middleware_class", + [SimpleRetryMiddleware, SmartRetryMiddleware], +) +async def test_per_task_exc_types_not_matching(middleware_class: type) -> None: + # per-task types_of_exceptions does not include the raised exception + broker = InMemoryBroker().with_middlewares( + middleware_class(no_result_on_retry=True, default_retry_label=True), + ) + runs = 0 + + @broker.task(max_retries=10, types_of_exceptions=(KeyError,)) + def run_task() -> None: + nonlocal runs + + runs += 1 + + raise ValueError(runs) + + task = await run_task.kiq() + resp = await task.wait_result(timeout=1) + with pytest.raises(ValueError): + resp.raise_for_error() + + assert runs == 1 + + +@pytest.mark.parametrize( + "middleware_class", + [SimpleRetryMiddleware, SmartRetryMiddleware], +) +async def test_per_task_exc_types_matching(middleware_class: type) -> None: + # per-task types_of_exceptions includes the raised exception + broker = InMemoryBroker().with_middlewares( + middleware_class(no_result_on_retry=True, default_retry_label=True), + ) + runs = 0 + + @broker.task(max_retries=10, types_of_exceptions=(ValueError,)) + def run_task() -> None: + nonlocal runs + + runs += 1 + + raise ValueError(runs) + + task = await run_task.kiq() + resp = await task.wait_result(timeout=1) + with pytest.raises(ValueError): + resp.raise_for_error() + + assert runs == 10 + + +@pytest.mark.parametrize( + "middleware_class", + [SimpleRetryMiddleware, SmartRetryMiddleware], +) +async def test_per_task_exc_types_override_global(middleware_class: type) -> None: + # per-task types_of_exceptions takes precedence over broker-wide value + broker = InMemoryBroker().with_middlewares( + middleware_class( + no_result_on_retry=True, + default_retry_label=True, + types_of_exceptions=(KeyError,), + ), + ) + runs = 0 + + @broker.task(max_retries=10, types_of_exceptions=(ValueError,)) + def run_task() -> None: + nonlocal runs + + runs += 1 + + raise ValueError(runs) + + task = await run_task.kiq() + resp = await task.wait_result(timeout=1) + with pytest.raises(ValueError): + resp.raise_for_error() + + assert runs == 10 + + +@pytest.mark.parametrize( + "middleware_class", + [SimpleRetryMiddleware, SmartRetryMiddleware], +) +async def test_global_exc_types_without_per_task(middleware_class: type) -> None: + # broker-wide types_of_exceptions still applies when no per-task value set + broker = InMemoryBroker().with_middlewares( + middleware_class( + no_result_on_retry=True, + default_retry_label=True, + types_of_exceptions=(KeyError,), + ), + ) + runs = 0 + + @broker.task(max_retries=10) + def run_task() -> None: + nonlocal runs + + runs += 1 + + raise ValueError(runs) + + task = await run_task.kiq() + resp = await task.wait_result(timeout=1) + with pytest.raises(ValueError): + resp.raise_for_error() + + assert runs == 1 + + async def test_retry_of_custom_exc_types_of_smart_middleware() -> None: # test that the passed error will be handled broker = InMemoryBroker().with_middlewares(