fix(health): honor allow_requests_on_db_unavailable in readiness probe (#37640)

* fix(health): honor allow_requests_on_db_unavailable in readiness probe

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(health): bound readiness DB check and pass reconnect timeout

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(health): bound whole readiness DB check with one deadline

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(health): keep readiness deadline fallback within lint budgets

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(health): suppress TQ008 for proxy-global readiness patches

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(db): release reconnect lock when a waiting reconnect is cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: milan <milan@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: yassin <yassin@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-08-29 10:17:12 -07:00 • committed by GitHub
parent cb7d41a5c6
commit 0de1825450
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 256 additions and 49 deletions

View file

@ -1372,8 +1372,25 @@ class DBHealthCache(TypedDict):
db_health_cache: DBHealthCache = {"status": "unknown", "last_updated": datetime.now()}
# Bounds each DB round-trip on the probe path so a hung connection during a
# failover cannot make the probe fail by timeout (k8s default timeoutSeconds: 5).
DB_READINESS_CHECK_TIMEOUT_SECONDS: Final = 2.0
# One deadline for the whole probe-path DB check (initial check + reconnect +
# re-check, including reconnect lock waits), kept under timeoutSeconds: 5.
DB_READINESS_PROBE_DEADLINE_SECONDS: Final = 4.0
async def _db_health_readiness_check():
async def _db_health_readiness_check() -> DBHealthCache:
try:
return await asyncio.wait_for(
_db_health_readiness_check_unbounded(),
timeout=DB_READINESS_PROBE_DEADLINE_SECONDS,
)
except asyncio.TimeoutError:
return {"status": "disconnected", "last_updated": db_health_cache["last_updated"]}
async def _db_health_readiness_check_unbounded() -> DBHealthCache:
from litellm.proxy.proxy_server import prisma_client
global db_health_cache
@ -1387,7 +1404,7 @@ async def _db_health_readiness_check():
db_health_cache = {"status": "disconnected", "last_updated": datetime.now()}
return db_health_cache
await prisma_client.health_check()
await asyncio.wait_for(prisma_client.health_check(), timeout=DB_READINESS_CHECK_TIMEOUT_SECONDS)
db_health_cache = {"status": "connected", "last_updated": datetime.now()}
return db_health_cache
except Exception as e:
@ -1395,8 +1412,15 @@ async def _db_health_readiness_check():
if PrismaDBExceptionHandler.is_database_transport_error(e):
try:
verbose_proxy_logger.warning("_db_health_readiness_check: health_check failed, attempting reconnect")
await prisma_client.attempt_db_reconnect(reason="health_readiness_check")
await prisma_client.health_check()
await prisma_client.attempt_db_reconnect(
reason="health_readiness_check",
timeout_seconds=DB_READINESS_CHECK_TIMEOUT_SECONDS,
lock_timeout_seconds=DB_READINESS_CHECK_TIMEOUT_SECONDS,
)
await asyncio.wait_for(
prisma_client.health_check(),
timeout=DB_READINESS_CHECK_TIMEOUT_SECONDS,
)
verbose_proxy_logger.info("_db_health_readiness_check: reconnect succeeded")
db_health_cache = {
"status": "connected",
@ -1580,7 +1604,14 @@ async def _get_health_readiness_details(
# serve requests that depend on persisted state (keys, budgets,
# spend logs). Return 503 so orchestrators take this pod out of
# rotation; "Not connected" (no DB configured at all) stays 200.
if response is not None and db_health_status["status"] != "connected":
# With allow_requests_on_db_unavailable the proxy keeps serving
# during a DB outage, so the pod must stay in rotation (200) and
# report the DB state through the body instead.
if (
response is not None
and db_health_status["status"] != "connected"
and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
):
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return {
"status": "healthy",
@ -1671,7 +1702,10 @@ async def _resolve_public_readiness_db(response: Response) -> str:
return "Not connected"
db_health_status: Final = await _db_health_readiness_check()
if db_health_status["status"] != "connected":
if (
db_health_status["status"] != "connected"
and not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
):
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return db_health_status["status"]

View file

@ -5580,12 +5580,8 @@ class PrismaClient:
return True
acquire_task: Final = asyncio.create_task(_acquire_reconnect_lock())
done, _pending = await asyncio.wait(
{acquire_task},
timeout=lock_timeout_seconds,
return_when=asyncio.FIRST_COMPLETED,
)
if acquire_task not in done:
async def _abandon_acquire_task() -> None:
acquire_task.cancel()
try:
await acquire_task
@ -5600,6 +5596,18 @@ class PrismaClient:
self._db_reconnect_lock.release()
except RuntimeError:
pass
try:
done, _pending = await asyncio.wait(
{acquire_task},
timeout=lock_timeout_seconds,
return_when=asyncio.FIRST_COMPLETED,
)
except asyncio.CancelledError:
await asyncio.shield(_abandon_acquire_task())
raise
if acquire_task not in done:
await _abandon_acquire_task()
verbose_proxy_logger.debug(
"Skipping DB reconnect attempt due to lock acquisition timeout. reason=%s timeout=%ss",
reason,

View file

@ -1,10 +1,10 @@
import asyncio
import json
import time
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
@ -15,7 +15,6 @@ from prisma.errors import ClientNotConnectedError, HTTPClientClosedError, Prisma
import litellm
import litellm.proxy.health_endpoints._health_endpoints as _health_endpoints_module
from litellm.litellm_core_utils.health_check_helpers import TEST_IMAGE_BASE64
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.health_endpoints._health_endpoints import (
@ -145,7 +144,11 @@ async def test_db_health_transport_error_never_raises(transport_error):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.attempt_db_reconnect.assert_called_once_with(reason="health_readiness_check")
mock_prisma.attempt_db_reconnect.assert_called_once_with(
reason="health_readiness_check",
timeout_seconds=_health_endpoints_module.DB_READINESS_CHECK_TIMEOUT_SECONDS,
lock_timeout_seconds=_health_endpoints_module.DB_READINESS_CHECK_TIMEOUT_SECONDS,
)
@pytest.mark.asyncio
@ -175,7 +178,11 @@ async def test_db_health_transport_error_reconnect_succeeds(transport_error):
result = await _db_health_readiness_check()
assert result["status"] == "connected"
mock_prisma.attempt_db_reconnect.assert_called_once_with(reason="health_readiness_check")
mock_prisma.attempt_db_reconnect.assert_called_once_with(
reason="health_readiness_check",
timeout_seconds=_health_endpoints_module.DB_READINESS_CHECK_TIMEOUT_SECONDS,
lock_timeout_seconds=_health_endpoints_module.DB_READINESS_CHECK_TIMEOUT_SECONDS,
)
assert mock_prisma.health_check.call_count == 2
@ -2276,6 +2283,159 @@ async def test_health_readiness_returns_503_when_db_disconnected():
assert result == {"status": "healthy", "db": "disconnected"}
@pytest.mark.asyncio
async def test_health_readiness_returns_200_when_db_down_and_allow_requests_on_db_unavailable():
"""
Regression test for https://github.com/BerriAI/litellm/issues/34934.
allow_requests_on_db_unavailable keeps the proxy serving through a DB
outage, so the readiness probe must keep the pod in rotation (200) and
report the DB state through the body, not the status code. Otherwise
K8s pulls every replica before the request-layer fail-open can run.
"""
from fastapi import Response
from litellm.proxy.health_endpoints._health_endpoints import health_readiness
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=PrismaError("nope"))
mock_prisma.attempt_db_reconnect = AsyncMock(side_effect=Exception("still nope"))
_health_endpoints_module.db_health_cache = {
"status": "unknown",
"last_updated": datetime.now() - timedelta(seconds=60),
}
response = Response()
with (
patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam
"litellm.proxy.proxy_server.prisma_client", mock_prisma
),
patch.dict( # test-quality-ok: the fail-open flag lives in the proxy-global general_settings; no injection seam
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
),
):
result = await health_readiness(response=response)
assert response.status_code == 200
assert result == {"status": "healthy", "db": "disconnected"}
@pytest.mark.asyncio
async def test_health_readiness_details_returns_200_when_db_down_and_allow_requests_on_db_unavailable():
"""
The detailed readiness payload (public via
allow_public_health_readiness_details, or /health/readiness/details)
must honor the same flag so probes pointed at it also stay 200.
"""
from fastapi import Response
from litellm.proxy.health_endpoints._health_endpoints import (
_get_health_readiness_details,
)
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=PrismaError("nope"))
mock_prisma.attempt_db_reconnect = AsyncMock(side_effect=Exception("still nope"))
_health_endpoints_module.db_health_cache = {
"status": "unknown",
"last_updated": datetime.now() - timedelta(seconds=60),
}
response = Response()
with (
patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam
"litellm.proxy.proxy_server.prisma_client", mock_prisma
),
patch.dict( # test-quality-ok: the fail-open flag lives in the proxy-global general_settings; no injection seam
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
),
):
result = await _get_health_readiness_details(response=response)
assert response.status_code == 200
assert result["db"] == "disconnected"
@pytest.mark.asyncio
async def test_db_health_readiness_check_bounds_hung_health_check():
"""
A connection that hangs mid-failover must not stall the probe past the
kubelet's timeoutSeconds; the DB round-trip is bounded and reported as
disconnected instead.
"""
from litellm.proxy.health_endpoints._health_endpoints import (
_db_health_readiness_check,
)
async def hang():
await asyncio.sleep(60)
mock_prisma = MagicMock()
mock_prisma.health_check = hang
mock_prisma.attempt_db_reconnect = AsyncMock(side_effect=Exception("still down"))
_health_endpoints_module.db_health_cache = {
"status": "unknown",
"last_updated": datetime.now() - timedelta(seconds=60),
}
with patch( # test-quality-ok: lowers the module-level probe timeout so the hung-call test finishes fast
"litellm.proxy.health_endpoints._health_endpoints.DB_READINESS_CHECK_TIMEOUT_SECONDS",
0.05,
):
start = time.monotonic()
with patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam
"litellm.proxy.proxy_server.prisma_client", mock_prisma
):
result = await _db_health_readiness_check()
elapsed = time.monotonic() - start
assert result["status"] == "disconnected"
assert elapsed < 5
@pytest.mark.asyncio
async def test_db_health_readiness_check_overall_deadline_bounds_hung_reconnect():
"""
The whole probe-path DB check (initial check + reconnect + re-check,
including reconnect lock waits) runs under one deadline, so a reconnect
that hangs on the lock still returns disconnected within the deadline.
"""
from litellm.proxy.health_endpoints._health_endpoints import (
_db_health_readiness_check,
)
async def hang(**kwargs):
await asyncio.sleep(60)
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=httpx.ConnectError("down"))
mock_prisma.attempt_db_reconnect = hang
_health_endpoints_module.db_health_cache = {
"status": "unknown",
"last_updated": datetime.now() - timedelta(seconds=60),
}
with patch( # test-quality-ok: lowers the module-level probe timeout so the hung-call test finishes fast
"litellm.proxy.health_endpoints._health_endpoints.DB_READINESS_PROBE_DEADLINE_SECONDS",
0.05,
):
start = time.monotonic()
with patch( # test-quality-ok: the readiness path reads the proxy-global DB client; it has no injection seam
"litellm.proxy.proxy_server.prisma_client", mock_prisma
):
result = await _db_health_readiness_check()
elapsed = time.monotonic() - start
assert result["status"] == "disconnected"
assert elapsed < 5
@pytest.mark.asyncio
async def test_health_readiness_returns_200_when_db_connected():
"""Happy path: connected DB keeps the legacy 200."""
@ -2746,13 +2906,13 @@ def test_test_model_connection_accepts_image_edit_mode(monkeypatch):
app = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
client = TestClient(app)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: the endpoint reads the proxy-global DB client and 500s when it is None; it has no injection seam
patch( # test-quality-ok: the endpoint reads the proxy-global DB client and 500s when it is None; it has no injection seam
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
respx.mock(assert_all_called=True) as respx_mock,
):
respx_mock.post(host="api.openai.com", path="/v1/images/edits").respond(

View file

@ -89,9 +89,7 @@ async def test_run_reconnect_cycle_direct_path_recreates_when_probe_fails(
prisma_client._cleanup_engine_watcher = MagicMock()
writer = MagicMock()
writer.query_raw = AsyncMock(
side_effect=[ConnectionError("probe failed"), [{"?column?": 1}]]
)
writer.query_raw = AsyncMock(side_effect=[ConnectionError("probe failed"), [{"?column?": 1}]])
monkeypatch.setattr(
PrismaClient,
"writer_db",
@ -171,9 +169,7 @@ async def test_run_reconnect_cycle_passes_writer_generation_to_recreate(
writer = MagicMock()
writer._engine_generation = 7
writer.query_raw = AsyncMock(
side_effect=[ConnectionError("probe failed"), [{"?column?": 1}]]
)
writer.query_raw = AsyncMock(side_effect=[ConnectionError("probe failed"), [{"?column?": 1}]])
monkeypatch.setattr(
PrismaClient,
"writer_db",
@ -229,9 +225,7 @@ async def test_attempt_reconnect_inside_lock_runs_cycle_and_resets_counter(
prisma_client._consecutive_reconnect_failures = 2
prisma_client._run_reconnect_cycle = AsyncMock()
ok = await prisma_client._attempt_reconnect_inside_lock(
force=True, reason="test", timeout_seconds=1
)
ok = await prisma_client._attempt_reconnect_inside_lock(force=True, reason="test", timeout_seconds=1)
pinned = {
"returned": ok,
"cycle_called": prisma_client._run_reconnect_cycle.await_count,
@ -254,9 +248,7 @@ async def test_attempt_reconnect_inside_lock_skips_when_in_cooldown(
prisma_client._db_last_reconnect_attempt_ts = time.time()
prisma_client._run_reconnect_cycle = AsyncMock()
ok = await prisma_client._attempt_reconnect_inside_lock(
force=False, reason="test", timeout_seconds=1
)
ok = await prisma_client._attempt_reconnect_inside_lock(force=False, reason="test", timeout_seconds=1)
assert ok is False
assert prisma_client._run_reconnect_cycle.await_count == 0
@ -269,9 +261,7 @@ async def test_attempt_reconnect_inside_lock_increments_failure_counter_on_error
prisma_client._consecutive_reconnect_failures = 0
prisma_client._run_reconnect_cycle = AsyncMock(side_effect=RuntimeError("boom"))
ok = await prisma_client._attempt_reconnect_inside_lock(
force=True, reason="failing_test", timeout_seconds=1
)
ok = await prisma_client._attempt_reconnect_inside_lock(force=True, reason="failing_test", timeout_seconds=1)
assert ok is False
assert prisma_client._consecutive_reconnect_failures == 1
@ -316,9 +306,7 @@ async def test_attempt_db_reconnect_lock_timeout_returns_false(
by replacing ``asyncio.wait`` with a callable that returns the loser
task as still-pending after it's already been completed elsewhere.
"""
completed_task: asyncio.Task[bool] = asyncio.get_running_loop().create_task(
_no_op_returning_true()
)
completed_task: asyncio.Task[bool] = asyncio.get_running_loop().create_task(_no_op_returning_true())
# Ensure the inner task has finished before attempt_db_reconnect sees it.
await completed_task
@ -329,7 +317,7 @@ async def test_attempt_db_reconnect_lock_timeout_returns_false(
monkeypatch.setattr(
asyncio,
"create_task",
lambda coro, *a, **kw: (coro.close() or completed_task),
lambda coro, *a, **kw: coro.close() or completed_task,
)
prisma_client._db_last_reconnect_attempt_ts = 0.0
@ -465,9 +453,7 @@ async def test_db_health_watchdog_loop_triggers_reconnect_on_timeout(
await prisma_client._db_health_watchdog_loop()
pinned = {
"reconnect_called": prisma_client.attempt_db_reconnect.await_count,
"reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs[
"reason"
],
"reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs["reason"],
"wait_for_calls": call_count["n"],
"loop_exited_clean": True,
}
@ -522,10 +508,7 @@ async def test_iam_refresh_racing_reconnect_recreates_engine_only_once(
from litellm.proxy.db.prisma_client import PrismaWrapper
def token_db_url(created: datetime) -> str:
token = (
f"host/?X-Amz-Date={created.strftime('%Y%m%dT%H%M%SZ')}"
f"&X-Amz-Expires=900&X-Amz-Signature=abc"
)
token = f"host/?X-Amz-Date={created.strftime('%Y%m%dT%H%M%SZ')}&X-Amz-Expires=900&X-Amz-Signature=abc"
return f"postgresql://user:{urllib.parse.quote(token, safe='')}@host:5432/db"
# Old engine (PID 111) carries an expired token; in-flight queries on it
@ -577,9 +560,7 @@ async def test_iam_refresh_racing_reconnect_recreates_engine_only_once(
# In-flight transport-error path fires while the refresh holds the
# wrapper's reconnection lock mid-recreate.
reconnect_task = asyncio.create_task(
prisma_client.attempt_db_reconnect(
reason="in_flight_transport_error", force=True
)
prisma_client.attempt_db_reconnect(reason="in_flight_transport_error", force=True)
)
await asyncio.sleep(0.05)
release_connect.set()
@ -1096,3 +1077,27 @@ async def test_unrelated_reconnect_failure_does_not_erase_the_burst_record(
"cycles_after": prisma_client._run_reconnect_cycle.await_count,
}
assert pinned == {"cycles_before": 2, "cycles_after": 2}
@pytest.mark.asyncio
async def test_attempt_db_reconnect_cancelled_while_waiting_does_not_strand_lock(
prisma_client: PrismaClient,
) -> None:
"""A reconnect cancelled while waiting on the lock (e.g. the readiness
probe deadline firing) must abandon its lock-acquisition task instead of
leaving it to grab the lock later with no owner to release it."""
prisma_client._db_last_reconnect_attempt_ts = 0.0
prisma_client._attempt_reconnect_inside_lock = AsyncMock(return_value=True)
await prisma_client._db_reconnect_lock.acquire()
waiting_reconnect: Final = asyncio.create_task(
prisma_client.attempt_db_reconnect(reason="probe_deadline", lock_timeout_seconds=30.0)
)
await asyncio.sleep(0.05)
waiting_reconnect.cancel()
with pytest.raises(asyncio.CancelledError):
await waiting_reconnect
prisma_client._db_reconnect_lock.release()
await asyncio.sleep(0.05)
assert prisma_client._db_reconnect_lock.locked() is False