From 0c9ab36172fe7e9d3e4f70c2dd025afcb161ebb8 Mon Sep 17 00:00:00 2001 From: RachelHuangZW Date: Tue, 29 Sep 2026 13:59:23 -0400 Subject: [PATCH] fix(scheduler): finish the queue removal when cancellation is delivered again during cleanup A second cancel, or an anyio cancel scope that re-cancels on every await, interrupted remove_request mid-write and left the entry in Redis. --- litellm/scheduler.py | 2 +- tests/unit/test_scheduler.py | 40 ++++++++++++++++++++++++++++++++++++ 2 files changed, 41 insertions(+), 1 deletion(-) diff --git a/litellm/scheduler.py b/litellm/scheduler.py index 66976310485..cffed3acce5 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -100,7 +100,7 @@ class Scheduler: return await asyncio.sleep(self.polling_interval) finally: - await self.remove_request(request_id=request.request_id, model_name=request.model_name) + await asyncio.shield(self.remove_request(request_id=request.request_id, model_name=request.model_name)) raise Timeout(message="Request timed out while polling queue", model=request.model_name, llm_provider="openai") async def remove_request(self, request_id: str, model_name: str) -> None: diff --git a/tests/unit/test_scheduler.py b/tests/unit/test_scheduler.py index 6b4c5f2fe80..2b643558089 100644 --- a/tests/unit/test_scheduler.py +++ b/tests/unit/test_scheduler.py @@ -60,3 +60,43 @@ async def test_wait_for_turn_removes_entry_when_cancelled_mid_enqueue(): await waiting assert await scheduler.get_queue("sched-model") == [] + + +class _HeldRemovalScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.removing: Final = asyncio.Event() + self.finish_removal: Final = asyncio.Event() + + async def remove_request(self, request_id: str, model_name: str) -> None: + self.removing.set() + await self.finish_removal.wait() + await super().remove_request(request_id=request_id, model_name=model_name) + + +@pytest.mark.asyncio +async def test_wait_for_turn_finishes_removal_when_cancelled_again_during_cleanup(): + scheduler: Final = _HeldRemovalScheduler() + await scheduler.add_request(FlowItem(priority=0, request_id="head", model_name="sched-model")) + polling: Final = asyncio.Event() + + async def no_healthy_deployments() -> Sequence[object]: + polling.set() + return () + + waiting: Final = asyncio.create_task( + scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="cancelled", model_name="sched-model"), + timeout=5, + get_healthy_deployments=no_healthy_deployments, + ) + ) + await polling.wait() + waiting.cancel() + await scheduler.removing.wait() + waiting.cancel() + scheduler.finish_removal.set() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [(0, "head")]