Skip to content

Commit 8733640

Browse files
committed
feat: allow middlewares to skip sending a task with SkipSendError
1 parent a00e4b0 commit 8733640

6 files changed

Lines changed: 91 additions & 13 deletions

File tree

‎docs/guide/architecture-overview.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -220,7 +220,7 @@ class MyMiddleware(TaskiqMiddleware):
220220

221221
Here are methods you can implement in the order they are executed:
222222

223-
- `pre_send` - executed on the client side before the message is sent. Here you can modify the message.
223+
- `pre_send` - executed on the client side before the message is sent. Here you can modify the message, or drop it by raising `SkipSendError`.
224224
- `post_send` - executed right after the message was sent.
225225
- `pre_execute` - executed on the worker side after the message was received by a worker and before its execution.
226226
- `on_error` - executed after the task was executed if an exception was found.

‎taskiq/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
ResultIsReadyError,
2222
SecurityError,
2323
SendTaskError,
24+
SkipSendError,
2425
TaskiqError,
2526
TaskiqResultTimeoutError,
2627
)
@@ -58,6 +59,7 @@
5859
"SecurityError",
5960
"SendTaskError",
6061
"SimpleRetryMiddleware",
62+
"SkipSendError",
6163
"SmartRetryMiddleware",
6264
"TaskiqDepends",
6365
"TaskiqError",

‎taskiq/abc/middleware.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,8 @@ def pre_send(
6060
This is a client-side hook, that executes right before
6161
the message is sent to broker.
6262
63+
This method may raise SkipSendError to drop the message.
64+
6365
:param message: message to send.
6466
:return: modified message.
6567
"""

‎taskiq/exceptions.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -102,6 +102,13 @@ class ScheduledTaskCancelledError(TaskiqError):
102102
__template__ = "Cannot send scheduled task to the queue."
103103

104104

105+
class SkipSendError(TaskiqError):
106+
"""Middleware asked to skip sending the task."""
107+
108+
__template__ = "Task was not sent to the queue"
109+
task_id: str | None = None
110+
111+
105112
class TaskBrokerMismatchError(TaskRejectedError):
106113
"""Task has a different broker than the one it was registered to."""
107114

‎taskiq/kicker.py‎

Lines changed: 19 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
from pydantic import BaseModel
1717

1818
from taskiq.abc.middleware import TaskiqMiddleware
19-
from taskiq.exceptions import SendTaskError
19+
from taskiq.exceptions import SendTaskError, SkipSendError
2020
from taskiq.labels import prepare_label
2121
from taskiq.message import TaskiqMessage
2222
from taskiq.scheduler.created_schedule import CreatedSchedule
@@ -145,6 +145,8 @@ async def kiq(
145145
It gets current broker and calls it's kick method,
146146
returning what it returns.
147147
148+
Returns without sending if a pre_send hook raises SkipSendError.
149+
148150
:param args: function's arguments.
149151
:param kwargs: function's key word arguments.
150152
@@ -159,20 +161,26 @@ async def kiq(
159161
kwargs,
160162
)
161163
message = self._prepare_message(*args, **kwargs)
162-
for middleware in self.broker.middlewares:
163-
if middleware.__class__.pre_send != TaskiqMiddleware.pre_send:
164-
message = await maybe_awaitable(middleware.pre_send(message))
165164
try:
166-
await self.broker.kick(self.broker.formatter.dumps(message))
167-
except Exception as exc:
168-
raise SendTaskError from exc
165+
for middleware in self.broker.middlewares:
166+
if middleware.__class__.pre_send != TaskiqMiddleware.pre_send:
167+
message = await maybe_awaitable(middleware.pre_send(message))
168+
except SkipSendError as exc:
169+
logger.debug("Task %s has been skipped.", self.task_name)
170+
task_id = exc.task_id or message.task_id
171+
else:
172+
try:
173+
await self.broker.kick(self.broker.formatter.dumps(message))
174+
except Exception as exc:
175+
raise SendTaskError from exc
169176

170-
for middleware in reversed(self.broker.middlewares):
171-
if middleware.__class__.post_send != TaskiqMiddleware.post_send:
172-
await maybe_awaitable(middleware.post_send(message))
177+
for middleware in reversed(self.broker.middlewares):
178+
if middleware.__class__.post_send != TaskiqMiddleware.post_send:
179+
await maybe_awaitable(middleware.post_send(message))
180+
task_id = message.task_id
173181

174182
return AsyncTaskiqTask(
175-
task_id=message.task_id,
183+
task_id=task_id,
176184
result_backend=self.broker.result_backend,
177185
return_type=self.return_type, # type: ignore # (pyright issue)
178186
)

‎tests/test_kicker.py‎

Lines changed: 60 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,10 @@
11
from typing import Any
22

3-
from taskiq import InMemoryBroker
3+
import pytest
4+
5+
from taskiq import InMemoryBroker, SkipSendError, TaskiqMessage, TaskiqMiddleware
46
from taskiq.kicker import AsyncKicker
7+
from tests.utils import AsyncQueueBroker
58

69

710
async def test_types_of_exceptions_not_serialized() -> None:
@@ -43,3 +46,59 @@ async def test_other_labels_still_serialized() -> None:
4346

4447
assert message.labels["retries"] == "3"
4548
assert message.labels["queue"] == "high_priority"
49+
50+
51+
async def test_skip_send_error_drops_task() -> None:
52+
"""SkipSendError in pre_send drops the task and returns the given task_id."""
53+
calls = []
54+
55+
class _BeforeMiddleware(TaskiqMiddleware):
56+
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
57+
calls.append("before.pre_send")
58+
return message
59+
60+
def post_send(self, message: TaskiqMessage) -> None:
61+
calls.append("before.post_send")
62+
63+
class _SkipMiddleware(TaskiqMiddleware):
64+
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
65+
raise SkipSendError(task_id="winner")
66+
67+
class _AfterMiddleware(TaskiqMiddleware):
68+
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
69+
calls.append("after.pre_send")
70+
return message
71+
72+
broker = AsyncQueueBroker().with_middlewares(
73+
_BeforeMiddleware(),
74+
_SkipMiddleware(),
75+
_AfterMiddleware(),
76+
)
77+
78+
@broker.task
79+
async def run_task() -> None:
80+
pass
81+
82+
task = await run_task.kiq()
83+
84+
assert task.task_id == "winner"
85+
assert broker.queue.empty()
86+
assert calls == ["before.pre_send"]
87+
88+
89+
async def test_other_pre_send_errors_propagate() -> None:
90+
"""Only SkipSendError is swallowed, other pre_send errors still propagate."""
91+
92+
class _FailingMiddleware(TaskiqMiddleware):
93+
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
94+
raise ValueError("boom")
95+
96+
broker = AsyncQueueBroker().with_middlewares(_FailingMiddleware())
97+
98+
@broker.task
99+
async def run_task() -> None:
100+
pass
101+
102+
with pytest.raises(ValueError, match="boom"):
103+
await run_task.kiq()
104+
assert broker.queue.empty()

0 commit comments

Comments
 (0)