Skip to content

Commit 5cae04c

Browse files
committed
feat: add per-task types_of_exceptions
1 parent ae2b788 commit 5cae04c

5 files changed

Lines changed: 193 additions & 4 deletions

File tree

docs/available-components/middlewares.md

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,25 @@ async def test():
3434

3535
`retry_on_error` enables retries for a task. `max_retries` is the maximum number of retry attempts.
3636

37+
### Retrying only specific exceptions
38+
39+
By default, all exceptions trigger a retry. You can limit retries to specific
40+
exception types either broker-wide via `types_of_exceptions`, or per task via
41+
the `types_of_exceptions` label in the task decorator. The per-task value
42+
overrides the broker-wide setting.
43+
44+
```python
45+
broker = ZeroMQBroker().with_middlewares(
46+
# Broker-wide default: retry only on ConnectionError.
47+
SimpleRetryMiddleware(types_of_exceptions=(ConnectionError,)),
48+
)
49+
50+
51+
@broker.task(retry_on_error=True, types_of_exceptions=(ValueError, KeyError))
52+
async def test():
53+
raise ValueError("retry only on ValueError or KeyError")
54+
```
55+
3756
## Smart retry middleware
3857

3958
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():
7897
* `retry_on_error`: Enables the retry mechanism for the specific task.
7998
* `max_retries`: Maximum number of retries (overrides middleware default).
8099
* `delay`: Initial delay before retrying the task, in seconds.
100+
* `types_of_exceptions`: Exception types that trigger a retry for this task. Overrides the broker-wide `types_of_exceptions` passed to the middleware.
101+
102+
```python
103+
@broker.task(retry_on_error=True, types_of_exceptions=(ConnectionError,))
104+
async def my_task():
105+
raise ConnectionError("retrying only on ConnectionError")
106+
```
81107

82108
### Usage Recommendations
83109

taskiq/middlewares/simple_retry_middleware.py

Lines changed: 25 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,28 @@ def __init__(
2626
self.no_result_on_retry = no_result_on_retry
2727
self.types_of_exceptions = types_of_exceptions
2828

29+
def _get_types_of_exceptions(
30+
self,
31+
message: "TaskiqMessage",
32+
) -> Iterable[type[BaseException]] | None:
33+
"""
34+
Resolve retryable exception types for a task.
35+
36+
Per-task ``types_of_exceptions`` set via the task decorator take
37+
precedence over the broker-wide value. Types are read from the
38+
registered task object, since label values are stringified when a
39+
message is serialized and cannot carry real exception types.
40+
41+
:param message: Original task message.
42+
:return: Effective exception types or None.
43+
"""
44+
task = self.broker.find_task(message.task_name)
45+
if task is not None:
46+
task_types = task.labels.get("types_of_exceptions")
47+
if task_types is not None:
48+
return task_types
49+
return self.types_of_exceptions
50+
2951
async def on_error(
3052
self,
3153
message: "TaskiqMessage",
@@ -45,9 +67,10 @@ async def on_error(
4567
:param result: execution result.
4668
:param exception: found exception.
4769
"""
48-
if self.types_of_exceptions is not None and not isinstance(
70+
types_of_exceptions = self._get_types_of_exceptions(message)
71+
if types_of_exceptions is not None and not isinstance(
4972
exception,
50-
tuple(self.types_of_exceptions),
73+
tuple(types_of_exceptions),
5174
):
5275
return
5376

taskiq/middlewares/smart_retry_middleware.py

Lines changed: 25 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -68,6 +68,28 @@ def __init__(
6868
"schedule_source must be an instance of ScheduleSource or None",
6969
)
7070

71+
def _get_types_of_exceptions(
72+
self,
73+
message: TaskiqMessage,
74+
) -> Iterable[type[BaseException]] | None:
75+
"""
76+
Resolve retryable exception types for a task.
77+
78+
Per-task ``types_of_exceptions`` set via the task decorator take
79+
precedence over the broker-wide value. Types are read from the
80+
registered task object, since label values are stringified when a
81+
message is serialized and cannot carry real exception types.
82+
83+
:param message: Original task message.
84+
:return: Effective exception types or None.
85+
"""
86+
task = self.broker.find_task(message.task_name)
87+
if task is not None:
88+
task_types = task.labels.get("types_of_exceptions")
89+
if task_types is not None:
90+
return task_types
91+
return self.types_of_exceptions
92+
7193
def is_retry_on_error(self, message: TaskiqMessage) -> bool:
7294
"""
7395
Check if retry is enabled for this task.
@@ -142,9 +164,10 @@ async def on_error(
142164
:param result: Execution result.
143165
:param exception: Caught exception.
144166
"""
145-
if self.types_of_exceptions is not None and not isinstance(
167+
types_of_exceptions = self._get_types_of_exceptions(message)
168+
if types_of_exceptions is not None and not isinstance(
146169
exception,
147-
tuple(self.types_of_exceptions),
170+
tuple(types_of_exceptions),
148171
):
149172
return
150173

tests/middlewares/test_simple_retry.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ def broker() -> AsyncMock:
1414
mocked_broker = AsyncMock()
1515
mocked_broker.id_generator = lambda: uuid.uuid4().hex
1616
mocked_broker.formatter = JSONFormatter()
17+
mocked_broker.find_task = lambda task_name: None
1718
return mocked_broker
1819

1920

tests/middlewares/test_task_retry.py

Lines changed: 116 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -200,6 +200,122 @@ def run_task2() -> None:
200200
assert runs == 1
201201

202202

203+
@pytest.mark.parametrize(
204+
"middleware_class",
205+
[SimpleRetryMiddleware, SmartRetryMiddleware],
206+
)
207+
async def test_per_task_exc_types_not_matching(middleware_class: type) -> None:
208+
# per-task types_of_exceptions does not include the raised exception
209+
broker = InMemoryBroker().with_middlewares(
210+
middleware_class(no_result_on_retry=True, default_retry_label=True),
211+
)
212+
runs = 0
213+
214+
@broker.task(max_retries=10, types_of_exceptions=(KeyError,))
215+
def run_task() -> None:
216+
nonlocal runs
217+
218+
runs += 1
219+
220+
raise ValueError(runs)
221+
222+
task = await run_task.kiq()
223+
resp = await task.wait_result(timeout=1)
224+
with pytest.raises(ValueError):
225+
resp.raise_for_error()
226+
227+
assert runs == 1
228+
229+
230+
@pytest.mark.parametrize(
231+
"middleware_class",
232+
[SimpleRetryMiddleware, SmartRetryMiddleware],
233+
)
234+
async def test_per_task_exc_types_matching(middleware_class: type) -> None:
235+
# per-task types_of_exceptions includes the raised exception
236+
broker = InMemoryBroker().with_middlewares(
237+
middleware_class(no_result_on_retry=True, default_retry_label=True),
238+
)
239+
runs = 0
240+
241+
@broker.task(max_retries=10, types_of_exceptions=(ValueError,))
242+
def run_task() -> None:
243+
nonlocal runs
244+
245+
runs += 1
246+
247+
raise ValueError(runs)
248+
249+
task = await run_task.kiq()
250+
resp = await task.wait_result(timeout=1)
251+
with pytest.raises(ValueError):
252+
resp.raise_for_error()
253+
254+
assert runs == 10
255+
256+
257+
@pytest.mark.parametrize(
258+
"middleware_class",
259+
[SimpleRetryMiddleware, SmartRetryMiddleware],
260+
)
261+
async def test_per_task_exc_types_override_global(middleware_class: type) -> None:
262+
# per-task types_of_exceptions takes precedence over broker-wide value
263+
broker = InMemoryBroker().with_middlewares(
264+
middleware_class(
265+
no_result_on_retry=True,
266+
default_retry_label=True,
267+
types_of_exceptions=(KeyError,),
268+
),
269+
)
270+
runs = 0
271+
272+
@broker.task(max_retries=10, types_of_exceptions=(ValueError,))
273+
def run_task() -> None:
274+
nonlocal runs
275+
276+
runs += 1
277+
278+
raise ValueError(runs)
279+
280+
task = await run_task.kiq()
281+
resp = await task.wait_result(timeout=1)
282+
with pytest.raises(ValueError):
283+
resp.raise_for_error()
284+
285+
assert runs == 10
286+
287+
288+
@pytest.mark.parametrize(
289+
"middleware_class",
290+
[SimpleRetryMiddleware, SmartRetryMiddleware],
291+
)
292+
async def test_global_exc_types_without_per_task(middleware_class: type) -> None:
293+
# broker-wide types_of_exceptions still applies when no per-task value set
294+
broker = InMemoryBroker().with_middlewares(
295+
middleware_class(
296+
no_result_on_retry=True,
297+
default_retry_label=True,
298+
types_of_exceptions=(KeyError,),
299+
),
300+
)
301+
runs = 0
302+
303+
@broker.task(max_retries=10)
304+
def run_task() -> None:
305+
nonlocal runs
306+
307+
runs += 1
308+
309+
raise ValueError(runs)
310+
311+
task = await run_task.kiq()
312+
resp = await task.wait_result(timeout=1)
313+
with pytest.raises(ValueError):
314+
resp.raise_for_error()
315+
316+
assert runs == 1
317+
318+
203319
async def test_retry_of_custom_exc_types_of_smart_middleware() -> None:
204320
# test that the passed error will be handled
205321
broker = InMemoryBroker().with_middlewares(

0 commit comments

Comments
 (0)