From aa5ac6ba3db6d5d776b302262e2d28071e0ee70f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 10 Mar 2025 20:03:19 -0700 Subject: [PATCH] can_team_access_model --- litellm/proxy/_types.py | 15 +++++++++ litellm/proxy/auth/auth_checks.py | 39 ++++++++++++++--------- tests/otel_tests/test_e2e_model_access.py | 10 ++---- 3 files changed, 41 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 6bf2ef9068..95931c06b8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2054,6 +2054,7 @@ class ProxyErrorTypes(str, enum.Enum): budget_exceeded = "budget_exceeded" key_model_access_denied = "key_model_access_denied" team_model_access_denied = "team_model_access_denied" + user_model_access_denied = "user_model_access_denied" expired_key = "expired_key" auth_error = "auth_error" internal_server_error = "internal_server_error" @@ -2062,6 +2063,20 @@ class ProxyErrorTypes(str, enum.Enum): validation_error = "bad_request_error" cache_ping_error = "cache_ping_error" + @classmethod + def get_model_access_error_type_for_object( + cls, object_type: Literal["key", "user", "team"] + ) -> "ProxyErrorTypes": + """ + Get the model access error type for object_type + """ + if object_type == "key": + return cls.key_model_access_denied + elif object_type == "team": + return cls.team_model_access_denied + elif object_type == "user": + return cls.user_model_access_denied + DB_CONNECTION_ERROR_TYPES = (httpx.ConnectError, httpx.ReadError, httpx.ReadTimeout) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3faf8c0107..f029511dd2 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -98,23 +98,19 @@ async def common_checks( ) # 2. If team can call model - if ( - team_object is not None - and _model is not None - and can_team_access_model( + if _model and team_object: + if not await can_team_access_model( model=_model, team_object=team_object, llm_router=llm_router, team_model_aliases=valid_token.team_model_aliases if valid_token else None, - ) - is False - ): - raise ProxyException( - message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", - type=ProxyErrorTypes.team_model_access_denied, - param="model", - code=status.HTTP_401_UNAUTHORIZED, - ) + ): + raise ProxyException( + message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", + type=ProxyErrorTypes.team_model_access_denied, + param="model", + code=status.HTTP_401_UNAUTHORIZED, + ) ## 2.1 If user can call model (if personal key) if team_object is None and user_object is not None: @@ -982,10 +978,18 @@ async def _can_object_call_model( llm_router: Optional[Router], models: List[str], team_model_aliases: Optional[Dict[str, str]] = None, + object_type: Literal["user", "team", "key"] = "user", ) -> Literal[True]: """ Checks if token can call a given model + Args: + - model: str + - llm_router: Optional[Router] + - models: List[str] + - team_model_aliases: Optional[Dict[str, str]] + - object_type: Literal["user", "team", "key"]. We use the object type to raise the correct exception type + Returns: - True: if token allowed to call model @@ -1034,8 +1038,10 @@ async def _can_object_call_model( if model is not None and model not in filtered_models and all_model_access is False: raise ProxyException( - message=f"API Key not allowed to access model. This token can only access models={models}. Tried to access {model}", - type=ProxyErrorTypes.key_model_access_denied, + message=f"{object_type} not allowed to access model. This {object_type} can only access models={models}. Tried to access {model}", + type=ProxyErrorTypes.get_model_access_error_type_for_object( + object_type=object_type + ), param="model", code=status.HTTP_401_UNAUTHORIZED, ) @@ -1086,6 +1092,7 @@ async def can_key_call_model( llm_router=llm_router, models=valid_token.models, team_model_aliases=valid_token.team_model_aliases, + object_type="key", ) @@ -1104,6 +1111,7 @@ async def can_team_access_model( llm_router=llm_router, models=team_object.models if team_object else [], team_model_aliases=team_model_aliases, + object_type="team", ) @@ -1128,6 +1136,7 @@ async def can_user_call_model( model=model, llm_router=llm_router, models=user_object.models, + object_type="user", ) diff --git a/tests/otel_tests/test_e2e_model_access.py b/tests/otel_tests/test_e2e_model_access.py index 8b633afefd..c4846c2478 100644 --- a/tests/otel_tests/test_e2e_model_access.py +++ b/tests/otel_tests/test_e2e_model_access.py @@ -159,12 +159,6 @@ async def test_model_access_update(): "team_models, test_model, expect_success", [ (["openai/*"], "anthropic/claude-2", False), # Non-matching model - (["gpt-4"], "gpt-4", True), # Exact model match - (["bedrock/*"], "bedrock/anthropic.claude-3", True), # Bedrock wildcard - (["bedrock/anthropic.*"], "bedrock/anthropic.claude-3", True), # Pattern match - (["bedrock/anthropic.*"], "bedrock/amazon.titan", False), # Pattern non-match - (None, "gpt-4", True), # No model restrictions - ([], "gpt-4", True), # Empty model list ], ) @pytest.mark.asyncio @@ -285,6 +279,6 @@ def _validate_model_access_exception( assert _error_body["param"] == "model" assert _error_body["code"] == "401" if expected_type == "key_model_access_denied": - assert "API Key not allowed to access model" in _error_body["message"] + assert "key not allowed to access model" in _error_body["message"] elif expected_type == "team_model_access_denied": - assert "Team not allowed to access model" in _error_body["message"] + assert "eam not allowed to access model" in _error_body["message"]