Skip to content

Commit a37ac15

Browse files
committed
fix(receiver): simplify prefetch listener lifecycle
- remove undocumented prefetch middleware hooks and metrics - preserve listener errors during task group failures - restore capacity when callback handoff fails - strengthen shutdown and cleanup coverage Refs: #528
1 parent b6f07dc commit a37ac15

5 files changed

Lines changed: 100 additions & 265 deletions

File tree

taskiq/middlewares/opentelemetry_middleware.py

Lines changed: 0 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -230,12 +230,6 @@ def __init__(
230230
unit="1",
231231
description="Number of tasks currently executing in the worker.",
232232
)
233-
# 9- Number of tasks executing
234-
self.number_of_broker_prefetched_tasks = self._meter.create_up_down_counter(
235-
"worker_prefetched_tasks",
236-
unit="1",
237-
description="Number of tasks currently prefetched in the worker.",
238-
)
239233

240234
def _observe_memory(self, options: Any) -> Generator[Observation, None, None]:
241235
if self.broker and self.broker.is_worker_process:
@@ -447,11 +441,3 @@ def post_execute(
447441
-1,
448442
attributes={"task_name": message.task_name},
449443
)
450-
451-
def on_prefetch_queue_add(self) -> None:
452-
"""This hook is called after task is added to the worker prefetch queue."""
453-
self.number_of_broker_prefetched_tasks.add(1)
454-
455-
def on_prefetch_queue_remove(self) -> None:
456-
"""This hook is called after task is removed from the worker prefetch queue."""
457-
self.number_of_broker_prefetched_tasks.add(-1)

taskiq/receiver/receiver.py

Lines changed: 11 additions & 66 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import asyncio
2+
import contextlib
23
import contextvars
34
import functools
45
import inspect
@@ -10,7 +11,7 @@
1011
from enum import Enum, auto
1112
from logging import getLogger
1213
from time import time
13-
from typing import Any, Literal, get_type_hints
14+
from typing import Any, get_type_hints
1415

1516
import anyio
1617
from taskiq_dependencies import DependencyGraph
@@ -442,7 +443,7 @@ async def listen(self, finish_event: asyncio.Event) -> None: # pragma: no cover
442443
logger.error(
443444
"A Receiver listener lifecycle error was recorded before "
444445
"the task group failed.",
445-
exc_info=(type(error), error, error.__traceback__),
446+
exc_info=error,
446447
)
447448
raise
448449

@@ -463,12 +464,7 @@ async def prefetcher(
463464
:param queue: queue for prefetched data.
464465
:param finish_event: event to indicate that we need to stop prefetching.
465466
"""
466-
try:
467-
state = _PrefetchState(iterator=self.broker.listen())
468-
except BaseException as exc:
469-
self._record_listen_error(exc)
470-
queue.put_nowait(_QueueSignal.DONE)
471-
return
467+
state = _PrefetchState(iterator=self.broker.listen())
472468

473469
fetched_tasks = 0
474470
finish_waiter = asyncio.create_task(finish_event.wait())
@@ -494,13 +490,6 @@ async def prefetcher(
494490
),
495491
)
496492
state.owns_delivery_slot = False
497-
try:
498-
await self._notify_prefetch_hook("on_prefetch_queue_add")
499-
except asyncio.CancelledError:
500-
raise
501-
except BaseException as exc:
502-
self._record_listen_error(exc)
503-
break
504493
finally:
505494
logger.info("Stopping prefetching messages...")
506495
with anyio.CancelScope(shield=True):
@@ -509,7 +498,8 @@ async def prefetcher(
509498
finally:
510499
queue.put_nowait(_QueueSignal.DONE)
511500
finish_waiter.cancel()
512-
await asyncio.gather(finish_waiter, return_exceptions=True)
501+
with contextlib.suppress(asyncio.CancelledError):
502+
await finish_waiter
513503

514504
def _should_stop_prefetch(
515505
self,
@@ -563,15 +553,15 @@ async def _acquire_delivery_slot(
563553
"""Acquire one delivery-admission slot unless shutdown wins the race."""
564554
acquire_task = asyncio.create_task(self.sem_prefetch.acquire())
565555
try:
566-
done, _ = await asyncio.wait(
556+
await asyncio.wait(
567557
{acquire_task, finish_waiter},
568558
return_when=asyncio.FIRST_COMPLETED,
569559
)
570560
except BaseException:
571561
await self._settle_delivery_acquire(acquire_task)
572562
raise
573563

574-
if finish_waiter in done or finish_event.is_set():
564+
if finish_waiter.done() or finish_event.is_set():
575565
await self._settle_delivery_acquire(acquire_task)
576566
return False
577567

@@ -605,10 +595,6 @@ async def _enqueue_late_prefetched_message(
605595
return
606596

607597
queue.put_nowait(late_delivery)
608-
try:
609-
await self._notify_prefetch_hook("on_prefetch_queue_add")
610-
except BaseException as exc:
611-
self._record_listen_error(exc)
612598

613599
async def _close_prefetch_state(
614600
self,
@@ -645,38 +631,11 @@ async def _close_prefetch_state(
645631
state.owns_delivery_slot = False
646632
return late_delivery
647633

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:
655-
"""Run all prefetch hooks and preserve the first failure."""
656-
first_error: BaseException | None = None
657-
for middleware in reversed(self.broker.middlewares):
658-
hook = getattr(middleware, hook_name, None)
659-
if hook is not None:
660-
try:
661-
await maybe_awaitable(hook())
662-
except BaseException as exc:
663-
if first_error is None:
664-
first_error = exc
665-
else:
666-
logger.error(
667-
"Additional error while running prefetch hook %s.",
668-
hook_name,
669-
exc_info=(type(exc), exc, exc.__traceback__),
670-
)
671-
672-
if first_error is not None:
673-
raise first_error
674-
675634
async def _discard_queued_messages(
676635
self,
677636
queue: "asyncio.Queue[_PrefetchedMessage | _QueueSignal]",
678637
) -> None:
679-
"""Release capacity and instrumentation for abandoned queue entries."""
638+
"""Release capacity for abandoned queue entries."""
680639
discarded_messages = 0
681640
while True:
682641
try:
@@ -694,18 +653,6 @@ async def _discard_queued_messages(
694653
discarded_messages,
695654
)
696655

697-
# Restore all capacity before cleanup hooks introduce suspension points.
698-
first_error: BaseException | None = None
699-
for _ in range(discarded_messages):
700-
try:
701-
await self._notify_prefetch_hook("on_prefetch_queue_remove")
702-
except BaseException as exc:
703-
if first_error is None:
704-
first_error = exc
705-
706-
if first_error is not None:
707-
self._record_listen_error(first_error)
708-
709656
async def runner(
710657
self,
711658
queue: "asyncio.Queue[_PrefetchedMessage | _QueueSignal]",
@@ -737,7 +684,7 @@ async def runner(
737684
except BaseException:
738685
queue.put_nowait(queued_message)
739686
raise
740-
started_callback = await self._start_callback(
687+
started_callback = self._start_callback(
741688
queued_message,
742689
owns_execution_slot=owns_execution_slot,
743690
)
@@ -761,7 +708,7 @@ async def runner(
761708
break
762709
logger.info("The runner is stopped.")
763710

764-
async def _start_callback(
711+
def _start_callback(
765712
self,
766713
message: _PrefetchedMessage,
767714
*,
@@ -770,8 +717,6 @@ async def _start_callback(
770717
"""Transfer execution and delivery capacity to a callback task."""
771718
owns_delivery_slot = message.owns_delivery_slot
772719
try:
773-
await self._notify_prefetch_hook("on_prefetch_queue_remove")
774-
775720
if self.sem is None and owns_delivery_slot:
776721
self.sem_prefetch.release()
777722
owns_delivery_slot = False

tests/opentelemetry/test_metrics.py

Lines changed: 0 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -189,19 +189,6 @@ async def test() -> None:
189189
"tests.opentelemetry.taskiq_test_tasks:task_add",
190190
)
191191

192-
def test_prefetch_queue_counter(self) -> None:
193-
middleware = next(
194-
m for m in broker.middlewares if isinstance(m, OpenTelemetryMiddleware)
195-
)
196-
middleware.on_prefetch_queue_add()
197-
middleware.on_prefetch_queue_add()
198-
middleware.on_prefetch_queue_add()
199-
middleware.on_prefetch_queue_remove()
200-
201-
points = self._get_data_points("worker_prefetched_tasks")
202-
self.assertEqual(len(points), 1)
203-
self.assertEqual(points[0].value, 2)
204-
205192
def test_worker_resource_metrics_when_worker_process(self) -> None:
206193
middleware = next(
207194
m for m in broker.middlewares if isinstance(m, OpenTelemetryMiddleware)

tests/receiver/receiver_listener_support.py

Lines changed: 1 addition & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -3,28 +3,11 @@
33
from typing import Literal, cast
44

55
from taskiq.abc.broker import AckableMessage, AsyncBroker
6-
from taskiq.abc.middleware import TaskiqMiddleware
76
from taskiq.message import BrokerMessage, TaskiqMessage
87

98

109
class ReceiverLifecycleError(RuntimeError):
11-
"""Marker error for listener and middleware lifecycle failures."""
12-
13-
14-
class PrefetchCounterMiddleware(TaskiqMiddleware):
15-
"""Track the observable number of messages in the prefetch queue."""
16-
17-
def __init__(self) -> None:
18-
super().__init__()
19-
self.queued_messages = 0
20-
21-
def on_prefetch_queue_add(self) -> None:
22-
"""Record one queued delivery."""
23-
self.queued_messages += 1
24-
25-
def on_prefetch_queue_remove(self) -> None:
26-
"""Record one removed or discarded delivery."""
27-
self.queued_messages -= 1
10+
"""Marker error for listener lifecycle failures."""
2811

2912

3013
class ObservedSemaphore(asyncio.Semaphore):

0 commit comments

Comments
 (0)