Skip to content

Commit d6af7d6

Browse files
committed
Add manual task acknowledgement
1 parent fe92683 commit d6af7d6

6 files changed

Lines changed: 189 additions & 20 deletions

File tree

‎docs/guide/cli.md‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -47,10 +47,11 @@ We have two options for this:
4747

4848
### Acknowledgements
4949

50-
The taskiq supports three types of acknowledgements:
50+
The taskiq supports four types of acknowledgements:
5151
* `when_received` - task is acknowledged when it is **received** by the worker.
5252
* `when_executed` - task is acknowledged right after it is **executed** by the worker.
5353
* `when_saved` - task is acknowledged when the result of execution is saved in the result backend.
54+
* `manual` - task is acknowledged by calling `Context.ack()` inside the task.
5455

5556
This can be configured using `--ack-type` parameter. For example:
5657

@@ -69,9 +70,16 @@ async def best_effort_task() -> None:
6970
@broker.task(ack_type="when_saved")
7071
async def critical_task() -> None:
7172
...
73+
74+
75+
@broker.task(ack_type="manual")
76+
async def manually_acked_task(context: Context = TaskiqDepends()) -> None:
77+
...
78+
await context.ack()
7279
```
7380

7481
If a task has no `ack_type` label, the worker-level `--ack-type` value is used.
82+
Manual acknowledgement requires a broker that yields `AckableMessage`.
7583

7684
### Type casts
7785

@@ -159,7 +167,7 @@ The number of signals before a hard kill can be configured with the `--hardkill-
159167
* `--no-propagate-errors` - if this parameter is enabled, exceptions won't be thrown in generator dependencies.
160168
* `--receiver` - python path to custom receiver class.
161169
* `--receiver_arg` - custom args for receiver.
162-
* `--ack-type` - Type of acknowledgement. This parameter is used to set when to acknowledge the task. Possible values are `when_received`, `when_executed`, `when_saved`. Default is `when_saved`.
170+
* `--ack-type` - Type of acknowledgement. This parameter is used to set when to acknowledge the task. Possible values are `when_received`, `when_executed`, `when_saved`, `manual`. Default is `when_saved`.
163171
* `--max-tasks-per-child` - maximum number of tasks to be executed by a single worker process before restart.
164172
* `--max-fails` - Maximum number of child process exits.
165173
* `--shutdown-timeout` - maximum amount of time for graceful broker's shutdown in seconds (default 5).

‎taskiq/acks.py‎

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,8 @@
44

55
from pydantic import BaseModel
66

7+
from taskiq.utils import maybe_awaitable
8+
79

810
@enum.unique
911
class AcknowledgeType(str, enum.Enum):
@@ -19,6 +21,9 @@ class AcknowledgeType(str, enum.Enum):
1921
# acknowledged when the task will be saved
2022
# only after it's saved in the result backend.
2123
WHEN_SAVED = "when_saved"
24+
# This option means that the task is responsible
25+
# for acknowledging the message through Context.ack.
26+
MANUAL = "manual"
2227

2328

2429
def parse_acknowledge_type(value: Any) -> AcknowledgeType:
@@ -50,3 +55,25 @@ class AckableMessage(BaseModel):
5055

5156
data: bytes
5257
ack: Callable[[], None | Awaitable[None]]
58+
59+
60+
class AckController:
61+
"""Controls acknowledgement state for a received message."""
62+
63+
def __init__(self, ack: Callable[[], None | Awaitable[None]] | None) -> None:
64+
self._ack = ack
65+
self.is_acked = False
66+
67+
@property
68+
def is_ackable(self) -> bool:
69+
"""Whether the current message supports acknowledgement."""
70+
return self._ack is not None
71+
72+
async def ack(self) -> None:
73+
"""Acknowledge the current message once."""
74+
if self._ack is None:
75+
raise RuntimeError("Current message is not ackable.")
76+
if self.is_acked:
77+
return
78+
await maybe_awaitable(self._ack())
79+
self.is_acked = True

‎taskiq/context.py‎

Lines changed: 28 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from typing import TYPE_CHECKING
22

33
from taskiq.abc.broker import AsyncBroker
4+
from taskiq.acks import AckController
45
from taskiq.exceptions import NoResultError, TaskRejectedError
56
from taskiq.message import TaskiqMessage
67

@@ -11,12 +12,38 @@
1112
class Context:
1213
"""Context class."""
1314

14-
def __init__(self, message: TaskiqMessage, broker: AsyncBroker) -> None:
15+
def __init__(
16+
self,
17+
message: TaskiqMessage,
18+
broker: AsyncBroker,
19+
ack_controller: AckController | None = None,
20+
) -> None:
1521
self.message = message
1622
self.broker = broker
23+
self._ack_controller = ack_controller
1724
self.state: TaskiqState = None # type: ignore
1825
self.state = broker.state
1926

27+
@property
28+
def is_ackable(self) -> bool:
29+
"""Whether the current message supports acknowledgement."""
30+
return self._ack_controller is not None and self._ack_controller.is_ackable
31+
32+
@property
33+
def is_acked(self) -> bool:
34+
"""Whether the current message has already been acknowledged."""
35+
return self._ack_controller is not None and self._ack_controller.is_acked
36+
37+
async def ack(self) -> None:
38+
"""
39+
Acknowledge current message.
40+
41+
:raises RuntimeError: if current broker message is not ackable.
42+
"""
43+
if self._ack_controller is None:
44+
raise RuntimeError("Current message is not ackable.")
45+
await self._ack_controller.ack()
46+
2047
async def requeue(self) -> None:
2148
"""
2249
Requeue task.

‎taskiq/receiver/receiver.py‎

Lines changed: 13 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@
1515

1616
from taskiq.abc.broker import AckableMessage, AsyncBroker
1717
from taskiq.abc.middleware import TaskiqMiddleware
18-
from taskiq.acks import AcknowledgeType, parse_acknowledge_type
18+
from taskiq.acks import AckController, AcknowledgeType, parse_acknowledge_type
1919
from taskiq.context import Context
2020
from taskiq.exceptions import NoResultError
2121
from taskiq.message import TaskiqMessage
@@ -116,6 +116,9 @@ async def callback( # noqa: C901, PLR0912
116116
:param raise_err: raise an error if cannot save result in
117117
result_backend.
118118
"""
119+
ack_controller = AckController(
120+
message.ack if isinstance(message, AckableMessage) else None,
121+
)
119122
message_data = message.data if isinstance(message, AckableMessage) else message
120123
try:
121124
taskiq_msg = self.broker.formatter.loads(message=message_data)
@@ -155,22 +158,17 @@ async def callback( # noqa: C901, PLR0912
155158
taskiq_msg.task_id,
156159
)
157160

158-
if ack_time == AcknowledgeType.WHEN_RECEIVED and isinstance(
159-
message,
160-
AckableMessage,
161-
):
162-
await maybe_awaitable(message.ack())
161+
if ack_time == AcknowledgeType.WHEN_RECEIVED and ack_controller.is_ackable:
162+
await ack_controller.ack()
163163

164164
result = await self.run_task(
165165
target=task.original_func,
166166
message=taskiq_msg,
167+
ack_controller=ack_controller,
167168
)
168169

169-
if ack_time == AcknowledgeType.WHEN_EXECUTED and isinstance(
170-
message,
171-
AckableMessage,
172-
):
173-
await maybe_awaitable(message.ack())
170+
if ack_time == AcknowledgeType.WHEN_EXECUTED and ack_controller.is_ackable:
171+
await ack_controller.ack()
174172

