mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
refactor(scheduler): move the wait loop into Scheduler.wait_for_turn
The router passes a healthy-deployments callable into the scheduler, so tests inject a Scheduler directly instead of replacing the router's scheduler attribute. The cancelled-mid-enqueue test moves to test_scheduler.py with the other scheduler tests, and poll() takes the deployments as a Sequence since it only checks whether any are healthy
This commit is contained in:
parent
629b9ef1de
commit
f5383a87dd
4 changed files with 73 additions and 44 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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") == []
|
||||
|
|
|
|||
|
|
@ -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") == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue