diff --git a/docs/extending-taskiq/broker.md b/docs/extending-taskiq/broker.md index 537227f9..9f2aecc3 100644 --- a/docs/extending-taskiq/broker.md +++ b/docs/extending-taskiq/broker.md @@ -36,13 +36,19 @@ async def listen(self) -> AsyncGenerator[AckableMessage, None]: for message in self.my_channel: yield AckableMessage( data=message.bytes, - # Ack is a function that takes no parameters. - # So you either set here method of a message, + # Ack can accept broker-specific keyword options. + # So you either set here a method of a message, # or you can make a closure. ack=message.ack, ) ``` +For manual acknowledgement, options passed to `Context.ack(**kwargs)` are +forwarded unchanged to this callback. Taskiq does not define or interpret these +options; they are specific to the broker. Automatic acknowledgement calls the +callback without options. Brokers that do not support options can continue to +provide a no-argument callback. + ## Conventions For brokers, we have several conventions. It's good if your broker implements them. diff --git a/taskiq/acks.py b/taskiq/acks.py index 894a9160..772aa615 100644 --- a/taskiq/acks.py +++ b/taskiq/acks.py @@ -1,11 +1,13 @@ import enum from collections.abc import Awaitable, Callable -from typing import Any +from typing import Any, TypeAlias from pydantic import BaseModel from taskiq.utils import maybe_awaitable +AckCallback: TypeAlias = Callable[..., None | Awaitable[None]] + @enum.unique class AcknowledgeType(str, enum.Enum): @@ -54,13 +56,13 @@ class AckableMessage(BaseModel): """ data: bytes - ack: Callable[[], None | Awaitable[None]] + ack: AckCallback class AckController: """Controls acknowledgement state for a received message.""" - def __init__(self, ack: Callable[[], None | Awaitable[None]] | None) -> None: + def __init__(self, ack: AckCallback | None) -> None: self._ack = ack self.is_acked = False @@ -69,11 +71,11 @@ def is_ackable(self) -> bool: """Whether the current message supports acknowledgement.""" return self._ack is not None - async def ack(self) -> None: - """Acknowledge the current message once.""" + async def ack(self, **kwargs: Any) -> None: + """Acknowledge the current message once with broker-specific options.""" if self._ack is None: raise RuntimeError("Current message is not ackable.") if self.is_acked: return - await maybe_awaitable(self._ack()) + await maybe_awaitable(self._ack(**kwargs)) self.is_acked = True diff --git a/taskiq/context.py b/taskiq/context.py index 48058ebb..ea9214fe 100644 --- a/taskiq/context.py +++ b/taskiq/context.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from taskiq.abc.broker import AsyncBroker from taskiq.acks import AckController @@ -34,15 +34,17 @@ def is_acked(self) -> bool: """Whether the current message has already been acknowledged.""" return self._ack_controller is not None and self._ack_controller.is_acked - async def ack(self) -> None: + async def ack(self, **kwargs: Any) -> None: """ Acknowledge current message. + :param kwargs: Broker-specific options forwarded to the acknowledgement + callback. :raises RuntimeError: if current broker message is not ackable. """ if self._ack_controller is None: raise RuntimeError("Current message is not ackable.") - await self._ack_controller.ack() + await self._ack_controller.ack(**kwargs) async def requeue(self) -> None: """ diff --git a/tests/receiver/test_receiver.py b/tests/receiver/test_receiver.py index d724b326..9cacd554 100644 --- a/tests/receiver/test_receiver.py +++ b/tests/receiver/test_receiver.py @@ -560,6 +560,42 @@ def ack_callback() -> None: assert events == ["task", "ack", "post_execute", "save", "post_save"] +async def test_manual_task_ack_forwards_broker_specific_options() -> None: + """Context.ack forwards options to the broker's acknowledgement callback.""" + events: list[str] = [] + received_options: dict[str, Any] = {} + broker = ( + InMemoryBroker() + .with_result_backend( + _EventResultBackend(events), + ) + .with_middlewares(_EventMiddleware(events)) + ) + + @broker.task(ack_type="manual") + async def my_task(context: Context = Depends()) -> int: + events.append("task") + await context.ack(delete_after_ack=True) + return 1 + + def ack_callback(**kwargs: Any) -> None: + received_options.update(kwargs) + events.append("ack") + + receiver = get_receiver(broker, ack_type=AcknowledgeType.WHEN_SAVED) + broker_message = broker.formatter.dumps(my_task.kicker()._prepare_message()) + + await receiver.callback( + AckableMessage( + data=broker_message.message, + ack=ack_callback, + ), + ) + + assert received_options == {"delete_after_ack": True} + assert events == ["task", "ack", "post_execute", "save", "post_save"] + + async def test_manual_task_ack_is_idempotent() -> None: """Calling Context.ack twice acknowledges the message once.""" events: list[str] = [] diff --git a/tests/test_acks.py b/tests/test_acks.py index fefd3f53..e498648c 100644 --- a/tests/test_acks.py +++ b/tests/test_acks.py @@ -1,6 +1,22 @@ +from typing import Any + import pytest -from taskiq.acks import AcknowledgeType, parse_acknowledge_type +from taskiq.acks import AckController, AcknowledgeType, parse_acknowledge_type + + +async def test_ack_forwards_broker_specific_options() -> None: + """AckController forwards explicit options to a supporting callback.""" + received_options: dict[str, Any] = {} + + async def ack_callback(**kwargs: Any) -> None: + received_options.update(kwargs) + + controller = AckController(ack_callback) + await controller.ack(delete_after_ack=True) + + assert controller.is_acked + assert received_options == {"delete_after_ack": True} def test_parse_acknowledge_type_from_enum() -> None: