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
2 changes: 1 addition & 1 deletion docs/guide/architecture-overview.md
Original file line number Diff line number Diff line change
Expand Up @@ -220,7 +220,7 @@ class MyMiddleware(TaskiqMiddleware):

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

- `pre_send` - executed on the client side before the message is sent. Here you can modify the message.
- `pre_send` - executed on the client side before the message is sent. Here you can modify the message, or drop it by raising `SkipSendError`.
- `post_send` - executed right after the message was sent.
- `pre_execute` - executed on the worker side after the message was received by a worker and before its execution.
- `on_error` - executed after the task was executed if an exception was found.
Expand Down
2 changes: 2 additions & 0 deletions taskiq/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
ResultIsReadyError,
SecurityError,
SendTaskError,
SkipSendError,
TaskiqError,
TaskiqResultTimeoutError,
)
Expand Down Expand Up @@ -58,6 +59,7 @@
"SecurityError",
"SendTaskError",
"SimpleRetryMiddleware",
"SkipSendError",
"SmartRetryMiddleware",
"TaskiqDepends",
"TaskiqError",
Expand Down
2 changes: 2 additions & 0 deletions taskiq/abc/middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,8 @@ def pre_send(
This is a client-side hook, that executes right before
the message is sent to broker.

This method may raise SkipSendError to drop the message.

:param message: message to send.
:return: modified message.
"""
Expand Down
7 changes: 7 additions & 0 deletions taskiq/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,13 @@ class ScheduledTaskCancelledError(TaskiqError):
__template__ = "Cannot send scheduled task to the queue."


class SkipSendError(TaskiqError):
"""Middleware asked to skip sending the task."""

__template__ = "Task was not sent to the queue"
task_id: str | None = None


class TaskBrokerMismatchError(TaskRejectedError):
"""Task has a different broker than the one it was registered to."""

Expand Down
30 changes: 19 additions & 11 deletions taskiq/kicker.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
from pydantic import BaseModel

from taskiq.abc.middleware import TaskiqMiddleware
from taskiq.exceptions import SendTaskError
from taskiq.exceptions import SendTaskError, SkipSendError
from taskiq.labels import prepare_label
from taskiq.message import TaskiqMessage
from taskiq.scheduler.created_schedule import CreatedSchedule
Expand Down Expand Up @@ -145,6 +145,8 @@ async def kiq(
It gets current broker and calls it's kick method,
returning what it returns.

Returns without sending if a pre_send hook raises SkipSendError.

:param args: function's arguments.
:param kwargs: function's key word arguments.

Expand All @@ -159,20 +161,26 @@ async def kiq(
kwargs,
)
message = self._prepare_message(*args, **kwargs)
for middleware in self.broker.middlewares:
if middleware.__class__.pre_send != TaskiqMiddleware.pre_send:
message = await maybe_awaitable(middleware.pre_send(message))
try:
await self.broker.kick(self.broker.formatter.dumps(message))
except Exception as exc:
raise SendTaskError from exc
for middleware in self.broker.middlewares:
if middleware.__class__.pre_send != TaskiqMiddleware.pre_send:
message = await maybe_awaitable(middleware.pre_send(message))
except SkipSendError as exc:
logger.debug("Task %s has been skipped.", self.task_name)
task_id = exc.task_id or message.task_id
else:
try:
await self.broker.kick(self.broker.formatter.dumps(message))
except Exception as exc:
raise SendTaskError from exc

for middleware in reversed(self.broker.middlewares):
if middleware.__class__.post_send != TaskiqMiddleware.post_send:
await maybe_awaitable(middleware.post_send(message))
for middleware in reversed(self.broker.middlewares):
if middleware.__class__.post_send != TaskiqMiddleware.post_send:
await maybe_awaitable(middleware.post_send(message))
task_id = message.task_id

return AsyncTaskiqTask(
task_id=message.task_id,
task_id=task_id,
result_backend=self.broker.result_backend,
return_type=self.return_type, # type: ignore # (pyright issue)
)
Expand Down
61 changes: 60 additions & 1 deletion tests/test_kicker.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,10 @@
from typing import Any

from taskiq import InMemoryBroker
import pytest

from taskiq import InMemoryBroker, SkipSendError, TaskiqMessage, TaskiqMiddleware
from taskiq.kicker import AsyncKicker
from tests.utils import AsyncQueueBroker


async def test_types_of_exceptions_not_serialized() -> None:
Expand Down Expand Up @@ -43,3 +46,59 @@ async def test_other_labels_still_serialized() -> None:

assert message.labels["retries"] == "3"
assert message.labels["queue"] == "high_priority"


async def test_skip_send_error_drops_task() -> None:
"""SkipSendError in pre_send drops the task and returns the given task_id."""
calls = []

class _BeforeMiddleware(TaskiqMiddleware):
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
calls.append("before.pre_send")
return message

def post_send(self, message: TaskiqMessage) -> None:
calls.append("before.post_send")

class _SkipMiddleware(TaskiqMiddleware):
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
raise SkipSendError(task_id="winner")

class _AfterMiddleware(TaskiqMiddleware):
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
calls.append("after.pre_send")
return message

broker = AsyncQueueBroker().with_middlewares(
_BeforeMiddleware(),
_SkipMiddleware(),
_AfterMiddleware(),
)

@broker.task
async def run_task() -> None:
pass

task = await run_task.kiq()

assert task.task_id == "winner"
assert broker.queue.empty()
assert calls == ["before.pre_send"]


async def test_other_pre_send_errors_propagate() -> None:
"""Only SkipSendError is swallowed, other pre_send errors still propagate."""

class _FailingMiddleware(TaskiqMiddleware):
def pre_send(self, message: TaskiqMessage) -> TaskiqMessage:
raise ValueError("boom")

broker = AsyncQueueBroker().with_middlewares(_FailingMiddleware())

@broker.task
async def run_task() -> None:
pass

with pytest.raises(ValueError, match="boom"):
await run_task.kiq()
assert broker.queue.empty()
Loading