diff --git a/litellm/router.py b/litellm/router.py index dcc9870565a..d5fb6457199 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4347,22 +4347,15 @@ class Router: raise e 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) - 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( - model=model, parent_otel_span=parent_otel_span - ) - if await self.scheduler.poll( - id=item.request_id, model_name=model, health_deployments=healthy_deployments - ): - return - await asyncio.sleep(self.scheduler.polling_interval) - finally: - await self.scheduler.remove_request(request_id=item.request_id, model_name=model) - raise litellm.Timeout(message="Request timed out while polling queue", model=model, llm_provider="openai") + async def healthy_deployments() -> Sequence[object]: + deployments, _ = await self._async_get_healthy_deployments(model=model, parent_otel_span=parent_otel_span) + return deployments + + await self.scheduler.wait_for_turn( + request=FlowItem(priority=priority, request_id=str(uuid.uuid4()), model_name=model), + timeout=self.timeout, + get_healthy_deployments=healthy_deployments, + ) def _is_prompt_management_model(self, model: str) -> bool: model_list: Final = self.get_model_list(model_name=model) diff --git a/litellm/scheduler.py b/litellm/scheduler.py index 942bd5ea395..66976310485 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -1,5 +1,8 @@ +import asyncio import enum import heapq +import time +from collections.abc import Awaitable, Callable, Sequence from typing import Final from pydantic import BaseModel @@ -7,6 +10,7 @@ from pydantic import BaseModel from litellm import print_verbose from litellm.caching.caching import DualCache, RedisCache from litellm.constants import DEFAULT_IN_MEMORY_TTL, DEFAULT_POLLING_INTERVAL +from litellm.exceptions import Timeout class SchedulerCacheKeys(enum.Enum): @@ -49,7 +53,7 @@ class Scheduler: # save the queue await self.save_queue(queue=queue, model_name=request.model_name) - async def poll(self, id: str, model_name: str, health_deployments: list) -> bool: + async def poll(self, id: str, model_name: str, health_deployments: Sequence[object]) -> bool: """ Return if request can be processed. @@ -78,6 +82,27 @@ class Scheduler: print_verbose(f"Popped id: {id}") return True + async def wait_for_turn( + self, + request: FlowItem, + timeout: float, + get_healthy_deployments: Callable[[], Awaitable[Sequence[object]]], + ) -> None: + try: + await self.add_request(request=request) + end_time: Final = time.monotonic() + timeout + while time.monotonic() < end_time: + if await self.poll( + id=request.request_id, + model_name=request.model_name, + health_deployments=await get_healthy_deployments(), + ): + return + await asyncio.sleep(self.polling_interval) + finally: + await 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: """ Remove a specific request from the priority queue for a model. diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index df466752951..b56b32b078d 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, Scheduler +from litellm.scheduler import FlowItem from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy @@ -17922,29 +17922,3 @@ async def test_prioritized_request_leaves_queue_when_it_stops_waiting(stop_waiti 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 index bf088642604..6b4c5f2fe80 100644 --- a/tests/test_litellm/test_scheduler.py +++ b/tests/test_litellm/test_scheduler.py @@ -1,3 +1,5 @@ +import asyncio +from collections.abc import Sequence from typing import Final import pytest @@ -23,3 +25,38 @@ async def test_poll_during_cooldown_admits_only_the_head_of_the_queue(): 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")] + + +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() + + +async def _no_healthy_deployments() -> Sequence[object]: + return () + + +@pytest.mark.asyncio +async def test_wait_for_turn_removes_entry_when_cancelled_mid_enqueue(): + scheduler: Final = _PausesAfterEnqueueScheduler() + 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 scheduler.enqueued.wait() + assert await scheduler.get_queue("sched-model") == [(1, "cancelled")] + + waiting.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == []