diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index 9bf7f6cab96..1b8e8f592c9 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -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 diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a6c0792a86f..7fc348acb54 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, ) diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index 38d6fb9b99d..4d16c0914a5 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index f03abe8f124..55e6f56ffb7 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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(): """ diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index 613ca847115..3aa53bb590a 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -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", [