mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(proxy): expand model access groups in UI model listing for team/user grants
This commit is contained in:
parent
0cd588ad10
commit
6f32b8f57a
2 changed files with 119 additions and 6 deletions
|
|
@ -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}}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue