[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.
This commit is contained in:
Ryan Crabbe 2026-03-28 12:32:12 -07:00
parent 666a31d47a
commit 8d9bbc6eb2
No known key found for this signature in database
2 changed files with 355 additions and 0 deletions

View file

@ -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():

View file

@ -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():
"""