mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
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:
parent
034998be71
commit
7c1186e880
4 changed files with 280 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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}]
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue