diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 14f198e0f12..054629a9c04 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3054,18 +3054,51 @@ async def can_key_call_model( Raises: - Exception: If token not allowed to call model """ + key_access_group_ids = valid_token.access_group_ids or [] + key_models = valid_token.models or [] + key_defers_to_team_models = ( + len(key_models) == 0 or SpecialModelNames.all_team_models.value in key_models + ) + + if key_access_group_ids and key_defers_to_team_models: + models_from_groups = await _get_models_from_access_groups( + access_group_ids=key_access_group_ids, + ) + if models_from_groups: + return _can_object_call_model( + model=model, + llm_router=llm_router, + models=models_from_groups, + team_model_aliases=valid_token.team_model_aliases, + team_id=valid_token.team_id, + object_type="key", + ) + verbose_proxy_logger.warning( + "Key has access_group_ids=%s, but those access groups resolved to no model permissions. " + "Denying model=%s for key_alias=%s, team_id=%s.", + key_access_group_ids, + model, + valid_token.key_alias, + valid_token.team_id, + ) + raise ProxyException( + message=f"key not allowed to access model. This key has access_group_ids={key_access_group_ids}, but those groups do not grant any models. Tried to access {model}", + type=ProxyErrorTypes.key_model_access_denied, + param="model", + code=status.HTTP_403_FORBIDDEN, + ) + try: return _can_object_call_model( model=model, llm_router=llm_router, - models=valid_token.models, + models=key_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, object_type="key", ) except ProxyException: # Fallback: check key's access_group_ids - key_access_group_ids = valid_token.access_group_ids or [] if key_access_group_ids: models_from_groups = await _get_models_from_access_groups( access_group_ids=key_access_group_ids, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 03278633928..737255d5b5b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2442,7 +2442,9 @@ async def _enforce_key_and_fallback_model_access( new_model_list = model_list verbose_proxy_logger.debug(f"\n new llm router model list {new_model_list}") elif ( - isinstance(valid_token.models, list) and "all-team-models" in valid_token.models + isinstance(valid_token.models, list) + and "all-team-models" in valid_token.models + and not valid_token.access_group_ids ): pass else: diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index d9f4a6e56b8..9535fe49ebe 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -1155,6 +1155,120 @@ async def test_can_key_call_model_via_access_group_ids(): ) +@pytest.mark.asyncio +async def test_can_key_call_model_restricts_empty_key_models_to_access_group_ids(): + """Keys with access_group_ids and no model list are restricted to those groups.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.auth_checks import can_key_call_model + + user_api_key_object = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=["ag-with-gpt4"], + ) + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "openai/gpt-4", "api_key": "test"}, + }, + { + "model_name": "claude-3", + "litellm_params": {"model": "anthropic/claude-3", "api_key": "test"}, + }, + ] + ) + + with patch( + "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + new_callable=AsyncMock, + return_value=["gpt-4"], + ): + with pytest.raises(ProxyException): + await can_key_call_model( + model="claude-3", + llm_model_list=[], + valid_token=user_api_key_object, + llm_router=router, + ) + + +@pytest.mark.asyncio +async def test_can_key_call_model_restricts_all_team_models_to_access_group_ids(): + """Keys with all-team-models and access_group_ids are restricted to those groups.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.auth_checks import can_key_call_model + + user_api_key_object = UserAPIKeyAuth( + token="test-token", + models=["all-team-models"], + access_group_ids=["ag-with-gpt4"], + team_models=["gpt-4", "claude-3"], + ) + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "openai/gpt-4", "api_key": "test"}, + }, + { + "model_name": "claude-3", + "litellm_params": {"model": "anthropic/claude-3", "api_key": "test"}, + }, + ] + ) + + with patch( + "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + new_callable=AsyncMock, + return_value=["gpt-4"], + ): + with pytest.raises(ProxyException): + await can_key_call_model( + model="claude-3", + llm_model_list=[], + valid_token=user_api_key_object, + llm_router=router, + ) + + +@pytest.mark.asyncio +async def test_can_key_call_model_denies_when_access_group_ids_resolve_no_models(): + """Keys with access_group_ids do not fall back to all models when groups are empty.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.auth_checks import can_key_call_model + + user_api_key_object = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=["empty-group"], + ) + + with ( + patch( + "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ), + patch("litellm.proxy.auth.auth_checks.verbose_proxy_logger.warning") as warning, + ): + with pytest.raises(ProxyException): + await can_key_call_model( + model="gpt-4", + llm_model_list=[], + valid_token=user_api_key_object, + llm_router=None, + ) + warning.assert_called_once() + assert "resolved to no model permissions" in warning.call_args.args[0] + + # --------------------------------------------------------------------------- # _key_access_group_grants_model (key access group overriding team restriction) # ---------------------------------------------------------------------------