mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Fix proxy auth error status codes
Co-authored-by: ishaan-berri <ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
3b78a3a545
commit
998620e27e
6 changed files with 80 additions and 7 deletions
|
|
@ -2849,7 +2849,7 @@ def _can_object_call_model(
|
|||
object_type=object_type
|
||||
),
|
||||
param="model",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -3082,7 +3082,7 @@ async def can_user_call_model(
|
|||
message=f"User not allowed to access model. No default model access, only team models allowed. Tried to access {model}",
|
||||
type=ProxyErrorTypes.key_model_access_denied,
|
||||
param="model",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
return _can_object_call_model(
|
||||
|
|
@ -3625,7 +3625,7 @@ async def _check_team_member_model_access(
|
|||
message=f"Team member not allowed to access model. User={valid_token.user_id}, Team={team_object.team_id}, Model={model}. Allowed member models = {member_allowed_models}",
|
||||
type=ProxyErrorTypes.team_model_access_denied,
|
||||
param="model",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -123,7 +123,7 @@ class UserAPIKeyAuthExceptionHandler:
|
|||
message=e.message,
|
||||
type=ProxyErrorTypes.budget_exceeded,
|
||||
param=None,
|
||||
code=400,
|
||||
code=getattr(e, "status_code", status.HTTP_429_TOO_MANY_REQUESTS),
|
||||
)
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
|
|
|
|||
|
|
@ -1107,7 +1107,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
raise ProxyException(
|
||||
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
||||
type=ProxyErrorTypes.expired_key,
|
||||
code=400,
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
param=abbreviate_api_key(api_key=api_key),
|
||||
)
|
||||
valid_token = update_valid_token_with_end_user_params(
|
||||
|
|
@ -1432,7 +1432,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
raise ProxyException(
|
||||
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
||||
type=ProxyErrorTypes.expired_key,
|
||||
code=400,
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
param=abbreviate_api_key(api_key=api_key),
|
||||
)
|
||||
|
||||
|
|
@ -2417,7 +2417,7 @@ async def _run_post_custom_auth_checks(
|
|||
raise ProxyException(
|
||||
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
||||
type=ProxyErrorTypes.expired_key,
|
||||
code=400,
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
param=(
|
||||
abbreviate_api_key(api_key=valid_token.token)
|
||||
if valid_token.token
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from datetime import datetime, timedelta
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import status
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -31,6 +32,7 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
ExperimentalUIJWTToken,
|
||||
_can_object_call_model,
|
||||
_can_object_call_vector_stores,
|
||||
_check_end_user_budget,
|
||||
_check_team_member_budget,
|
||||
|
|
@ -206,6 +208,52 @@ def test_get_key_object_from_ui_hash_key_invalid():
|
|||
assert key_object is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"object_type,expected_error_type",
|
||||
[
|
||||
("key", ProxyErrorTypes.key_model_access_denied),
|
||||
("team", ProxyErrorTypes.team_model_access_denied),
|
||||
("user", ProxyErrorTypes.user_model_access_denied),
|
||||
("org", ProxyErrorTypes.org_model_access_denied),
|
||||
("project", ProxyErrorTypes.project_model_access_denied),
|
||||
],
|
||||
)
|
||||
def test_can_object_call_model_denials_return_forbidden(
|
||||
object_type, expected_error_type
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
_can_object_call_model(
|
||||
model="restricted-model",
|
||||
llm_router=None,
|
||||
models=["allowed-model"],
|
||||
object_type=object_type,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == expected_error_type
|
||||
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_can_user_call_model_no_default_models_returns_forbidden():
|
||||
from litellm.proxy._types import SpecialModelNames
|
||||
from litellm.proxy.auth.auth_checks import can_user_call_model
|
||||
|
||||
user_object = LiteLLM_UserTable(
|
||||
user_id="test-user",
|
||||
models=[SpecialModelNames.no_default_models.value],
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await can_user_call_model(
|
||||
model="restricted-model",
|
||||
llm_router=None,
|
||||
user_object=user_object,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_should_reconnect_once_on_db_connection_error():
|
||||
mock_prisma_client = MagicMock()
|
||||
|
|
@ -1144,6 +1192,7 @@ async def test_check_team_member_model_access_denied_model():
|
|||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
|
||||
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -140,6 +140,7 @@ async def test_handle_authentication_error_budget_exceeded():
|
|||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
|
||||
assert int(exc_info.value.code) == status.HTTP_429_TOO_MANY_REQUESTS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -9,6 +10,7 @@ sys.path.insert(
|
|||
) # Adds the parent directory to the system path
|
||||
|
||||
import pytest
|
||||
from fastapi import status
|
||||
|
||||
import litellm
|
||||
import litellm.proxy.proxy_server
|
||||
|
|
@ -178,6 +180,26 @@ async def test_custom_auth_does_not_enforce_key_model_access_by_default():
|
|||
mock_can_key.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_custom_auth_expired_key_returns_unauthorized():
|
||||
expired_token = UserAPIKeyAuth(
|
||||
token="test_token",
|
||||
expires=datetime.now() - timedelta(minutes=1),
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _run_post_custom_auth_checks(
|
||||
valid_token=expired_token,
|
||||
request=MagicMock(),
|
||||
request_data={},
|
||||
route="/v1/chat/completions",
|
||||
parent_otel_span=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.expired_key
|
||||
assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_auth_honors_key_level_model_access_restriction_allowed_with_opt_in():
|
||||
valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"])
|
||||
|
|
@ -934,6 +956,7 @@ async def test_proxy_admin_expired_key_from_cache():
|
|||
assert (
|
||||
exc_info.value.type == ProxyErrorTypes.expired_key
|
||||
), f"Expected expired_key error type, got {exc_info.value.type}"
|
||||
assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED
|
||||
assert "Expired Key" in str(
|
||||
exc_info.value.message
|
||||
), f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue