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:
Clement 2026-09-10 01:47:34 +08:00 committed by GitHub
parent 7c6e33ef70
commit 699ae63b2a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 225 additions and 18 deletions

View file

@ -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)

View file

@ -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

View file

@ -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)
]

View file

@ -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