fix(proxy): honor allow_requests_on_db_unavailable in /health/readiness

This commit is contained in:
Devin AI 2026-07-28 14:10:11 +00:00
parent daf22ec871
commit ec40a2bfb3
2 changed files with 181 additions and 127 deletions

View file

@ -1406,6 +1406,21 @@ def callback_name(callback):
return str(callback)
def _db_unavailable_should_flip_readiness_to_503() -> bool:
"""
Whether an unreachable-but-configured DB should mark this worker NotReady.
Defaults to True so orchestrators pull the pod out of rotation when the DB
is down. Operators who set general_settings.allow_requests_on_db_unavailable
are explicitly opting to keep serving during a DB outage (the request layer
already fails open via PrismaDBExceptionHandler), so readiness must stay 200
and keep the pod in the Service endpoints; flipping to 503 here would pull
every replica out of rotation and defeat the very flag meant to survive the
outage.
"""
return not PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
async def _get_health_readiness_details(
response: Optional[Response] = None,
) -> Dict[str, Any]:
@ -1454,7 +1469,13 @@ 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":
# When allow_requests_on_db_unavailable is set the operator wants
# the pod to keep serving through the outage, so readiness stays 200.
if (
response is not None
and db_health_status["status"] != "connected"
and _db_unavailable_should_flip_readiness_to_503()
):
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return {
"status": "healthy",
@ -1534,8 +1555,10 @@ def _authorize_drain_request(request: Request) -> None:
async def _resolve_public_readiness_db(response: Response) -> str:
"""
Return the db status string for the public probe and flip the response to
503 when a configured DB is unreachable. Mirrors the legacy values:
"Not connected" (no DB configured), "connected", "disconnected".
503 when a configured DB is unreachable, unless the operator opted into
allow_requests_on_db_unavailable (see _db_unavailable_should_flip_readiness_to_503).
Mirrors the legacy values: "Not connected" (no DB configured), "connected",
"disconnected".
"""
from litellm.proxy.proxy_server import prisma_client
@ -1543,7 +1566,7 @@ async def _resolve_public_readiness_db(response: Response) -> str:
return "Not connected"
db_health_status = await _db_health_readiness_check()
if db_health_status["status"] != "connected":
if db_health_status["status"] != "connected" and _db_unavailable_should_flip_readiness_to_503():
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return db_health_status["status"]

View file

@ -5,9 +5,7 @@ from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
import httpx
import pytest
@ -145,9 +143,7 @@ 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")
@pytest.mark.asyncio
@ -177,9 +173,7 @@ 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")
assert mock_prisma.health_check.call_count == 2
@ -199,9 +193,7 @@ async def test_db_health_transport_error_reconnect_fails(transport_error):
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=transport_error)
mock_prisma.attempt_db_reconnect = AsyncMock(
side_effect=RuntimeError("reconnect failed")
)
mock_prisma.attempt_db_reconnect = AsyncMock(side_effect=RuntimeError("reconnect failed"))
_health_endpoints_module.db_health_cache = {
"status": "connected",
@ -253,9 +245,7 @@ async def test_health_services_endpoint_sqs(status, error_message):
"""
with patch("litellm.integrations.sqs.SQSLogger") as MockSQSLogger:
mock_instance = MagicMock()
mock_instance.async_health_check = AsyncMock(
return_value={"status": status, "error_message": error_message}
)
mock_instance.async_health_check = AsyncMock(return_value={"status": status, "error_message": error_message})
MockSQSLogger.return_value = mock_instance
result = await health_services_endpoint(service="sqs")
@ -452,14 +442,9 @@ async def test_test_model_connection_loads_config_from_router():
# Verify that config params were loaded and merged
# Note: request params override config params, so model from request is used
assert model_params.get("api_key") == "resolved-api-key-from-env"
assert (
model_params.get("api_base")
== "https://resolved-endpoint.openai.azure.com/"
)
assert model_params.get("api_base") == "https://resolved-endpoint.openai.azure.com/"
assert model_params.get("api_version") == "2024-10-21"
assert (
model_params.get("model") == "gpt-4o"
) # Request param overrides config param
assert model_params.get("model") == "gpt-4o" # Request param overrides config param
# Verify result
assert result["status"] == "success"
@ -595,9 +580,7 @@ async def test_test_model_connection_uses_model_info_id_to_disambiguate_duplicat
assert ahealth_check_call_args is not None
model_params = ahealth_check_call_args.kwargs.get("model_params", {})
assert model_params.get("api_base") == (
"https://deployment-B-base.invalid/v1"
), (
assert model_params.get("api_base") == ("https://deployment-B-base.invalid/v1"), (
"Expected /health/test_connection to probe deployment B's "
"api_base when model_info.id='deployment-B-id' was provided. "
f"Got: {model_params.get('api_base')!r}. This means the "
@ -772,14 +755,10 @@ async def test_test_model_connection_uses_loaded_deployment_team_id():
"can_user_make_model_call",
wraps=ModelManagementAuthChecks.can_user_make_model_call,
) as spy_auth_check,
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.TeamRepository"
) as MockTeamRepo,
patch("litellm.proxy.management_endpoints.model_management_endpoints.TeamRepository") as MockTeamRepo,
):
mock_team_repo_instance = MagicMock()
mock_team_repo_instance.table.find_unique = AsyncMock(
side_effect=fake_find_unique
)
mock_team_repo_instance.table.find_unique = AsyncMock(side_effect=fake_find_unique)
MockTeamRepo.return_value = mock_team_repo_instance
with pytest.raises(HTTPException) as exc_info:
@ -874,14 +853,10 @@ async def test_test_model_connection_uses_loaded_deployment_team_id_via_model_na
"can_user_make_model_call",
wraps=ModelManagementAuthChecks.can_user_make_model_call,
) as spy_auth_check,
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.TeamRepository"
) as MockTeamRepo,
patch("litellm.proxy.management_endpoints.model_management_endpoints.TeamRepository") as MockTeamRepo,
):
mock_team_repo_instance = MagicMock()
mock_team_repo_instance.table.find_unique = AsyncMock(
side_effect=fake_find_unique
)
mock_team_repo_instance.table.find_unique = AsyncMock(side_effect=fake_find_unique)
MockTeamRepo.return_value = mock_team_repo_instance
with pytest.raises(HTTPException) as exc_info:
@ -950,9 +925,7 @@ async def test_test_model_connection_authorized_team_admin_passes_real_auth():
return SimpleNamespace(
model_dump=lambda: LiteLLM_TeamTable(
team_id=owner_team_id,
members_with_roles=[
{"user_id": owner_admin_user_id, "role": "admin"}
],
members_with_roles=[{"user_id": owner_admin_user_id, "role": "admin"}],
).model_dump()
)
return None
@ -968,9 +941,7 @@ async def test_test_model_connection_authorized_team_admin_passes_real_auth():
"can_user_make_model_call",
wraps=ModelManagementAuthChecks.can_user_make_model_call,
) as spy_auth_check,
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.TeamRepository"
) as MockTeamRepo,
patch("litellm.proxy.management_endpoints.model_management_endpoints.TeamRepository") as MockTeamRepo,
patch(
"litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check",
AsyncMock(return_value=health_result),
@ -981,9 +952,7 @@ async def test_test_model_connection_authorized_team_admin_passes_real_auth():
),
):
mock_team_repo_instance = MagicMock()
mock_team_repo_instance.table.find_unique = AsyncMock(
side_effect=fake_find_unique
)
mock_team_repo_instance.table.find_unique = AsyncMock(side_effect=fake_find_unique)
MockTeamRepo.return_value = mock_team_repo_instance
result = await health_test_model_connection(
@ -1010,9 +979,7 @@ async def test_test_model_connection_authorized_team_admin_passes_real_auth():
async def test_health_services_endpoint_galileo(status, error_message):
with patch("litellm.integrations.galileo.GalileoObserve") as MockGalileoObserve:
mock_instance = MagicMock()
mock_instance.async_health_check = AsyncMock(
return_value={"status": status, "error_message": error_message}
)
mock_instance.async_health_check = AsyncMock(return_value={"status": status, "error_message": error_message})
MockGalileoObserve.return_value = mock_instance
result = await health_services_endpoint(service="galileo")
@ -1085,13 +1052,9 @@ async def test_health_services_endpoint_newrelic_blocks_non_admin(role):
user_role=role,
)
with patch(
"litellm.integrations.newrelic.newrelic.NewRelicLogger"
) as MockNewRelicLogger:
with patch("litellm.integrations.newrelic.newrelic.NewRelicLogger") as MockNewRelicLogger:
mock_instance = MagicMock()
mock_instance.async_health_check = AsyncMock(
return_value={"status": "healthy", "error_message": ""}
)
mock_instance.async_health_check = AsyncMock(return_value={"status": "healthy", "error_message": ""})
MockNewRelicLogger.return_value = mock_instance
with pytest.raises(ProxyException) as exc_info:
@ -1120,13 +1083,9 @@ async def test_health_services_endpoint_newrelic_allows_proxy_admin(admin_role):
user_role=admin_role,
)
with patch(
"litellm.integrations.newrelic.newrelic.NewRelicLogger"
) as MockNewRelicLogger:
with patch("litellm.integrations.newrelic.newrelic.NewRelicLogger") as MockNewRelicLogger:
mock_instance = MagicMock()
mock_instance.async_health_check = AsyncMock(
return_value={"status": "healthy", "error_message": ""}
)
mock_instance.async_health_check = AsyncMock(return_value={"status": "healthy", "error_message": ""})
MockNewRelicLogger.return_value = mock_instance
result = await health_services_endpoint(
@ -1177,20 +1136,14 @@ def test_health_liveliness_endpoint(proxy_client):
duration_ms = (end_time - start_time) * 1000
# Assert response status
assert (
response.status_code == 200
), f"Expected 200 OK, got {response.status_code}: {response.text}"
assert response.status_code == 200, f"Expected 200 OK, got {response.status_code}: {response.text}"
# Assert response content (FastAPI JSON-encodes the string)
assert (
response.json() == "I'm alive!"
), f"Expected 'I'm alive!' message, got: {response.json()}"
assert response.json() == "I'm alive!", f"Expected 'I'm alive!' message, got: {response.json()}"
# Verify response is fast (should be < 100ms for a simple endpoint)
# This is critical for orchestration systems that poll frequently
assert (
duration_ms < 100
), f"Health check took {duration_ms:.2f}ms, expected < 100ms for a simple endpoint"
assert duration_ms < 100, f"Health check took {duration_ms:.2f}ms, expected < 100ms for a simple endpoint"
# Log the duration for visibility (useful for CI/CD monitoring)
print(f"\n/health/liveliness response time: {duration_ms:.2f}ms")
@ -1210,19 +1163,13 @@ def test_health_liveness_endpoint(proxy_client):
duration_ms = (end_time - start_time) * 1000
# Assert response status
assert (
response.status_code == 200
), f"Expected 200 OK, got {response.status_code}: {response.text}"
assert response.status_code == 200, f"Expected 200 OK, got {response.status_code}: {response.text}"
# Assert response content (FastAPI JSON-encodes the string)
assert (
response.json() == "I'm alive!"
), f"Expected 'I'm alive!' message, got: {response.json()}"
assert response.json() == "I'm alive!", f"Expected 'I'm alive!' message, got: {response.json()}"
# Verify response is fast (should be < 100ms for a simple endpoint)
assert (
duration_ms < 100
), f"Health check took {duration_ms:.2f}ms, expected < 100ms for a simple endpoint"
assert duration_ms < 100, f"Health check took {duration_ms:.2f}ms, expected < 100ms for a simple endpoint"
# Log the duration for visibility (useful for CI/CD monitoring)
print(f"\n/health/liveness response time: {duration_ms:.2f}ms")
@ -1243,15 +1190,11 @@ def test_health_readiness(proxy_client):
duration_ms = (end_time - start_time) * 1000
# Assert response status
assert (
response.status_code == 200
), f"Expected 200 OK, got {response.status_code}: {response.text}"
assert response.status_code == 200, f"Expected 200 OK, got {response.status_code}: {response.text}"
# Verify response is fast (readiness may include DB check if available, so < 500ms is reasonable)
# This is critical for orchestration systems (Kubernetes) that poll frequently
assert (
duration_ms < 500
), f"Health check took {duration_ms:.2f}ms, expected < 500ms for readiness endpoint"
assert duration_ms < 500, f"Health check took {duration_ms:.2f}ms, expected < 500ms for readiness endpoint"
# Assert response contains only low-detail public probe fields. `db` is
# included so unauthenticated probes can distinguish "DB unreachable"
@ -1270,9 +1213,7 @@ def test_health_readiness_details_returns_diagnostic_fields(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)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
@ -1424,9 +1365,7 @@ def test_get_callback_identifier_custom_logger_registry_and_fallback():
unregistered = UnregisteredCallback()
# Mock registry to return empty list (not registered)
with patch.object(
CustomLoggerRegistry, "get_all_callback_strs_from_class_type", return_value=[]
):
with patch.object(CustomLoggerRegistry, "get_all_callback_strs_from_class_type", return_value=[]):
result = get_callback_identifier(unregistered)
# Should fall back to callback_name() which returns __class__.__name__
assert result == "UnregisteredCallback"
@ -1515,13 +1454,9 @@ async def test_health_endpoint_filters_model_list_by_user_access():
await health_endpoint(response=Response(), user_api_key_dict=user_api_key_dict)
assert (
"model_list" in captured
), "health_endpoint did not call _perform_health_check_and_save"
assert "model_list" in captured, "health_endpoint did not call _perform_health_check_and_save"
returned_names = {m["model_name"] for m in captured["model_list"]}
assert returned_names == {
"model-a"
}, f"health_endpoint did not scope model_list to caller access: {returned_names}"
assert returned_names == {"model-a"}, f"health_endpoint did not scope model_list to caller access: {returned_names}"
@pytest.mark.asyncio
@ -1651,9 +1586,7 @@ async def test_health_endpoint_resolves_all_team_models_to_team_allowlist():
await health_endpoint(response=Response(), user_api_key_dict=user_api_key_dict)
returned_names = {m["model_name"] for m in captured["model_list"]}
assert returned_names == {
"model-b"
}, f"all-team-models key should health-check the team's models: {returned_names}"
assert returned_names == {"model-b"}, f"all-team-models key should health-check the team's models: {returned_names}"
@pytest.mark.asyncio
@ -1735,15 +1668,13 @@ async def test_health_endpoint_filters_background_cache_by_user_access():
# vacuously when the cache filter drops everything because cached
# entries lack the model_id key — both entries carry model_id above.)
assert len(cached_results["healthy_endpoints"]) == 2
assert all(
ep.get("model_id") for ep in cached_results["healthy_endpoints"]
), "test fixture invariant: every cached entry must carry a model_id"
assert all(ep.get("model_id") for ep in cached_results["healthy_endpoints"]), (
"test fixture invariant: every cached entry must carry a model_id"
)
# The non-admin caller must not see api_base on the returned cache entries.
returned = result.get("healthy_endpoints", [])
assert (
len(returned) == 1
), f"expected exactly one cached entry after scoping, got {len(returned)}"
assert len(returned) == 1, f"expected exactly one cached entry after scoping, got {len(returned)}"
assert returned[0]["model_id"] == "id-a"
assert "api_base" not in returned[0]
assert result["healthy_count"] == 1
@ -1834,13 +1765,12 @@ async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not():
non_admin_eps = non_admin_result.get("healthy_endpoints", [])
assert len(admin_eps) == 1
assert (
admin_eps[0]["api_base"]
== "https://us-central1-aiplatform.googleapis.com/v1/projects/p"
), "admin must see the full api_base so they can identify the region"
assert (
admin_eps[0]["api_version"] == "2024-10-21"
), "admin must see api_version so they can distinguish provider deployments"
assert admin_eps[0]["api_base"] == "https://us-central1-aiplatform.googleapis.com/v1/projects/p", (
"admin must see the full api_base so they can identify the region"
)
assert admin_eps[0]["api_version"] == "2024-10-21", (
"admin must see api_version so they can distinguish provider deployments"
)
assert len(non_admin_eps) == 1
assert "api_base" not in non_admin_eps[0]
@ -1857,10 +1787,7 @@ async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not():
# Stripping must produce a copy — the shared cache must still carry the
# routing fields so the next admin caller can read them.
cached_first = cached_results["healthy_endpoints"][0]
assert (
cached_first["api_base"]
== "https://us-central1-aiplatform.googleapis.com/v1/projects/p"
)
assert cached_first["api_base"] == "https://us-central1-aiplatform.googleapis.com/v1/projects/p"
assert cached_first["api_version"] == "2024-10-21"
@ -2005,9 +1932,7 @@ async def test_health_endpoint_blocks_cross_scope_model_id_under_background_cach
leaked_ids = {ep.get("model_id") for ep in result.get("healthy_endpoints", [])}
leaked_ids |= {ep.get("model_id") for ep in result.get("unhealthy_endpoints", [])}
assert (
"id-b" not in leaked_ids
), "background cache leaked an out-of-scope deployment to a scoped caller"
assert "id-b" not in leaked_ids, "background cache leaked an out-of-scope deployment to a scoped caller"
assert result["healthy_count"] == 0
assert response.status_code == 503
@ -2234,9 +2159,7 @@ async def test_health_endpoint_no_model_param_returns_200_even_when_zero_healthy
async def fake_perform(**kwargs):
return {
"healthy_endpoints": [],
"unhealthy_endpoints": [
{"model": "openai/gpt-4o", "model_id": "id-a", "error": "boom"}
],
"unhealthy_endpoints": [{"model": "openai/gpt-4o", "model_id": "id-a", "error": "boom"}],
"healthy_count": 0,
"unhealthy_count": 1,
}
@ -2299,6 +2222,114 @@ async def test_health_readiness_returns_503_when_db_disconnected():
assert result == {"status": "healthy", "db": "disconnected"}
@pytest.mark.asyncio
async def test_health_readiness_stays_200_when_db_down_and_allow_requests_on_db_unavailable():
"""
Regression for #34934: with general_settings.allow_requests_on_db_unavailable
set, an unreachable-but-configured DB must NOT flip readiness to 503. The
operator explicitly opted to keep serving through the outage (the request
layer fails open), so the probe has to stay 200 and report db disconnected;
otherwise K8s pulls every replica out of rotation and defeats the flag.
"""
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("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"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_stays_200_when_db_down_and_allow_requests_on_db_unavailable():
"""
Same #34934 regression on the authenticated/detailed payload path
(allow_public_health_readiness_details): the flag must keep it 200.
"""
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("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.proxy_server.general_settings",
{
"allow_requests_on_db_unavailable": True,
"allow_public_health_readiness_details": True,
},
),
):
result = await health_readiness(response=response)
assert response.status_code == 200
assert result["status"] == "healthy"
assert result["db"] == "disconnected"
@pytest.mark.asyncio
async def test_health_readiness_details_returns_503_when_db_down_without_flag():
"""
Detailed payload path keeps the default: no allow_requests_on_db_unavailable
flag means an unreachable DB still flips the pod to 503.
"""
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("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_public_health_readiness_details": True},
),
):
result = await health_readiness(response=response)
assert response.status_code == 503
assert result["db"] == "disconnected"
@pytest.mark.asyncio
async def test_health_readiness_returns_200_when_db_connected():
"""Happy path: connected DB keeps the legacy 200."""