diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 6ddf2cfeb20..570b6742f57 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2949,6 +2949,26 @@ async def _get_agent_ids_from_access_groups( ) +def _resolve_all_team_model_sentinel_for_auth_check( + models: List[str], + llm_router: Optional[Router], + team_id: Optional[str], +) -> List[str]: + if ( + SpecialModelNames.all_team_models.value not in models + or team_id is None + or llm_router is None + ): + return models + proxy_models = llm_router.get_model_names() + non_sentinel_models = [ + model for model in models if model != SpecialModelNames.all_team_models.value + ] + if not proxy_models: + return non_sentinel_models or models + return list(dict.fromkeys(non_sentinel_models + proxy_models)) + + def _check_model_access_helper( model: str, llm_router: Optional[Router], @@ -2966,6 +2986,12 @@ def _check_model_access_helper( model_name=model, team_id=team_id ) + models = _resolve_all_team_model_sentinel_for_auth_check( + models=models, + llm_router=llm_router, + team_id=team_id, + ) + if ( len(access_groups) > 0 and llm_router is not None ): # check if token contains any model access groups diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 6e477ac10a4..aa53954da8f 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -123,11 +123,13 @@ def get_key_models( and user_api_key_dict.team_id is not None ): all_models = list(user_api_key_dict.team_models) - # GH#30619: if team_models also contains all-team-models, - # expand to actual proxy model list instead of leaking - # the sentinel string into /model/info if SpecialModelNames.all_team_models.value in all_models: - all_models = list(proxy_model_list) + all_models = [ + model + for model in all_models + if model != SpecialModelNames.all_team_models.value + ] + all_models.extend(proxy_model_list) if include_model_access_groups: all_models.extend(model_access_groups.keys()) if SpecialModelNames.all_proxy_models.value in all_models: diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 52634cc25fe..2dcf040a556 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -351,6 +351,43 @@ async def test_can_key_call_model_all_team_models_no_team_id_is_denied(): assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied +@pytest.mark.asyncio +async def test_can_team_access_model_all_team_models_expands_router_models(): + from litellm import Router + from litellm.proxy._types import SpecialModelNames + from litellm.proxy.auth.auth_checks import can_team_access_model + + team_object = LiteLLM_TeamTable( + team_id="team-123", + models=[SpecialModelNames.all_team_models.value], + ) + router = Router( + model_list=[ + { + "model_name": "allowed-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}, + } + ] + ) + + assert ( + await can_team_access_model( + model="allowed-model", + team_object=team_object, + llm_router=router, + ) + is True + ) + with pytest.raises(ProxyException) as exc_info: + await can_team_access_model( + model="blocked-model", + team_object=team_object, + llm_router=router, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + @pytest.mark.asyncio async def test_get_key_object_should_reconnect_once_on_db_connection_error(): mock_prisma_client = MagicMock() diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index d35e4848ef7..261485e8965 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -565,6 +565,27 @@ def test_get_key_models_all_team_models_recursive_team(): assert set(result) == {"model-a", "model-b"} +def test_get_key_models_all_team_models_keeps_mixed_team_entries(): + from litellm.proxy.auth.model_checks import get_key_models + from litellm.proxy._types import SpecialModelNames + + user_api_key_dict = type( + "obj", + (object,), + { + "models": [SpecialModelNames.all_team_models.value], + "team_id": "team-1", + "team_models": [ + SpecialModelNames.all_team_models.value, + "restricted-model", + ], + }, + )() + result = get_key_models(user_api_key_dict, ["model-a", "model-b"], {}) + assert SpecialModelNames.all_team_models.value not in result + assert set(result) == {"model-a", "model-b", "restricted-model"} + + def test_get_team_models_all_team_models_expands(): """GH#30619: all-team-models in team_models should expand.""" from litellm.proxy.auth.model_checks import get_team_models