Skip to content

Commit f9d4935

Browse files
committed
feat: add worker-side task batching
Allow grouping same-type tasks into batches that are executed together, reducing per-call overhead for tasks like DB writes, external API calls, ML inference, event publishing and indexing. Usage: @broker.task(batch=True, batch_size=100, batch_timeout=3) async def process_items(items: list[Item]) -> None: ... await process_items.kiq(item) # one item per kiq Each .kiq sends a single item; the worker buffers messages per task name and invokes the function once with the collected list. A batch is flushed when batch_size is reached or batch_timeout seconds pass since the first buffered item (whichever comes first), and on graceful shutdown. Details: - AsyncBatchedTaskiqDecoratedTask types .kiq to accept a single element while the function body receives list[item]. - Batch config is validated at registration (TaskiqBatchConfigError). - The whole batch shares one result, stored under every task_id, and every message is acked per the configured AcknowledgeType. pre_execute, post_execute, post_save and on_error middlewares run per message. - InMemoryBroker supports batching too (flush on size, on wait_all, or immediately when await_inplace=True) for testing.
1 parent 9f8db96 commit f9d4935

14 files changed

Lines changed: 1289 additions & 63 deletions

File tree

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
# broker.py
2+
import asyncio
3+
4+
from taskiq import InMemoryBroker
5+
6+
broker = InMemoryBroker()
7+
8+
9+
@broker.task(batch=True, batch_size=100, batch_timeout=3)
10+
async def process_items(items: list[int]) -> int:
11+
# The worker collects many `.kiq` calls and invokes this
12+
# function once with the accumulated list of items.
13+
print(f"Processing a batch of {len(items)} items.")
14+
return sum(items)
15+
16+
17+
async def main() -> None:
18+
await broker.startup()
19+
# Each `.kiq` sends a single item. They are buffered and run
20+
# together once the batch is flushed.
21+
tasks = [await process_items.kiq(i) for i in range(10)]
22+
# In tests, `wait_all` flushes any pending batches and waits
23+
# for them to finish before we read the results.
24+
await broker.wait_all()
25+
for task in tasks:
26+
result = await task.wait_result(timeout=5)
27+
# Every item in the batch shares the same batch result.
28+
print(f"Returned value: {result.return_value}")
29+
await broker.shutdown()
30+
31+
32+
if __name__ == "__main__":
33+
asyncio.run(main())
Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
# broker.py
2+
import asyncio
3+
4+
from taskiq_redis import RedisAsyncResultBackend, RedisStreamBroker
5+
6+
broker = RedisStreamBroker(url="redis://localhost:6379").with_result_backend(
7+
RedisAsyncResultBackend(redis_url="redis://localhost:6379"),
8+
)
9+
10+
11+
@broker.task(batch=True, batch_size=100, batch_timeout=3)
12+
async def process_items(items: list[int]) -> int:
13+
# The worker collects many `.kiq` calls and invokes this
14+
# function once with the accumulated list of items.
15+
print(f"Processing a batch of {len(items)} items.")
16+
return sum(items)
17+
18+
19+
async def main() -> None:
20+
await broker.startup()
21+
# Each `.kiq` sends a single item. The worker buffers them and
22+
# runs `process_items` once with the whole batch.
23+
tasks = [await process_items.kiq(i) for i in range(10)]
24+
for task in tasks:
25+
result = await task.wait_result(timeout=5)
26+
# Every item in the batch shares the same batch result.
27+
print(f"Returned value: {result.return_value}")
28+
await broker.shutdown()
29+
30+
31+
if __name__ == "__main__":
32+
asyncio.run(main())

‎docs/guide/batching-tasks.md‎

