mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(routing_groups): build every group selector before swapping router state
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
d835aedddb
commit
eac63695d5
2 changed files with 45 additions and 10 deletions
|
|
@ -28,6 +28,7 @@ import anyio
|
|||
import httpx
|
||||
import openai
|
||||
from openai import AsyncOpenAI
|
||||
from pydantic import ValidationError
|
||||
from typing_extensions import overload
|
||||
|
||||
import litellm
|
||||
|
|
@ -1030,6 +1031,12 @@ class Router:
|
|||
known_model_names=known_model_names,
|
||||
)
|
||||
|
||||
built: Final = tuple((group, self._try_build_group_selector(group)) for group in groups)
|
||||
failures: Final = tuple(outcome for _, outcome in built if isinstance(outcome, Exception))
|
||||
if failures:
|
||||
self._unregister_router_selectors([outcome for _, outcome in built if not isinstance(outcome, Exception)])
|
||||
raise failures[0]
|
||||
|
||||
self._unregister_router_selectors(
|
||||
[sel for selectors in getattr(self, "_group_selectors", {}).values() for sel in selectors.values()]
|
||||
)
|
||||
|
|
@ -1040,19 +1047,20 @@ class Router:
|
|||
}
|
||||
self._group_selectors: dict[str, dict[str, Any]] = {
|
||||
group.group_name: (
|
||||
{}
|
||||
if (
|
||||
selector := self._build_strategy_selector(
|
||||
strategy=group.routing_strategy,
|
||||
routing_strategy_args=group.routing_strategy_args or {},
|
||||
)
|
||||
)
|
||||
is None
|
||||
else {self._normalize_strategy(group.routing_strategy) or "": selector}
|
||||
{} if selector is None else {self._normalize_strategy(group.routing_strategy) or "": selector}
|
||||
)
|
||||
for group in groups
|
||||
for group, selector in built
|
||||
}
|
||||
|
||||
def _try_build_group_selector(self, group: RoutingGroup) -> object | None:
|
||||
try:
|
||||
return self._build_strategy_selector(
|
||||
strategy=group.routing_strategy,
|
||||
routing_strategy_args=group.routing_strategy_args or {},
|
||||
)
|
||||
except ValidationError as build_error:
|
||||
return build_error
|
||||
|
||||
_OVERRIDABLE_ROUTING_STRATEGIES: frozenset[str] = frozenset({"simple-shuffle", *_DEFAULT_SELECTOR_ATTR_BY_STRATEGY})
|
||||
|
||||
def _get_request_routing_strategy_override(self, request_kwargs: dict | None) -> str | None:
|
||||
|
|
|
|||
|
|
@ -788,3 +788,30 @@ def test_invalid_group_strategy_does_not_leak_a_selector():
|
|||
|
||||
assert list(router._routing_groups) == ["g1"]
|
||||
assert any(id(cb) == id(selector) for cb in litellm.callbacks)
|
||||
|
||||
|
||||
def test_unbuildable_group_selector_keeps_previous_groups():
|
||||
router = _build_router(
|
||||
routing_groups=[
|
||||
{"group_name": "g1", "models": ["filtered-model"], "routing_strategy": "latency-based-routing"},
|
||||
],
|
||||
)
|
||||
selector = router._group_selectors["g1"]["latency-based-routing"]
|
||||
|
||||
with pytest.raises(Exception):
|
||||
router.update_settings(
|
||||
routing_groups=[
|
||||
{"group_name": "g1", "models": ["filtered-model"], "routing_strategy": "latency-based-routing"},
|
||||
{
|
||||
"group_name": "g2",
|
||||
"models": ["other-model"],
|
||||
"routing_strategy": "latency-based-routing",
|
||||
"routing_strategy_args": {"ttl": "not-a-number"},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
assert list(router._routing_groups) == ["g1"]
|
||||
assert router._model_to_group == {"filtered-model": "g1"}
|
||||
assert router._group_selectors["g1"]["latency-based-routing"] is selector
|
||||
assert sum(1 for cb in litellm.callbacks if type(cb) is type(selector)) == 1
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue