fix(proxy): prevent custom auth fail-open fallback

This commit is contained in:
jeevan6996 2026-09-01 18:34:28 +01:00
parent 2d19f0f65b
commit 5f4828b65b
4 changed files with 81 additions and 5 deletions

View file

@ -59,6 +59,7 @@ class UserAPIKeyAuthExceptionHandler:
parent_otel_span: Span | None,
api_key: str,
resolved_identity: UserAPIKeyAuth | None = None,
is_custom_auth_error: bool = False,
) -> UserAPIKeyAuth:
"""
Handles Connection Errors when reading a Virtual Key from LiteLLM DB
@ -81,6 +82,7 @@ class UserAPIKeyAuthExceptionHandler:
if (
PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
and not is_custom_auth_error
and PrismaDBExceptionHandler.is_database_connection_error(e)
):
# log this as a DB failure on prometheus

View file

@ -1231,6 +1231,7 @@ async def _user_api_key_auth_builder(
route: Final[str] = get_request_route(request=request)
valid_token: UserAPIKeyAuth | None = None
custom_auth_api_key: bool = False
custom_auth_error: bool = False
try:
with tracer.trace("litellm.proxy.auth.pre_db_read_auth_checks"):
@ -1269,10 +1270,14 @@ async def _user_api_key_auth_builder(
### USER-DEFINED AUTH FUNCTION ###
if enterprise_custom_auth is not None:
with tracer.trace("litellm.proxy.auth.enterprise_custom_auth"):
response = await enterprise_custom_auth(
request=request, api_key=api_key, user_custom_auth=user_custom_auth
)
try:
with tracer.trace("litellm.proxy.auth.enterprise_custom_auth"):
response = await enterprise_custom_auth(
request=request, api_key=api_key, user_custom_auth=user_custom_auth
)
except Exception:
custom_auth_error = True
raise
if response is not None and isinstance(response, UserAPIKeyAuth):
validated = UserAPIKeyAuth.model_validate(response)
if getattr(litellm, "enable_post_custom_auth_checks", False):
@ -1288,7 +1293,11 @@ async def _user_api_key_auth_builder(
api_key = response
custom_auth_api_key = True
elif user_custom_auth is not None:
response = await user_custom_auth(request=request, api_key=api_key)
try:
response = await user_custom_auth(request=request, api_key=api_key)
except Exception:
custom_auth_error = True
raise
validated = UserAPIKeyAuth.model_validate(response)
if getattr(litellm, "enable_post_custom_auth_checks", False):
validated = await _run_post_custom_auth_checks(
@ -2281,6 +2290,7 @@ async def _user_api_key_auth_builder(
parent_otel_span=parent_otel_span,
api_key=api_key,
resolved_identity=valid_token,
is_custom_auth_error=custom_auth_error,
)

View file

@ -66,6 +66,30 @@ async def test_handle_authentication_error_db_unavailable_connectivity(db_error)
assert result.token == "failed-to-connect-to-db"
@pytest.mark.asyncio
async def test_handle_custom_auth_transport_error_does_not_fall_back():
"""Custom-authenticator outages must not issue the DB fail-open identity."""
handler = UserAPIKeyAuthExceptionHandler()
with patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
):
with pytest.raises(ProxyException) as exc_info:
await handler._handle_authentication_error(
httpx.ConnectError("custom authenticator unavailable"),
MagicMock(),
{},
"/test",
None,
"test-key",
is_custom_auth_error=True,
)
assert exc_info.value.type == ProxyErrorTypes.no_db_connection
assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"prisma_error",

View file

@ -6,6 +6,7 @@ from types import SimpleNamespace
from unittest.mock import ANY, AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import status
@ -537,6 +538,45 @@ async def test_user_custom_auth_skips_post_custom_auth_checks_by_default():
litellm.enable_post_custom_auth_checks = original_flag
@pytest.mark.asyncio
async def test_user_custom_auth_transport_error_does_not_use_db_fallback():
"""An unavailable custom authenticator must not receive the DB fallback identity."""
from fastapi import Request
from starlette.datastructures import URL
import litellm.proxy.proxy_server as _proxy_server_mod
mock_user_custom_auth = AsyncMock(side_effect=httpx.ConnectError("auth service unavailable"))
attrs = _proxy_server_attrs_for_custom_auth(user_custom_auth=mock_user_custom_auth)
attrs["general_settings"] = {"allow_requests_on_db_unavailable": True}
originals = {attr: getattr(_proxy_server_mod, attr, None) for attr in attrs}
try:
for attr, val in attrs.items():
setattr(_proxy_server_mod, attr, val)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with pytest.raises(ProxyException) as exc_info:
await _user_api_key_auth_builder(
request=request,
api_key="Bearer sk-custom-auth-unavailable",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
assert exc_info.value.type == ProxyErrorTypes.no_db_connection
assert int(exc_info.value.code) == status.HTTP_503_SERVICE_UNAVAILABLE
mock_user_custom_auth.assert_awaited_once()
finally:
for attr, val in originals.items():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
async def test_user_custom_auth_runs_post_custom_auth_checks_when_opt_in():
"""