From 8733640d73dca55501b08e7d015402e4ce7c9eab Mon Sep 17 00:00:00 2001 From: d3vyce Date: Tue, 29 Sep 2026 08:47:40 -0400 Subject: [PATCH] feat: allow middlewares to skip sending a task with SkipSendError --- docs/guide/architecture-overview.md | 2 +- taskiq/__init__.py | 2 + taskiq/abc/middleware.py | 2 + taskiq/exceptions.py | 7 ++++ taskiq/kicker.py | 30 ++++++++------ tests/test_kicker.py | 61 ++++++++++++++++++++++++++++- 6 files changed, 91 insertions(+), 13 deletions(-) diff --git a/docs/guide/architecture-overview.md b/docs/guide/architecture-overview.md index 10d6b113..729f666a 100644 --- a/docs/guide/architecture-overview.md +++ b/docs/guide/architecture-overview.md @@ -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. diff --git a/taskiq/__init__.py b/taskiq/__init__.py index 1a764bc4..cdf0a687 100644 --- a/taskiq/__init__.py +++ b/taskiq/__init__.py @@ -21,6 +21,7 @@ ResultIsReadyError, SecurityError, SendTaskError, + SkipSendError, TaskiqError, TaskiqResultTimeoutError, ) @@ -58,6 +59,7 @@ "SecurityError", "SendTaskError", "SimpleRetryMiddleware", + "SkipSendError", "SmartRetryMiddleware", "TaskiqDepends", "TaskiqError", diff --git a/taskiq/abc/middleware.py b/taskiq/abc/middleware.py index cba52e16..8e4d7586 100644 --- a/taskiq/abc/middleware.py +++ b/taskiq/abc/middleware.py @@ -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. """ diff --git a/taskiq/exceptions.py b/taskiq/exceptions.py index 1e6fbb83..aed50984 100644 --- a/taskiq/exceptions.py +++ b/taskiq/exceptions.py @@ -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.""" diff --git a/taskiq/kicker.py b/taskiq/kicker.py index 07f4e51a..fda533c8 100644 --- a/taskiq/kicker.py +++ b/taskiq/kicker.py @@ -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 @@ -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. @@ -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) ) diff --git a/tests/test_kicker.py b/tests/test_kicker.py index 05afc0cb..4fa7c41f 100644 --- a/tests/test_kicker.py +++ b/tests/test_kicker.py @@ -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: @@ -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()