mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(proxy): direct-access models honor "all-proxy-models" sentinel and empty list
This commit is contained in:
parent
643989989f
commit
7ee3d373ce
2 changed files with 113 additions and 1 deletions
|
|
@ -10773,8 +10773,18 @@ def get_direct_access_models(
|
|||
llm_router: Router,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get all models that user has direct access to
|
||||
Get all models that user has direct access to.
|
||||
|
||||
If user.models is empty or contains the "all-proxy-models" sentinel,
|
||||
return all non-team model IDs — same semantic as the PROXY_ADMIN branch
|
||||
in get_all_team_and_direct_access_models and the "all model access"
|
||||
handling in auth_checks._check_model_access_helper.
|
||||
"""
|
||||
if (
|
||||
not user_db_object.models
|
||||
or SpecialModelNames.all_proxy_models.value in user_db_object.models
|
||||
):
|
||||
return llm_router.get_model_ids(exclude_team_models=True)
|
||||
|
||||
direct_access_models: List[str] = []
|
||||
for model in user_db_object.models:
|
||||
|
|
|
|||
|
|
@ -1145,6 +1145,108 @@ async def test_get_all_team_models():
|
|||
assert result == {"gpt-4-model-1": ["team1"], "gpt-4-model-2": ["team1"]}
|
||||
|
||||
|
||||
def test_get_direct_access_models_sentinel_returns_all_non_team_models():
|
||||
"""models == ["all-proxy-models"] returns all non-team model IDs (issue #22791)."""
|
||||
from litellm.proxy.proxy_server import get_direct_access_models
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.models = ["all-proxy-models"]
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_ids.return_value = ["m1", "m2", "m3"]
|
||||
|
||||
result = get_direct_access_models(user_db_object=mock_user, llm_router=mock_router)
|
||||
|
||||
assert result == ["m1", "m2", "m3"]
|
||||
mock_router.get_model_ids.assert_called_once_with(exclude_team_models=True)
|
||||
mock_router.get_model_list.assert_not_called()
|
||||
|
||||
|
||||
def test_get_direct_access_models_empty_list_returns_all_non_team_models():
|
||||
"""models == [] returns all non-team model IDs (Pfizer's case — the bug PR #22875 missed)."""
|
||||
from litellm.proxy.proxy_server import get_direct_access_models
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.models = []
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_ids.return_value = ["m1", "m2"]
|
||||
|
||||
result = get_direct_access_models(user_db_object=mock_user, llm_router=mock_router)
|
||||
|
||||
assert result == ["m1", "m2"]
|
||||
mock_router.get_model_ids.assert_called_once_with(exclude_team_models=True)
|
||||
mock_router.get_model_list.assert_not_called()
|
||||
|
||||
|
||||
def test_get_direct_access_models_specific_model_preserved():
|
||||
"""Explicit model names continue to resolve via get_model_list (existing behaviour)."""
|
||||
from litellm.proxy.proxy_server import get_direct_access_models
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.models = ["gpt-4o-mini"]
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_list.return_value = [
|
||||
{"model_info": {"id": "deploy-1"}},
|
||||
{"model_info": {"id": "deploy-2"}},
|
||||
]
|
||||
|
||||
result = get_direct_access_models(user_db_object=mock_user, llm_router=mock_router)
|
||||
|
||||
assert result == ["deploy-1", "deploy-2"]
|
||||
mock_router.get_model_list.assert_called_once_with(model_name="gpt-4o-mini")
|
||||
mock_router.get_model_ids.assert_not_called()
|
||||
|
||||
|
||||
def test_get_direct_access_models_unknown_model_returns_empty():
|
||||
"""Unknown model name returns [] (existing behaviour preserved)."""
|
||||
from litellm.proxy.proxy_server import get_direct_access_models
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.models = ["nonexistent-model"]
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_list.return_value = None
|
||||
|
||||
result = get_direct_access_models(user_db_object=mock_user, llm_router=mock_router)
|
||||
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_get_direct_access_models_sentinel_with_extra_names_wins():
|
||||
"""If "all-proxy-models" appears alongside other names, sentinel wins — no per-name lookups."""
|
||||
from litellm.proxy.proxy_server import get_direct_access_models
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.models = ["all-proxy-models", "gpt-4"]
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_ids.return_value = ["m1", "m2", "m3"]
|
||||
|
||||
result = get_direct_access_models(user_db_object=mock_user, llm_router=mock_router)
|
||||
|
||||
assert result == ["m1", "m2", "m3"]
|
||||
mock_router.get_model_ids.assert_called_once_with(exclude_team_models=True)
|
||||
mock_router.get_model_list.assert_not_called()
|
||||
|
||||
|
||||
def test_get_direct_access_models_excludes_team_models_for_all_access():
|
||||
"""exclude_team_models=True is passed so team-only models do not get a phantom direct_access flag."""
|
||||
from litellm.proxy.proxy_server import get_direct_access_models
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.models = []
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_ids.return_value = ["public-1", "public-2"]
|
||||
|
||||
result = get_direct_access_models(user_db_object=mock_user, llm_router=mock_router)
|
||||
|
||||
assert result == ["public-1", "public-2"]
|
||||
mock_router.get_model_ids.assert_called_once_with(exclude_team_models=True)
|
||||
|
||||
|
||||
def test_add_team_models_to_all_models():
|
||||
"""
|
||||
Test add_team_models_to_all_models function
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue