mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge fe169152a9 into 9e9c29f404
This commit is contained in:
commit
b0063978b6
5 changed files with 228 additions and 118 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
102
tests/unit/test_scheduler.py
Normal file
102
tests/unit/test_scheduler.py
Normal file
|
|
@ -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")]
|
||||
Loading…
Add table
Reference in a new issue