Skip to content

Commit b6f07dc

Browse files
committed
fix(receiver): tighten prefetch ownership checks
Reject invalid prefetch values before receiver retries and retain delivery accounting until execution capacity is acquired. Add deterministic coverage for CLI defaults, cancellation races, and listener lifecycle cleanup. Refs #528
1 parent 4321e09 commit b6f07dc

6 files changed

Lines changed: 208 additions & 10 deletions

File tree

taskiq/api/receiver.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,11 @@ async def run_receiver_task(
5252
:param ack_time: acknowledge type to use.
5353
:param use_process_pool: whether to use process pool or threadpool.
5454
:raises asyncio.CancelledError: if the task was cancelled.
55+
:raises ValueError: if max_prefetch is negative.
5556
"""
57+
if max_prefetch < 0:
58+
raise ValueError("max_prefetch cannot be negative.")
59+
5660
finish_event = asyncio.Event()
5761

5862
def on_exit(_: Receiver) -> None:

taskiq/cli/worker/args.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -286,6 +286,8 @@ def from_cli(
286286
args,
287287
namespace=None if defaults is None else Namespace(**defaults),
288288
)
289+
if namespace.max_prefetch < 0:
290+
parser.error("argument --max-prefetch: max_prefetch cannot be negative.")
289291
# If there are any patterns specified, remove default.
290292
# This is an argparse limitation.
291293
if len(namespace.tasks_pattern) > 1:

taskiq/receiver/receiver.py

Lines changed: 25 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
from enum import Enum, auto
1111
from logging import getLogger
1212
from time import time
13-
from typing import Any, get_type_hints
13+
from typing import Any, Literal, get_type_hints
1414

1515
import anyio
1616
from taskiq_dependencies import DependencyGraph
@@ -645,7 +645,13 @@ async def _close_prefetch_state(
645645
state.owns_delivery_slot = False
646646
return late_delivery
647647

648-
async def _notify_prefetch_hook(self, hook_name: str) -> None:
648+
async def _notify_prefetch_hook(
649+
self,
650+
hook_name: Literal[
651+
"on_prefetch_queue_add",
652+
"on_prefetch_queue_remove",
653+
],
654+
) -> None:
649655
"""Run all prefetch hooks and preserve the first failure."""
650656
first_error: BaseException | None = None
651657
for middleware in reversed(self.broker.middlewares):
@@ -723,7 +729,18 @@ async def runner(
723729
await asyncio.wait(tasks, timeout=self.wait_tasks_timeout)
724730
logger.info("No more tasks to wait for. Shutting down.")
725731
break
726-
started_callback = await self._start_callback(queued_message)
732+
execution_semaphore = self.sem
733+
owns_execution_slot = execution_semaphore is not None
734+
if execution_semaphore is not None:
735+
try:
736+
await execution_semaphore.acquire()
737+
except BaseException:
738+
queue.put_nowait(queued_message)
739+
raise
740+
started_callback = await self._start_callback(
741+
queued_message,
742+
owns_execution_slot=owns_execution_slot,
743+
)
727744
tasks.add(started_callback.task)
728745

729746
# We want the task to remove itself from the set when it's done.
@@ -747,15 +764,13 @@ async def runner(
747764
async def _start_callback(
748765
self,
749766
message: _PrefetchedMessage,
767+
*,
768+
owns_execution_slot: bool,
750769
) -> _StartedCallback:
751770
"""Transfer execution and delivery capacity to a callback task."""
752771
owns_delivery_slot = message.owns_delivery_slot
753-
owns_execution_slot = False
754772
try:
755773
await self._notify_prefetch_hook("on_prefetch_queue_remove")
756-
if self.sem is not None:
757-
await self.sem.acquire()
758-
owns_execution_slot = True
759774

760775
if self.sem is None and owns_delivery_slot:
761776
self.sem_prefetch.release()
@@ -767,6 +782,9 @@ async def _start_callback(
767782
owns_delivery_slot=owns_delivery_slot,
768783
)
769784
except BaseException:
785+
logger.warning(
786+
"Discarding 1 prefetched delivery during Receiver cleanup.",
787+
)
770788
if owns_delivery_slot:
771789
self.sem_prefetch.release()
772790
if owns_execution_slot and self.sem is not None:

tests/api/test_receiver_task.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,30 @@
11
import asyncio
22
import contextlib
3+
from typing import Any
34

45
import pytest
56

67
from taskiq.api import run_receiver_task
8+
from taskiq.receiver import Receiver
79
from tests.utils import AsyncQueueBroker
810

911

12+
class _UnexpectedReceiverRetry(BaseException):
13+
"""Signal that invalid configuration reached the Receiver retry loop."""
14+
15+
16+
class _ValidationProbeReceiver(Receiver):
17+
"""Fail deterministically if invalid configuration is retried."""
18+
19+
construction_attempts = 0
20+
21+
def __init__(self, *args: Any, **kwargs: Any) -> None:
22+
type(self).construction_attempts += 1
23+
if type(self).construction_attempts > 1:
24+
raise _UnexpectedReceiverRetry
25+
super().__init__(*args, **kwargs)
26+
27+
1028
async def test_successful() -> None:
1129
broker = AsyncQueueBroker()
1230
kicked = 0
@@ -52,3 +70,18 @@ def test_func() -> None:
5270
with pytest.raises(asyncio.TimeoutError):
5371
await asyncio.wait_for(broker.wait_tasks(), 0.2)
5472
assert kicked == 1
73+
74+
75+
async def test_negative_prefetch_is_rejected_before_receiver_retry() -> None:
76+
broker = AsyncQueueBroker()
77+
_ValidationProbeReceiver.construction_attempts = 0
78+
79+
with pytest.raises(ValueError, match="max_prefetch cannot be negative"):
80+
await run_receiver_task(
81+
broker,
82+
receiver_cls=_ValidationProbeReceiver,
83+
max_prefetch=-1,
84+
)
85+
86+
assert _ValidationProbeReceiver.construction_attempts == 0
87+
assert not broker.is_worker_process

tests/cli/worker/test_args.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
import pytest
2+
3+
from taskiq.cli.worker.args import WorkerArgs
4+
5+
6+
@pytest.mark.parametrize("max_prefetch", [0, 3])
7+
def test_max_prefetch_accepts_non_negative_values(max_prefetch: int) -> None:
8+
args = WorkerArgs.from_cli(
9+
["example:broker", "--max-prefetch", str(max_prefetch)],
10+
)
11+
12+
assert args.max_prefetch == max_prefetch
13+
14+
15+
def test_max_prefetch_rejects_negative_value(
16+
capsys: pytest.CaptureFixture[str],
17+
) -> None:
18+
with pytest.raises(SystemExit) as exc_info:
19+
WorkerArgs.from_cli(
20+
["example:broker", "--max-prefetch", "-1"],
21+
)
22+
23+
assert exc_info.value.code == 2
24+
assert "max_prefetch cannot be negative" in capsys.readouterr().err
25+
26+
27+
def test_max_prefetch_rejects_negative_default(
28+
capsys: pytest.CaptureFixture[str],
29+
) -> None:
30+
with pytest.raises(SystemExit) as exc_info:
31+
WorkerArgs.from_cli(
32+
["example:broker"],
33+
defaults={"max_prefetch": -1},
34+
)
35+
36+
assert exc_info.value.code == 2
37+
assert "max_prefetch cannot be negative" in capsys.readouterr().err

tests/receiver/test_receiver_listener.py

Lines changed: 107 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -215,6 +215,74 @@ async def task() -> None:
215215
await assert_semaphore_capacity(receiver.sem, 2)
216216

217217

218+
async def test_runner_retains_prefetch_ownership_until_callback_handoff(
219+
caplog: pytest.LogCaptureFixture,
220+
) -> None:
221+
broker = ControlledBroker()
222+
prefetch_counter = PrefetchCounterMiddleware()
223+
three_messages_added = asyncio.Event()
224+
callback_started = asyncio.Event()
225+
callback_finished = asyncio.Event()
226+
release_callback = asyncio.Event()
227+
added_messages = 0
228+
229+
class AddBarrierMiddleware(TaskiqMiddleware):
230+
def on_prefetch_queue_add(self) -> None:
231+
nonlocal added_messages
232+
added_messages += 1
233+
if added_messages == 3:
234+
three_messages_added.set()
235+
236+
broker.with_middlewares(prefetch_counter, AddBarrierMiddleware())
237+
238+
@broker.task(task_name="receiver.prefetch.runner-owned")
239+
async def task() -> None:
240+
callback_started.set()
241+
try:
242+
await release_callback.wait()
243+
finally:
244+
callback_finished.set()
245+
246+
receiver = Receiver(
247+
broker,
248+
max_async_tasks=1,
249+
max_prefetch=2,
250+
run_startup=False,
251+
)
252+
execution_capacity = ObservedSemaphore(1)
253+
receiver.sem = execution_capacity
254+
for _ in range(3):
255+
await task.kiq()
256+
257+
listen_task = asyncio.create_task(receiver.listen(asyncio.Event()))
258+
try:
259+
await asyncio.wait_for(callback_started.wait(), timeout=1)
260+
assert await execution_capacity.acquire_attempts.get() == 1
261+
assert await execution_capacity.acquire_attempts.get() == 2
262+
await asyncio.wait_for(three_messages_added.wait(), timeout=1)
263+
264+
assert prefetch_counter.queued_messages == 2
265+
266+
with caplog.at_level(logging.WARNING, logger="taskiq.receiver.receiver"):
267+
listen_task.cancel()
268+
with pytest.raises(asyncio.CancelledError):
269+
await listen_task
270+
271+
assert (
272+
"Discarding 2 prefetched deliveries during Receiver cleanup" in caplog.text
273+
)
274+
assert prefetch_counter.queued_messages == 0
275+
finally:
276+
if not listen_task.done():
277+
listen_task.cancel()
278+
await asyncio.gather(listen_task, return_exceptions=True)
279+
release_callback.set()
280+
await asyncio.wait_for(callback_finished.wait(), timeout=1)
281+
282+
await assert_semaphore_capacity(receiver.sem_prefetch, 3)
283+
await assert_semaphore_capacity(execution_capacity, 1)
284+
285+
218286
async def test_unlimited_execution_keeps_zero_prefetch_handoff_progress() -> None:
219287
broker = ControlledBroker()
220288
two_callbacks_started = asyncio.Event()
@@ -255,7 +323,7 @@ async def task() -> None:
255323
await asyncio.wait_for(callbacks_finished.wait(), timeout=1)
256324

257325

258-
def test_delivery_capacity_uses_jittered_execution_limit() -> None:
326+
async def test_delivery_capacity_uses_jittered_execution_limit() -> None:
259327
with unittest.mock.patch(
260328
"taskiq.receiver.receiver.random.randint",
261329
return_value=3,
@@ -269,8 +337,8 @@ def test_delivery_capacity_uses_jittered_execution_limit() -> None:
269337
)
270338

271339
assert receiver.sem is not None
272-
assert receiver.sem._value == 8
273-
assert receiver.sem_prefetch._value == 10
340+
await assert_semaphore_capacity(receiver.sem, 8)
341+
await assert_semaphore_capacity(receiver.sem_prefetch, 10)
274342

275343

276344
def test_negative_prefetch_is_rejected_before_listener_startup() -> None:
@@ -297,6 +365,23 @@ async def test_finish_wakes_prefetcher_blocked_on_capacity() -> None:
297365
assert observed_semaphore.locked()
298366

299367

368+
async def test_finish_returns_concurrently_acquired_capacity_once() -> None:
369+
broker = ControlledBroker()
370+
receiver = Receiver(broker, max_prefetch=0, run_startup=False)
371+
observed_semaphore = ObservedSemaphore(0)
372+
receiver.sem_prefetch = observed_semaphore
373+
finish_event = asyncio.Event()
374+
listen_task = asyncio.create_task(receiver.listen(finish_event))
375+
376+
await observed_semaphore.acquire_started.wait()
377+
observed_semaphore.release()
378+
finish_event.set()
379+
await asyncio.wait_for(listen_task, timeout=1)
380+
381+
assert broker.read_started.empty()
382+
await assert_semaphore_capacity(observed_semaphore, 1)
383+
384+
300385
async def test_pending_read_is_cancelled_and_iterator_closed() -> None:
301386
broker = ControlledBroker()
302387
receiver = Receiver(broker, max_prefetch=1, run_startup=False)
@@ -571,6 +656,25 @@ def opening_failure() -> AsyncGenerator[bytes | AckableMessage, None]:
571656
assert exc_info.value is listener_error
572657

573658

659+
async def test_listener_exhaustion_stops_cleanly_and_releases_capacity() -> None:
660+
listener_closed = asyncio.Event()
661+
662+
async def exhausted_listener() -> AsyncGenerator[bytes | AckableMessage, None]:
663+
try:
664+
if False: # pragma: no branch
665+
yield b""
666+
finally:
667+
listener_closed.set()
668+
669+
broker = ListenerBroker(exhausted_listener)
670+
receiver = Receiver(broker, max_prefetch=0, run_startup=False)
671+
672+
await asyncio.wait_for(receiver.listen(asyncio.Event()), timeout=1)
673+
674+
assert listener_closed.is_set()
675+
await assert_semaphore_capacity(receiver.sem_prefetch, 1)
676+
677+
574678
async def test_prefetch_add_hook_failure_preserves_delivery_and_capacity() -> None:
575679
hook_error = ReceiverLifecycleError("prefetch add hook failed")
576680
executed = asyncio.Event()

0 commit comments

Comments
 (0)