11import asyncio
2+ import contextlib
23import contextvars
34import functools
45import inspect
1011from enum import Enum , auto
1112from logging import getLogger
1213from time import time
13- from typing import Any , Literal , get_type_hints
14+ from typing import Any , get_type_hints
1415
1516import anyio
1617from 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
0 commit comments