Fix proxy auth error status codes

Co-authored-by: ishaan-berri <ishaan-berri@users.noreply.github.com>
This commit is contained in:
oss-agent-shin 2026-05-07 20:56:09 +00:00
parent 3b78a3a545
commit 998620e27e
No known key found for this signature in database
6 changed files with 80 additions and 7 deletions

View file

@ -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,
)

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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}"