diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 63109916ab1..e4db02758a1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8139,17 +8139,26 @@ def _add_team_models_to_all_models( if can_add_model: team_models.setdefault(model_id, set()).add(team_object.team_id) else: + model_access_groups = ( + llm_router.get_model_access_groups(team_id=team_object.team_id) + if llm_router + else {} + ) for model_name in team_object.models: - _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: - model_id = model.get("model_info", {}).get("id", None) - if model_id is not None: - team_models.setdefault(model_id, set()).add( - team_object.team_id - ) + names_to_resolve = [model_name] + if model_name in model_access_groups: + names_to_resolve = model_access_groups[model_name] + for name in names_to_resolve: + _models = llm_router.get_model_list( + model_name=name, team_id=team_object.team_id + ) + if _models is not None: + for model in _models: + model_id = model.get("model_info", {}).get("id", None) + if model_id is not None: + team_models.setdefault(model_id, set()).add( + team_object.team_id + ) return team_models @@ -8702,18 +8711,25 @@ async def _filter_models_by_team_id( if can_add_model: team_accessible_model_ids.add(model_id) else: - # Team has access to specific models + # Team has access to specific models (or model access group names like "on-prem") + model_access_groups = ( + llm_router.get_model_access_groups(team_id=team_id) if llm_router else {} + ) for model_name in team_object.models: - _models = ( - llm_router.get_model_list(model_name=model_name, team_id=team_id) - if llm_router - else [] - ) - if _models is not None: - for model in _models: - model_id = model.get("model_info", {}).get("id", None) - if model_id is not None: - team_accessible_model_ids.add(model_id) + names_to_resolve = [model_name] + if model_name in model_access_groups: + names_to_resolve = model_access_groups[model_name] + for name in names_to_resolve: + _models = ( + llm_router.get_model_list(model_name=name, team_id=team_id) + if llm_router + else [] + ) + if _models is not None: + for model in _models: + model_id = model.get("model_info", {}).get("id", None) + if model_id is not None: + team_accessible_model_ids.add(model_id) # Also search database for models accessible to this team # This complements the config search done above diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index d65df0087ad..c3e8d4429cd 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -896,6 +896,37 @@ def test_add_team_models_to_all_models(): assert result == {"gpt-4-model-2": {"team1"}} +def test_add_team_models_to_all_models_resolves_access_groups(): + """ + Team with models = ["on-prem"] (access group name) should get model ids + for models in that access group. + """ + 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 = "team-1" + team.models = ["on-prem"] + + llm_router = MagicMock() + llm_router.get_model_access_groups.return_value = { + "on-prem": ["gemma3-4b", "qwen3-embedding-8b"], + } + llm_router.get_model_list.side_effect = lambda model_name, team_id: [ + {"model_info": {"id": f"id-{model_name}", "team_id": None}}, + ] + + result = _add_team_models_to_all_models( + team_db_objects_typed=[team], + llm_router=llm_router, + ) + + assert result == { + "id-gemma3-4b": {"team-1"}, + "id-qwen3-embedding-8b": {"team-1"}, + } + + @pytest.mark.asyncio async def test_delete_deployment_type_mismatch(): """