diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 44d41097833..18a25e52cf0 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 1dde9c9a6eb..cbaa71b835e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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: diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index f3c89003105..48c9c2359bd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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