This commit is contained in:
Jeevan Mohan Pawar 2026-09-23 14:49:34 +00:00 • committed by GitHub
commit b3fc1c3619
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 129 additions and 25 deletions

View file

@ -114,6 +114,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
@ -136,6 +137,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

@ -1499,6 +1499,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"):
@ -1537,10 +1538,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):
@ -1556,7 +1561,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(
@ -2540,6 +2549,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

@ -102,13 +102,17 @@ class PrismaDBExceptionHandler:
Reporting decisions want the opposite breadth; use
``is_database_infrastructure_error`` for those.
"""
import prisma.engine.errors
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
return True
if isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection:
return True
try:
import prisma.engine.errors
except ImportError:
return False
if isinstance(e, _exception_types(prisma.engine.errors.EngineConnectionError)):
return True
return isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection
return False
@staticmethod
def is_database_infrastructure_error(e: Exception) -> bool:
@ -126,8 +130,14 @@ class PrismaDBExceptionHandler:
``RecordNotFoundError``, etc.) are excluded — the DB IS reachable and
the request itself is what failed.
"""
import prisma
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
return True
if isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection:
return True
try:
import prisma
except ImportError:
return False
data_layer_errors: Final = _exception_types(
prisma.errors.DataError,
prisma.errors.UniqueViolationError,
@ -139,12 +149,8 @@ class PrismaDBExceptionHandler:
)
if isinstance(e, data_layer_errors):
return False
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
return True
if isinstance(e, _exception_types(prisma.errors.PrismaError)):
return True
if isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection:
return True
return False
@staticmethod
@ -166,8 +172,10 @@ class PrismaDBExceptionHandler:
per-row data rejection has to additionally consult
``is_database_service_unavailable_error`` before acting on a True here.
"""
import prisma
try:
import prisma
except ImportError:
return False
return type(e) is prisma.errors.DataError
@staticmethod
@ -179,10 +187,14 @@ class PrismaDBExceptionHandler:
Use this for reconnect logic — data-layer errors like UniqueViolationError
mean the DB IS reachable, so reconnecting would be pointless.
"""
import prisma
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
return True
if isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection:
return True
try:
import prisma
except ImportError:
return False
if isinstance(
e,
_exception_types(
@ -209,21 +221,24 @@ class PrismaDBExceptionHandler:
)
if any(keyword in error_message for keyword in connection_keywords):
return True
if isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection:
return True
return False
@staticmethod
def is_prisma_error(e: Exception) -> bool:
import prisma
try:
import prisma
except ImportError:
return False
return isinstance(e, _exception_types(prisma.errors.PrismaError))
@staticmethod
def is_deadlock_error(e: Exception) -> bool:
"""True iff ``e`` is a Postgres deadlock (P2034 / 40P01) surfaced through prisma."""
import prisma
try:
import prisma
except ImportError:
return False
if not isinstance(e, _exception_types(prisma.errors.PrismaError)):
return False
if getattr(e, "code", None) == "P2034":
@ -279,8 +294,10 @@ class PrismaDBExceptionHandler:
are already classified by type/keyword above, and data-layer ones
(the DB IS reachable) must stay 401.
"""
import prisma
try:
import prisma
except ImportError:
return False
if isinstance(e, _exception_types(prisma.errors.PrismaError)):
return False
tb = e.__traceback__ if hasattr(e, "__traceback__") else None

View file

@ -14,6 +14,7 @@ from types import SimpleNamespace
from unittest.mock import ANY, AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException, status
@ -586,6 +587,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():
"""

View file

@ -78,6 +78,18 @@ def test_is_database_transport_error_non_connection_prisma_errors(prisma_error):
assert PrismaDBExceptionHandler.is_database_transport_error(prisma_error) == False
@pytest.mark.parametrize(
"transport_error",
[
httpx.ConnectError("connection refused"),
ClientNotConnectedError(),
HTTPClientClosedError(),
],
)
def test_is_database_transport_error_connection_errors(transport_error):
assert PrismaDBExceptionHandler.is_database_transport_error(transport_error) is True
def test_is_database_connection_generic_errors():
"""
Test non-Prisma error cases for database connection checking
@ -99,6 +111,8 @@ def test_is_database_connection_generic_errors():
PrismaDBExceptionHandler.is_database_connection_error(db_proxy_exception)
== True
)
assert PrismaDBExceptionHandler.is_database_infrastructure_error(db_proxy_exception)
assert PrismaDBExceptionHandler.is_database_transport_error(db_proxy_exception)
# Test with non-DB error
regular_exception = Exception("Regular error")
@ -108,6 +122,27 @@ def test_is_database_connection_generic_errors():
)
def test_db_less_proxy_auth_error_does_not_require_prisma():
"""A missing Prisma dependency must not mask a normal authentication error."""
original_import = __import__
def import_without_prisma(name, globals=None, locals=None, fromlist=(), level=0):
if name == "prisma" or name.startswith("prisma."):
raise ModuleNotFoundError("No module named 'prisma'")
return original_import(name, globals, locals, fromlist, level)
with patch("builtins.__import__", side_effect=import_without_prisma):
error = Exception("No api key passed in.")
assert PrismaDBExceptionHandler.is_database_connection_error(error) is False
assert PrismaDBExceptionHandler.is_database_infrastructure_error(error) is False
assert PrismaDBExceptionHandler.is_prisma_data_error(error) is False
assert PrismaDBExceptionHandler.is_prisma_error(error) is False
assert PrismaDBExceptionHandler.is_database_transport_error(error) is False
assert PrismaDBExceptionHandler.is_deadlock_error(error) is False
assert PrismaDBExceptionHandler.is_prisma_engine_internal_error(error) is False
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is False
@pytest.mark.parametrize(
"error",
[