mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
fbc4baebc4
commit
1245ab49bb
3 changed files with 183 additions and 79 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue