@@ -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+
218286async 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
276344def 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+
300385async 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+
574678async 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