Lines changed: 120 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,120 @@
1+
---
2+
title: Batching tasks
3+
order: 6
4+
---
5+
6+
# Batching tasks
7+
8+
Some tasks have a high fixed cost per call but become much cheaper when processed
9+
together. Think of database writes, calls to an external API, ML inference,
10+
event publishing or search indexing — running them one by one wastes most of the
11+
time on overhead. If a single call takes ~1 second, then 10 separate calls take
12+
~10 seconds, but processing all 10 at once might take only ~3 seconds.
13+
14+
Taskiq can collect many task invocations into a single batched call. Instead of
15+
running every message on its own, the worker buffers messages of the same task
16+
and executes the function once with the whole list.
17+
18+
## Defining a batched task
19+
20+
Pass `batch=True` to the `task` decorator and declare the function with a single
21+
parameter that receives the list of items.
22+
23+
```python
24+
@broker.task(batch=True, batch_size=100, batch_timeout=3)
25+
async def process_items(items: list[int]) -> int:
26+
return sum(items)
27+
```
28+
29+
Each `.kiq` call sends a single item, exactly like a normal task:
30+
31+
```python
32+
await process_items.kiq(1)
33+
await process_items.kiq(2)
34+
```
35+
36+
The worker accumulates these items and calls `process_items` once with the
37+
collected list (e.g. `[1, 2, ...]`).
38+
39+
::: tip Typed by design
40+
41+
`.kiq` accepts a single element, while the function body receives `list[item]`.
42+
Both sides are correctly typed: `process_items.kiq(1)` type-checks, but
43+
`process_items.kiq([1, 2])` is reported as a type error.
44+
45+
:::
46+
47+
## When a batch is flushed
48+
49+
A batch is sent for execution as soon as **either** condition is met:
50+
51+
- **`batch_size`** — the buffer reaches this number of items, or
52+
- **`batch_timeout`** — this many seconds pass since the first item entered the
53+
buffer.
54+
55+
Whichever happens first wins. You must set at least one of the two; you can set
56+
both. The timer starts with the first item of a fresh buffer and resets after
57+
each flush. When a worker shuts down gracefully, any buffered items are flushed
58+
so nothing is lost.
59+
60+
Each worker buffers independently and keeps a separate buffer per task name.
61+
62+
## Results and acknowledgement
63+
64+
A batch produces a single result. That same return value (or error) is stored
65+
for **every** task in the batch, so each `.kiq` call can still await its own
66+
result. If the batched function raises, every task in the batch receives that
67+
error. Every message is acknowledged according to the configured
68+
[acknowledgement type](./cli.md).
69+
70+
::: caution Per-item granularity
71+
72+
Batching trades per-item isolation for throughput. The whole batch shares one
73+
result and one fate — there are no per-item results or per-item error handling.
74+
A batched task must take exactly one positional argument (the list); keyword
75+
arguments are not part of the batched call.
76+
77+
:::
78+
79+
## Trying it locally
80+
81+
Batching is a worker-side feature, but the `InMemoryBroker` supports it too, so
82+
you can try it without setting up a real broker. Call `wait_all` to flush any
83+
pending batches and wait for them to finish before reading results.
84+
85+
@[code python](../examples/batching/inmemory_batch.py)
86+
87+
Running this prints a single batch execution and the shared result:
88+
89+
```bash:no-line-numbers
90+
$ python broker.py
91+
Processing a batch of 10 items.
92+
Returned value: 45
93+
... (10 times)
94+
```
95+
96+
::: warning InMemoryBroker behavior
97+
98+
The `InMemoryBroker` executes tasks inplace, so batches are flushed by
99+
`batch_size`, by `wait_all`, or — with `await_inplace=True` — immediately as
100+
one-item batches. This is convenient for tests, but to see real batching across
101+
processes you need a distributed broker and a worker.
102+
103+
:::
104+
105+
## Running with a worker
106+
107+
In production, batching happens inside the worker. Using
108+
[taskiq-redis](https://pypi.org/project/taskiq-redis/) as an example:
109+
110+
@[code python](../examples/batching/redis_batch.py)
111+
112+
Start one or more workers:
113+
114+
```bash:no-line-numbers
115+
taskiq worker broker:broker
116+
```
117+
118+
Then run the script to send items. The worker collects them and runs
119+
`process_items` once per batch. With several workers, each one batches the
120+
messages it receives independently, so the load is spread across all of them.

‎taskiq/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
from taskiq.brokers.shared_broker import async_shared_broker
1515
from taskiq.brokers.zmq_broker import ZeroMQBroker
1616
from taskiq.context import Context
17+
from taskiq.decor import AsyncBatchedTaskiqDecoratedTask
1718
from taskiq.events import TaskiqEvents
1819
from taskiq.exceptions import (
1920
NoResultError,
@@ -41,6 +42,7 @@
4142

4243
__all__ = [
4344
"AckableMessage",
45+
"AsyncBatchedTaskiqDecoratedTask",
4446
"AsyncBroker",
4547
"AsyncResultBackend",
4648
"AsyncTaskiqDecoratedTask",

‎taskiq/abc/broker.py‎

Lines changed: 43 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
TYPE_CHECKING,
1111
Any,
1212
ClassVar,
13+
Literal,
1314
ParamSpec,
1415
TypeAlias,
1516
TypeVar,
@@ -21,9 +22,9 @@
2122
from taskiq.abc.middleware import TaskiqMiddleware
2223
from taskiq.abc.serializer import TaskiqSerializer
2324
from taskiq.acks import AckableMessage
24-
from taskiq.decor import AsyncTaskiqDecoratedTask
25+
from taskiq.decor import AsyncBatchedTaskiqDecoratedTask, AsyncTaskiqDecoratedTask
2526
from taskiq.events import TaskiqEvents
26-
from taskiq.exceptions import TaskBrokerMismatchError
27+
from taskiq.exceptions import TaskBrokerMismatchError, TaskiqBatchConfigError
2728
from taskiq.formatters.proxy_formatter import ProxyFormatter
2829
from taskiq.message import BrokerMessage
2930
from taskiq.result_backends.dummy import DummyResultBackend
@@ -43,6 +44,7 @@
4344
from taskiq.abc.result_backend import AsyncResultBackend
4445

4546
_T = TypeVar("_T")
47+
_Item = TypeVar("_Item")
4648
_FuncParams = ParamSpec("_FuncParams")
4749
_ReturnType = TypeVar("_ReturnType")
4850

@@ -255,6 +257,40 @@ def listen(self) -> AsyncGenerator[bytes | AckableMessage, None]:
255257
:return: nothing.
256258
"""
257259

260+
@staticmethod
261+
def _validate_batch_labels(labels: dict[str, Any]) -> None:
262+
"""
263+
Validate batch related labels.
264+
265+
:param labels: labels passed to the task decorator.
266+
:raises TaskiqBatchConfigError: if batch configuration is invalid.
267+
"""
268+
if not labels.get("batch"):
269+
return
270+
batch_size = labels.get("batch_size")
271+
batch_timeout = labels.get("batch_timeout")
272+
invalid = (
273+
(batch_size is None and batch_timeout is None)
274+
or (batch_size is not None and batch_size < 1)
275+
or (batch_timeout is not None and batch_timeout <= 0)
276+
)
277+
if invalid:
278+
raise TaskiqBatchConfigError
279+
280+
@overload
281+
def task(
282+
self,
283+
*,
284+
batch: Literal[True],
285+
batch_size: int | None = None,
286+
batch_timeout: float | None = None,
287+
**labels: Any,
288+
) -> Callable[
289+
[Callable[[list[_Item]], Awaitable[_ReturnType]]],
290+
AsyncBatchedTaskiqDecoratedTask[_Item, ..., _ReturnType],
291+
]: # pragma: no cover
292+
...
293+
258294
@overload
259295
def task(
260296
self,
@@ -301,6 +337,7 @@ def task( # type: ignore[misc]
301337
302338
:returns: decorator function or AsyncTaskiqDecoratedTask.
303339
"""
340+
self._validate_batch_labels(labels)
304341

305342
def make_decorated_task(
306343
inner_labels: dict[str, str | int],
@@ -332,8 +369,11 @@ def inner(
332369
if "return" in sign:
333370
return_type = sign["return"]
334371

372+
decorator_cls = self.decorator_class
373+
if inner_labels.get("batch"):
374+
decorator_cls = AsyncBatchedTaskiqDecoratedTask
335375
decorated_task = wrapper(
336-
self.decorator_class(
376+
decorator_cls(
337377
broker=self,
338378
original_func=func,
339379
labels=inner_labels,

‎taskiq/brokers/inmemory_broker.py‎

Lines changed: 42 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -153,6 +153,10 @@ async def kick(self, message: BrokerMessage) -> None:
153153
154154
This method just executes given task.
155155
156+
For batched tasks the message is buffered and the batch is executed
157+
once the batch size is reached (or immediately when `await_inplace`
158+
is enabled). Call `wait_all` to flush any remaining buffered batches.
159+
156160
:param message: incoming message.
157161
158162
:raises TaskiqError: if someone wants to kick unknown task.
@@ -161,6 +165,11 @@ async def kick(self, message: BrokerMessage) -> None:
161165
if target_task is None:
162166
raise UnknownTaskError(task_name=message.task_name)
163167

168+
batch_config = self.receiver.get_batch_config(message.task_name)
169+
if batch_config is not None:
170+
await self._kick_batched(message, batch_config)
171+
return
172+
164173
receiver_cb = self.receiver.callback(message=message.message)
165174
if self.await_inplace:
166175
await receiver_cb
@@ -170,6 +179,35 @@ async def kick(self, message: BrokerMessage) -> None:
170179
self._running_tasks.add(task)
171180
task.add_done_callback(self._running_tasks.discard)
172181

182+
async def _kick_batched(
183+
self,
184+
message: BrokerMessage,
185+
batch_config: "tuple[int | None, float | None]",
186+
) -> None:
187+
"""
188+
Buffer a batched task message and flush according to its config.
189+
190+
With `await_inplace` each message is flushed immediately as a
191+
one-item batch so that results are available right after `kiq`.
192+
193+
:param message: incoming message.
194+
:param batch_config: (batch_size, batch_timeout) of the task.
195+
"""
196+
size, timeout = batch_config
197+
if self.await_inplace:
198+
# Flush immediately: one message becomes a one-item batch.
199+
await self.receiver.batched_callback(
200+
task_name=message.task_name,
201+
messages=[message.message],
202+
)
203+
return
204+
await self.receiver.batcher.add(
205+
message.task_name,
206+
message.message,
207+
size,
208+
timeout,
209+
)
210+
173211
def listen(self) -> AsyncGenerator[bytes, None]:
174212
"""
175213
Inmemory broker cannot listen.
@@ -186,11 +224,14 @@ async def wait_all(self) -> None:
186224
Wait for all currently running tasks to complete.
187225
188226
Useful when used in testing and you need to await all sent tasks
189-
before asserting results.
227+
before asserting results. Any buffered batched tasks are flushed
228+
first so their results become available too.
190229
"""
230+
await self.receiver.batcher.flush_all()
191231
to_await = list(self._running_tasks)
192232
for task in to_await:
193233
await task
234+
await self.receiver.wait_for_batch_tasks()
194235

195236
async def startup(self) -> None:
196237
"""Runs startup events for client and worker side."""

0 commit comments

Comments
 (0)