mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
feat(router): support percentile-based TTFT routing (#40352)
* feat(router): support percentile-based TTFT routing * fix(router): apply routing_strategy_args updates to the live selector Runtime routing_strategy_args updates (config reload, update_settings) only rebuilt the strategy selector when routing_strategy itself changed, so a newly added ttft_percentile sat unused until the proxy restarted. Also drops a comment that only restated the code it sat above. Claude-Session: https://claude.ai/code/session_01PmqjhFYcUh6vA72d8W9gdB * refactor(router): drop unreachable empty-samples guard in percentile latency _percentile_latency is only called behind use_ttft, which already requires a non-empty ttft sample list, so the early return was dead code and the one line Codecov flagged as uncovered on this patch. Claude-Session: https://claude.ai/code/session_01PmqjhFYcUh6vA72d8W9gdB * test(router): cover the no-selector path of a routing_strategy_args update simple-shuffle has no selector attribute to re-link, so the early return guards a setattr with a None attribute name. Dropping the guard makes the new test fail with "attribute name must be string, not 'NoneType'". Claude-Session: https://claude.ai/code/session_01PmqjhFYcUh6vA72d8W9gdB * fix(test): assert ValidationError on out-of-range ttft_percentile pytest.raises(ValueError) tripped PT011 for being too broad. Pydantic raises ValidationError for the gt/le constraint, so naming it satisfies the rule and pins the assertion to the constraint under test. Claude-Session: https://claude.ai/code/session_01PmqjhFYcUh6vA72d8W9gdB * fix(router): drop Final from a per-deployment loop variable basedpyright rejects "A Final variable cannot be assigned within a loop", which pushed reportGeneralTypeIssues one over its budget. selected_latency is rebound each iteration, so it matches its unannotated neighbours in the same loop. Claude-Session: https://claude.ai/code/session_01PmqjhFYcUh6vA72d8W9gdB * test(router): exempt _apply_updated_routing_strategy_args from the name scan The scan only reads test files with "router" in the filename, so it cannot see the update_settings tests in router_strategy/test_lowest_latency.py. Calling the private helper directly would test structure rather than behaviour, so it joins the existing entries ignored for the same reason. Claude-Session: https://claude.ai/code/session_01PmqjhFYcUh6vA72d8W9gdB
This commit is contained in:
parent
7c6e33ef70
commit
699ae63b2a
4 changed files with 225 additions and 18 deletions
|
|
@ -1317,6 +1317,43 @@ class Router:
|
|||
if isinstance(litellm.input_callback, list):
|
||||
litellm.input_callback = [c for c in litellm.input_callback if id(c) not in selector_ids]
|
||||
|
||||
def _apply_updated_routing_strategy_args(self) -> None:
|
||||
"""
|
||||
Re-link the default group's selector to the current `routing_strategy_args`.
|
||||
|
||||
Selectors freeze their `RoutingArgs` at construction, so a runtime args
|
||||
update would otherwise keep serving the boot-time values until restart.
|
||||
Latency/usage state survives the rebuild: it lives in the shared router
|
||||
cache, not on the selector.
|
||||
"""
|
||||
strategy: Final = self._normalize_strategy(self.routing_strategy)
|
||||
if strategy == "lar1":
|
||||
from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
|
||||
|
||||
apply_lar1_routing_strategy(self, self.routing_strategy_args)
|
||||
return
|
||||
|
||||
attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy or "")
|
||||
current: Final = getattr(self, attr, None) if attr is not None else None
|
||||
if attr is None or current is None:
|
||||
return
|
||||
|
||||
try:
|
||||
rebuilt: Final = self._build_strategy_selector(
|
||||
strategy=strategy or "",
|
||||
routing_strategy_args=self.routing_strategy_args,
|
||||
)
|
||||
except (TypeError, ValidationError):
|
||||
verbose_router_logger.exception(
|
||||
"Invalid routing_strategy_args %s for '%s'; keeping the previous ones",
|
||||
self.routing_strategy_args,
|
||||
strategy,
|
||||
)
|
||||
return
|
||||
|
||||
self._unregister_router_selectors((current,))
|
||||
setattr(self, attr, rebuilt)
|
||||
|
||||
def routing_strategy_init(self, routing_strategy: RoutingStrategy | str, routing_strategy_args: dict):
|
||||
verbose_router_logger.info("Routing strategy: %s", routing_strategy)
|
||||
self._validate_routing_strategy(routing_strategy)
|
||||
|
|
@ -11847,7 +11884,7 @@ class Router:
|
|||
|
||||
_existing_router_settings: Final = self.get_settings()
|
||||
rebuild_routing_groups = False
|
||||
relink_lar1_from_args = False
|
||||
routing_args_updated = False
|
||||
for var in kwargs:
|
||||
if var in RUNTIME_UPDATABLE_ROUTER_SETTINGS:
|
||||
if var in _int_settings:
|
||||
|
|
@ -11886,15 +11923,13 @@ class Router:
|
|||
)
|
||||
rebuild_routing_groups = True
|
||||
elif var == "routing_strategy_args":
|
||||
relink_lar1_from_args = True
|
||||
routing_args_updated = True
|
||||
setattr(self, var, value)
|
||||
else:
|
||||
verbose_router_logger.debug("Setting %s is not allowed", var)
|
||||
|
||||
if relink_lar1_from_args and self._normalize_strategy(self.routing_strategy) == "lar1":
|
||||
from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
|
||||
|
||||
apply_lar1_routing_strategy(self, self.routing_strategy_args)
|
||||
if routing_args_updated:
|
||||
self._apply_updated_routing_strategy_args()
|
||||
|
||||
if rebuild_routing_groups:
|
||||
self._init_routing_groups(self._routing_groups_input)
|
||||
|
|
|
|||
|
|
@ -3,8 +3,11 @@
|
|||
import random
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime, timedelta
|
||||
from math import ceil
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
import litellm
|
||||
from litellm import ModelResponse, token_counter, verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -24,6 +27,7 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
|
|||
ttl: float = 1 * 60 * 60 # 1 hour
|
||||
lowest_latency_buffer: float = 0
|
||||
max_latency_list_size: int = 10
|
||||
ttft_percentile: float | None = Field(default=None, gt=0, le=1)
|
||||
|
||||
|
||||
def _average_latency(samples: Sequence[float]) -> float:
|
||||
|
|
@ -32,6 +36,12 @@ def _average_latency(samples: Sequence[float]) -> float:
|
|||
return sum(samples) / len(samples)
|
||||
|
||||
|
||||
def _percentile_latency(samples: Sequence[float], percentile: float) -> float:
|
||||
values: Final = sorted(samples)
|
||||
index: Final = ceil(len(values) * percentile) - 1
|
||||
return values[index]
|
||||
|
||||
|
||||
def _ttft_seconds(elapsed: timedelta | float) -> float:
|
||||
if isinstance(elapsed, timedelta):
|
||||
return elapsed.total_seconds()
|
||||
|
|
@ -427,14 +437,17 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
item_rpm = item_map.get(precise_minute, {}).get("rpm", 0)
|
||||
item_tpm = item_map.get(precise_minute, {}).get("tpm", 0)
|
||||
|
||||
# get average latency or average ttft (depending on streaming/non-streaming)
|
||||
use_ttft = (
|
||||
request_kwargs is not None
|
||||
and request_kwargs.get("stream", None) is not None
|
||||
and request_kwargs["stream"] is True
|
||||
and len(item_ttft_latency) > 0
|
||||
)
|
||||
average_latency = _average_latency(item_ttft_latency if use_ttft else item_latency)
|
||||
selected_latency = (
|
||||
_percentile_latency(item_ttft_latency, self.routing_args.ttft_percentile)
|
||||
if use_ttft and self.routing_args.ttft_percentile is not None
|
||||
else _average_latency(item_ttft_latency if use_ttft else item_latency)
|
||||
)
|
||||
|
||||
# -------------- #
|
||||
# Debugging Logic
|
||||
|
|
@ -443,7 +456,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
# this helps a user to debug why the router picked a specfic deployment #
|
||||
_deployment_api_base = _deployment.get("litellm_params", {}).get("api_base", "")
|
||||
if _deployment_api_base is not None:
|
||||
_latency_per_deployment[_deployment_api_base] = average_latency
|
||||
_latency_per_deployment[_deployment_api_base] = selected_latency
|
||||
# -------------- #
|
||||
# End of Debugging Logic
|
||||
# -------------- #
|
||||
|
|
@ -453,7 +466,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
): # if user passed in tpm / rpm in the model_list
|
||||
continue
|
||||
else:
|
||||
potential_deployments.append((_deployment, average_latency))
|
||||
potential_deployments.append((_deployment, selected_latency))
|
||||
|
||||
if len(potential_deployments) == 0:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -87,6 +87,7 @@ ignored_function_names = [
|
|||
"_delete_claude_code_session_router_binding", # Tested through Redis cleanup failure in test_router.py
|
||||
"_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)
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8,11 +8,12 @@ import json
|
|||
from datetime import datetime, timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
|
||||
from litellm.router import Router
|
||||
from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler, RoutingArgs
|
||||
|
||||
DEPLOYMENT_ID = "9876"
|
||||
KWARGS = {
|
||||
|
|
@ -58,9 +59,9 @@ def test_sync_embedding_latency_is_json_serializable():
|
|||
|
||||
latencies = _recorded_latencies(cache)
|
||||
assert latencies, "expected a latency entry to be recorded"
|
||||
assert all(
|
||||
not isinstance(value, timedelta) for value in latencies
|
||||
), f"raw timedelta leaked into latency list: {latencies}"
|
||||
assert all(not isinstance(value, timedelta) for value in latencies), (
|
||||
f"raw timedelta leaked into latency list: {latencies}"
|
||||
)
|
||||
assert latencies[-1] == pytest.approx(2.0)
|
||||
# the exact failure mode from production: redis cache sync json.dumps
|
||||
json.dumps({"latency": latencies})
|
||||
|
|
@ -84,9 +85,9 @@ async def test_async_embedding_latency_is_json_serializable():
|
|||
|
||||
latencies = _recorded_latencies(cache)
|
||||
assert latencies, "expected a latency entry to be recorded"
|
||||
assert all(
|
||||
not isinstance(value, timedelta) for value in latencies
|
||||
), f"raw timedelta leaked into latency list: {latencies}"
|
||||
assert all(not isinstance(value, timedelta) for value in latencies), (
|
||||
f"raw timedelta leaked into latency list: {latencies}"
|
||||
)
|
||||
assert latencies[-1] == pytest.approx(3.0)
|
||||
json.dumps({"latency": latencies})
|
||||
|
||||
|
|
@ -292,6 +293,85 @@ async def test_streaming_routing_ignores_per_token_ttft_samples_from_older_worke
|
|||
assert picked["model_info"]["id"] == FAST_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sync_mode", [True, False], ids=["sync", "async"])
|
||||
@pytest.mark.parametrize(
|
||||
("ttft_percentile", "first_samples", "second_samples", "expected_id"),
|
||||
[
|
||||
(None, [0.1, 0.1, 1.0], [0.3, 0.3, 0.3], SLOW_TTFT_ID),
|
||||
(0.5, [0.1, 0.1, 1.0], [0.3, 0.3, 0.3], FAST_TTFT_ID),
|
||||
(0.9, [0.1, 0.1, 0.1, 0.1, 1.5], [0.3, 0.3, 0.3, 0.3, 0.3], SLOW_TTFT_ID),
|
||||
],
|
||||
ids=["default_average", "p50", "p90"],
|
||||
)
|
||||
async def test_streaming_ttft_ranking_percentile(
|
||||
sync_mode: bool,
|
||||
ttft_percentile: float | None,
|
||||
first_samples: list[float],
|
||||
second_samples: list[float],
|
||||
expected_id: str,
|
||||
):
|
||||
cache = DualCache()
|
||||
routing_args = {} if ttft_percentile is None else {"ttft_percentile": ttft_percentile}
|
||||
handler = LowestLatencyLoggingHandler(router_cache=cache, routing_args=routing_args)
|
||||
cache.set_cache(
|
||||
key=f"{MODEL_GROUP}_map",
|
||||
value={
|
||||
FAST_TTFT_ID: {"time_to_first_token_seconds": first_samples},
|
||||
SLOW_TTFT_ID: {"time_to_first_token_seconds": second_samples},
|
||||
},
|
||||
)
|
||||
|
||||
if sync_mode:
|
||||
picked = handler.get_available_deployments(
|
||||
model_group=MODEL_GROUP,
|
||||
healthy_deployments=STREAMING_DEPLOYMENTS,
|
||||
request_kwargs={"stream": True, "metadata": {}},
|
||||
)
|
||||
else:
|
||||
picked = await handler.async_get_available_deployments(
|
||||
model_group=MODEL_GROUP,
|
||||
healthy_deployments=STREAMING_DEPLOYMENTS,
|
||||
request_kwargs={"stream": True, "metadata": {}},
|
||||
)
|
||||
|
||||
assert picked is not None
|
||||
assert picked["model_info"]["id"] == expected_id
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ttft_percentile", [0, -0.1, 1.1])
|
||||
def test_ttft_percentile_validation(ttft_percentile: float):
|
||||
with pytest.raises(ValidationError):
|
||||
RoutingArgs(ttft_percentile=ttft_percentile)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ttft_percentile", [0.5, 0.9, 0.95, 1.0])
|
||||
def test_ttft_percentile_accepts_valid_values(ttft_percentile: float):
|
||||
assert RoutingArgs(ttft_percentile=ttft_percentile).ttft_percentile == ttft_percentile
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ttft_percentile_does_not_change_non_streaming_routing():
|
||||
cache = DualCache()
|
||||
handler = LowestLatencyLoggingHandler(router_cache=cache, routing_args={"ttft_percentile": 0.9})
|
||||
cache.set_cache(
|
||||
key=f"{MODEL_GROUP}_map",
|
||||
value={
|
||||
FAST_TTFT_ID: {"latency": [1.0], "time_to_first_token_seconds": [0.1]},
|
||||
SLOW_TTFT_ID: {"latency": [0.2], "time_to_first_token_seconds": [1.5]},
|
||||
},
|
||||
)
|
||||
|
||||
picked = await handler.async_get_available_deployments(
|
||||
model_group=MODEL_GROUP,
|
||||
healthy_deployments=STREAMING_DEPLOYMENTS,
|
||||
request_kwargs={"stream": False, "metadata": {}},
|
||||
)
|
||||
|
||||
assert picked is not None
|
||||
assert picked["model_info"]["id"] == SLOW_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"cached_entry",
|
||||
|
|
@ -318,3 +398,81 @@ async def test_async_get_available_deployments_treats_missing_samples_as_zero_la
|
|||
|
||||
assert picked is not None
|
||||
assert picked["model_info"]["id"] == DEPLOYMENT_ID
|
||||
|
||||
|
||||
def _latency_router(routing_strategy_args: dict) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": MODEL_GROUP,
|
||||
"litellm_params": {"model": f"openai/{MODEL_GROUP}", "api_key": "sk-fake"},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
for deployment_id in (FAST_TTFT_ID, SLOW_TTFT_ID)
|
||||
],
|
||||
routing_strategy="latency-based-routing",
|
||||
routing_strategy_args=routing_strategy_args,
|
||||
)
|
||||
|
||||
|
||||
def _seed_streaming_ttft(router: Router) -> None:
|
||||
router.cache.set_cache(
|
||||
key=f"{MODEL_GROUP}_map",
|
||||
value={
|
||||
FAST_TTFT_ID: {"time_to_first_token_seconds": [0.1, 0.1, 1.0]},
|
||||
SLOW_TTFT_ID: {"time_to_first_token_seconds": [0.3, 0.3, 0.3]},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _pick_streaming(router: Router) -> str:
|
||||
picked = await router.async_get_available_deployment(
|
||||
model=MODEL_GROUP,
|
||||
request_kwargs={"stream": True, "metadata": {}},
|
||||
)
|
||||
return picked["model_info"]["id"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_routing_strategy_args_update_applies_ttft_percentile():
|
||||
"""A config reload that adds ttft_percentile must reach the live selector,
|
||||
not sit unused until the proxy restarts."""
|
||||
router = _latency_router({"max_latency_list_size": 50})
|
||||
_seed_streaming_ttft(router)
|
||||
|
||||
assert await _pick_streaming(router) == SLOW_TTFT_ID
|
||||
|
||||
router.update_settings(routing_strategy_args={"max_latency_list_size": 50, "ttft_percentile": 0.5})
|
||||
|
||||
assert await _pick_streaming(router) == FAST_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_routing_strategy_args_update_keeps_previous_args_when_invalid():
|
||||
router = _latency_router({"ttft_percentile": 0.5})
|
||||
_seed_streaming_ttft(router)
|
||||
|
||||
router.update_settings(routing_strategy_args={"ttft_percentile": 5})
|
||||
|
||||
assert await _pick_streaming(router) == FAST_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_routing_strategy_args_update_is_a_noop_without_a_selector():
|
||||
"""simple-shuffle has no selector to re-link, so an args update must leave
|
||||
the router alone instead of blowing up on a missing selector attribute."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": MODEL_GROUP,
|
||||
"litellm_params": {"model": f"openai/{MODEL_GROUP}", "api_key": "sk-fake"},
|
||||
"model_info": {"id": FAST_TTFT_ID},
|
||||
}
|
||||
],
|
||||
routing_strategy="simple-shuffle",
|
||||
)
|
||||
|
||||
router.update_settings(routing_strategy_args={"ttl": 5})
|
||||
|
||||
assert router.routing_strategy_args == {"ttl": 5}
|
||||
assert await _pick_streaming(router) == FAST_TTFT_ID
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue