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" 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 183b5609d03..f029511dd23 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -98,12 +98,19 @@ 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 _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, + ): + 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: @@ -971,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 @@ -1018,10 +1033,15 @@ 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.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: 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, ) @@ -1072,6 +1092,26 @@ 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", + ) + + +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, + object_type="team", ) @@ -1096,6 +1136,7 @@ async def can_user_call_model( model=model, llm_router=llm_router, models=user_object.models, + object_type="user", ) @@ -1248,53 +1289,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. 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, diff --git a/tests/otel_tests/test_e2e_model_access.py b/tests/otel_tests/test_e2e_model_access.py index 73c93212bf9..4628dc7e9c8 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 = {} @@ -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 @@ -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"] diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 8e9618297e4..0eb1a387558 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, expect_to_work): """ - 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 @@ -438,16 +438,16 @@ async def test_team_model_access_check(model, team_models, expect_to_work): 6. Empty model list 7. None model list """ - team_object = LiteLLM_TeamTable( - team_id="test-team", - models=team_models, - ) - try: - _team_model_access_check( + 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( 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