fix(router): bind per-request routing_strategy override selectors to the request's callbacks

Backport of #41178 to rc/1.102.0.
Cherry-picked from merge commit 1ca4579375 (main), originally by app/devin-ai-integration.
This commit is contained in:
mateo-berri 2026-09-18 11:10:28 -07:00
parent 034998be71
commit 7c1186e880
4 changed files with 280 additions and 2 deletions

View file

@ -633,6 +633,24 @@ class Logging(LiteLLMLoggingBaseClass):
"""Keep ``_response_ms`` / ``litellm_overhead_time_ms`` for a result that has no ``_hidden_params``."""
self.response_timing_metrics = dict(timing_metrics) # mutable-ok: kept deep-copyable
def add_dynamic_callback(self, callback: CustomLogger) -> None:
self.dynamic_input_callbacks = self._with_dynamic_callback(self.dynamic_input_callbacks, callback)
self.dynamic_success_callbacks = self._with_dynamic_callback(self.dynamic_success_callbacks, callback)
self.dynamic_async_success_callbacks = self._with_dynamic_callback(
self.dynamic_async_success_callbacks, callback
)
self.dynamic_failure_callbacks = self._with_dynamic_callback(self.dynamic_failure_callbacks, callback)
self.dynamic_async_failure_callbacks = self._with_dynamic_callback(
self.dynamic_async_failure_callbacks, callback
)
@staticmethod
def _with_dynamic_callback(
callbacks: Sequence[str | Callable | CustomLogger] | None, callback: CustomLogger
) -> list[str | Callable | CustomLogger]:
existing: Final = tuple(callbacks or ())
return [*existing, *(() if callback in existing else (callback,))]
def process_dynamic_callbacks(self):
"""
Initializes CustomLogger compatible callbacks in self.dynamic_* callbacks

View file

@ -1625,6 +1625,24 @@ class Router:
return
await selector.async_pre_call_check(deployment, parent_otel_span)
def _bind_override_selector_to_request(
self, strategy: str, selector: RouterStrategySelector | None, request_kwargs: Mapping[str, object] | None
) -> None:
if selector is None or request_kwargs is None or strategy in self._globally_registered_strategies():
return
logging_obj: Final = request_kwargs.get("litellm_logging_obj")
if isinstance(logging_obj, LiteLLMLogging):
logging_obj.add_dynamic_callback(selector)
def _globally_registered_strategies(self) -> frozenset[str]:
configured: Final = (
self.routing_strategy,
*(group.routing_strategy for group in self._routing_groups.values()),
)
return frozenset(
normalized for normalized in map(self._normalize_strategy, configured) if normalized is not None
)
def _get_routing_context(
self, model: str, request_kwargs: dict | None = None
) -> tuple[str | None, RouterStrategySelector | None]:
@ -1650,7 +1668,9 @@ class Router:
override: Final = self._get_request_routing_strategy_override(request_kwargs)
if override is not None:
verbose_router_logger.debug("routing_group=request-override model=%s strategy=%s", model, override)
return override, self._get_override_strategy_selector(override)
override_selector: Final = self._get_override_strategy_selector(override)
self._bind_override_selector_to_request(override, override_selector, request_kwargs)
return override, override_selector
group_name: Final = model if self.get_routing_group(model) is not None else self._model_to_group.get(model)
if group_name is None:

View file

@ -7004,3 +7004,21 @@ def test_get_additional_headers_survives_a_thread_growing_headers_mid_copy():
assert copied["llm_provider-x-custom-1999"] == "1999"
_run_while_a_thread_grows(headers, read, reads=300)
def test_add_dynamic_callback_registers_once_per_list_without_touching_the_callers_list(logging_obj: LitellmLogging):
callback: Final = CustomLogger()
caller_owned: Final = ["langfuse"]
logging_obj.dynamic_success_callbacks = caller_owned
logging_obj.add_dynamic_callback(callback)
logging_obj.add_dynamic_callback(callback)
assert caller_owned == ["langfuse"]
assert logging_obj.dynamic_success_callbacks == ["langfuse", callback]
assert logging_obj.dynamic_input_callbacks == [callback]
assert logging_obj.dynamic_async_success_callbacks == [callback]
assert logging_obj.dynamic_failure_callbacks == [callback]
assert logging_obj.dynamic_async_failure_callbacks == [callback]
assert LitellmLogging._with_dynamic_callback(None, callback) == [callback]
assert LitellmLogging._with_dynamic_callback((callback,), callback) == [callback]

View file

@ -5,15 +5,20 @@ the implicit `"default"` group driven by the router's top-level
`routing_strategy` / `routing_strategy_args`.
"""
import asyncio
import datetime
import time
import uuid
from collections.abc import Callable
from unittest.mock import patch
import pytest
import litellm
from litellm import Router
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.router import RoutingGroup, RoutingStrategy
from litellm.utils import Rules, function_setup
def _model_list():
@ -954,6 +959,223 @@ def test_sync_pass_through_specific_deployment_runs_the_override_pre_call_check(
assert plain["model_info"]["id"] == "deploy-3"
def _two_deployment_model_list(**d1_params: object) -> list[dict[str, object]]:
return [
{
"model_name": "grp",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test-1", "mock_response": "ok", **d1_params},
"model_info": {"id": "d1"},
},
{
"model_name": "grp",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test-2", "mock_response": "ok"},
"model_info": {"id": "d2"},
},
]
def _proxy_shaped_request(**data: object) -> dict[str, object]:
"""The proxy builds the request's `Logging` object before it hands the call to the router."""
logging_obj, kwargs = function_setup(
"acompletion",
Rules(),
datetime.datetime.now(),
litellm_call_id=str(uuid.uuid4()),
messages=[{"role": "user", "content": "hi"}],
**data,
)
return {**kwargs, "litellm_logging_obj": logging_obj}
async def _async_override_pick(router: Router, strategy: str) -> str:
deployment = await router.async_get_available_deployment(
"grp", request_kwargs=_proxy_shaped_request(model="grp", routing_strategy=strategy)
)
return deployment["model_info"]["id"]
def _sync_override_pick(router: Router, strategy: str) -> str:
deployment = router.get_available_deployment(
"grp", request_kwargs=_proxy_shaped_request(model="grp", routing_strategy=strategy)
)
return deployment["model_info"]["id"]
def _in_flight(router: Router, deployment_id: str) -> int | None:
return router.cache.get_cache(f"grp_request_count:{deployment_id}")
async def _async_wait_until(predicate: Callable[[], bool]) -> None:
for _ in range(100):
if predicate():
return
await asyncio.sleep(0.02)
raise AssertionError("lifecycle callback never reached the override selector")
def _sync_wait_until(predicate: Callable[[], bool]) -> None:
for _ in range(100):
if predicate():
return
time.sleep(0.02)
raise AssertionError("lifecycle callback never reached the override selector")
def _selector_is_not_global(selector: CustomLogger) -> bool:
global_lists = (
litellm.callbacks,
litellm.input_callback,
litellm.success_callback,
litellm.failure_callback,
litellm._async_success_callback,
litellm._async_failure_callback,
)
return not any(cb is selector for cbs in global_lists for cb in cbs)
@pytest.mark.asyncio
async def test_least_busy_override_sees_the_overriding_request_in_flight():
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle", num_retries=0)
stream = await router.acompletion(**_proxy_shaped_request(model="grp", routing_strategy="least-busy", stream=True))
busy = stream._hidden_params["model_id"]
idle = "d2" if busy == "d1" else "d1"
assert [await _async_override_pick(router, "least-busy") for _ in range(3)] == [idle, idle, idle]
async for _ in stream:
pass
await _async_wait_until(lambda: _in_flight(router, busy) == 0)
assert await _async_override_pick(router, "least-busy") == "d1"
assert _selector_is_not_global(router._override_selectors["least-busy"])
def test_sync_least_busy_override_sees_the_overriding_request_in_flight():
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle", num_retries=0)
stream = router.completion(**_proxy_shaped_request(model="grp", routing_strategy="least-busy", stream=True))
busy = stream._hidden_params["model_id"]
idle = "d2" if busy == "d1" else "d1"
assert [_sync_override_pick(router, "least-busy") for _ in range(3)] == [idle, idle, idle]
for _ in stream:
pass
_sync_wait_until(lambda: _in_flight(router, busy) == 0)
assert _sync_override_pick(router, "least-busy") == "d1"
assert _selector_is_not_global(router._override_selectors["least-busy"])
@pytest.mark.asyncio
async def test_least_busy_override_releases_the_slot_when_the_overriding_request_fails():
router = Router(
model_list=_two_deployment_model_list(mock_response="litellm.InternalServerError"),
routing_strategy="simple-shuffle",
num_retries=0,
)
with pytest.raises(litellm.InternalServerError):
await router.acompletion(**_proxy_shaped_request(model="grp", routing_strategy="least-busy"))
await _async_wait_until(lambda: _in_flight(router, "d1") == 0)
assert await _async_override_pick(router, "least-busy") == "d1"
assert _selector_is_not_global(router._override_selectors["least-busy"])
@pytest.mark.asyncio
async def test_latency_based_override_learns_from_the_overriding_requests():
router = Router(
model_list=_two_deployment_model_list(mock_delay=0.05), routing_strategy="simple-shuffle", num_retries=0
)
def samples(deployment_id: str) -> list[float]:
recorded = (router.cache.get_cache("grp_map") or {}).get(deployment_id, {}).get("latency", [])
return [latency for latency in recorded if latency > 0]
async def overriding_call() -> str:
sampled_before = {"d1": len(samples("d1")), "d2": len(samples("d2"))}
response = await router.acompletion(
**_proxy_shaped_request(model="grp", routing_strategy="latency-based-routing")
)
deployment_id = response._hidden_params["model_id"]
await _async_wait_until(lambda: len(samples(deployment_id)) > sampled_before[deployment_id])
return deployment_id
served = [await overriding_call() for _ in range(6)]
assert "d1" in served
assert served[2:] == ["d2"] * 4
assert _selector_is_not_global(router._override_selectors["latency-based-routing"])
def test_override_selector_is_bound_only_to_the_request_that_asked_for_it():
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle")
overriding = _proxy_shaped_request(model="grp", routing_strategy="least-busy")
plain = _proxy_shaped_request(model="grp")
router.get_available_deployment("grp", request_kwargs=overriding)
router.get_available_deployment("grp", request_kwargs=overriding)
router.get_available_deployment("grp", request_kwargs=plain)
selector = router._override_selectors["least-busy"]
bound = overriding["litellm_logging_obj"]
for callbacks in (
bound.dynamic_input_callbacks,
bound.dynamic_success_callbacks,
bound.dynamic_async_success_callbacks,
bound.dynamic_failure_callbacks,
bound.dynamic_async_failure_callbacks,
):
assert callbacks == [selector]
unbound = plain["litellm_logging_obj"]
assert unbound.dynamic_input_callbacks is None and unbound.dynamic_success_callbacks is None
assert unbound.dynamic_failure_callbacks is None and unbound.dynamic_async_failure_callbacks is None
def test_override_matching_the_router_strategy_is_not_bound_twice():
router = Router(model_list=_two_deployment_model_list(), routing_strategy="least-busy")
request = _proxy_shaped_request(model="grp", routing_strategy="least-busy")
router.get_available_deployment("grp", request_kwargs=request)
assert request["litellm_logging_obj"].dynamic_input_callbacks is None
@pytest.mark.asyncio
async def test_override_matching_a_routing_group_strategy_records_each_request_once():
router = Router(
model_list=_two_deployment_model_list(),
routing_strategy="simple-shuffle",
routing_groups=[RoutingGroup(group_name="lat", models=["grp"], routing_strategy="latency-based-routing")],
num_retries=0,
)
request = _proxy_shaped_request(model="grp", routing_strategy="latency-based-routing")
assert router._globally_registered_strategies() == {"simple-shuffle", "latency-based-routing"}
response = await router.acompletion(**request)
deployment_id = response._hidden_params["model_id"]
await _async_wait_until(lambda: (router.cache.get_cache("grp_map") or {}).get(deployment_id) is not None)
assert len(router.cache.get_cache("grp_map")[deployment_id]["latency"]) == 1
assert request["litellm_logging_obj"].dynamic_success_callbacks is None
def test_bind_override_selector_to_request_binds_once_and_ignores_requests_without_logging():
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle")
selector = router._get_override_strategy_selector("least-busy")
request = _proxy_shaped_request(model="grp", routing_strategy="least-busy")
request["litellm_logging_obj"].dynamic_success_callbacks = ["langfuse"]
router._bind_override_selector_to_request("least-busy", selector, request)
router._bind_override_selector_to_request("least-busy", selector, request)
router._bind_override_selector_to_request("least-busy", selector, None)
router._bind_override_selector_to_request("least-busy", selector, {"model": "grp"})
logging_obj = request["litellm_logging_obj"]
assert logging_obj.dynamic_success_callbacks == ["langfuse", selector]
assert logging_obj.dynamic_input_callbacks == [selector]
assert logging_obj.dynamic_async_failure_callbacks == [selector]
assert _selector_is_not_global(selector)
def _quality_group(strategy="latency-based-routing"):
return [{"group_name": "quality", "models": ["filtered-model", "other-model"], "routing_strategy": strategy}]