Skip to content
Open
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
26 changes: 26 additions & 0 deletions docs/available-components/middlewares.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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

Expand Down
27 changes: 25 additions & 2 deletions taskiq/middlewares/simple_retry_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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

Expand Down
27 changes: 25 additions & 2 deletions taskiq/middlewares/smart_retry_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A regular label belongs to a specific message. It is produced by the producer, passed through the broker, and can be overridden via .kicker().with_labels(...).

In contrast, types_of_exceptions is effectively part of the task definition on the worker side. The value passed with the message is not used for any decision-making.

It seems to me that this results in the task labels being silently lost before the retry.

The same issue exists in simple_retry_middleware.py:44

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

Expand Down
1 change: 1 addition & 0 deletions tests/middlewares/test_simple_retry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
116 changes: 116 additions & 0 deletions tests/middlewares/test_task_retry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading