Skip to content

Commit 9ce2069

Browse files
authored
feat: add per-task types_of_exceptions (#651)
1 parent 3ff8b12 commit 9ce2069

7 files changed

Lines changed: 244 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/kicker.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -322,6 +322,12 @@ def _prepare_message(
322322
formatted_kwargs[kwarg_name] = self._prepare_arg(kwarg_val)
323323

324324
for label, label_val in self.labels.items():
325+
# `types_of_exceptions` is only ever read back from the
326+
# locally registered task's labels (see retry middlewares),
327+
# never from the wire. It holds exception classes, which
328+
# can't be faithfully serialized, so don't ship it at all.
329+
if label == "types_of_exceptions":
330+
continue
325331
labels[label], labels_types[label] = prepare_label(label_val)
326332

327333
task_id = self.custom_task_id

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(

tests/test_kicker.py

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,45 @@
1+
from typing import Any
2+
3+
from taskiq import InMemoryBroker
4+
from taskiq.kicker import AsyncKicker
5+
6+
7+
async def test_types_of_exceptions_not_serialized() -> None:
8+
"""`types_of_exceptions` label should never be sent over the wire."""
9+
broker = InMemoryBroker()
10+
11+
@broker.task(types_of_exceptions=(ValueError, TypeError))
12+
async def run_task() -> None:
13+
pass
14+
15+
kicker = run_task.kicker()
16+
message = kicker._prepare_message()
17+
18+
assert "types_of_exceptions" not in message.labels
19+
assert "types_of_exceptions" not in (message.labels_types or {})
20+
21+
22+
async def test_types_of_exceptions_still_local_on_task() -> None:
23+
"""The registered task still keeps the real exception types locally."""
24+
broker = InMemoryBroker()
25+
26+
@broker.task(types_of_exceptions=(ValueError, TypeError))
27+
async def run_task() -> None:
28+
pass
29+
30+
task = broker.find_task(run_task.task_name)
31+
assert task is not None
32+
assert task.labels["types_of_exceptions"] == (ValueError, TypeError)
33+
34+
35+
async def test_other_labels_still_serialized() -> None:
36+
"""Unrelated labels are unaffected by the fix."""
37+
kicker: AsyncKicker[Any, Any] = AsyncKicker(
38+
task_name="some_task",
39+
broker=InMemoryBroker(),
40+
labels={"retries": 3, "queue": "high_priority"},
41+
)
42+
message = kicker._prepare_message()
43+
44+
assert message.labels["retries"] == "3"
45+
assert message.labels["queue"] == "high_priority"

0 commit comments

Comments
 (0)