Prisma DB Failure Detection and Self-Healing (#21059)

* fix(proxy): readiness check returns 200 when database is unreachable

_db_health_readiness_check() catches health_check() exceptions but
never updates db_health_cache to "disconnected" and never re-raises.
The caller health_readiness() always returns 200 with "db": "connected"
hardcoded, regardless of actual DB state.

In Kubernetes, this means pods with dead database connections stay in
the Service endpoints and continue receiving traffic they cannot serve.

Changes:
- Set db_health_cache to "disconnected" and re-raise the exception on
  health_check failure so health_readiness() returns 503
- Use actual db_health_status["status"] in the response instead of
  hardcoding "db": "connected"
- Reduce cache TTL from 2 minutes to 15 seconds. The 2-minute window
  is too wide for readiness probes (typically 10-15s intervals) and
  means a pod can report healthy for up to 2 minutes after the DB dies
- Only serve cached results when status is "connected". The previous
  condition (status != "unknown") would also cache "disconnected" for
  2 minutes, delaying recovery detection after a DB comes back

* fix(proxy): add DB connection self-healing to readiness check

When the Prisma query engine's internal TCP connection pool holds dead
connections (caused by network blips, Cloud SQL proxy restarts, or
node-level issues), health_check() fails with httpx.ConnectError.
The engine never recovers on its own because nothing triggers a
disconnect/connect cycle to restart the subprocess with fresh
connections.

This leaves pods permanently failing readiness checks until they are
manually restarted, even after the underlying DB becomes reachable
again.

Add a reconnect attempt to _db_health_readiness_check() when
health_check() fails:
1. disconnect() - kills the query engine subprocess and closes all
   connections (has built-in backoff retry: 3 tries, 10s max)
2. connect() - starts a new engine with fresh TCP connections (has
   built-in backoff retry: 3 tries, 10s max)
3. health_check() - verifies the new connection works (has built-in
   backoff retry: 3 tries, 10s max)

If reconnect succeeds, the pod immediately returns to service (200).
If it fails, the original exception is re-raised (503). Reconnect
attempts are rate-limited by probe frequency (~10-15s), so a
permanently unreachable DB gets one attempt per cycle with no retry
loops.

This uses the same disconnect/connect mechanism that
PrismaWrapper.recreate_prisma_client() uses for IAM token refresh,
and aligns with the community-documented pattern for Prisma connection
recovery in long-running processes (prisma/prisma#24718, #27024).

* Add poetry lock and modify test_health_endpoints

* Address allow_requests_on_db_unavailable regression

* Address comments

* resolve greptile issue

* Restore accidentally deleted UI HTML files

These were removed in an earlier commit but still exist on main.
Restoring to keep the PR diff clean.

* Guard reconnect with is_database_transport_error

Only attempt disconnect/connect/health_check cycle for transport-level
failures (unreachable DB, dropped connection). Data-layer errors like
UniqueViolationError indicate the DB is reachable, so reconnecting
would be pointless churn.

* Address greptile's comments

* Fix module alias after rebase and add adversarial test coverage

- Unify module alias to _health_endpoints_module after rebase conflict
- Add test for non-transport error with flag on (exercises is_database_transport_error guard)
- Add test for disconnect() failure during reconnect cycle
- Split non-transport error test into flag-off (re-raises) and flag-on (skips reconnect) variants

* Remove stale UI HTML files reintroduced during rebase
This commit is contained in:
Curtis 2026-03-05 13:44:49 -08:00 • committed by GitHub
parent 503eb2fd4c
commit 725c0c158f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 270 additions and 53 deletions

View file

@ -1141,11 +1141,9 @@ async def _db_health_readiness_check():
global db_health_cache
# Note - Intentionally don't try/except this so it raises an exception when it fails
try:
# if timedelta is less than 2 minutes return DB Status
time_diff = datetime.now() - db_health_cache["last_updated"]
if db_health_cache["status"] != "unknown" and time_diff < timedelta(minutes=2):
if db_health_cache["status"] == "connected" and time_diff < timedelta(seconds=15):
return db_health_cache
if prisma_client is None:
@ -1156,7 +1154,25 @@ async def _db_health_readiness_check():
db_health_cache = {"status": "connected", "last_updated": datetime.now()}
return db_health_cache
except Exception as e:
db_health_cache = {"status": "disconnected", "last_updated": datetime.now()}
PrismaDBExceptionHandler.handle_db_exception(e)
if PrismaDBExceptionHandler.is_database_transport_error(e):
try:
verbose_proxy_logger.warning(
"_db_health_readiness_check: health_check failed, attempting reconnect"
)
await prisma_client.disconnect()
await prisma_client.connect()
await prisma_client.health_check()
verbose_proxy_logger.info(
"_db_health_readiness_check: reconnect succeeded"
)
db_health_cache = {"status": "connected", "last_updated": datetime.now()}
return db_health_cache
except Exception:
verbose_proxy_logger.error(
"_db_health_readiness_check: reconnect failed"
)
return db_health_cache
@ -1302,14 +1318,13 @@ async def health_readiness():
db_health_status = await _db_health_readiness_check()
return {
"status": "healthy",
"db": "connected",
"db": db_health_status["status"],
"cache": cache_type,
"litellm_version": version,
"success_callbacks": success_callback_names,
"use_aiohttp_transport": AsyncHTTPHandler._should_use_aiohttp_transport(),
"log_level": log_level_name,
"is_detailed_debug": is_detailed_debug,
**db_health_status,
}
else:
return {

View file

@ -13,9 +13,9 @@ import pytest
from prisma.errors import ClientNotConnectedError, HTTPClientClosedError, PrismaError
import litellm.proxy.health_endpoints._health_endpoints as _health_endpoints_module
from litellm.proxy.health_endpoints._health_endpoints import (
_db_health_readiness_check,
db_health_cache,
get_callback_identifier,
health_license_endpoint,
health_services_endpoint,
@ -29,45 +29,119 @@ from tests.test_litellm.proxy.conftest import create_proxy_test_client
@pytest.mark.asyncio
@pytest.mark.parametrize(
"prisma_error",
[
PrismaError("Can't reach database server"),
ClientNotConnectedError(),
HTTPClientClosedError(),
],
)
async def test_db_health_readiness_check_with_prisma_error(prisma_error):
async def test_db_health_cache_hit_returns_cached():
"""
Test that when prisma_client.health_check() raises a PrismaError and
allow_requests_on_db_unavailable is True, the function should not raise an error
and return the cached health status.
When cache is 'connected' and within the 15s TTL, return the cache
without calling health_check.
"""
# Mock the prisma client
mock_prisma_client = MagicMock()
mock_prisma_client.health_check.side_effect = prisma_error
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock()
# Reset the health cache in the source module so _db_health_readiness_check
# sees the updated value (assigning to a test-module global doesn't work).
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now(),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await _db_health_readiness_check()
assert result["status"] == "connected"
mock_prisma.health_check.assert_not_called()
@pytest.mark.asyncio
async def test_db_health_cache_expired_calls_health_check():
"""
When cache is 'connected' but older than 15s, call health_check
to re-validate the connection.
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await _db_health_readiness_check()
assert result["status"] == "connected"
mock_prisma.health_check.assert_called_once()
@pytest.mark.asyncio
async def test_db_health_non_connected_ignores_cache_ttl():
"""
When cache status is not 'connected' (e.g. 'disconnected', 'unknown'),
always call health_check regardless of how fresh the cache is.
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "disconnected",
"last_updated": datetime.now(),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await _db_health_readiness_check()
assert result["status"] == "connected"
mock_prisma.health_check.assert_called_once()
@pytest.mark.asyncio
async def test_db_health_prisma_client_none():
"""
When prisma_client is None, return 'disconnected' without attempting
a health_check call.
"""
_health_endpoints_module.db_health_cache = {
"status": "unknown",
"last_updated": datetime.now() - timedelta(minutes=5),
}
# Patch the imports and general_settings
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
):
# Call the function
with patch("litellm.proxy.proxy_server.prisma_client", None):
result = await _db_health_readiness_check()
# Verify that the function called health_check
mock_prisma_client.health_check.assert_called_once()
assert result["status"] == "disconnected"
# Verify that the function returned the cache
assert result is not None
assert result["status"] == "unknown" # Should retain the status from the cache
@pytest.mark.asyncio
@pytest.mark.parametrize(
"prisma_error",
[
PrismaError(),
ClientNotConnectedError(),
HTTPClientClosedError(),
],
)
async def test_db_health_error_flag_off_raises_no_reconnect(prisma_error):
"""
When health_check raises and allow_requests_on_db_unavailable is False,
handle_db_exception re-raises immediately. The reconnect path is never
reached, so disconnect/connect are never called.
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=prisma_error)
mock_prisma.disconnect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": False},
):
with pytest.raises(Exception) as exc_info:
await _db_health_readiness_check()
assert exc_info.value is prisma_error
mock_prisma.disconnect.assert_not_called()
assert _health_endpoints_module.db_health_cache["status"] == "disconnected"
@pytest.mark.asyncio
@ -79,32 +153,161 @@ async def test_db_health_readiness_check_with_prisma_error(prisma_error):
HTTPClientClosedError(),
],
)
async def test_db_health_readiness_check_with_error_and_flag_off(prisma_error):
async def test_db_health_error_flag_on_reconnect_succeeds(prisma_error):
"""
Test that when prisma_client.health_check() raises a DB error but
allow_requests_on_db_unavailable is False, the exception should be raised.
When health_check raises, allow_requests_on_db_unavailable is True,
and the reconnect cycle (disconnect -> connect -> health_check) succeeds,
return 'connected' and update the cache.
"""
# Mock the prisma client
mock_prisma_client = MagicMock()
mock_prisma_client.health_check.side_effect = prisma_error
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(
side_effect=[prisma_error, None]
)
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
# Reset the health cache in the source module
_health_endpoints_module.db_health_cache = {
"status": "unknown",
"last_updated": datetime.now() - timedelta(minutes=5),
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
# Patch the imports and general_settings where the flag is False
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), patch(
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
):
result = await _db_health_readiness_check()
assert result["status"] == "connected"
mock_prisma.disconnect.assert_called_once()
mock_prisma.connect.assert_called_once()
assert mock_prisma.health_check.call_count == 2
@pytest.mark.asyncio
@pytest.mark.parametrize(
"prisma_error",
[
PrismaError("Can't reach database server"),
ClientNotConnectedError(),
HTTPClientClosedError(),
],
)
async def test_db_health_error_flag_on_reconnect_fails(prisma_error):
"""
When health_check raises, allow_requests_on_db_unavailable is True,
but the reconnect also fails, return 'disconnected' instead of raising.
This respects the flag's intent: keep serving even without a DB.
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=prisma_error)
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.disconnect.assert_called_once()
mock_prisma.connect.assert_called_once()
@pytest.mark.asyncio
async def test_db_health_non_transport_error_flag_off_raises():
"""
When health_check raises a non-transport error and
allow_requests_on_db_unavailable is False, handle_db_exception
re-raises before reaching the is_database_transport_error guard.
Cache is still invalidated before the re-raise.
"""
non_transport_error = PrismaError("UniqueViolationError")
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=non_transport_error)
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": False},
):
# The function should raise the exception
with pytest.raises(Exception) as excinfo:
with pytest.raises(PrismaError):
await _db_health_readiness_check()
# Verify that the raised exception is the same
assert excinfo.value == prisma_error
assert _health_endpoints_module.db_health_cache["status"] == "disconnected"
mock_prisma.disconnect.assert_not_called()
mock_prisma.connect.assert_not_called()
@pytest.mark.asyncio
async def test_db_health_non_transport_error_flag_on_skips_reconnect():
"""
When health_check raises a non-transport error (e.g. data-layer) and
allow_requests_on_db_unavailable is True, handle_db_exception swallows
the exception, then is_database_transport_error returns False so the
reconnect cycle is skipped. Returns 'disconnected' without calling
disconnect/connect.
"""
non_transport_error = PrismaError("UniqueViolationError")
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=non_transport_error)
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.disconnect.assert_not_called()
mock_prisma.connect.assert_not_called()
@pytest.mark.asyncio
async def test_db_health_reconnect_disconnect_fails():
"""
When disconnect() itself raises during the reconnect cycle,
the inner except catches it and returns 'disconnected'.
connect() and the second health_check() are never called.
"""
transport_error = ClientNotConnectedError()
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=transport_error)
mock_prisma.disconnect = AsyncMock(side_effect=RuntimeError("already closed"))
mock_prisma.connect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.disconnect.assert_called_once()
mock_prisma.connect.assert_not_called()
@pytest.mark.asyncio
@ -374,7 +577,6 @@ def proxy_client(monkeypatch):
yield client
@pytest.mark.xdist_group("proxy_health")
def test_health_liveliness_endpoint(proxy_client):
"""
Test that /health/liveliness endpoint returns 200 OK with "I'm alive!" message.
@ -492,7 +694,7 @@ def test_get_callback_identifier_string_and_object_with_callback_name():
- Object with empty/None callback_name (should fall through to other checks)
"""
from litellm.proxy.health_endpoints._health_endpoints import get_callback_identifier
# Test 1: String callback should be returned as-is
assert get_callback_identifier("datadog") == "datadog"
assert get_callback_identifier("langfuse") == "langfuse"
@ -523,9 +725,9 @@ def test_get_callback_identifier_custom_logger_registry_and_fallback():
- Object with callback_name that matches registry entry
- Fallback to callback_name() helper function
"""
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
from litellm.proxy.health_endpoints._health_endpoints import get_callback_identifier
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
# Test 1: Object registered in CustomLoggerRegistry (without callback_name attribute)
# Mock a class that's registered in the registry
class MockRegisteredLogger: