From 3eb6e1d91716994df7f43e2032a5258f138bca31 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Sat, 16 May 2026 19:52:11 +0000 Subject: [PATCH] feat(proxy): implement reconnect circuit breaker for database connections - Added a reconnect circuit breaker mechanism to manage database connection failures, including configurable thresholds for attempts, failures, and engine deaths. - Introduced methods to handle reconnect attempts and prune timestamps for tracking failures. - Updated the PrismaClient class to include new attributes for circuit breaker configuration and state management. - Enhanced logging for circuit breaker actions and reasons for state changes. - Added unit tests to validate the behavior of the reconnect circuit breaker under various scenarios. --- litellm/proxy/utils.py | 211 +++++++++++++++++- .../proxy/test_prisma_engine_watchdog.py | 145 ++++++++++++ 2 files changed, 353 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 2998555db1a..62de50a87c8 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1,9 +1,11 @@ import asyncio +from collections import deque import copy import hashlib import inspect import json import os +import signal import smtplib import sys import threading @@ -17,6 +19,7 @@ from typing import ( Any, AsyncGenerator, Awaitable, + Deque, Dict, List, Literal, @@ -2698,6 +2701,48 @@ class PrismaClient: self._reconnect_escalation_threshold: int = max( 1, int(os.getenv("PRISMA_RECONNECT_ESCALATION_THRESHOLD", "3")) ) + self._db_reconnect_circuit_breaker_enabled: bool = ( + str_to_bool( + os.getenv("PRISMA_RECONNECT_CIRCUIT_BREAKER_ENABLED", "true") + ) + is True + ) + self._db_reconnect_circuit_breaker_window_seconds: int = max( + 1, + int(os.getenv("PRISMA_RECONNECT_CIRCUIT_BREAKER_WINDOW_SECONDS", "900")), + ) + self._db_reconnect_circuit_breaker_max_attempts: int = max( + 1, int(os.getenv("PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_ATTEMPTS", "8")) + ) + self._db_reconnect_circuit_breaker_max_failures: int = max( + 1, int(os.getenv("PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_FAILURES", "3")) + ) + self._db_reconnect_circuit_breaker_max_engine_deaths: int = max( + 1, + int(os.getenv("PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_ENGINE_DEATHS", "1")), + ) + self._db_reconnect_circuit_breaker_action: str = os.getenv( + "PRISMA_RECONNECT_CIRCUIT_BREAKER_ACTION", "exit" + ).lower() + if self._db_reconnect_circuit_breaker_action not in {"exit", "log"}: + verbose_proxy_logger.warning( + "Invalid PRISMA_RECONNECT_CIRCUIT_BREAKER_ACTION=%s; defaulting to exit", + self._db_reconnect_circuit_breaker_action, + ) + self._db_reconnect_circuit_breaker_action = "exit" + self._db_reconnect_circuit_breaker_opened: bool = False + self._db_reconnect_breaker_attempts: Deque[float] = deque( + maxlen=self._db_reconnect_circuit_breaker_max_attempts + 1 + ) + self._db_reconnect_breaker_failures: Deque[float] = deque( + maxlen=self._db_reconnect_circuit_breaker_max_failures + ) + self._db_reconnect_breaker_engine_deaths: Deque[float] = deque( + maxlen=self._db_reconnect_circuit_breaker_max_engine_deaths + ) + self._db_reconnect_breaker_last_event_type: Optional[str] = None + self._db_reconnect_breaker_last_reason: Optional[str] = None + self._db_reconnect_breaker_last_engine_pid: Optional[int] = None self._engine_pidfd: int = -1 self._engine_pid: int = 0 self._watching_engine: bool = False @@ -4081,6 +4126,7 @@ class PrismaClient: self.attempt_db_reconnect( reason="engine_process_death", force=True, + engine_pid=pid, ) ) return True @@ -4135,6 +4181,7 @@ class PrismaClient: self.attempt_db_reconnect( reason="engine_process_death", force=True, + engine_pid=dead_pid, ) ) @@ -4189,6 +4236,7 @@ class PrismaClient: self.attempt_db_reconnect( reason="engine_process_death", force=True, + engine_pid=dead_pid, ) ) @@ -4201,9 +4249,10 @@ class PrismaClient: try: os.kill(self._engine_pid, 0) except ProcessLookupError: + dead_pid = self._engine_pid verbose_proxy_logger.error( "prisma-query-engine PID %s gone; triggering reconnect.", - self._engine_pid, + dead_pid, ) self._engine_confirmed_dead = True self._reap_all_zombies() @@ -4211,6 +4260,7 @@ class PrismaClient: await self.attempt_db_reconnect( reason="engine_process_death", force=True, + engine_pid=dead_pid, ) return except (PermissionError, OSError): @@ -4290,6 +4340,135 @@ class PrismaClient: self._engine_confirmed_dead = False verbose_proxy_logger.debug("Stopped engine process watcher.") + @staticmethod + def _prune_reconnect_breaker_timestamps( + timestamps: Deque[float], + cutoff: float, + ) -> None: + while timestamps and timestamps[0] < cutoff: + timestamps.popleft() + + def _prune_reconnect_breaker_events(self, now: float) -> None: + cutoff = now - self._db_reconnect_circuit_breaker_window_seconds + self._prune_reconnect_breaker_timestamps( + self._db_reconnect_breaker_attempts, cutoff + ) + self._prune_reconnect_breaker_timestamps( + self._db_reconnect_breaker_failures, cutoff + ) + self._prune_reconnect_breaker_timestamps( + self._db_reconnect_breaker_engine_deaths, cutoff + ) + + def _get_reconnect_breaker_counts(self) -> Dict[str, int]: + return { + "attempts": len(self._db_reconnect_breaker_attempts), + "failures": len(self._db_reconnect_breaker_failures), + "engine_deaths": len(self._db_reconnect_breaker_engine_deaths), + } + + def _get_reconnect_breaker_open_reason( + self, counts: Dict[str, int] + ) -> Optional[str]: + if ( + counts["engine_deaths"] + >= self._db_reconnect_circuit_breaker_max_engine_deaths + ): + return "engine_process_death_threshold_exceeded" + if counts["failures"] >= self._db_reconnect_circuit_breaker_max_failures: + return "reconnect_failure_threshold_exceeded" + if counts["attempts"] > self._db_reconnect_circuit_breaker_max_attempts: + return "reconnect_attempt_threshold_exceeded" + return None + + def _terminate_for_reconnect_breaker(self) -> None: + if self._db_reconnect_circuit_breaker_action == "log": + return + + def _hard_exit_if_sigterm_did_not_stop_process() -> None: + time.sleep(30) + os._exit(1) + + threading.Thread( + target=_hard_exit_if_sigterm_did_not_stop_process, + daemon=True, + name="prisma-reconnect-circuit-breaker-hard-exit", + ).start() + os.kill(os.getpid(), signal.SIGTERM) + + def _open_reconnect_breaker( + self, + *, + open_reason: str, + current_reason: str, + counts: Dict[str, int], + last_error: Optional[BaseException], + ) -> None: + if self._db_reconnect_circuit_breaker_opened: + return + self._db_reconnect_circuit_breaker_opened = True + verbose_proxy_logger.critical( + "Prisma DB reconnect circuit breaker opened. " + "open_reason=%s current_reason=%s attempts=%s failures=%s engine_deaths=%s " + "window_seconds=%s max_attempts=%s max_failures=%s max_engine_deaths=%s " + "last_event_type=%s last_event_reason=%s last_error_type=%s last_error=%s " + "action=%s worker_pid=%s engine_pid=%s", + open_reason, + current_reason, + counts["attempts"], + counts["failures"], + counts["engine_deaths"], + self._db_reconnect_circuit_breaker_window_seconds, + self._db_reconnect_circuit_breaker_max_attempts, + self._db_reconnect_circuit_breaker_max_failures, + self._db_reconnect_circuit_breaker_max_engine_deaths, + self._db_reconnect_breaker_last_event_type, + self._db_reconnect_breaker_last_reason, + type(last_error).__name__ if last_error is not None else None, + last_error, + self._db_reconnect_circuit_breaker_action, + os.getpid(), + self._db_reconnect_breaker_last_engine_pid, + ) + self._terminate_for_reconnect_breaker() + + def _record_reconnect_breaker_event( + self, + *, + event_type: Literal["attempt", "failure", "success"], + reason: str, + error: Optional[BaseException] = None, + engine_pid: Optional[int] = None, + ) -> bool: + if self._db_reconnect_circuit_breaker_enabled is not True: + return False + now = time.time() + self._prune_reconnect_breaker_events(now=now) + self._db_reconnect_breaker_last_event_type = event_type + self._db_reconnect_breaker_last_reason = reason + if engine_pid is not None: + self._db_reconnect_breaker_last_engine_pid = engine_pid + if event_type == "attempt": + self._db_reconnect_breaker_attempts.append(now) + if reason == "engine_process_death": + self._db_reconnect_breaker_engine_deaths.append(now) + elif event_type == "failure": + self._db_reconnect_breaker_failures.append(now) + else: + return False + + counts = self._get_reconnect_breaker_counts() + open_reason = self._get_reconnect_breaker_open_reason(counts) + if open_reason is None: + return False + self._open_reconnect_breaker( + open_reason=open_reason, + current_reason=reason, + counts=counts, + last_error=error, + ) + return self._db_reconnect_circuit_breaker_action == "exit" + async def _run_reconnect_cycle( self, timeout_seconds: Optional[float] = None ) -> None: @@ -4372,6 +4551,7 @@ class PrismaClient: force: bool, reason: str, timeout_seconds: Optional[float], + engine_pid: Optional[int] = None, ) -> bool: now = time.time() if ( @@ -4398,6 +4578,13 @@ class PrismaClient: ) self._engine_confirmed_dead = True + if self._record_reconnect_breaker_event( + event_type="attempt", + reason=reason, + engine_pid=engine_pid, + ): + return False + verbose_proxy_logger.warning( "Attempting Prisma DB reconnect. reason=%s", reason ) @@ -4407,6 +4594,11 @@ class PrismaClient: await self._run_reconnect_cycle(timeout_seconds=timeout_seconds) reconnect_succeeded = True self._consecutive_reconnect_failures = 0 + self._record_reconnect_breaker_event( + event_type="success", + reason=reason, + engine_pid=engine_pid, + ) verbose_proxy_logger.info( "Prisma DB reconnect succeeded. reason=%s", reason ) @@ -4418,6 +4610,12 @@ class PrismaClient: reason, reconnect_err, ) + self._record_reconnect_breaker_event( + event_type="failure", + reason=reason, + error=reconnect_err, + engine_pid=engine_pid, + ) finally: self._db_last_reconnect_attempt_ts = time.time() @@ -4429,6 +4627,7 @@ class PrismaClient: force: bool = False, timeout_seconds: Optional[float] = None, lock_timeout_seconds: Optional[float] = None, + engine_pid: Optional[int] = None, ) -> bool: """ Attempt to reconnect the Prisma client in a singleflight manner. @@ -4451,7 +4650,7 @@ class PrismaClient: if lock_timeout_seconds is None: async with self._db_reconnect_lock: return await self._attempt_reconnect_inside_lock( - force, reason, timeout_seconds + force, reason, timeout_seconds, engine_pid ) lock_acquired_by_timeout_task = False @@ -4502,7 +4701,7 @@ class PrismaClient: try: return await self._attempt_reconnect_inside_lock( - force, reason, timeout_seconds + force, reason, timeout_seconds, engine_pid ) finally: self._db_reconnect_lock.release() @@ -4558,6 +4757,12 @@ class PrismaClient: if isinstance( e, asyncio.TimeoutError ) or PrismaDBExceptionHandler.is_database_connection_error(e): + verbose_proxy_logger.warning( + "Prisma DB health watchdog probe failed; attempting reconnect. " + "error_type=%s error=%s", + type(e).__name__, + e, + ) await self.attempt_db_reconnect( reason="db_health_watchdog_connection_error", timeout_seconds=self._db_watchdog_reconnect_timeout_seconds, diff --git a/tests/litellm/proxy/test_prisma_engine_watchdog.py b/tests/litellm/proxy/test_prisma_engine_watchdog.py index 0d241f75749..00774dfde9d 100644 --- a/tests/litellm/proxy/test_prisma_engine_watchdog.py +++ b/tests/litellm/proxy/test_prisma_engine_watchdog.py @@ -109,6 +109,7 @@ async def test_poll_missing_process_triggers_reconnect(engine_client) -> None: engine_client.attempt_db_reconnect.assert_awaited_once_with( reason="engine_process_death", force=True, + engine_pid=1234, ) @@ -174,6 +175,7 @@ async def test_pidfd_readable_schedules_reconnect(engine_client) -> None: engine_client.attempt_db_reconnect.assert_awaited_once_with( reason="engine_process_death", force=True, + engine_pid=1234, ) @@ -459,6 +461,7 @@ async def test_on_engine_death_from_thread_triggers_reconnect(engine_client) -> engine_client.attempt_db_reconnect.assert_awaited_once_with( reason="engine_process_death", force=True, + engine_pid=1234, ) @@ -493,6 +496,7 @@ async def test_escalation_after_consecutive_direct_reconnect_failures(engine_cli """After N consecutive direct reconnect failures, _engine_confirmed_dead is set to True so _run_reconnect_cycle takes the heavy reconnect path.""" engine_client._reconnect_escalation_threshold = 3 + engine_client._db_reconnect_circuit_breaker_enabled = False engine_client._consecutive_reconnect_failures = 0 engine_client._db_reconnect_cooldown_seconds = 0 # disable cooldown for test engine_client._start_engine_watcher = AsyncMock(return_value=None) @@ -560,3 +564,144 @@ def test_escalation_threshold_min_guard(mock_proxy_logging): database_url="mock://test", proxy_logging_obj=mock_proxy_logging ) assert client._reconnect_escalation_threshold == 1 + + +# --------------------------------------------------------------------------- +# Reconnect circuit breaker +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_reconnect_circuit_breaker_opens_after_repeated_failures( + engine_client, +): + engine_client._db_reconnect_circuit_breaker_enabled = True + engine_client._db_reconnect_circuit_breaker_max_attempts = 100 + engine_client._db_reconnect_circuit_breaker_max_failures = 2 + engine_client._db_reconnect_circuit_breaker_max_engine_deaths = 100 + engine_client._db_reconnect_circuit_breaker_window_seconds = 900 + engine_client._db_reconnect_cooldown_seconds = 0 + engine_client._start_engine_watcher = AsyncMock(return_value=None) + engine_client.db.recreate_prisma_client = AsyncMock( + side_effect=RuntimeError("recreate failed") + ) + engine_client._terminate_for_reconnect_breaker = MagicMock() + + with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}): + first_result = await engine_client._attempt_reconnect_inside_lock( + force=True, + reason="db_health_watchdog_connection_error", + timeout_seconds=5.0, + ) + second_result = await engine_client._attempt_reconnect_inside_lock( + force=True, + reason="db_health_watchdog_connection_error", + timeout_seconds=5.0, + ) + + assert first_result is False + assert second_result is False + assert engine_client._db_reconnect_circuit_breaker_opened is True + engine_client._terminate_for_reconnect_breaker.assert_called_once() + + +@pytest.mark.asyncio +async def test_reconnect_circuit_breaker_opens_immediately_on_engine_death( + engine_client, +): + engine_client._db_reconnect_circuit_breaker_enabled = True + engine_client._db_reconnect_circuit_breaker_max_attempts = 100 + engine_client._db_reconnect_circuit_breaker_max_failures = 100 + engine_client._db_reconnect_circuit_breaker_max_engine_deaths = 1 + engine_client._db_reconnect_cooldown_seconds = 0 + engine_client._run_reconnect_cycle = AsyncMock(return_value=None) + engine_client._terminate_for_reconnect_breaker = MagicMock() + + result = await engine_client._attempt_reconnect_inside_lock( + force=True, + reason="engine_process_death", + timeout_seconds=5.0, + engine_pid=5678, + ) + + assert result is False + engine_client._run_reconnect_cycle.assert_not_awaited() + assert engine_client._db_reconnect_circuit_breaker_opened is True + assert engine_client._db_reconnect_breaker_last_engine_pid == 5678 + engine_client._terminate_for_reconnect_breaker.assert_called_once() + + +@pytest.mark.asyncio +async def test_reconnect_circuit_breaker_stays_closed_on_transient_success( + engine_client, +): + engine_client._db_reconnect_circuit_breaker_enabled = True + engine_client._db_reconnect_circuit_breaker_max_attempts = 3 + engine_client._db_reconnect_circuit_breaker_max_failures = 2 + engine_client._db_reconnect_circuit_breaker_max_engine_deaths = 1 + engine_client._db_reconnect_cooldown_seconds = 0 + engine_client._start_engine_watcher = AsyncMock(return_value=None) + engine_client.db.recreate_prisma_client = AsyncMock(return_value=None) + engine_client.db.query_raw = AsyncMock(return_value=[{"result": 1}]) + engine_client._terminate_for_reconnect_breaker = MagicMock() + + with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}): + result = await engine_client._attempt_reconnect_inside_lock( + force=True, + reason="db_health_watchdog_connection_error", + timeout_seconds=5.0, + ) + + assert result is True + assert engine_client._db_reconnect_circuit_breaker_opened is False + engine_client._terminate_for_reconnect_breaker.assert_not_called() + + +@pytest.mark.asyncio +async def test_reconnect_circuit_breaker_log_action_does_not_skip_reconnect( + engine_client, +): + engine_client._db_reconnect_circuit_breaker_enabled = True + engine_client._db_reconnect_circuit_breaker_action = "log" + engine_client._db_reconnect_circuit_breaker_max_attempts = 100 + engine_client._db_reconnect_circuit_breaker_max_failures = 100 + engine_client._db_reconnect_circuit_breaker_max_engine_deaths = 1 + engine_client._db_reconnect_cooldown_seconds = 0 + engine_client._run_reconnect_cycle = AsyncMock(return_value=None) + + result = await engine_client._attempt_reconnect_inside_lock( + force=True, + reason="engine_process_death", + timeout_seconds=5.0, + ) + + assert result is True + assert engine_client._db_reconnect_circuit_breaker_opened is True + engine_client._run_reconnect_cycle.assert_awaited_once_with(timeout_seconds=5.0) + + +def test_reconnect_circuit_breaker_env_vars_are_respected(mock_proxy_logging): + with patch.dict( + os.environ, + { + "PRISMA_RECONNECT_CIRCUIT_BREAKER_ENABLED": "false", + "PRISMA_RECONNECT_CIRCUIT_BREAKER_WINDOW_SECONDS": "120", + "PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_ATTEMPTS": "4", + "PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_FAILURES": "2", + "PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_ENGINE_DEATHS": "1", + "PRISMA_RECONNECT_CIRCUIT_BREAKER_ACTION": "log", + }, + ): + client = PrismaClient( + database_url="mock://test", proxy_logging_obj=mock_proxy_logging + ) + + assert client._db_reconnect_circuit_breaker_enabled is False + assert client._db_reconnect_circuit_breaker_window_seconds == 120 + assert client._db_reconnect_circuit_breaker_max_attempts == 4 + assert client._db_reconnect_circuit_breaker_max_failures == 2 + assert client._db_reconnect_circuit_breaker_max_engine_deaths == 1 + assert client._db_reconnect_circuit_breaker_action == "log" + assert client._db_reconnect_breaker_attempts.maxlen == 5 + assert client._db_reconnect_breaker_failures.maxlen == 2 + assert client._db_reconnect_breaker_engine_deaths.maxlen == 1