175173
for middleware in reversed(self.broker.middlewares):
176174
if middleware.__class__.post_execute != TaskiqMiddleware.post_execute:
@@ -193,11 +191,8 @@ async def callback( # noqa: C901, PLR0912
193191
if raise_err:
194192
raise exc
195193

196-
if ack_time == AcknowledgeType.WHEN_SAVED and isinstance(
197-
message,
198-
AckableMessage,
199-
):
200-
await maybe_awaitable(message.ack())
194+
if ack_time == AcknowledgeType.WHEN_SAVED and ack_controller.is_ackable:
195+
await ack_controller.ack()
201196

202197
def _get_ack_time(self, message: TaskiqMessage) -> AcknowledgeType:
203198
"""
@@ -219,6 +214,7 @@ async def run_task( # noqa: C901, PLR0912, PLR0915
219214
self,
220215
target: Callable[..., Any],
221216
message: TaskiqMessage,
217+
ack_controller: AckController | None = None,
222218
) -> TaskiqResult[Any]:
223219
"""
224220
This function actually executes functions.
@@ -259,7 +255,7 @@ async def run_task( # noqa: C901, PLR0912, PLR0915
259255
broker_ctx = self.broker.custom_dependency_context
260256
broker_ctx.update(
261257
{
262-
Context: Context(message, self.broker),
258+
Context: Context(message, self.broker, ack_controller),
263259
TaskiqState: self.broker.state,
264260
},
265261
)

‎tests/receiver/test_receiver.py‎

Lines changed: 107 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
from taskiq.abc.result_backend import AsyncResultBackend
1717
from taskiq.acks import AcknowledgeType
1818
from taskiq.brokers.inmemory_broker import InMemoryBroker
19+
from taskiq.context import Context
1920
from taskiq.exceptions import NoResultError, TaskiqResultTimeoutError
2021
from taskiq.message import TaskiqMessage
2122
from taskiq.receiver import Receiver
@@ -523,6 +524,112 @@ def ack_callback() -> None:
523524
assert events == ["ack", "task", "post_execute", "save", "post_save"]
524525

525526

527+
async def test_manual_task_ack_from_context() -> None:
528+
"""Manual ack lets a task acknowledge the message through context."""
529+
events: list[str] = []
530+
broker = (
531+
InMemoryBroker()
532+
.with_result_backend(
533+
_EventResultBackend(events),
534+
)
535+
.with_middlewares(_EventMiddleware(events))
536+
)
537+
538+
@broker.task(ack_type="manual")
539+
async def my_task(context: Context = Depends()) -> int:
540+
events.append("task")
541+
assert context.is_ackable
542+
assert not context.is_acked
543+
await context.ack()
544+
assert context.is_acked
545+
return 1
546+
547+
def ack_callback() -> None:
548+
events.append("ack")
549+
550+
receiver = get_receiver(broker, ack_type=AcknowledgeType.WHEN_SAVED)
551+
broker_message = broker.formatter.dumps(my_task.kicker()._prepare_message())
552+
553+
await receiver.callback(
554+
AckableMessage(
555+
data=broker_message.message,
556+
ack=ack_callback,
557+
),
558+
)
559+
560+
assert events == ["task", "ack", "post_execute", "save", "post_save"]
561+
562+
563+
async def test_manual_task_ack_is_idempotent() -> None:
564+
"""Calling Context.ack twice acknowledges the message once."""
565+
events: list[str] = []
566+
broker = (
567+
InMemoryBroker()
568+
.with_result_backend(
569+
_EventResultBackend(events),
570+
)
571+
.with_middlewares(_EventMiddleware(events))
572+
)
573+
574+
@broker.task(ack_type="manual")
575+
async def my_task(context: Context = Depends()) -> int:
576+
events.append("task")
577+
await context.ack()
578+
await context.ack()
579+
return 1
580+
581+
def ack_callback() -> None:
582+
events.append("ack")
583+
584+
receiver = get_receiver(broker, ack_type=AcknowledgeType.WHEN_SAVED)
585+
broker_message = broker.formatter.dumps(my_task.kicker()._prepare_message())
586+
587+
await receiver.callback(
588+
AckableMessage(
589+
data=broker_message.message,
590+
ack=ack_callback,
591+
),
592+
)
593+
594+
assert events == ["task", "ack", "post_execute", "save", "post_save"]
595+
596+
597+
async def test_manual_task_ack_without_ackable_message_saves_error() -> None:
598+
"""Context.ack fails clearly when broker messages do not support acking."""
599+
events: list[str] = []
600+
result_backend = _EventResultBackend(events)
601+
broker = (
602+
InMemoryBroker()
603+
.with_result_backend(result_backend)
604+
.with_middlewares(_EventMiddleware(events))
605+
)
606+
607+
@broker.task(ack_type="manual")
608+
async def my_task(context: Context = Depends()) -> int:
609+
assert not context.is_ackable
610+
await context.ack()
611+
return 1
612+
613+
receiver = get_receiver(broker)
614+
broker_message = broker.formatter.dumps(
615+
TaskiqMessage(
616+
task_id="task_id",
617+
task_name=my_task.task_name,
618+
labels={"ack_type": "manual"},
619+
args=[],
620+
kwargs={},
621+
),
622+
)
623+
624+
await receiver.callback(broker_message.message)
625+
626+
result = result_backend.results["task_id"]
627+
assert result.is_err
628+
assert isinstance(result.error, RuntimeError)
629+
assert str(result.error) == "Current message is not ackable."
630+
assert events == ["post_execute", "save", "post_save"]
631+
632+
526633
async def test_invalid_task_ack_type_raises_before_execution() -> None:
527634
"""Invalid task ack_type fails before task execution and acknowledgement."""
528635
events: list[str] = []

‎tests/test_acks.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,10 @@ def test_parse_acknowledge_type_from_string() -> None:
1414
assert parse_acknowledge_type("when_executed") == AcknowledgeType.WHEN_EXECUTED
1515

1616

17+
def test_parse_manual_acknowledge_type_from_string() -> None:
18+
assert parse_acknowledge_type("manual") == AcknowledgeType.MANUAL
19+
20+
1721
def test_parse_acknowledge_type_unknown_string() -> None:
1822
with pytest.raises(ValueError, match="Unknown acknowledge type value"):
1923
parse_acknowledge_type("when_save")

0 commit comments

Comments
 (0)