From 8d9bbc6eb2f004538a6291bf38ed88c0f304a9f2 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Sat, 28 Mar 2026 12:32:12 -0700 Subject: [PATCH] [Fix] Include access group models in UI model listing Models associated with a team only through access groups (not directly in team.models) were not appearing on the /ui/?page=models page. The API authorization path already resolved access groups correctly, but the /v2/model/info listing endpoint only checked team.models. Add _add_access_group_models_to_team_models() which batch-fetches all distinct access groups in a single find_many query, then resolves each team's access group models into deployments and merges them into the team_models dict. --- litellm/proxy/proxy_server.py | 77 +++++ tests/test_litellm/proxy/test_proxy_server.py | 278 ++++++++++++++++++ 2 files changed, 355 insertions(+) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 853ceb63c9d..4f0a25296b2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9249,6 +9249,75 @@ def _add_team_models_to_all_models( return team_models +async def _add_access_group_models_to_team_models( + team_db_objects_typed: List[LiteLLM_TeamTable], + llm_router: Router, + prisma_client: PrismaClient, + team_models: Dict[str, Set[str]], +) -> Dict[str, Set[str]]: + """ + Resolve models reachable via team access groups and merge them into team_models. + + Batch-fetches all distinct access groups in a single DB query, then resolves + each eligible team's access group models via the pre-fetched map. + + This ensures models associated with a team only through access groups + (not directly in team.models) are included in the UI model listing. + """ + # First pass: identify eligible teams and collect all distinct access group IDs + eligible_teams: List[LiteLLM_TeamTable] = [] + all_access_group_ids: Set[str] = set() + + for team_object in team_db_objects_typed: + if not team_object.access_group_ids: + continue + + # Skip teams with empty models list — they already have access to everything + # (handled by _add_team_models_to_all_models) + if ( + len(team_object.models) == 0 + or SpecialModelNames.all_proxy_models.value in team_object.models + ): + continue + + eligible_teams.append(team_object) + all_access_group_ids.update(team_object.access_group_ids) + + if not eligible_teams: + return team_models + + # Single batch fetch for all access groups + access_group_rows = ( + await prisma_client.db.litellm_accessgrouptable.find_many( + where={"access_group_id": {"in": list(all_access_group_ids)}} + ) + ) + ag_model_map: Dict[str, List[str]] = { + row.access_group_id: getattr(row, "access_model_names", []) or [] + for row in access_group_rows + } + + # Second pass: resolve deployments for each eligible team + for team_object in eligible_teams: + model_names: Set[str] = set() + for ag_id in team_object.access_group_ids: + model_names.update(ag_model_map.get(ag_id, [])) + + for model_name in model_names: + deployments = llm_router.get_model_list( + model_name=model_name, team_id=team_object.team_id + ) + if deployments is not None: + for deployment in deployments: + model_id = deployment.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 + + async def get_all_team_models( user_teams: Union[List[str], Literal["*"]], prisma_client: PrismaClient, @@ -9285,6 +9354,14 @@ async def get_all_team_models( llm_router=llm_router, ) + # Also resolve models reachable via team access groups + team_models = await _add_access_group_models_to_team_models( + team_db_objects_typed=team_db_objects_typed, + llm_router=llm_router, + prisma_client=prisma_client, + team_models=team_models, + ) + # convert set to list returned_team_models: Dict[str, List[str]] = {} for model_id, team_ids in team_models.items(): diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 11bae9009a2..daabed0def1 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -988,6 +988,7 @@ async def test_get_all_team_models(): mock_instance = MagicMock() mock_instance.team_id = kwargs["team_id"] mock_instance.models = kwargs["models"] + mock_instance.access_group_ids = kwargs.get("access_group_ids") return mock_instance mock_team_table_class.side_effect = mock_team_table_constructor @@ -1108,6 +1109,283 @@ def test_add_team_models_to_all_models(): assert result == {"gpt-4-model-2": {"team1"}} +@pytest.mark.asyncio +async def test_add_access_group_models_to_team_models(): + """ + Test that models reachable via team access groups are included in team_models. + + Scenario: A team has models=["gpt-4"] and access_group_ids=["premium"]. + The "premium" access group contains ["claude-3", "gemini"]. + After resolution, the team should see gpt-4 (direct) + claude-3/gemini (via access group). + """ + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.proxy_server import _add_access_group_models_to_team_models + + # Team with specific models AND access groups + team_with_access_groups = MagicMock(spec=LiteLLM_TeamTable) + team_with_access_groups.team_id = "team1" + team_with_access_groups.models = ["gpt-4"] # non-empty = specific models + team_with_access_groups.access_group_ids = ["premium"] + + # Team with no access groups — should be skipped + team_without_access_groups = MagicMock(spec=LiteLLM_TeamTable) + team_without_access_groups.team_id = "team2" + team_without_access_groups.models = ["gpt-4"] + team_without_access_groups.access_group_ids = None + + # Team with empty access_group_ids list — should be skipped + team_empty_access_groups = MagicMock(spec=LiteLLM_TeamTable) + team_empty_access_groups.team_id = "team2b" + team_empty_access_groups.models = ["gpt-4"] + team_empty_access_groups.access_group_ids = [] + + # Team with empty models (all access) — should be skipped + team_all_access = MagicMock(spec=LiteLLM_TeamTable) + team_all_access.team_id = "team3" + team_all_access.models = [] + team_all_access.access_group_ids = ["premium"] + + # Team with all-proxy-models sentinel (all access) — should be skipped + team_all_proxy = MagicMock(spec=LiteLLM_TeamTable) + team_all_proxy.team_id = "team4" + team_all_proxy.models = ["all-proxy-models"] + team_all_proxy.access_group_ids = ["premium"] + + # Mock router + mock_router = MagicMock() + + def mock_get_model_list(model_name, team_id=None): + if model_name == "claude-3": + return [{"model_info": {"id": "claude-3-id"}}] + elif model_name == "gemini": + return [{"model_info": {"id": "gemini-id"}}] + return None + + mock_router.get_model_list.side_effect = mock_get_model_list + + # Pre-existing team_models (e.g., from _add_team_models_to_all_models) + existing_team_models = { + "gpt-4-id": {"team1"}, + } + + # Mock prisma client with batch find_many returning access group rows + mock_ag_row = MagicMock() + mock_ag_row.access_group_id = "premium" + mock_ag_row.access_model_names = ["claude-3", "gemini"] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock( + return_value=[mock_ag_row] + ) + + result = await _add_access_group_models_to_team_models( + team_db_objects_typed=[ + team_with_access_groups, + team_without_access_groups, + team_empty_access_groups, + team_all_access, + team_all_proxy, + ], + llm_router=mock_router, + prisma_client=mock_prisma_client, + team_models=existing_team_models, + ) + + # Single batch query with only the eligible team's access group IDs + mock_prisma_client.db.litellm_accessgrouptable.find_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_accessgrouptable.find_many.call_args + queried_ids = call_args[1]["where"]["access_group_id"]["in"] + assert set(queried_ids) == {"premium"} + + # Original model still present + assert "gpt-4-id" in result + assert "team1" in result["gpt-4-id"] + + # Access group models added for team1 + assert "claude-3-id" in result + assert "team1" in result["claude-3-id"] + assert "gemini-id" in result + assert "team1" in result["gemini-id"] + + # Skipped teams should NOT have added these models + for skipped_team in ["team2", "team2b", "team3", "team4"]: + assert skipped_team not in result.get("claude-3-id", set()) + assert skipped_team not in result.get("gemini-id", set()) + + +@pytest.mark.asyncio +async def test_add_access_group_models_multiple_teams_shared_group(): + """ + Test that multiple teams sharing the same access group each get the models, + and only one batch DB query is made. + """ + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.proxy_server import _add_access_group_models_to_team_models + + team_a = MagicMock(spec=LiteLLM_TeamTable) + team_a.team_id = "team-a" + team_a.models = ["gpt-4"] + team_a.access_group_ids = ["shared-group"] + + team_b = MagicMock(spec=LiteLLM_TeamTable) + team_b.team_id = "team-b" + team_b.models = ["gpt-3.5"] + team_b.access_group_ids = ["shared-group", "extra-group"] + + mock_router = MagicMock() + + def mock_get_model_list(model_name, team_id=None): + if model_name == "claude-3": + return [{"model_info": {"id": "claude-3-id"}}] + elif model_name == "gemini": + return [{"model_info": {"id": "gemini-id"}}] + return None + + mock_router.get_model_list.side_effect = mock_get_model_list + + mock_shared_row = MagicMock() + mock_shared_row.access_group_id = "shared-group" + mock_shared_row.access_model_names = ["claude-3"] + + mock_extra_row = MagicMock() + mock_extra_row.access_group_id = "extra-group" + mock_extra_row.access_model_names = ["gemini"] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock( + return_value=[mock_shared_row, mock_extra_row] + ) + + result = await _add_access_group_models_to_team_models( + team_db_objects_typed=[team_a, team_b], + llm_router=mock_router, + prisma_client=mock_prisma_client, + team_models={}, + ) + + # Single batch query for both groups + mock_prisma_client.db.litellm_accessgrouptable.find_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_accessgrouptable.find_many.call_args + queried_ids = set(call_args[1]["where"]["access_group_id"]["in"]) + assert queried_ids == {"shared-group", "extra-group"} + + # Both teams get claude-3 from the shared group + assert "claude-3-id" in result + assert "team-a" in result["claude-3-id"] + assert "team-b" in result["claude-3-id"] + + # Only team-b gets gemini (from extra-group) + assert "gemini-id" in result + assert "team-b" in result["gemini-id"] + assert "team-a" not in result["gemini-id"] + + +@pytest.mark.asyncio +async def test_add_access_group_models_no_eligible_teams(): + """ + When no teams have access groups, find_many should not be called at all. + """ + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.proxy_server import _add_access_group_models_to_team_models + + team = MagicMock(spec=LiteLLM_TeamTable) + team.team_id = "team1" + team.models = ["gpt-4"] + team.access_group_ids = None + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock() + + result = await _add_access_group_models_to_team_models( + team_db_objects_typed=[team], + llm_router=MagicMock(), + prisma_client=mock_prisma_client, + team_models={"existing-id": {"team1"}}, + ) + + # No DB call made + mock_prisma_client.db.litellm_accessgrouptable.find_many.assert_not_called() + + # Original data unchanged + assert result == {"existing-id": {"team1"}} + + +@pytest.mark.asyncio +async def test_get_all_team_models_with_access_groups(): + """ + End-to-end test: get_all_team_models includes models from access groups. + + Scenario: User is on team1 which has models=["gpt-4"] and + access_group_ids=["premium"]. The "premium" group has ["claude-3"]. + The result should include both gpt-4 and claude-3 deployments for team1. + """ + from litellm.proxy.proxy_server import get_all_team_models + + mock_team1 = MagicMock() + mock_team1.model_dump.return_value = { + "team_id": "team1", + "models": ["gpt-4"], + "team_alias": "Team 1", + "access_group_ids": ["premium"], + } + + # Mock access group row returned by batch find_many + mock_ag_row = MagicMock() + mock_ag_row.access_group_id = "premium" + mock_ag_row.access_model_names = ["claude-3"] + + mock_prisma_client = MagicMock() + mock_db = MagicMock() + mock_litellm_teamtable = MagicMock() + mock_prisma_client.db = mock_db + mock_db.litellm_teamtable = mock_litellm_teamtable + mock_litellm_teamtable.find_many = AsyncMock(return_value=[mock_team1]) + mock_db.litellm_accessgrouptable = MagicMock() + mock_db.litellm_accessgrouptable.find_many = AsyncMock( + return_value=[mock_ag_row] + ) + + mock_router = MagicMock() + + def mock_get_model_list(model_name, team_id=None): + if model_name == "gpt-4": + return [{"model_info": {"id": "gpt-4-deploy-1"}}] + elif model_name == "claude-3": + return [{"model_info": {"id": "claude-3-deploy-1"}}] + return None + + mock_router.get_model_list.side_effect = mock_get_model_list + + with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_tt_class: + + def mock_team_table_constructor(**kwargs): + mock_instance = MagicMock() + mock_instance.team_id = kwargs["team_id"] + mock_instance.models = kwargs["models"] + mock_instance.access_group_ids = kwargs.get("access_group_ids") + return mock_instance + + mock_tt_class.side_effect = mock_team_table_constructor + + result = await get_all_team_models( + user_teams=["team1"], + prisma_client=mock_prisma_client, + llm_router=mock_router, + ) + + # gpt-4 from direct team.models + assert "gpt-4-deploy-1" in result + assert "team1" in result["gpt-4-deploy-1"] + + # claude-3 from access group + assert "claude-3-deploy-1" in result + assert "team1" in result["claude-3-deploy-1"] + + # Return type is Dict[str, List[str]] + assert isinstance(result["gpt-4-deploy-1"], list) + assert isinstance(result["claude-3-deploy-1"], list) + + @pytest.mark.asyncio async def test_delete_deployment_type_mismatch(): """