diff --git a/litellm/router.py b/litellm/router.py index b93a4abdf03..5ce49e7becb 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4419,57 +4419,17 @@ class Router: stream=False, **kwargs, ): - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### - - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) - - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) - _response._hidden_params.setdefault("additional_headers", {}) - _response._hidden_params["additional_headers"].update({"x-litellm-request-prioritization-used": True}) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + try: + _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) + _response._hidden_params.setdefault("additional_headers", {}) + _response._hidden_params["additional_headers"].update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e async def _schedule_factory( self, @@ -4479,60 +4439,29 @@ class Router: args: tuple[object, ...], kwargs: dict[str, object], ): - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### + try: + _response: Final = await original_function(*args, **kwargs) + if isinstance(_response._hidden_params, dict): + _response._hidden_params.setdefault("additional_headers", {}) + _response._hidden_params["additional_headers"].update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) + async def _wait_for_scheduler_turn(self, model: str, priority: int, parent_otel_span: Span | None) -> None: + async def healthy_deployments() -> Sequence[object]: + deployments, _ = await self._async_get_healthy_deployments(model=model, parent_otel_span=parent_otel_span) + return deployments - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await original_function(*args, **kwargs) - if isinstance(_response._hidden_params, dict): - _response._hidden_params.setdefault("additional_headers", {}) - _response._hidden_params["additional_headers"].update( - {"x-litellm-request-prioritization-used": True} - ) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + 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 e19e386f527..8cc5c08e7d7 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 @@ -8,6 +11,7 @@ from litellm import print_verbose from litellm._internal_context import with_service_target from litellm.caching.caching import DualCache, RedisCache from litellm.constants import DEFAULT_IN_MEMORY_TTL, DEFAULT_POLLING_INTERVAL +from litellm.exceptions import Timeout SCHEDULER_QUEUE_TARGET: Final = "scheduler_queue" @@ -52,7 +56,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. @@ -64,29 +68,44 @@ 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 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 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: """ Remove a specific request from the priority queue for a model. diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index f55f415b76c..b9237c4e1e4 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -90,6 +90,7 @@ ignored_function_names = [ "_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name) + "_wait_for_scheduler_turn", # Tested through prioritized acompletion and atext_completion in test_router.py "arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name) "_async_get_available_deployment", # Body of the `route {model}` phase wrapper, exercised through async_get_available_deployment in test_router.py "_async_get_available_deployment_for_pass_through", # Same, through async_get_available_deployment_for_pass_through in test_router.py diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 924375a18b7..e241d1a8b74 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -57,6 +57,7 @@ from litellm.router_utils.client_initalization_utils import MaxParallelRequestsL from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute +from litellm.scheduler import FlowItem from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy @@ -18683,6 +18684,64 @@ def test_access_windows_filter_reserved_deployments_method(): ] == ["reserved-deployment", "open-deployment"] +def _scheduled_router(timeout: float) -> Router: + return Router( + model_list=[ + { + "model_name": "sched-model", + "litellm_params": {"model": "openai/sched-model", "api_key": "sk-fake", "mock_response": "hi"}, + "model_info": {"id": "sched-deployment"}, + } + ], + timeout=timeout, + ) + + +async def _send_scheduled_chat(router: Router, priority: int) -> object: + return await router.acompletion( + model="sched-model", messages=[{"role": "user", "content": "hi"}], priority=priority + ) + + +async def _send_scheduled_text(router: Router, priority: int) -> object: + return await router.atext_completion(model="sched-model", prompt="hi", priority=priority) + + +@pytest.mark.parametrize( + "send", [_send_scheduled_chat, _send_scheduled_text], ids=["schedule_acompletion", "schedule_factory"] +) +@pytest.mark.asyncio +async def test_admitted_prioritized_request_does_not_block_later_request_during_cooldown( + send: Callable[[Router, int], Awaitable[object]], +): + from litellm.types.router import RouterRateLimitError + + router: Final = _scheduled_router(timeout=1) + await send(router, 1) + _cool_down(router, "sched-deployment") + + with pytest.raises(RouterRateLimitError, match="cooldown"): + await send(router, 2) + + +@pytest.mark.parametrize("stop_waiting", ["cancel", "timeout"]) +@pytest.mark.asyncio +async def test_prioritized_request_leaves_queue_when_it_stops_waiting(stop_waiting: Literal["cancel", "timeout"]): + router: Final = _scheduled_router(timeout=0.5) + _cool_down(router, "sched-deployment") + await router.scheduler.add_request(FlowItem(priority=0, request_id="head-of-queue", model_name="sched-model")) + waiting: Final = asyncio.create_task(_send_scheduled_chat(router, 5)) + await asyncio.sleep(0.05) + assert len(await router.scheduler.get_queue("sched-model")) == 2 + + if stop_waiting == "cancel": + waiting.cancel() + with pytest.raises(asyncio.CancelledError if stop_waiting == "cancel" else litellm.Timeout): + await waiting + + assert await router.scheduler.get_queue("sched-model") == [(0, "head-of-queue")] + + @pytest.mark.asyncio async def test_bare_model_group_served_by_wildcard_deployment_uses_provider_prefixed_fallback_key() -> None: """Claude Code sends the bare "claude-sonnet-4-6" to /v1/messages; routing serves it through the diff --git a/tests/unit/test_scheduler.py b/tests/unit/test_scheduler.py new file mode 100644 index 00000000000..2b643558089 --- /dev/null +++ b/tests/unit/test_scheduler.py @@ -0,0 +1,102 @@ +import asyncio +from collections.abc import Sequence +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")] + + +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") == [] + + +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")]