diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a20b557e38b..ba122fdf1d3 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11071,6 +11071,7 @@ def _add_team_models_to_all_models( Add team models to all models """ team_models: Dict[str, Set[str]] = {} + access_groups = llm_router.get_model_access_groups() for team_object in team_db_objects_typed: if ( @@ -11094,7 +11095,7 @@ def _add_team_models_to_all_models( if can_add_model: team_models.setdefault(model_id, set()).add(team_object.team_id) else: - for model_name in team_object.models: + for model_name in _resolve_models_with_access_groups(team_object.models, access_groups): _models = llm_router.get_model_list(model_name=model_name, team_id=team_object.team_id) if _models is not None: for model in _models: @@ -11220,9 +11221,13 @@ def get_direct_access_models( if SpecialModelNames.all_proxy_models.value in user_db_object.models: return llm_router.get_model_ids(exclude_team_models=True) + resolved_models = _resolve_models_with_access_groups( + user_db_object.models, + llm_router.get_model_access_groups(), + ) return [ model_id - for model in user_db_object.models + for model in resolved_models for deployment in (llm_router.get_model_list(model_name=model) or []) if (model_id := deployment.get("model_info", {}).get("id", None)) is not None ] @@ -11775,10 +11780,10 @@ def _paginate_models_response( } -def _team_models_resolve_to_names(team_models: List[str], access_groups: Dict[str, Any]) -> List[str]: - """Expand team model entries (including access group names) to concrete model names.""" +def _resolve_models_with_access_groups(models: List[str], access_groups: Dict[str, Any]) -> List[str]: + """Expand model entries (including access group names) to concrete model names.""" resolved: List[str] = [] - for name in team_models: + for name in models: if name in access_groups: resolved.extend(access_groups[name]) else: @@ -11837,7 +11842,7 @@ async def _gather_team_accessible_model_ids( try: if team_object.models and SpecialModelNames.all_proxy_models.value not in team_object.models: - _resolved_names = _team_models_resolve_to_names(team_object.models, access_groups) + _resolved_names = _resolve_models_with_access_groups(team_object.models, access_groups) db_models = await ModelRepository(prisma_client).table.find_many( where={"model_name": {"in": _resolved_names}} ) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index bad76864ca7..17700428b22 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1740,6 +1740,114 @@ def test_add_team_models_to_all_models(): assert result == {"gpt-4-model-2": {"team1"}} +def test_add_team_models_to_all_models_expands_model_access_groups(): + """ + Regression test for #34998 / #17475: a team whose `models` field holds a + config-level model access group name (e.g. `beta-models`) saw a blank + Models + Endpoints list, because the access group name was passed straight + to `llm_router.get_model_list()` and matched no deployment. + """ + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.proxy_server import _add_team_models_to_all_models + + team = MagicMock(spec=LiteLLM_TeamTable) + team.team_id = "team1" + team.models = ["beta-models"] + + llm_router = MagicMock() + llm_router.get_model_access_groups.return_value = {"beta-models": ["gpt-4", "claude-3"]} + + def get_model_list(model_name=None, team_id=None): + return { + "gpt-4": [{"model_info": {"id": "gpt-4-deploy"}}], + "claude-3": [{"model_info": {"id": "claude-3-deploy"}}], + }.get(model_name) + + llm_router.get_model_list.side_effect = get_model_list + + result = _add_team_models_to_all_models( + team_db_objects_typed=[team], + llm_router=llm_router, + ) + + assert result == { + "gpt-4-deploy": {"team1"}, + "claude-3-deploy": {"team1"}, + } + + +def test_get_direct_access_models_expands_model_access_groups(): + """ + Same regression as above, for a user granted models only through a + config-level model access group on their own `models` field. + """ + from litellm.proxy.proxy_server import get_direct_access_models + + user = MagicMock() + user.models = ["beta-models"] + + llm_router = MagicMock() + llm_router.get_model_access_groups.return_value = {"beta-models": ["gpt-4"]} + llm_router.get_model_list.side_effect = lambda model_name=None, team_id=None: ( + [{"model_info": {"id": "gpt-4-deploy"}}] if model_name == "gpt-4" else None + ) + + assert get_direct_access_models(user_db_object=user, llm_router=llm_router) == ["gpt-4-deploy"] + + +@pytest.mark.asyncio +async def test_get_all_team_and_direct_access_models_with_team_access_group(): + """ + End-to-end regression for the reported symptom: an internal user whose only + model access comes from a model access group on their team must still see + those deployments in the `include_team_models=true` listing. + """ + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.proxy_server import get_all_team_and_direct_access_models + + mock_user_row = MagicMock() + mock_user_row.model_dump.return_value = { + "user_id": "u1", + "models": [], + "teams": ["team1"], + } + mock_team_row = MagicMock() + mock_team_row.model_dump.return_value = { + "team_id": "team1", + "models": ["beta-models"], + "access_group_ids": None, + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user_row) + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team_row]) + + llm_router = MagicMock() + llm_router.get_model_access_groups.return_value = {"beta-models": ["gpt-4"]} + llm_router.get_model_list.side_effect = lambda model_name=None, team_id=None: ( + [{"model_info": {"id": "gpt-4-deploy"}}] if model_name == "gpt-4" else None + ) + + all_models = [ + {"model_name": "gpt-4", "model_info": {"id": "gpt-4-deploy"}}, + {"model_name": "gpt-5", "model_info": {"id": "gpt-5-deploy"}}, + ] + + result = await get_all_team_and_direct_access_models( + user_api_key_dict=UserAPIKeyAuth( + user_id="u1", + user_role=LitellmUserRoles.INTERNAL_USER, + team_id="litellm-dashboard", + ), + prisma_client=mock_prisma_client, + llm_router=llm_router, + all_models=all_models, + ) + + assert [m["model_info"]["id"] for m in result] == ["gpt-4-deploy"] + assert result[0]["model_info"]["access_via_team_ids"] == ["team1"] + + @pytest.mark.asyncio async def test_apply_search_filter_matches_team_public_model_name(): """