diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 76612fdd178..b401528f64d 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -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 { diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index 0e108ea5817..bc3aec58991 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -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: