From eeee61db658de558a914569567c048fb51278c0f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 25 Feb 2025 14:50:10 -0800 Subject: [PATCH 01/10] can_team_access_model --- litellm/proxy/auth/auth_checks.py | 92 +++++++++++++------------------ 1 file changed, 39 insertions(+), 53 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 0590bcb50a4..c922599f86e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -38,6 +38,7 @@ from litellm.proxy._types import ( ProxyErrorTypes, ProxyException, RoleBasedPermissions, + SpecialModelNames, UserAPIKeyAuth, ) from litellm.proxy.auth.route_checks import RouteChecks @@ -97,12 +98,23 @@ async def common_checks( ) # 2. If team can call model - _team_model_access_check( - team_object=team_object, - model=_model, - llm_router=llm_router, - team_model_aliases=valid_token.team_model_aliases if valid_token else None, - ) + if ( + team_object is not None + and _model is not None + and 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, + ) ## 2.1 If user can call model (if personal key) if team_object is None and user_object is not None: @@ -1017,6 +1029,9 @@ async def _can_object_call_model( if (len(filtered_models) == 0 and len(models) == 0) or "*" in filtered_models: all_model_access = True + if SpecialModelNames.all_proxy_models in filtered_models: + all_model_access = True + 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}", @@ -1074,6 +1089,24 @@ async def can_key_call_model( ) +async def can_team_access_model( + model: str, + team_object: Optional[LiteLLM_TeamTable], + llm_router: Optional[Router], + team_model_aliases: Optional[Dict[str, str]] = None, +) -> Literal[True]: + """ + Returns True if the team can access a specific model. + + """ + return await _can_object_call_model( + model=model, + llm_router=llm_router, + models=team_object.models if team_object else [], + team_model_aliases=team_model_aliases, + ) + + async def can_user_call_model( model: str, llm_router: Optional[Router], @@ -1239,53 +1272,6 @@ async def _team_max_budget_check( ) -def _team_model_access_check( - model: Optional[str], - team_object: Optional[LiteLLM_TeamTable], - llm_router: Optional[Router], - team_model_aliases: Optional[Dict[str, str]] = None, -): - """ - Access check for team models - Raises: - Exception if the team is not allowed to call the`model` - """ - if ( - model is not None - and team_object is not None - and team_object.models is not None - and len(team_object.models) > 0 - and model not in team_object.models - ): - # this means the team has access to all models on the proxy - if "all-proxy-models" in team_object.models or "*" in team_object.models: - # this means the team has access to all models on the proxy - pass - # check if the team model is an access_group - elif ( - model_in_access_group( - model=model, team_models=team_object.models, llm_router=llm_router - ) - is True - ): - pass - elif model and "*" in model: - pass - elif _model_in_team_aliases(model=model, team_model_aliases=team_model_aliases): - pass - elif _model_matches_any_wildcard_pattern_in_list( - model=model, allowed_model_list=team_object.models - ): - pass - else: - 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, - ) - - def is_model_allowed_by_pattern(model: str, allowed_model_pattern: str) -> bool: """ Check if a model matches an allowed pattern. From b6d6e270b49e72a49c3f8a6496a8e965e8eb55c6 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 25 Feb 2025 14:51:57 -0800 Subject: [PATCH 02/10] can_team_access_model --- litellm/proxy/auth/handle_jwt.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 29f4b31f9cd..248d553662b 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -33,6 +33,7 @@ from litellm.proxy._types import ( ScopeMapping, Span, ) +from litellm.proxy.auth.auth_checks import can_team_access_model from litellm.proxy.utils import PrismaClient, ProxyLogging from .auth_checks import ( @@ -723,8 +724,12 @@ class JWTAuthManager: team_models = team_object.models if isinstance(team_models, list) and ( not requested_model - or requested_model in team_models - or "*" in team_models + or can_team_access_model( + model=requested_model, + team_object=team_object, + llm_router=None, + team_model_aliases=None, + ) ): is_allowed = allowed_routes_check( user_role=LitellmUserRoles.TEAM, From 3d0b56e8a34b73baf5057e8b45fd6dcb2a558920 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 25 Feb 2025 14:57:13 -0800 Subject: [PATCH 03/10] test_can_team_access_model --- tests/proxy_unit_tests/test_auth_checks.py | 28 ++++++++-------------- 1 file changed, 10 insertions(+), 18 deletions(-) diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 0a8ebbe0185..5b79ace1b9d 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -27,7 +27,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.utils import PrismaClient from litellm.proxy.auth.auth_checks import ( - _team_model_access_check, + can_team_access_model, _virtual_key_soft_budget_check, ) from litellm.proxy.utils import ProxyLogging @@ -427,9 +427,9 @@ async def test_virtual_key_max_budget_check( ], ) @pytest.mark.asyncio -async def test_team_model_access_check(model, team_models, expect_to_work): +async def test_can_team_access_model(model, team_models, expected_result): """ - Test cases for _team_model_access_check: + Test cases for can_team_access_model: 1. Exact model match 2. all-proxy-models access 3. Wildcard (*) access @@ -443,21 +443,13 @@ async def test_team_model_access_check(model, team_models, expect_to_work): models=team_models, ) - try: - _team_model_access_check( - model=model, - team_object=team_object, - llm_router=None, - ) - if not expect_to_work: - pytest.fail( - f"Expected model access check to fail for model={model}, team_models={team_models}" - ) - except Exception as e: - if expect_to_work: - pytest.fail( - f"Expected model access check to work for model={model}, team_models={team_models}. Got error: {str(e)}" - ) + result = await can_team_access_model( + model=model, + team_object=team_object, + llm_router=None, + team_model_aliases=None, + ) + assert result == expected_result @pytest.mark.parametrize( From 7eaf0039193363fd562349d08df2a7744646a7f2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 25 Feb 2025 15:25:51 -0800 Subject: [PATCH 04/10] expected_result --- tests/proxy_unit_tests/test_auth_checks.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 5b79ace1b9d..a5782653ad7 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -394,7 +394,7 @@ async def test_virtual_key_max_budget_check( @pytest.mark.parametrize( - "model, team_models, expect_to_work", + "model, team_models, expected_result", [ ("gpt-4", ["gpt-4"], True), # exact match ("gpt-4", ["all-proxy-models"], True), # all-proxy-models access From 5ead81786dea2322b614118d0f4eac7c7843e3f0 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 1 Mar 2025 17:42:50 -0800 Subject: [PATCH 05/10] test_can_team_access_model --- tests/proxy_unit_tests/test_auth_checks.py | 36 +++++++++++++--------- 1 file changed, 22 insertions(+), 14 deletions(-) diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index ec36823633b..0eb1a387558 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -394,7 +394,7 @@ async def test_virtual_key_max_budget_check( @pytest.mark.parametrize( - "model, team_models, expected_result", + "model, team_models, expect_to_work", [ ("gpt-4", ["gpt-4"], True), # exact match ("gpt-4", ["all-proxy-models"], True), # all-proxy-models access @@ -427,7 +427,7 @@ async def test_virtual_key_max_budget_check( ], ) @pytest.mark.asyncio -async def test_can_team_access_model(model, team_models, expected_result): +async def test_can_team_access_model(model, team_models, expect_to_work): """ Test cases for can_team_access_model: 1. Exact model match @@ -438,18 +438,26 @@ async def test_can_team_access_model(model, team_models, expected_result): 6. Empty model list 7. None model list """ - team_object = LiteLLM_TeamTable( - team_id="test-team", - models=team_models, - ) - - result = await can_team_access_model( - model=model, - team_object=team_object, - llm_router=None, - team_model_aliases=None, - ) - assert result == expected_result + try: + team_object = LiteLLM_TeamTable( + team_id="test-team", + models=team_models, + ) + result = await can_team_access_model( + model=model, + team_object=team_object, + llm_router=None, + team_model_aliases=None, + ) + if not expect_to_work: + pytest.fail( + f"Expected model access check to fail for model={model}, team_models={team_models}" + ) + except Exception as e: + if expect_to_work: + pytest.fail( + f"Expected model access check to work for model={model}, team_models={team_models}. Got error: {str(e)}" + ) @pytest.mark.parametrize( From 8811d4cd1197e90175b272c3b976169f58c174be Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 10 Mar 2025 19:03:50 -0700 Subject: [PATCH 06/10] generate_key --- tests/otel_tests/test_e2e_model_access.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/otel_tests/test_e2e_model_access.py b/tests/otel_tests/test_e2e_model_access.py index 73c93212bf9..8b633afefd7 100644 --- a/tests/otel_tests/test_e2e_model_access.py +++ b/tests/otel_tests/test_e2e_model_access.py @@ -9,7 +9,7 @@ from typing import Any, Optional, List, Literal async def generate_key( session, models: Optional[List[str]] = None, team_id: Optional[str] = None ): - """Helper function to generate a key with specific model access""" + """Helper function to generate a key with specific model access controls""" url = "http://0.0.0.0:4000/key/generate" headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} data = {} From 0d6df360bfab29144e972af6a5fae136988c61a4 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 10 Mar 2025 19:09:50 -0700 Subject: [PATCH 07/10] test_can_team_access_model fix --- litellm/proxy/auth/auth_checks.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 4df84fb9b12..3faf8c0107b 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1029,7 +1029,7 @@ async def _can_object_call_model( if (len(filtered_models) == 0 and len(models) == 0) or "*" in filtered_models: all_model_access = True - if SpecialModelNames.all_proxy_models in filtered_models: + if SpecialModelNames.all_proxy_models.value in filtered_models: all_model_access = True if model is not None and model not in filtered_models and all_model_access is False: From aa5ac6ba3db6d5d776b302262e2d28071e0ee70f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 10 Mar 2025 20:03:19 -0700 Subject: [PATCH 08/10] 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 6bf2ef90683..95931c06b8d 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 3faf8c0107b..f029511dd23 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 8b633afefd7..c4846c2478e 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"] From 91a7fb0a2353f4b14ce89e35e412104b2a23f5b5 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 10 Mar 2025 20:04:18 -0700 Subject: [PATCH 09/10] test string checked for model access control --- tests/otel_tests/test_e2e_model_access.py | 2 +- tests/test_openai_endpoints.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/otel_tests/test_e2e_model_access.py b/tests/otel_tests/test_e2e_model_access.py index c4846c2478e..4628dc7e9c8 100644 --- a/tests/otel_tests/test_e2e_model_access.py +++ b/tests/otel_tests/test_e2e_model_access.py @@ -94,7 +94,7 @@ async def test_model_access_patterns(key_models, test_model, expect_success): assert _error_body["type"] == "key_model_access_denied" assert _error_body["param"] == "model" assert _error_body["code"] == "401" - assert "API Key not allowed to access model" in _error_body["message"] + assert "key not allowed to access model" in _error_body["message"] @pytest.mark.asyncio diff --git a/tests/test_openai_endpoints.py b/tests/test_openai_endpoints.py index 45fd29721f7..16b9838d80b 100644 --- a/tests/test_openai_endpoints.py +++ b/tests/test_openai_endpoints.py @@ -308,7 +308,7 @@ async def test_chat_completion(): model="gpt-4", messages=[{"role": "user", "content": "Hello!"}], ) - assert "API Key not allowed to access model." in str(e) + assert "key not allowed to access model." in str(e) @pytest.mark.asyncio From 24b8dcff1fe6c795890fa8bc21b74ea920ed2469 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 10 Mar 2025 20:45:58 -0700 Subject: [PATCH 10/10] get_complete_url --- litellm/llms/triton/completion/transformation.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/llms/triton/completion/transformation.py b/litellm/llms/triton/completion/transformation.py index 0a65e216dfe..4037c32365e 100644 --- a/litellm/llms/triton/completion/transformation.py +++ b/litellm/llms/triton/completion/transformation.py @@ -69,11 +69,13 @@ class TritonConfig(BaseConfig): def get_complete_url( self, - api_base: str, + api_base: Optional[str], model: str, optional_params: dict, stream: Optional[bool] = None, ) -> str: + if api_base is None: + raise ValueError("api_base is required") llm_type = self._get_triton_llm_type(api_base) if llm_type == "generate" and stream: return api_base + "_stream"