mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fix: resolve model access groups when filtering team models in UI
This commit is contained in:
parent
9a2410be71
commit
82dcb6a042
2 changed files with 68 additions and 21 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue