fix(team-routing): keep team model routing on public names

Remove team model_alias rewrites and resolve team deployments by team_public_model_name with team_id so sibling deployments stay in the routing candidate pool, with explicit logs showing candidate selection before load balancing.

Made-with: Cursor
This commit is contained in:
Sameer Kankute 2026-03-23 16:12:27 +05:30 • committed by Yuneng Jiang
parent fbc4baebc4
commit 1245ab49bb
No known key found for this signature in database
3 changed files with 183 additions and 79 deletions

View file

@ -344,17 +344,6 @@ async def _add_team_model_to_db(
prisma_client=prisma_client,
)
## CREATE MODEL ALIAS IN DB ##
await update_team(
data=UpdateTeamRequest(
team_id=_team_id,
model_aliases={original_model_name: unique_model_name},
),
user_api_key_dict=user_api_key_dict,
http_request=Request(scope={"type": "http"}),
)
# add model to team object
await team_model_add(
data=TeamModelAddRequest(
team_id=_team_id,
@ -490,8 +479,8 @@ async def _update_existing_team_model_assignment(
# Update alias only if public name changed
if old_public_name and public_model_name != old_public_name:
await update_team(
data=UpdateTeamRequest(
await team_model_add(
data=TeamModelAddRequest(
team_id=team_id,
model_aliases={public_model_name: db_model.model_name},
),
@ -499,7 +488,6 @@ async def _update_existing_team_model_assignment(
http_request=Request(scope={"type": "http"}),
)
# Keep existing unique model_name
patch_data.model_name = None

View file

@ -5288,63 +5288,6 @@ class Router:
if "fallback_depth" not in input_kwargs:
input_kwargs["fallback_depth"] = 0
# ORDER-BASED FALLBACKS: prepend higher order levels to the fallback list
# Skip for error types that have their own dedicated fallback handlers
_skip_order_fallback = isinstance(
e,
(litellm.ContextWindowExceededError, litellm.ContentPolicyViolationError),
)
all_deployments = self._get_all_deployments(model_name=original_model_group)
_order_set: set = {
d.get("litellm_params", {}).get("order")
for d in all_deployments
if d.get("litellm_params", {}).get("order") is not None
}
order_values: list = sorted(_order_set)
if len(order_values) > 1 and not _skip_order_fallback:
# Determine which order levels have already been tried
current_target = kwargs.get("_target_order")
skip_up_to = (
current_target if current_target is not None else order_values[0]
)
# Build order-based fallback entries (skip already-tried levels)
order_fallback_entries: List = [
{"model": original_model_group, "_target_order": o}
for o in order_values
if o > skip_up_to
]
# Get external fallbacks — handle both standard and non-standard formats
external_fallback_group: Optional[List] = None
if fallbacks is not None and model_group is not None:
if _check_non_standard_fallback_format(fallbacks=fallbacks):
# Non-standard formats (e.g. ["claude-3-haiku"] or
# [{"model": "...", "messages": [...]}]) are passed through directly
external_fallback_group = fallbacks
else:
external_fallback_group, generic_idx = get_fallback_model_group(
fallbacks=fallbacks,
model_group=cast(str, model_group),
)
if external_fallback_group is None and generic_idx is not None:
external_fallback_group = fallbacks[generic_idx]["*"]
# Combined list: order fallbacks first, then external
combined_fallbacks = order_fallback_entries + (
external_fallback_group or []
)
if combined_fallbacks:
input_kwargs.update(
{
"fallback_model_group": combined_fallbacks,
"original_model_group": original_model_group,
}
)
response = await run_async_fallback(
*args,
**input_kwargs,
)
return response
try:
verbose_router_logger.info("Trying to fallback b/w models")
@ -8218,7 +8161,6 @@ class Router:
if model.get("model_info", {}).get("team_id") == team_id:
return model.get("model_name")
## wildcard models
return None
def should_include_deployment(
@ -8924,6 +8866,38 @@ class Router:
model = _model_from_alias
if model not in self.model_names:
# Check for team-specific deployments by team_public_model_name
if request_team_id is not None:
team_deployments = self._get_all_deployments(
model_name=model, team_id=request_team_id
)
if team_deployments:
candidate_details = []
for deployment in team_deployments:
deployment_info = deployment.get("model_info", {}) or {}
deployment_params = deployment.get("litellm_params", {}) or {}
candidate_details.append(
{
"model_name": deployment.get("model_name"),
"model_id": deployment_info.get("id"),
"team_public_model_name": deployment_info.get(
"team_public_model_name"
),
"api_base": deployment_params.get("api_base"),
}
)
verbose_router_logger.info(
"🔥 routing_candidates_before_lb "
f"model={model} count={len(team_deployments)} "
f"candidates={candidate_details}"
)
if len(team_deployments) > 1:
verbose_router_logger.info(
"🔥 load_balancer_candidate_pool "
f"model={model} candidate_count={len(team_deployments)}"
)
return model, team_deployments
# check if provider/ specific wildcard routing use pattern matching
pattern_deployments = self.pattern_router.get_deployments_by_pattern(
model=model,
@ -8956,14 +8930,38 @@ class Router:
## get healthy deployments
### get all deployments
healthy_deployments = self._get_all_deployments(
model_name=model, team_id=request_team_id
)
healthy_deployments = self._get_all_deployments(model_name=model)
if len(healthy_deployments) == 0:
# check if the user sent in a deployment name instead
healthy_deployments = self._get_deployment_by_litellm_model(model=model)
if isinstance(healthy_deployments, list) and len(healthy_deployments) > 0:
candidate_details = []
for deployment in healthy_deployments:
deployment_info = deployment.get("model_info", {}) or {}
deployment_params = deployment.get("litellm_params", {}) or {}
candidate_details.append(
{
"model_name": deployment.get("model_name"),
"model_id": deployment_info.get("id"),
"team_public_model_name": deployment_info.get(
"team_public_model_name"
),
"api_base": deployment_params.get("api_base"),
}
)
verbose_router_logger.info(
"🔥 routing_candidates_before_lb "
f"model={model} count={len(healthy_deployments)} "
f"candidates={candidate_details}"
)
if len(healthy_deployments) > 1:
verbose_router_logger.info(
"🔥 load_balancer_candidate_pool "
f"model={model} candidate_count={len(healthy_deployments)}"
)
if verbose_router_logger.isEnabledFor(logging.DEBUG):
verbose_router_logger.debug(
f"initial list of deployments: {healthy_deployments}"
@ -8979,9 +8977,7 @@ class Router:
)
# Re-assign model to the fallback and try to get deployments again
model = fallback_model
healthy_deployments = self._get_all_deployments(
model_name=model, team_id=request_team_id
)
healthy_deployments = self._get_all_deployments(model_name=model)
# If still no deployments after checking for fallbacks, raise an error
if len(healthy_deployments) == 0:

View file

@ -558,6 +558,126 @@ class TestUpdatePublicModelGroups:
litellm.public_model_groups_links = original_value
class TestTeamModelSiblingRouting:
"""
Verify that sibling team deployments (same public model name, different
api_base) are all reachable through routing — no alias overwrite, no
collapse to a single deployment.
"""
@pytest.mark.asyncio
async def test_no_model_aliases_written_for_team_models(self):
"""
_add_team_model_to_db must NOT write model_aliases (which caused
the second sibling to overwrite the first). It should only call
team_model_add to register the public name on the team's models list.
"""
from litellm.proxy.management_endpoints.model_management_endpoints import (
_add_team_model_to_db,
)
from litellm.types.router import ModelInfo
team_id = "team_no_alias"
public_name = "gpt-4.1-mini"
mock_update_team = AsyncMock()
async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client):
return MagicMock(model_id=str(uuid.uuid4()))
mock_team_model_add = AsyncMock()
user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
prisma_client = MockPrismaClient(team_exists=True)
for api_base in ["https://eastus.example.com", "https://westus.example.com"]:
dep = Deployment(
model_name=public_name,
litellm_params=LiteLLM_Params(
model="azure/gpt-4o-mini",
api_key="key",
api_base=api_base,
),
model_info=ModelInfo(team_id=team_id),
)
with patch(
"litellm.proxy.management_endpoints.model_management_endpoints.update_team",
mock_update_team,
), patch(
"litellm.proxy.management_endpoints.model_management_endpoints._add_model_to_db",
side_effect=mock_add_model_to_db,
), patch(
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add",
mock_team_model_add,
):
await _add_team_model_to_db(
model_params=dep,
user_api_key_dict=user,
prisma_client=prisma_client,
)
mock_update_team.assert_not_called()
assert mock_team_model_add.call_count == 2
@pytest.mark.asyncio
async def test_router_finds_all_sibling_team_deployments(self):
"""
When two team deployments share team_public_model_name="gpt-4.1-mini",
the router's _common_checks_available_deployment must return BOTH as
healthy_deployments (not collapse to one).
"""
import litellm
team_id = "teamA"
public_name = "gpt-4.1-mini"
router = litellm.Router(
model_list=[
{
"model_name": f"model_name_{team_id}_uuid1",
"litellm_params": {
"model": "azure/gpt-4o-mini",
"api_key": "key-1",
"api_base": "https://eastus.openai.azure.com",
},
"model_info": {
"team_id": team_id,
"team_public_model_name": public_name,
},
},
{
"model_name": f"model_name_{team_id}_uuid2",
"litellm_params": {
"model": "azure/gpt-4o-mini",
"api_key": "key-2",
"api_base": "https://westus.openai.azure.com",
},
"model_info": {
"team_id": team_id,
"team_public_model_name": public_name,
},
},
],
)
# map_team_model should return the public name (not an internal UUID)
result = router.map_team_model(public_name, team_id)
assert result == public_name
# _common_checks_available_deployment should return both deployments
model, healthy = router._common_checks_available_deployment(
model=public_name,
request_kwargs={"metadata": {"user_api_key_team_id": team_id}},
)
assert isinstance(healthy, list)
assert len(healthy) == 2
api_bases = {d["litellm_params"]["api_base"] for d in healthy}
assert api_bases == {
"https://eastus.openai.azure.com",
"https://westus.openai.azure.com",
}
class TestTeamModelUpdate:
"""Test team model update handles team_id consistently with model creation"""
@ -604,7 +724,7 @@ class TestTeamModelUpdate:
assert result.get("model_name", "").startswith("model_name_test_team_123_")
assert "team_public_model_name" in str(result.get("model_info", ""))
mock_update_team.assert_called_once()
mock_update_team.assert_not_called()
mock_team_model_add.assert_called_once()
@pytest.mark.asyncio