litellm/tests/test_litellm/proxy/db/test_exception_handler.py
devin-ai-integration[bot] 52403d7a8d
fix(jwt): retry JWKS fetches, serve stale keys, and return 503 when the IdP is unreachable (#37690)
A JWKS fetch had no retry, so a single connect timeout to the identity provider
failed authentication outright, and once the cached copy expired there was
nothing to fall back on. How that surfaced depended on the outage shape:
httpx.ConnectTimeout was missing from DB_CONNECTION_ERROR_TYPES so it fell
through to the generic auth handler as a 401 with an empty detail, while a read
timeout took the database path and reported a healthy database as unreachable.

Transport failures are now retried three times with a short backoff, and the
last-known-good JWKS stays usable for a bounded window past public_key_ttl.
That window is public_key_stale_ttl, a new config field defaulting to 3600s and
settable to 0 to fail closed. It is checked on every read against the current
setting rather than baked into the cache entry when it is written, so lowering
it binds immediately instead of waiting for entries written under the old value
to age out, which matters because a shared cache survives the restart an
operator performs to make the change take effect. A copy whose write time
cannot be established is not servable. Only httpx.TransportError unlocks the
stale copy, so an identity provider that answers at all, including with a
narrowed key set, revokes on the next refresh. Every stale serve logs the kid
it authenticated, how long ago that copy was refreshed, and how long until it
stops being trusted.

A sustained outage is remembered for 30s per key url, so it costs one fetch per
window instead of three timeouts per request serialised behind the refresh lock.
Non-200 JWKS responses now raise instead of being cached as the key set, which
previously let an error body overwrite the last-known-good copy. An unreachable
identity provider with no cached copy left returns 503 auth_provider_unavailable.

Resolves LIT-5524

Co-authored-by: Yassin Kortam <yassin@berri.ai>
2026-08-20 16:41:06 -07:00

586 lines
23 KiB
Python

import asyncio
import json
import os
import sys
from unittest.mock import MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException, Request, status
from prisma import errors as prisma_errors
from prisma.errors import (
ClientNotConnectedError,
DataError,
ForeignKeyViolationError,
HTTPClientClosedError,
MissingRequiredValueError,
PrismaError,
RawQueryError,
RecordNotFoundError,
TableNotFoundError,
UniqueViolationError,
)
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ProxyErrorTypes, ProxyException
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
# Test is_database_connection_error method
@pytest.mark.parametrize(
"prisma_error",
[
HTTPClientClosedError(),
ClientNotConnectedError(),
PrismaError("can't reach database server"),
PrismaError("connection refused"),
PrismaError("timed out while connecting"),
],
)
def test_is_database_infrastructure_error_prisma_connection_errors(prisma_error):
"""
Test that Prisma failures originating below the request are reported as
infrastructure faults, so a caller is told the service failed rather than
that its credentials did.
"""
assert PrismaDBExceptionHandler.is_database_infrastructure_error(prisma_error) == True
@pytest.mark.parametrize(
"prisma_error",
[
PrismaError(),
PrismaError("validation failed on query"),
DataError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
UniqueViolationError(
data={"user_facing_error": {"meta": {"table": "test_table"}}}
),
ForeignKeyViolationError(
data={"user_facing_error": {"meta": {"table": "test_table"}}}
),
MissingRequiredValueError(
data={"user_facing_error": {"meta": {"table": "test_table"}}}
),
RawQueryError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
TableNotFoundError(
data={"user_facing_error": {"meta": {"table": "test_table"}}}
),
RecordNotFoundError(
data={"user_facing_error": {"meta": {"table": "test_table"}}}
),
],
)
def test_is_database_transport_error_non_connection_prisma_errors(prisma_error):
"""Data-layer errors should not trigger reconnect — DB is reachable when these occur."""
assert PrismaDBExceptionHandler.is_database_transport_error(prisma_error) == False
def test_is_database_connection_generic_errors():
"""
Test non-Prisma error cases for database connection checking
"""
assert (
PrismaDBExceptionHandler.is_database_connection_error(
Exception("Regular error")
)
== False
)
# Test with ProxyException (DB connection)
db_proxy_exception = ProxyException(
message="DB Connection Error",
type=ProxyErrorTypes.no_db_connection,
param="test-param",
)
assert (
PrismaDBExceptionHandler.is_database_connection_error(db_proxy_exception)
== True
)
# Test with non-DB error
regular_exception = Exception("Regular error")
assert (
PrismaDBExceptionHandler.is_database_connection_error(regular_exception)
== False
)
@pytest.mark.parametrize(
"error",
[
ConnectionError("connection refused"),
TimeoutError("timed out"),
OSError("network is unreachable"),
asyncio.TimeoutError(),
httpx.ConnectError("connection refused"),
httpx.ConnectTimeout("connect timed out"),
HTTPClientClosedError(),
ClientNotConnectedError(),
PrismaError("can't reach database server"),
PrismaError(),
],
)
def test_is_database_service_unavailable_error_infra_failures(error):
"""Infrastructure-level failures (socket/connection/timeout, prisma
transport, unknown PrismaError) mean the DB could not answer, so auth
must surface 503 instead of treating a valid key as invalid."""
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is True
def test_is_database_service_unavailable_error_prisma_p1001_masquerades_as_dataerror():
"""Real-world regression: prisma-client-py raises the P1001 "can't reach
database server" connectivity failure as a DataError (a data-layer type).
A type-only check would miss it and return 401 during a genuine outage;
the message keyword must still classify it as service-unavailable -> 503."""
p1001_as_dataerror = DataError(
data={
"user_facing_error": {
"message": "Can't reach database server at `127.0.0.1`:`5499`",
"meta": {"table": "t"},
}
}
)
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(
p1001_as_dataerror
)
is True
)
def test_is_prisma_data_error_only_true_for_dataerror():
"""The spend-log poison-row isolation gates on this: only a prisma
``DataError`` (the DB refused the data, e.g. a NUL byte) may be bisected
into a per-row drop. A connectivity failure or any non-prisma exception
must not be treated as a data rejection, so the whole batch surfaces."""
import httpx
data_error = DataError(data={"user_facing_error": {"message": "invalid byte sequence for encoding UTF8: 0x00"}})
assert PrismaDBExceptionHandler.is_prisma_data_error(data_error) is True
for non_data in (
httpx.ConnectError("conn refused"),
PrismaError("can't reach database server"),
UniqueViolationError(data={"user_facing_error": {"meta": {"table": "t"}}}),
RuntimeError("boom"),
):
assert PrismaDBExceptionHandler.is_prisma_data_error(non_data) is False
def test_is_prisma_data_error_true_for_connection_masquerade_dataerror():
"""The P1001 outage prisma mislabels as a ``DataError`` is still a
``DataError`` by type, so this returns True; the spend-log helper relies on
``is_database_service_unavailable_error`` (not this check) to keep that
outage on the retry path instead of dropping rows."""
p1001_as_dataerror = DataError(
data={"user_facing_error": {"message": "Can't reach database server at `127.0.0.1`:`5499`"}}
)
assert PrismaDBExceptionHandler.is_prisma_data_error(p1001_as_dataerror) is True
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(p1001_as_dataerror) is True
def test_is_database_service_unavailable_error_cached_plan_escapes_as_503():
"""Composes with the cached-plan retry: when that recovery fails and the
Postgres "cached plan must not change result type" error escapes (raised by
prisma as a data-layer RawQueryError), it is a transient stale-DB-state
condition, not an invalid key, so it must classify as service-unavailable
-> 503 rather than fall through to 401."""
cached_plan_error = RawQueryError(
data={
"user_facing_error": {
"message": "cached plan must not change result type",
"meta": {"table": "t"},
}
}
)
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(
cached_plan_error
)
is True
)
def test_is_database_service_unavailable_error_prisma_engine_malformed_payload():
"""Real-world regression: at the instant the DB socket drops, the prisma
query engine returns a malformed error payload (``user_facing_error.meta``
is ``null``). prisma-client-py's ``handle_response_errors`` then crashes
with ``AttributeError: 'NoneType' object has no attribute 'get'`` before it
can raise the proper P1001 error. That bare AttributeError has no
connection keyword, so without the prisma-engine-origin check it falls
through to 401 on the first request of an outage. Reproduce the exact
prisma crash and assert it classifies as service-unavailable -> 503."""
from prisma.engine import utils as prisma_engine_utils
malformed_payload = [
{
"error": "Can't reach database server",
"user_facing_error": {
"error_code": "P1001",
"message": "Can't reach database server at `localhost`:`5503`",
"meta": None,
},
}
]
with pytest.raises(AttributeError) as exc_info:
prisma_engine_utils.handle_response_errors(None, malformed_payload)
assert "no attribute 'get'" in str(exc_info.value)
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value)
is True
)
def test_is_prisma_engine_internal_error_excludes_application_attributeerror():
"""The prisma-engine-origin check must stay narrow: a genuine AttributeError
raised by application code (a real bug) must NOT be classified as
service-unavailable, otherwise real bugs would silently become 503s."""
def application_bug():
none_value = None
return none_value.get("oops")
with pytest.raises(AttributeError) as exc_info:
application_bug()
assert (
PrismaDBExceptionHandler.is_prisma_engine_internal_error(exc_info.value)
is False
)
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value)
is False
)
def test_is_prisma_engine_internal_error_excludes_data_layer_prisma_error():
"""A data-layer ``PrismaError`` (the DB IS reachable and rejected the data)
must stay 401. These are always raised from prisma internals, so the check
excludes any ``PrismaError`` by type before inspecting the traceback."""
data_layer_error = UniqueViolationError(
data={"user_facing_error": {"meta": {"table": "t"}}}
)
try:
raise data_layer_error
except UniqueViolationError as e:
assert PrismaDBExceptionHandler.is_prisma_engine_internal_error(e) is False
@pytest.mark.parametrize(
"error",
[
DataError(data={"user_facing_error": {"meta": {"table": "t"}}}),
UniqueViolationError(data={"user_facing_error": {"meta": {"table": "t"}}}),
RecordNotFoundError(data={"user_facing_error": {"meta": {"table": "t"}}}),
Exception("some unrelated error"),
ValueError("bad value"),
],
)
def test_is_database_service_unavailable_error_excludes_non_infra(error):
"""Data-layer errors (the DB IS reachable and answered) and generic
non-DB errors must NOT be classified as service-unavailable, otherwise a
genuine 401 would be masked as a transient 503."""
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is False
)
def _wrapped_like_get_user_object(original):
"""Reproduce get_user_object's exception contract (litellm/proxy/auth/auth_checks.py): it catches
every DB failure in a broad ``except`` and re-raises a bare ``ValueError``, so the original error
survives only as ``__context__``. Building it by raising inside an ``except`` sets ``__context__``
exactly as production does."""
try:
raise original
except BaseException:
try:
raise ValueError("User doesn't exist in db. Got error - x")
except ValueError as wrapped:
return wrapped
def test_is_database_service_unavailable_error_in_chain_sees_through_wrapping():
"""The chain-aware classifier must see a real outage that a caller wrapped in a different type.
get_user_object turns a connection error into a bare ValueError whose type check reads as non-infra,
so the single-exception check returns False and only the chain walk recovers the outage. A missing
user (whose wrapped cause is a plain Exception) must stay non-infra on both."""
outage = _wrapped_like_get_user_object(ConnectionError("can't reach database server"))
missing_user = _wrapped_like_get_user_object(Exception())
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(outage) is False
assert PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(outage) is True
assert PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(missing_user) is False
# parity: a raw outage with no wrapper is still an outage, and a plain ValueError is not
assert PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(ConnectionError("boom")) is True
assert PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(ValueError("nope")) is False
def test_is_database_service_unavailable_error_in_chain_terminates_on_a_cause_cycle():
"""The walk must terminate on a pathological __cause__ cycle rather than hang. Neither link is an
outage, so the bounded walk returns False instead of looping forever."""
first = ValueError("first")
second = ValueError("second")
first.__cause__ = second
second.__cause__ = first
assert PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(first) is False
def test_is_database_service_unavailable_error_asyncpg(monkeypatch):
"""asyncpg connection/interface errors map to service-unavailable. asyncpg
is not a hard dependency, so inject a stand-in module to exercise the
branch deterministically regardless of the install environment."""
import sys
import types
fake_asyncpg = types.ModuleType("asyncpg")
fake_exceptions = types.ModuleType("asyncpg.exceptions")
class PostgresConnectionError(Exception):
pass
class InterfaceError(Exception):
pass
class UniqueViolationError(Exception): # data-layer, must stay False
pass
fake_exceptions.PostgresConnectionError = PostgresConnectionError
fake_exceptions.InterfaceError = InterfaceError
fake_exceptions.UniqueViolationError = UniqueViolationError
fake_asyncpg.exceptions = fake_exceptions
monkeypatch.setitem(sys.modules, "asyncpg", fake_asyncpg)
monkeypatch.setitem(sys.modules, "asyncpg.exceptions", fake_exceptions)
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(
PostgresConnectionError("connection reset")
)
is True
)
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(
InterfaceError("connection was closed")
)
is True
)
assert (
PrismaDBExceptionHandler.is_database_service_unavailable_error(
UniqueViolationError("duplicate key")
)
is False
)
# Test should_allow_request_on_db_unavailable method
@patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
)
def test_should_allow_request_on_db_unavailable_true():
assert PrismaDBExceptionHandler.should_allow_request_on_db_unavailable() == True
@patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": False},
)
def test_should_allow_request_on_db_unavailable_false():
assert PrismaDBExceptionHandler.should_allow_request_on_db_unavailable() == False
@patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
)
def test_handle_db_exception_with_connection_error():
"""
Test that DB connection errors are handled gracefully when allow_requests_on_db_unavailable is True
"""
db_error = httpx.ConnectError("All connection attempts failed")
result = PrismaDBExceptionHandler.handle_db_exception(db_error)
assert result is None
@patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": False},
)
def test_handle_db_exception_raises_error():
"""
Test that DB connection errors are raised when allow_requests_on_db_unavailable is False
"""
db_error = httpx.ConnectError("All connection attempts failed")
with pytest.raises(httpx.ConnectError):
PrismaDBExceptionHandler.handle_db_exception(db_error)
def test_handle_db_exception_with_non_db_error():
"""
Test that non-DB errors are always raised regardless of allow_requests_on_db_unavailable setting
"""
regular_error = litellm.BudgetExceededError(
current_cost=10,
max_budget=10,
)
with pytest.raises(litellm.BudgetExceededError):
PrismaDBExceptionHandler.handle_db_exception(regular_error)
def _permanent_prisma_faults():
"""Every prisma error class that is not a transient outage and not a
data-layer error, built by enumeration so the list cannot drift out of sync
with the installed prisma version."""
import inspect
from prisma import engine as prisma_engine
payload = {"user_facing_error": {"meta": {"target": ["x"]}}, "error_message": "boom"}
class _Response:
status = 422
def build(cls):
attempts = (
((), {}),
((payload,), {}),
(("x",), {}),
(("x", "y"), {}),
((_Response(),), {}),
((_Response(), "body"), {}),
((), {"expected": "1", "got": "2"}),
)
for args, kwargs in attempts:
try:
return cls(*args, **kwargs)
except Exception:
continue
return None
discovered = {}
for module in (prisma_errors, prisma_engine.errors):
for name, obj in vars(module).items():
if inspect.isclass(obj) and issubclass(obj, BaseException) and obj.__module__.startswith("prisma"):
discovered[name] = obj
faults = []
for name, cls in sorted(discovered.items()):
if issubclass(cls, prisma_engine.errors.EngineConnectionError):
continue
instance = build(cls)
if instance is None or not PrismaDBExceptionHandler.is_database_infrastructure_error(instance):
continue
faults.append(pytest.param(instance, id=name))
return faults
PERMANENT_PRISMA_FAULTS = _permanent_prisma_faults()
def test_permanent_fault_enumeration_is_not_empty():
"""Guards the parametrized tests below: if the enumeration silently found
nothing, those tests would pass without asserting anything."""
assert len(PERMANENT_PRISMA_FAULTS) >= 15
@pytest.mark.parametrize("prisma_error", PERMANENT_PRISMA_FAULTS)
def test_permanent_prisma_faults_are_not_transient(prisma_error):
"""A fault that cannot resolve on its own must not qualify a request to be
served without a verified database.
``allow_requests_on_db_unavailable`` trades verification for availability on
the assumption the database returns. A missing or version-skewed query
engine, a malformed generated query, or a misused transaction never returns,
so absorbing one turns a bounded degraded window into a permanent one."""
assert PrismaDBExceptionHandler.is_database_connection_error(prisma_error) is False
@pytest.mark.parametrize("prisma_error", PERMANENT_PRISMA_FAULTS)
def test_permanent_prisma_faults_are_still_reported_as_service_problems(prisma_error):
"""Narrowing what may be served without a database must not change what the
caller is told.
A permanently faulted engine is still the service's fault, so it has to keep
reaching the reporting predicate that renders it as unavailable rather than
as a rejected credential."""
assert PrismaDBExceptionHandler.is_database_infrastructure_error(prisma_error) is True
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(prisma_error) is True
@pytest.mark.parametrize(
"transient_error",
[
pytest.param(httpx.ConnectError("All connection attempts failed"), id="ConnectError"),
pytest.param(httpx.ReadError("read failed"), id="ReadError"),
pytest.param(httpx.ReadTimeout("timed out"), id="ReadTimeout"),
pytest.param(prisma_errors.HTTPClientClosedError(), id="HTTPClientClosedError_is_not_transient"),
],
)
def test_genuine_outage_still_qualifies_for_the_fallback(transient_error):
"""The httpx transport errors are how a real outage actually reaches the
caller: the query engine is a local HTTP server, so an unreachable database
surfaces as a transport failure against it rather than as a prisma type.
Narrowing the classifier must leave that path intact."""
expected = not isinstance(transient_error, prisma_errors.HTTPClientClosedError)
assert PrismaDBExceptionHandler.is_database_connection_error(transient_error) is expected
def test_engine_connection_error_is_the_transient_prisma_type():
"""``EngineConnectionError`` is the one prisma class that means the engine
could not reach the database and may succeed later."""
from prisma.engine.errors import EngineConnectionError
assert PrismaDBExceptionHandler.is_database_connection_error(EngineConnectionError()) is True
@patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
)
def test_handle_db_exception_surfaces_a_permanent_fault_even_when_degraded_mode_is_enabled():
"""The startup gate must not boot a proxy whose engine is permanently
faulted. Absorbing it produces a process that reports healthy and persists
nothing, indefinitely, with no error for an operator to find."""
from prisma.engine.errors import BinaryNotFoundError
with pytest.raises(BinaryNotFoundError):
PrismaDBExceptionHandler.handle_db_exception(BinaryNotFoundError("query engine binary not found"))
@pytest.mark.parametrize(
"error",
[
RawQueryError(data={"user_facing_error": {"error_code": "P2034", "meta": {"table": "t"}}}),
PrismaError("Transaction failed due to a write conflict or a deadlock. Please retry your transaction"),
RawQueryError(data={"user_facing_error": {"message": "deadlock detected", "meta": {"table": "t"}}}),
RawQueryError(
data={"user_facing_error": {"message": "ERROR: 40P01: deadlock detected", "meta": {"table": "t"}}}
),
],
)
def test_is_deadlock_error_matches_postgres_deadlock(error):
"""A Postgres deadlock surfaced through prisma (P2034 or 40P01 / "deadlock detected" text) is recognized."""
assert PrismaDBExceptionHandler.is_deadlock_error(error) is True
@pytest.mark.parametrize(
"error",
[
UniqueViolationError(data={"user_facing_error": {"error_code": "P2002", "meta": {"table": "t"}}}),
RecordNotFoundError(data={"user_facing_error": {"meta": {"table": "t"}}}),
PrismaError("validation failed on query"),
PrismaError("can't reach database server"),
httpx.ConnectError("connection refused"),
RuntimeError("deadlock detected"),
ValueError("40P01"),
],
)
def test_is_deadlock_error_excludes_non_deadlocks(error):
"""Non-deadlock prisma errors, connectivity failures, and non-prisma exceptions are not treated as deadlocks."""
assert PrismaDBExceptionHandler.is_deadlock_error(error) is False