mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
feat(router): make routing groups callable as virtual models and list them in /v1/models (#36519)
* feat(router): make routing groups callable as virtual models and list them in /v1/models * fix(router): traffic-scoped cooldown exemption, live model_names on delete, group-info cache invalidation * fix(router): share one recognized-model predicate across proxy gates, resolve aliases in group cooldown, read metadata via the dual-bucket owner * fix(router): close the gate and cache families for callable groups, strip member access_groups from group rows, prove cooldown wiring end to end * refactor(router): cache materialized group rows under the model-group cache owner and drop the redundant wiring test * fix(router): warn-and-shadow on group name collisions, name-level test coverage for group helpers, faithful router doubles in a2a and cursor tests * test(router): pin group cooldown metadata across the retry path
This commit is contained in:
parent
4725cb4661
commit
06943b6468
10 changed files with 646 additions and 34 deletions
|
|
@ -233,15 +233,7 @@ async def chat_completion_pass_through_endpoint(
|
|||
# skip router if user passed their key
|
||||
if "api_key" in data:
|
||||
llm_response = asyncio.create_task(litellm.aadapter_completion(**data))
|
||||
elif llm_router is not None and data["model"] in router_model_names: # model in router model list
|
||||
llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
|
||||
elif (
|
||||
llm_router is not None
|
||||
and llm_router.model_group_alias is not None
|
||||
and data["model"] in llm_router.model_group_alias
|
||||
): # model set in model_group_alias
|
||||
llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
|
||||
elif llm_router is not None and llm_router.has_model_id(data["model"]): # model in router model list
|
||||
elif llm_router is not None and llm_router.is_recognized_model(data["model"]):
|
||||
llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
|
||||
elif (
|
||||
llm_router is not None
|
||||
|
|
|
|||
|
|
@ -121,7 +121,7 @@ def _parse_cursor_model_variant(model: str) -> _CursorModelVariant:
|
|||
def _router_can_serve(model: str, llm_router: "Router | None") -> bool:
|
||||
if llm_router is None:
|
||||
return False
|
||||
if model in llm_router.model_names or model in llm_router.model_group_alias:
|
||||
if llm_router.is_recognized_model(model):
|
||||
return True
|
||||
if model in llm_router.team_public_model_names:
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -587,16 +587,10 @@ async def route_request(
|
|||
return getattr(llm_router, f"{route_type}")(**data)
|
||||
|
||||
elif (
|
||||
(
|
||||
is_proxy_admin_without_team
|
||||
and data["model"] not in router_model_names
|
||||
and data["model"] in llm_router.team_public_model_names
|
||||
)
|
||||
or data["model"] in router_model_names
|
||||
or llm_router.has_model_id(data["model"])
|
||||
or llm_router.model_group_alias is not None
|
||||
and data["model"] in llm_router.model_group_alias
|
||||
):
|
||||
is_proxy_admin_without_team
|
||||
and data["model"] not in router_model_names
|
||||
and data["model"] in llm_router.team_public_model_names
|
||||
) or llm_router.is_recognized_model(data["model"]):
|
||||
return getattr(llm_router, f"{route_type}")(**data)
|
||||
|
||||
elif data["model"] not in router_model_names:
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ from litellm.litellm_core_utils.asyncify import run_async_function
|
|||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
coerce_token_limit,
|
||||
get_litellm_metadata_from_kwargs,
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_or_create_metadata_bucket,
|
||||
)
|
||||
|
|
@ -608,8 +609,10 @@ class Router:
|
|||
self.team_public_model_names: frozenset[str] = frozenset()
|
||||
|
||||
# Initialize cache attributes that ``_invalidate_model_group_info_cache``
|
||||
# touches *before* the first ``set_model_list`` below (which calls
|
||||
# that invalidation as part of building the model index).
|
||||
# and ``_invalidate_access_groups_cache`` touch *before* the first
|
||||
# ``set_model_list`` below (which calls those invalidations as part of
|
||||
# building the model index) and before ``_init_routing_groups(None)``
|
||||
# (which calls them on every group rebuild).
|
||||
self._access_groups_cache: dict[str, list[str]] | None = None
|
||||
# Per-router cache for the proxy auth-layer "is this model explicitly
|
||||
# zero-cost?" check. Lives on the router so it is invalidated alongside
|
||||
|
|
@ -617,6 +620,8 @@ class Router:
|
|||
# ``id()``-reuse risk after GC). See
|
||||
# ``litellm.proxy.auth.auth_checks._is_model_cost_zero``.
|
||||
self._zero_cost_cache: dict[str, bool] = {}
|
||||
self._routing_group_rows: tuple[DeploymentTypedDict, ...] | None = None
|
||||
self._init_routing_groups(None)
|
||||
|
||||
self.deployment_affinity_ttl_seconds = deployment_affinity_ttl_seconds
|
||||
self.model_group_affinity_config = model_group_affinity_config
|
||||
|
|
@ -1039,6 +1044,8 @@ class Router:
|
|||
self._routing_groups: dict[str, RoutingGroup] = {}
|
||||
self._model_to_group: dict[str, str] = {}
|
||||
self._group_selectors: dict[str, dict[str, RouterStrategySelector]] = {}
|
||||
self._invalidate_model_group_info_cache()
|
||||
self._invalidate_access_groups_cache()
|
||||
|
||||
if not groups_input:
|
||||
return
|
||||
|
|
@ -1053,6 +1060,12 @@ class Router:
|
|||
raise ValueError("routing_groups: group_name must be non-empty.")
|
||||
if group.group_name == "default":
|
||||
raise ValueError("routing_groups: 'default' is reserved for the implicit fallback group.")
|
||||
if group.group_name in known_model_names or group.group_name in (self.model_group_alias or {}):
|
||||
verbose_router_logger.warning(
|
||||
"routing_groups: group_name '%s' is shadowed by an existing model_name or model_group_alias; "
|
||||
"the group's strategy still applies to its members, but the name is not callable until renamed.",
|
||||
group.group_name,
|
||||
)
|
||||
if group.group_name in seen_group_names:
|
||||
raise ValueError(
|
||||
f"routing_groups: group names must be unique, duplicate group_name '{group.group_name}'."
|
||||
|
|
@ -1089,6 +1102,82 @@ class Router:
|
|||
{strategy_value: group_selector} if group_selector is not None else {}
|
||||
)
|
||||
|
||||
def get_routing_group(self, model_name: str) -> RoutingGroup | None:
|
||||
"""
|
||||
The routing group callable as `model_name`, or None. A real deployment
|
||||
`model_name` added after init shadows a same-named group (mirroring
|
||||
`_try_early_resolve_deployments_for_model_not_in_names`, where concrete
|
||||
models win over indirection); config-time collisions are rejected by
|
||||
`_init_routing_groups`.
|
||||
"""
|
||||
if not self._routing_groups:
|
||||
return None
|
||||
group: Final = self._routing_groups.get(model_name)
|
||||
if (
|
||||
group is None
|
||||
or model_name in self.model_name_to_deployment_indices
|
||||
or model_name in (self.model_group_alias or {})
|
||||
):
|
||||
return None
|
||||
return group
|
||||
|
||||
def _get_routing_group_deployments(
|
||||
self, model: str, team_id: str | None = None
|
||||
) -> list[DeploymentTypedDict] | None: # mutable-ok: list matches _get_all_deployments' contract for callers
|
||||
"""
|
||||
The union of member deployments for a routing group called as `model`,
|
||||
or None when `model` is not a callable group. The requested name stays
|
||||
the group name so strategy selectors key their state by it.
|
||||
|
||||
`_common_checks_available_deployment` consults this BEFORE its
|
||||
early-resolve step so a wildcard `default_deployment` or pattern route
|
||||
cannot hijack a group call. Overall resolution precedence there:
|
||||
specific deployment > model id > model_group_alias > routing group >
|
||||
model_name > team/pattern/default fallbacks.
|
||||
"""
|
||||
if not self._routing_groups:
|
||||
return None
|
||||
routing_group: Final = self.get_routing_group(model)
|
||||
if routing_group is None:
|
||||
return None
|
||||
return [ # mutable-ok: matches _get_all_deployments' list contract expected by downstream filters
|
||||
deployment
|
||||
for member in routing_group.models
|
||||
for deployment in self._get_all_deployments(model_name=member, team_id=team_id)
|
||||
]
|
||||
|
||||
def is_recognized_model(self, model: str) -> bool:
|
||||
"""
|
||||
Whether `model` names something this router serves directly: a
|
||||
deployment model_name, a deployment id, a `model_group_alias`, or a
|
||||
callable routing group. Proxy request gates share this predicate so a
|
||||
new virtual-model kind cannot be forgotten at one of them; wildcard,
|
||||
default-deployment, and deployment-name fallbacks stay caller policy.
|
||||
"""
|
||||
return (
|
||||
model in self.model_names
|
||||
or self.has_model_id(model)
|
||||
or (self.model_group_alias is not None and model in self.model_group_alias)
|
||||
or self.get_routing_group(model) is not None
|
||||
)
|
||||
|
||||
def routing_group_has_alternatives(self, model_group: str | None) -> bool:
|
||||
"""
|
||||
True when `model_group` names a callable routing group whose member
|
||||
union spans more than one deployment. Cooldown handling passes the
|
||||
FAILING REQUEST's model group here: a 429 on a group call cools the
|
||||
member down so selection moves to the group's alternatives, while a
|
||||
direct call to a single-deployment member keeps the
|
||||
single-deployment-model-group cooldown exemption.
|
||||
"""
|
||||
if model_group is None:
|
||||
return False
|
||||
resolved: Final = self._get_model_from_alias(model=model_group) or model_group
|
||||
group: Final = self.get_routing_group(resolved)
|
||||
if group is None:
|
||||
return False
|
||||
return sum(len(self.model_name_to_deployment_indices.get(member) or ()) for member in group.models) > 1
|
||||
|
||||
_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:
|
||||
|
|
@ -1149,8 +1238,10 @@ class Router:
|
|||
the most specific expression of caller intent.
|
||||
|
||||
Otherwise every model belongs to exactly one group: an explicit entry
|
||||
from `routing_groups`, or the implicit `"default"` group driven by the
|
||||
router's top-level `routing_strategy` / `routing_strategy_args`.
|
||||
from `routing_groups` (either because `model` IS a callable group name,
|
||||
or because it is a member of one), or the implicit `"default"` group
|
||||
driven by the router's top-level `routing_strategy` /
|
||||
`routing_strategy_args`.
|
||||
|
||||
`self.routing_strategy` may be either a string or a `RoutingStrategy`
|
||||
enum member (the constructor accepts both), so it is normalized to a
|
||||
|
|
@ -1162,7 +1253,7 @@ class Router:
|
|||
verbose_router_logger.debug("routing_group=request-override model=%s strategy=%s", model, override)
|
||||
return override, self._get_override_strategy_selector(override)
|
||||
|
||||
group_name: Final = self._model_to_group.get(model)
|
||||
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:
|
||||
strategy = self._normalize_strategy(self.routing_strategy)
|
||||
attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy or "")
|
||||
|
|
@ -7143,6 +7234,7 @@ class Router:
|
|||
original_exception=exception,
|
||||
deployment=deployment_id,
|
||||
time_to_cooldown=_time_to_cooldown,
|
||||
requested_model_group=(get_litellm_metadata_from_kwargs(kwargs) or {}).get("model_group"),
|
||||
) # setting deployment_id in cooldown deployments
|
||||
|
||||
return result
|
||||
|
|
@ -8326,6 +8418,7 @@ class Router:
|
|||
self.model_name_to_deployment_indices[model_name] = updated_indices
|
||||
else:
|
||||
del self.model_name_to_deployment_indices[model_name]
|
||||
self.model_names.discard(model_name)
|
||||
|
||||
# Update team_model_to_deployment_indices
|
||||
for key, indices in list(self.team_model_to_deployment_indices.items()):
|
||||
|
|
@ -9981,6 +10074,52 @@ class Router:
|
|||
|
||||
return returned_models
|
||||
|
||||
def get_model_list_from_routing_groups(self, model_name: str | None = None) -> Sequence[DeploymentTypedDict]:
|
||||
"""
|
||||
Callable routing groups materialized as model-list rows, mirroring
|
||||
`get_model_list_from_model_alias`: each member deployment is emitted
|
||||
under the group's name (via `_get_all_deployments`' `model_alias`
|
||||
rewrite), which is what surfaces groups in `get_model_names`,
|
||||
`/v1/models` discovery, `get_model_group_usage`, and the
|
||||
blocked/unhealthy hiding that all read `get_model_list`.
|
||||
"""
|
||||
if model_name is not None:
|
||||
group: Final = self.get_routing_group(model_name)
|
||||
return self._materialize_routing_group_rows((group,)) if group is not None else ()
|
||||
cached: Final = self._routing_group_rows
|
||||
if cached is not None:
|
||||
return cached
|
||||
rows: Final = self._materialize_routing_group_rows(
|
||||
tuple(
|
||||
callable_group
|
||||
for name in self._routing_groups
|
||||
if (callable_group := self.get_routing_group(name)) is not None
|
||||
)
|
||||
)
|
||||
self._routing_group_rows = rows
|
||||
return rows
|
||||
|
||||
def _materialize_routing_group_rows(self, groups: tuple[RoutingGroup, ...]) -> tuple[DeploymentTypedDict, ...]:
|
||||
return tuple(
|
||||
self._as_routing_group_row(deployment)
|
||||
for group in groups
|
||||
for member in group.models
|
||||
for deployment in self._get_all_deployments(model_name=member, model_alias=group.group_name)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _as_routing_group_row(deployment: DeploymentTypedDict) -> DeploymentTypedDict:
|
||||
"""
|
||||
A member deployment re-emitted under its group's name must not carry
|
||||
the member's `access_groups`: access groups grant member names, never
|
||||
the group, so inheriting them here would let a key holding a member's
|
||||
access group list and call the whole group.
|
||||
"""
|
||||
model_info: Final = { # mutable-ok: DeploymentTypedDict rows are plain dicts
|
||||
k: v for k, v in (deployment.get("model_info") or {}).items() if k != "access_groups"
|
||||
}
|
||||
return {**deployment, "model_info": model_info} # mutable-ok: DeploymentTypedDict rows are plain dicts
|
||||
|
||||
def get_model_list(
|
||||
self, model_name: str | None = None, team_id: str | None = None
|
||||
) -> list[DeploymentTypedDict] | None:
|
||||
|
|
@ -9997,6 +10136,7 @@ class Router:
|
|||
returned_models.extend(self._get_all_deployments(model_name=model_name, team_id=team_id))
|
||||
|
||||
returned_models.extend(self.get_model_list_from_model_alias(model_name=model_name))
|
||||
returned_models.extend(self.get_model_list_from_routing_groups(model_name=model_name))
|
||||
|
||||
if len(returned_models) == 0: # check if wildcard route
|
||||
potential_wildcard_models: Final = self.pattern_router.route(model_name) or []
|
||||
|
|
@ -10028,6 +10168,7 @@ class Router:
|
|||
"""
|
||||
self._cached_get_model_group_info.cache_clear()
|
||||
self._zero_cost_cache.clear()
|
||||
self._routing_group_rows = None
|
||||
|
||||
def _invalidate_access_groups_cache(self) -> None:
|
||||
"""Invalidate the cached access groups.
|
||||
|
|
@ -10598,17 +10739,23 @@ class Router:
|
|||
if _model_from_alias is not None:
|
||||
model = _model_from_alias
|
||||
|
||||
early: Final = self._try_early_resolve_deployments_for_model_not_in_names(
|
||||
model=model,
|
||||
request_team_id=request_team_id,
|
||||
include_team_models=_is_proxy_admin_request(request_kwargs),
|
||||
)
|
||||
if early is not None:
|
||||
return early
|
||||
_routing_group_deployments: Final = self._get_routing_group_deployments(model=model, team_id=request_team_id)
|
||||
if _routing_group_deployments is None:
|
||||
early: Final = self._try_early_resolve_deployments_for_model_not_in_names(
|
||||
model=model,
|
||||
request_team_id=request_team_id,
|
||||
include_team_models=_is_proxy_admin_request(request_kwargs),
|
||||
)
|
||||
if early is not None:
|
||||
return early
|
||||
|
||||
## get healthy deployments
|
||||
### get all deployments
|
||||
healthy_deployments = self._get_all_deployments(model_name=model, team_id=request_team_id)
|
||||
healthy_deployments = (
|
||||
_routing_group_deployments
|
||||
if _routing_group_deployments is not None
|
||||
else self._get_all_deployments(model_name=model, team_id=request_team_id)
|
||||
)
|
||||
_pre_model_access_group_filter_len: Final = len(healthy_deployments)
|
||||
healthy_deployments = self._filter_deployments_by_model_access_groups(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -319,6 +319,7 @@ def _should_cooldown_deployment(
|
|||
deployment: str,
|
||||
exception_status: str | int,
|
||||
original_exception: Any,
|
||||
requested_model_group: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Helper that decides if a deployment should be put in cooldown
|
||||
|
|
@ -341,7 +342,9 @@ def _should_cooldown_deployment(
|
|||
model_group: Final = litellm_router_instance.get_model_group(id=deployment)
|
||||
is_single_deployment_model_group = False
|
||||
if model_group is not None and len(model_group) == 1:
|
||||
is_single_deployment_model_group = True
|
||||
is_single_deployment_model_group = not litellm_router_instance.routing_group_has_alternatives(
|
||||
requested_model_group
|
||||
)
|
||||
|
||||
## CHECK DEPLOYMENT-LEVEL POLICY FIRST (overrides router-level)
|
||||
dep_policy, dep_allowed_fails = _get_deployment_cooldown_policy(litellm_router_instance, deployment)
|
||||
|
|
@ -413,6 +416,7 @@ def _set_cooldown_deployments(
|
|||
exception_status: str | int,
|
||||
deployment: str | None = None,
|
||||
time_to_cooldown: float | None = None,
|
||||
requested_model_group: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Add a model to the list of models being cooled down for that minute, if it exceeds the allowed fails / minute
|
||||
|
|
@ -449,6 +453,7 @@ def _set_cooldown_deployments(
|
|||
deployment=deployment,
|
||||
exception_status=exception_status,
|
||||
original_exception=original_exception,
|
||||
requested_model_group=requested_model_group,
|
||||
):
|
||||
litellm_router_instance.cooldown_cache.add_deployment_to_cooldown(
|
||||
model_id=deployment,
|
||||
|
|
|
|||
|
|
@ -1480,6 +1480,12 @@ def _router_serving_only(base_model: str) -> MagicMock:
|
|||
mock_router.model_names = set()
|
||||
mock_router.model_group_alias = {}
|
||||
mock_router.team_public_model_names = frozenset()
|
||||
mock_router.is_recognized_model.side_effect = lambda model: (
|
||||
model in mock_router.model_names or model in mock_router.model_group_alias
|
||||
)
|
||||
mock_router.router_general_settings.pass_through_all_models = False
|
||||
mock_router.default_deployment = None
|
||||
mock_router.pattern_router.patterns = {base_model: ["anthropic/*"]}
|
||||
mock_router.pattern_router.get_pattern.side_effect = (
|
||||
lambda model: [{"model_name": "anthropic/*"}] if model == base_model else None
|
||||
)
|
||||
|
|
@ -1723,3 +1729,22 @@ class TestCursorVariantResolvedBeforeAuth:
|
|||
)
|
||||
assert auth_body["model"] == "claude-opus-5-thinking-high"
|
||||
assert "reasoning_effort" not in auth_body
|
||||
|
||||
|
||||
class TestCursorGateRecognizesRoutingGroups:
|
||||
def test_group_name_variant_is_not_mangled(self):
|
||||
from litellm import Router
|
||||
from litellm.proxy.response_api_endpoints.endpoints import _resolve_cursor_model_variant
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{"model_name": "member-fast", "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}}
|
||||
],
|
||||
routing_groups=[
|
||||
{"group_name": "grouped-thinking-high", "models": ["member-fast"], "routing_strategy": "simple-shuffle"}
|
||||
],
|
||||
)
|
||||
body = {"model": "grouped-thinking-high", "messages": [{"role": "user", "content": "hi"}]}
|
||||
resolved = _resolve_cursor_model_variant(body, router)
|
||||
assert resolved["model"] == "grouped-thinking-high"
|
||||
assert "reasoning_effort" not in resolved
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ async def test_route_a2a_model_bypasses_router():
|
|||
mock_router.model_names = ["gpt-4", "gpt-3.5-turbo"]
|
||||
mock_router.deployment_names = []
|
||||
mock_router.has_model_id = Mock(return_value=False)
|
||||
mock_router.is_recognized_model = Mock(return_value=False)
|
||||
mock_router.model_group_alias = None
|
||||
mock_router.router_general_settings = Mock(pass_through_all_models=False)
|
||||
mock_router.default_deployment = None
|
||||
|
|
@ -88,6 +89,7 @@ async def test_route_non_a2a_model_raises_error_if_not_in_router():
|
|||
mock_router.model_names = ["gpt-4", "gpt-3.5-turbo"]
|
||||
mock_router.deployment_names = []
|
||||
mock_router.has_model_id = Mock(return_value=False)
|
||||
mock_router.is_recognized_model = Mock(return_value=False)
|
||||
mock_router.model_group_alias = None
|
||||
mock_router.router_general_settings = Mock(pass_through_all_models=False)
|
||||
mock_router.default_deployment = None
|
||||
|
|
|
|||
|
|
@ -1091,3 +1091,27 @@ async def test_route_request_rejects_chat_completion_without_messages():
|
|||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.param == "messages"
|
||||
llm_router.acompletion.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_routing_group_name_passes_model_gate():
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{"model_name": "member-a", "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}},
|
||||
{"model_name": "member-b", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"}},
|
||||
],
|
||||
routing_groups=[
|
||||
{"group_name": "grouped-quality", "models": ["member-a", "member-b"], "routing_strategy": "simple-shuffle"}
|
||||
],
|
||||
)
|
||||
data = {"model": "grouped-quality", "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
with patch.object(router, "acompletion", new=AsyncMock(return_value="group_response")) as spy:
|
||||
response = await (await route_request(data, router, None, "acompletion"))
|
||||
|
||||
assert response == "group_response"
|
||||
spy.assert_called_once_with(**data)
|
||||
|
|
|
|||
|
|
@ -726,3 +726,345 @@ def test_strategy_reinit_unregisters_override_selectors():
|
|||
assert router._override_selectors == {}
|
||||
assert not any(id(cb) == id(override_selector) for cb in litellm.callbacks)
|
||||
assert router._get_override_strategy_selector("latency-based-routing") is router.lowestlatency_logger
|
||||
|
||||
|
||||
def _quality_group(strategy="latency-based-routing"):
|
||||
return [{"group_name": "quality", "models": ["filtered-model", "other-model"], "routing_strategy": strategy}]
|
||||
|
||||
|
||||
def test_group_name_is_callable_and_unions_member_deployments():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
model, deployments = router._common_checks_available_deployment(model="quality")
|
||||
assert model == "quality"
|
||||
assert sorted(d["model_info"]["id"] for d in deployments) == ["deploy-1", "deploy-2", "deploy-3"]
|
||||
|
||||
|
||||
def test_group_name_appears_in_model_names_and_model_list():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
assert "quality" in router.get_model_names()
|
||||
rows = router.get_model_list(model_name="quality")
|
||||
assert {r["model_name"] for r in rows} == {"quality"}
|
||||
assert sorted(r["model_info"]["id"] for r in rows) == ["deploy-1", "deploy-2", "deploy-3"]
|
||||
|
||||
|
||||
def test_get_routing_context_for_group_name_uses_group_strategy():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
strategy, selector = router._get_routing_context("quality")
|
||||
assert strategy == "latency-based-routing"
|
||||
assert selector is router._group_selectors["quality"]["latency-based-routing"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_call_dispatches_via_group_selector():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
group_selector = router._group_selectors["quality"]["latency-based-routing"]
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
group_selector,
|
||||
"async_get_available_deployments",
|
||||
wraps=group_selector.async_get_available_deployments,
|
||||
) as latency_spy,
|
||||
patch("litellm.router.simple_shuffle", wraps=litellm.router.simple_shuffle) as shuffle_spy,
|
||||
):
|
||||
deployment = await router.async_get_available_deployment(model="quality", request_kwargs={})
|
||||
|
||||
assert latency_spy.called
|
||||
assert not shuffle_spy.called
|
||||
assert deployment["model_name"] in {"filtered-model", "other-model"}
|
||||
|
||||
|
||||
def test_group_name_colliding_with_model_name_is_shadowed_with_warning(caplog):
|
||||
import logging
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Router"):
|
||||
router = _build_router(
|
||||
routing_groups=[
|
||||
{"group_name": "filtered-model", "models": ["other-model"], "routing_strategy": "latency-based-routing"}
|
||||
]
|
||||
)
|
||||
assert any("shadowed" in record.getMessage() for record in caplog.records)
|
||||
assert router.get_routing_group("filtered-model") is None
|
||||
assert router._get_routing_context("other-model")[0] == "latency-based-routing"
|
||||
|
||||
model, deployments = router._common_checks_available_deployment(model="filtered-model")
|
||||
assert sorted(d["model_info"]["id"] for d in deployments) == ["deploy-1", "deploy-2"]
|
||||
|
||||
|
||||
def test_group_name_colliding_with_model_group_alias_is_shadowed_with_warning(caplog):
|
||||
import logging
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Router"):
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
model_group_alias={"quality": "filtered-model"},
|
||||
routing_groups=_quality_group(),
|
||||
)
|
||||
assert any("shadowed" in record.getMessage() for record in caplog.records)
|
||||
assert router.get_routing_group("quality") is None
|
||||
|
||||
model, deployments = router._common_checks_available_deployment(model="quality")
|
||||
assert model == "filtered-model"
|
||||
assert sorted(d["model_info"]["id"] for d in deployments) == ["deploy-1", "deploy-2"]
|
||||
|
||||
|
||||
def test_real_model_added_later_shadows_group():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
assert router.get_routing_group("quality") is not None
|
||||
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
router.add_deployment(
|
||||
Deployment(
|
||||
model_name="quality",
|
||||
litellm_params={"model": "openai/gpt-4o", "api_key": "sk-test-4", "api_base": "https://example.invalid"},
|
||||
model_info={"id": "deploy-shadow"},
|
||||
)
|
||||
)
|
||||
assert router.get_routing_group("quality") is None
|
||||
model, deployments = router._common_checks_available_deployment(model="quality")
|
||||
assert [d["model_info"]["id"] for d in deployments] == ["deploy-shadow"]
|
||||
|
||||
router.delete_deployment(id="deploy-shadow")
|
||||
assert "quality" not in router.model_names
|
||||
assert router.get_routing_group("quality") is not None
|
||||
_, restored = router._common_checks_available_deployment(model="quality")
|
||||
assert sorted(d["model_info"]["id"] for d in restored) == ["deploy-1", "deploy-2", "deploy-3"]
|
||||
|
||||
|
||||
def test_group_with_no_member_deployments_raises_no_healthy():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
routing_groups=[{"group_name": "empty-group", "models": ["ghost-model"], "routing_strategy": "simple-shuffle"}],
|
||||
)
|
||||
with pytest.raises(litellm.BadRequestError):
|
||||
router._common_checks_available_deployment(model="empty-group")
|
||||
|
||||
|
||||
def test_alias_pointing_at_group_composes():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
model_group_alias={"quality-alias": "quality"},
|
||||
routing_groups=_quality_group(),
|
||||
)
|
||||
model, deployments = router._common_checks_available_deployment(model="quality-alias")
|
||||
assert model == "quality"
|
||||
assert sorted(d["model_info"]["id"] for d in deployments) == ["deploy-1", "deploy-2", "deploy-3"]
|
||||
|
||||
|
||||
def test_model_group_info_reports_group():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
info = router.get_model_group_info("quality")
|
||||
assert info is not None
|
||||
assert info.model_group == "quality"
|
||||
assert "openai" in info.providers
|
||||
|
||||
|
||||
def test_routing_group_has_alternatives():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
assert router.routing_group_has_alternatives("quality") is True
|
||||
assert router.routing_group_has_alternatives("filtered-model") is False
|
||||
assert router.routing_group_has_alternatives(None) is False
|
||||
|
||||
solo_router = Router(
|
||||
model_list=_model_list(),
|
||||
routing_groups=[{"group_name": "solo-group", "models": ["other-model"], "routing_strategy": "simple-shuffle"}],
|
||||
)
|
||||
assert solo_router.routing_group_has_alternatives("solo-group") is False
|
||||
|
||||
|
||||
def test_member_direct_call_unchanged_by_callable_groups():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
model, deployments = router._common_checks_available_deployment(model="other-model")
|
||||
assert model == "other-model"
|
||||
assert [d["model_info"]["id"] for d in deployments] == ["deploy-3"]
|
||||
|
||||
|
||||
def test_update_settings_group_change_invalidates_model_group_info():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
assert router.get_model_group_info("quality") is not None
|
||||
assert router.get_model_group_info("renamed-group") is None
|
||||
|
||||
router.update_settings(
|
||||
routing_groups=[
|
||||
{"group_name": "renamed-group", "models": ["filtered-model"], "routing_strategy": "simple-shuffle"}
|
||||
]
|
||||
)
|
||||
assert router.get_model_group_info("quality") is None
|
||||
info = router.get_model_group_info("renamed-group")
|
||||
assert info is not None
|
||||
assert info.model_group == "renamed-group"
|
||||
|
||||
|
||||
def test_is_recognized_model_covers_every_virtual_model_kind():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
model_group_alias={"my-alias": "filtered-model"},
|
||||
routing_groups=_quality_group(),
|
||||
)
|
||||
assert router.is_recognized_model("filtered-model") is True
|
||||
assert router.is_recognized_model("deploy-1") is True
|
||||
assert router.is_recognized_model("my-alias") is True
|
||||
assert router.is_recognized_model("quality") is True
|
||||
assert router.is_recognized_model("ghost") is False
|
||||
|
||||
|
||||
def test_routing_group_has_alternatives_resolves_aliases():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
model_group_alias={"quality-alias": "quality"},
|
||||
routing_groups=_quality_group(),
|
||||
)
|
||||
assert router.routing_group_has_alternatives("quality-alias") is True
|
||||
assert router.routing_group_has_alternatives("quality") is True
|
||||
|
||||
|
||||
def test_group_rows_cache_invalidated_on_model_list_change():
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
assert sum(1 for row in router.get_model_list() if row["model_name"] == "quality") == 3
|
||||
|
||||
router.add_deployment(
|
||||
Deployment(
|
||||
model_name="filtered-model",
|
||||
litellm_params={"model": "openai/gpt-4o", "api_key": "sk-test-5", "api_base": "https://example.invalid"},
|
||||
model_info={"id": "deploy-4"},
|
||||
)
|
||||
)
|
||||
assert sum(1 for row in router.get_model_list() if row["model_name"] == "quality") == 4
|
||||
|
||||
|
||||
def _pin_choice_to(deployment_id):
|
||||
def _pick(seq):
|
||||
for candidate in seq:
|
||||
if candidate["model_info"]["id"] == deployment_id:
|
||||
return candidate
|
||||
return seq[0]
|
||||
|
||||
return _pick
|
||||
|
||||
|
||||
async def _call_and_get_cooldowns(router, model):
|
||||
from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments
|
||||
|
||||
with (
|
||||
patch("litellm.router_strategy.simple_shuffle.random.choice", side_effect=_pin_choice_to("deploy-3")),
|
||||
pytest.raises(litellm.RateLimitError),
|
||||
):
|
||||
await router.acompletion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="litellm.RateLimitError",
|
||||
)
|
||||
return await _async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_call_429_registers_cooldown_end_to_end():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
routing_groups=_quality_group("simple-shuffle"),
|
||||
num_retries=0,
|
||||
cooldown_time=60,
|
||||
)
|
||||
cooldown_ids = await _call_and_get_cooldowns(router, "quality")
|
||||
assert "deploy-3" in cooldown_ids
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_alias_to_group_429_registers_cooldown_end_to_end():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
model_group_alias={"quality-alias": "quality"},
|
||||
routing_groups=_quality_group("simple-shuffle"),
|
||||
num_retries=0,
|
||||
cooldown_time=60,
|
||||
)
|
||||
cooldown_ids = await _call_and_get_cooldowns(router, "quality-alias")
|
||||
assert "deploy-3" in cooldown_ids
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_direct_single_deployment_member_429_keeps_exemption_end_to_end():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
routing_groups=_quality_group("simple-shuffle"),
|
||||
num_retries=0,
|
||||
cooldown_time=60,
|
||||
)
|
||||
cooldown_ids = await _call_and_get_cooldowns(router, "other-model")
|
||||
assert "deploy-3" not in cooldown_ids
|
||||
|
||||
|
||||
def test_group_rows_do_not_inherit_member_access_groups():
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gated-member",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
|
||||
"model_info": {"id": "gated-1", "access_groups": ["restricted-team"]},
|
||||
}
|
||||
],
|
||||
routing_groups=[
|
||||
{"group_name": "gated-group", "models": ["gated-member"], "routing_strategy": "simple-shuffle"}
|
||||
],
|
||||
)
|
||||
access_groups = router.get_model_access_groups()
|
||||
assert "gated-group" not in access_groups.get("restricted-team", [])
|
||||
assert all("access_groups" not in (row.get("model_info") or {}) for row in router.get_model_list(model_name="gated-group"))
|
||||
assert "access_groups" in router.get_model_list(model_name="gated-member")[0]["model_info"]
|
||||
|
||||
|
||||
def test_group_rebuild_invalidates_access_groups_cache():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
router.get_model_access_groups()
|
||||
assert router._access_groups_cache is not None
|
||||
|
||||
router.update_settings(routing_groups=[])
|
||||
assert router._access_groups_cache is None
|
||||
|
||||
|
||||
def test_get_model_list_from_routing_groups_materializes_rows():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
rows = router.get_model_list_from_routing_groups()
|
||||
assert {row["model_name"] for row in rows} == {"quality"}
|
||||
assert router.get_model_list_from_routing_groups() is rows
|
||||
|
||||
named = router.get_model_list_from_routing_groups(model_name="quality")
|
||||
assert sorted(row["model_info"]["id"] for row in named) == ["deploy-1", "deploy-2", "deploy-3"]
|
||||
assert router.get_model_list_from_routing_groups(model_name="filtered-model") == ()
|
||||
|
||||
|
||||
def test_get_routing_group_deployments_unions_members():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
union = router._get_routing_group_deployments("quality")
|
||||
assert sorted(d["model_info"]["id"] for d in union) == ["deploy-1", "deploy-2", "deploy-3"]
|
||||
assert router._get_routing_group_deployments("filtered-model") is None
|
||||
|
||||
|
||||
def test_materialize_routing_group_rows_labels_members_with_group_name():
|
||||
router = _build_router(routing_groups=_quality_group())
|
||||
group = router.get_routing_group("quality")
|
||||
rows = router._materialize_routing_group_rows((group,))
|
||||
assert {row["model_name"] for row in rows} == {"quality"}
|
||||
assert len(rows) == 3
|
||||
|
||||
|
||||
def test_as_routing_group_row_strips_access_groups():
|
||||
source = {"model_name": "member", "model_info": {"id": "d1", "access_groups": ["restricted"]}}
|
||||
row = Router._as_routing_group_row(source)
|
||||
assert row["model_info"] == {"id": "d1"}
|
||||
assert source["model_info"]["access_groups"] == ["restricted"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_call_429_cools_down_member_across_retries():
|
||||
router = Router(
|
||||
model_list=_model_list(),
|
||||
routing_groups=_quality_group("simple-shuffle"),
|
||||
num_retries=1,
|
||||
cooldown_time=60,
|
||||
)
|
||||
cooldown_ids = await _call_and_get_cooldowns(router, "quality")
|
||||
assert "deploy-3" in cooldown_ids
|
||||
|
|
|
|||
|
|
@ -296,3 +296,84 @@ class TestShouldCooldownBasedOnAllowedFailsPolicy:
|
|||
assert set_cache_call[1]["ttl"] == 0.0, (
|
||||
"cooldown_time_override=0 should be used as TTL, not the router-level 60.0"
|
||||
)
|
||||
|
||||
|
||||
class TestRoutingGroupCooldownAlternatives:
|
||||
def _router(self, routing_groups=None):
|
||||
from litellm import Router
|
||||
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "solo-member",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
|
||||
"model_info": {"id": "cg-deploy-1"},
|
||||
},
|
||||
{
|
||||
"model_name": "other-member",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"},
|
||||
"model_info": {"id": "cg-deploy-2"},
|
||||
},
|
||||
],
|
||||
routing_groups=routing_groups,
|
||||
)
|
||||
|
||||
def test_group_call_429_cools_down_member_with_alternatives(self):
|
||||
from litellm.router_utils.cooldown_handlers import _should_cooldown_deployment
|
||||
|
||||
router = self._router(
|
||||
routing_groups=[
|
||||
{
|
||||
"group_name": "grouped",
|
||||
"models": ["solo-member", "other-member"],
|
||||
"routing_strategy": "simple-shuffle",
|
||||
}
|
||||
]
|
||||
)
|
||||
assert (
|
||||
_should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="cg-deploy-1",
|
||||
exception_status=429,
|
||||
original_exception=Exception("rate limited"),
|
||||
requested_model_group="grouped",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_direct_member_429_keeps_single_deployment_exemption(self):
|
||||
from litellm.router_utils.cooldown_handlers import _should_cooldown_deployment
|
||||
|
||||
router = self._router(
|
||||
routing_groups=[
|
||||
{
|
||||
"group_name": "grouped",
|
||||
"models": ["solo-member", "other-member"],
|
||||
"routing_strategy": "simple-shuffle",
|
||||
}
|
||||
]
|
||||
)
|
||||
assert (
|
||||
_should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="cg-deploy-1",
|
||||
exception_status=429,
|
||||
original_exception=Exception("rate limited"),
|
||||
requested_model_group="solo-member",
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_429_without_request_context_keeps_exemption(self):
|
||||
from litellm.router_utils.cooldown_handlers import _should_cooldown_deployment
|
||||
|
||||
router = self._router(routing_groups=None)
|
||||
assert (
|
||||
_should_cooldown_deployment(
|
||||
litellm_router_instance=router,
|
||||
deployment="cg-deploy-1",
|
||||
exception_status=429,
|
||||
original_exception=Exception("rate limited"),
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue