diff --git a/litellm/router.py b/litellm/router.py index 7a4da72a2b8..dcc9870565a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4348,8 +4348,8 @@ class Router: async def _wait_for_scheduler_turn(self, model: str, priority: int, parent_otel_span: Span | None) -> None: item: Final = FlowItem(priority=priority, request_id=str(uuid.uuid4()), model_name=model) - await self.scheduler.add_request(request=item) try: + await self.scheduler.add_request(request=item) end_time: Final = time.monotonic() + self.timeout while time.monotonic() < end_time: healthy_deployments, _ = await self._async_get_healthy_deployments( diff --git a/litellm/scheduler.py b/litellm/scheduler.py index 028b5d085e2..942bd5ea395 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -61,27 +61,21 @@ class Scheduler: * If no healthy deployments available * AND request not at the top of queue """ + print_verbose(f"len(health_deployments): {len(health_deployments)}") + if len(health_deployments) > 0: + return True + queue: Final = await self.get_queue(model_name=model_name) if not queue: raise Exception(f"Incorrectly setup. Queue is invalid. Queue={queue}") - # ------------ - # Setup values - # ------------ - - print_verbose(f"len(health_deployments): {len(health_deployments)}") - if len(health_deployments) == 0: - print_verbose(f"queue: {queue}, seeking id={id}") - # Check if the id is at the top of the heap - if queue[0][1] == id: - # Remove the item from the queue - heapq.heappop(queue) - await self.save_queue(queue=queue, model_name=model_name) - print_verbose(f"Popped id: {id}") - return True - else: - return False + print_verbose(f"queue: {queue}, seeking id={id}") + if queue[0][1] != id: + return False + heapq.heappop(queue) + await self.save_queue(queue=queue, model_name=model_name) + print_verbose(f"Popped id: {id}") return True async def remove_request(self, request_id: str, model_name: str) -> None: diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 69f499a020c..df466752951 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -50,7 +50,7 @@ from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute -from litellm.scheduler import FlowItem +from litellm.scheduler import FlowItem, Scheduler from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy @@ -17921,3 +17921,30 @@ async def test_prioritized_request_leaves_queue_when_it_stops_waiting(stop_waiti await waiting assert await router.scheduler.get_queue("sched-model") == [(0, "head-of-queue")] + + +class _PausesAfterEnqueueScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.enqueued: Final = asyncio.Event() + + async def add_request(self, request: FlowItem) -> None: + await super().add_request(request) + self.enqueued.set() + await asyncio.Event().wait() + + +@pytest.mark.asyncio +async def test_prioritized_request_cancelled_while_enqueueing_leaves_queue(): + router: Final = _scheduled_router(timeout=5) + scheduler: Final = _PausesAfterEnqueueScheduler() + router.scheduler = scheduler + waiting: Final = asyncio.create_task(_send_scheduled_chat(router, 1)) + await scheduler.enqueued.wait() + assert len(await scheduler.get_queue("sched-model")) == 1 + + waiting.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [] diff --git a/tests/test_litellm/test_scheduler.py b/tests/test_litellm/test_scheduler.py new file mode 100644 index 00000000000..bf088642604 --- /dev/null +++ b/tests/test_litellm/test_scheduler.py @@ -0,0 +1,25 @@ +from typing import Final + +import pytest + +from litellm.scheduler import FlowItem, Scheduler + + +@pytest.mark.asyncio +async def test_poll_admits_request_missing_from_queue_while_a_deployment_is_healthy(): + scheduler: Final = Scheduler() + + assert await scheduler.poll( + id="erased-by-concurrent-write", model_name="sched-model", health_deployments=[{"model_info": {"id": "a"}}] + ) + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_admits_only_the_head_of_the_queue(): + scheduler: Final = Scheduler() + await scheduler.add_request(FlowItem(priority=2, request_id="later", model_name="sched-model")) + await scheduler.add_request(FlowItem(priority=1, request_id="head", model_name="sched-model")) + + assert not await scheduler.poll(id="later", model_name="sched-model", health_deployments=[]) + assert await scheduler.poll(id="head", model_name="sched-model", health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [(2, "later")]