fix(router): resolve routing group names to member deployment ids

Signed-off-by: pjdurden <prajjwalchittori1@gmail.com>
This commit is contained in:
pjdurden 2026-08-18 10:27:23 -05:00
parent 4fd7a73ef5
commit 7b73dc0950
No known key found for this signature in database
2 changed files with 64 additions and 4 deletions

View file

@ -1167,6 +1167,18 @@ class Router:
or self.get_routing_group(model) is not None
)
def _resolve_to_deployment_model_names(self, model: str) -> tuple[str, ...]:
"""
The deployment `model_name`s behind a requested name: a
`model_group_alias` resolves to its target and a callable routing group
to its members, so lookups keyed by model group name reach the same
deployments the router would route the request to. Names that are
neither resolve to themselves.
"""
resolved: Final = self._get_model_from_alias(model=model) or model
group: Final = self.get_routing_group(resolved)
return tuple(group.models) if group is not None else (resolved,)
def routing_group_has_alternatives(self, model_group: str | None) -> bool:
"""
True when `model_group` names a callable routing group whose member
@ -9764,6 +9776,11 @@ class Router:
Returns list of model id's.
`model_name` may be any name the router serves: a deployment
`model_name`, a `model_group_alias`, or a callable routing group, which
resolve to the deployments they route to (see
`_resolve_to_deployment_model_names`).
Optimized with O(1) or O(k) index lookup when model_name provided,
instead of O(n) linear scan.
"""
@ -9771,9 +9788,8 @@ class Router:
if model_name is not None:
# O(1) lookup in model_name index, then O(k) iteration where k = deployments for this model_name
if model_name in self.model_name_to_deployment_indices:
indices: Final = self.model_name_to_deployment_indices[model_name]
for idx in indices:
for deployment_model_name in self._resolve_to_deployment_model_names(model=model_name):
for idx in self.model_name_to_deployment_indices.get(deployment_model_name) or ():
model = self.model_list[idx]
if "model_info" in model and "id" in model["model_info"]:
if exclude_team_models and model["model_info"].get("team_id"):

View file

@ -15,7 +15,7 @@ sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import Router
from litellm.types.router import RoutingGroup, RoutingStrategy
from litellm.types.router import RouterRateLimitError, RoutingGroup, RoutingStrategy
def _model_list():
@ -1068,3 +1068,47 @@ async def test_group_call_429_cools_down_member_across_retries():
)
cooldown_ids = await _call_and_get_cooldowns(router, "quality")
assert "deploy-3" in cooldown_ids
def test_get_model_ids_resolves_group_and_alias_to_member_deployments():
router = Router(
model_list=_model_list(),
model_group_alias={"quality-alias": "quality"},
routing_groups=_quality_group("simple-shuffle"),
)
assert sorted(router.get_model_ids(model_name="quality")) == ["deploy-1", "deploy-2", "deploy-3"]
assert sorted(router.get_model_ids(model_name="quality-alias")) == ["deploy-1", "deploy-2", "deploy-3"]
assert sorted(router.get_model_ids(model_name="filtered-model")) == ["deploy-1", "deploy-2"]
assert router.get_model_ids(model_name="not-a-model") == []
@pytest.mark.asyncio
async def test_group_call_reports_member_cooldown_time_when_every_member_is_cooling_down():
"""
Regression: the group name has to resolve to its members' deployment ids, else the
429 raised for a fully cooled-down group reports the router's default cooldown
instead of the members' own, and the proxy sends that as `retry-after`.
"""
deployment_cooldown_time = 120
router = Router(
model_list=[
{**deployment, "model_info": {**deployment["model_info"], "cooldown_time": deployment_cooldown_time}}
for deployment in _model_list()
],
routing_groups=_quality_group("simple-shuffle"),
num_retries=0,
)
assert router.cooldown_cache.default_cooldown_time != deployment_cooldown_time
for deployment_id in ("deploy-1", "deploy-2", "deploy-3"):
router.cooldown_cache.add_deployment_to_cooldown(
model_id=deployment_id,
original_exception=litellm.RateLimitError(message="rate limited", llm_provider="openai", model="gpt-4o"),
exception_status=429,
cooldown_time=deployment_cooldown_time,
)
with pytest.raises(RouterRateLimitError) as exc_info:
await router.async_get_available_deployment(model="quality", request_kwargs={})
assert exc_info.value.cooldown_time == deployment_cooldown_time