mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
503eb2fd4c
commit
725c0c158f
2 changed files with 270 additions and 53 deletions
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue