diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 6f140266fca..79be325bd44 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2722,15 +2722,16 @@ class PrismaClient: 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" + "PRISMA_RECONNECT_CIRCUIT_BREAKER_ACTION", "log" ).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", + "Invalid PRISMA_RECONNECT_CIRCUIT_BREAKER_ACTION=%s; defaulting to log", self._db_reconnect_circuit_breaker_action, ) - self._db_reconnect_circuit_breaker_action = "exit" + self._db_reconnect_circuit_breaker_action = "log" self._db_reconnect_circuit_breaker_opened: bool = False + self._db_reconnect_circuit_breaker_termination_started: bool = False self._db_reconnect_breaker_attempts: Deque[float] = deque( maxlen=self._db_reconnect_circuit_breaker_max_attempts ) @@ -4382,8 +4383,12 @@ class PrismaClient: return None def _terminate_for_reconnect_breaker(self) -> None: - if self._db_reconnect_circuit_breaker_action == "log": + if ( + self._db_reconnect_circuit_breaker_action == "log" + or self._db_reconnect_circuit_breaker_termination_started + ): return + self._db_reconnect_circuit_breaker_termination_started = True def _hard_exit_if_sigterm_did_not_stop_process() -> None: time.sleep(30) @@ -4396,6 +4401,11 @@ class PrismaClient: ).start() os.kill(os.getpid(), signal.SIGTERM) + def _should_terminate_for_reconnect_breaker(self, open_reason: str) -> bool: + if self._db_reconnect_circuit_breaker_action != "exit": + return False + return open_reason == "engine_process_death_threshold_exceeded" + def _open_reconnect_breaker( self, *, @@ -4404,33 +4414,37 @@ class PrismaClient: counts: Dict[str, int], last_error: Optional[BaseException], ) -> None: - if self._db_reconnect_circuit_breaker_opened: + should_terminate = self._should_terminate_for_reconnect_breaker(open_reason) + if self._db_reconnect_circuit_breaker_opened and not should_terminate: 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() + if not self._db_reconnect_circuit_breaker_opened or should_terminate: + 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 should_terminate=%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, + should_terminate, + os.getpid(), + self._db_reconnect_breaker_last_engine_pid, + ) + if should_terminate: + self._terminate_for_reconnect_breaker() def _record_reconnect_breaker_event( self, @@ -4469,7 +4483,7 @@ class PrismaClient: counts=counts, last_error=error, ) - return self._db_reconnect_circuit_breaker_action == "exit" + return self._should_terminate_for_reconnect_breaker(open_reason) async def _run_reconnect_cycle( self, timeout_seconds: Optional[float] = None diff --git a/tests/litellm/proxy/test_prisma_engine_watchdog.py b/tests/litellm/proxy/test_prisma_engine_watchdog.py index 11b68e2b7a4..15eff561ae7 100644 --- a/tests/litellm/proxy/test_prisma_engine_watchdog.py +++ b/tests/litellm/proxy/test_prisma_engine_watchdog.py @@ -573,10 +573,11 @@ def test_escalation_threshold_min_guard(mock_proxy_logging): @pytest.mark.asyncio -async def test_reconnect_circuit_breaker_opens_after_repeated_failures( +async def test_reconnect_circuit_breaker_exit_action_does_not_terminate_on_reconnect_failures( engine_client, ): engine_client._db_reconnect_circuit_breaker_enabled = True + engine_client._db_reconnect_circuit_breaker_action = "exit" 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 @@ -603,7 +604,7 @@ async def test_reconnect_circuit_breaker_opens_after_repeated_failures( 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() + engine_client._terminate_for_reconnect_breaker.assert_not_called() @pytest.mark.asyncio @@ -611,6 +612,7 @@ 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_action = "exit" 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 @@ -663,6 +665,7 @@ async def test_reconnect_circuit_breaker_opens_on_exact_attempt_threshold( engine_client, ): engine_client._db_reconnect_circuit_breaker_enabled = True + engine_client._db_reconnect_circuit_breaker_action = "exit" engine_client._db_reconnect_circuit_breaker_max_attempts = 2 engine_client._db_reconnect_circuit_breaker_max_failures = 100 engine_client._db_reconnect_circuit_breaker_max_engine_deaths = 100 @@ -685,10 +688,10 @@ async def test_reconnect_circuit_breaker_opens_on_exact_attempt_threshold( ) assert first_result is True - assert second_result is False + assert second_result is True assert len(engine_client._db_reconnect_breaker_attempts) == 2 assert engine_client._db_reconnect_circuit_breaker_opened is True - engine_client._terminate_for_reconnect_breaker.assert_called_once() + engine_client._terminate_for_reconnect_breaker.assert_not_called() def test_reconnect_circuit_breaker_success_clears_failure_state(engine_client): @@ -754,3 +757,25 @@ def test_reconnect_circuit_breaker_env_vars_are_respected(mock_proxy_logging): assert client._db_reconnect_breaker_attempts.maxlen == 4 assert client._db_reconnect_breaker_failures.maxlen == 2 assert client._db_reconnect_breaker_engine_deaths.maxlen == 1 + + +def test_reconnect_circuit_breaker_action_defaults_to_log(mock_proxy_logging): + with patch.dict(os.environ, {}, clear=True): + client = PrismaClient( + database_url="mock://test", proxy_logging_obj=mock_proxy_logging + ) + + assert client._db_reconnect_circuit_breaker_enabled is True + assert client._db_reconnect_circuit_breaker_action == "log" + + +def test_reconnect_circuit_breaker_invalid_action_defaults_to_log(mock_proxy_logging): + with patch.dict( + os.environ, + {"PRISMA_RECONNECT_CIRCUIT_BREAKER_ACTION": "invalid"}, + ): + client = PrismaClient( + database_url="mock://test", proxy_logging_obj=mock_proxy_logging + ) + + assert client._db_reconnect_circuit_breaker_action == "log" diff --git a/tests/test_litellm/proxy/db/test_prisma_reconnect_circuit_breaker_e2e.py b/tests/test_litellm/proxy/db/test_prisma_reconnect_circuit_breaker_e2e.py new file mode 100644 index 00000000000..401ed3c3bd9 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_prisma_reconnect_circuit_breaker_e2e.py @@ -0,0 +1,144 @@ +import os +import subprocess +import sys +from pathlib import Path + +import pytest + + +REPO_ROOT = Path(__file__).resolve().parents[4] + +REPRO_CODE = r""" +import asyncio +import os +import signal +import sys +import types +from pathlib import Path + + +async def noop_failure_handler(*_args, **_kwargs) -> None: + return None + + +def add_repo_root_to_python_path() -> None: + repo_root = Path(os.environ["LITELLM_REPO_ROOT"]) + sys.path.insert(0, str(repo_root)) + + +def configure_breaker_for_fast_repro() -> None: + breaker_action = os.environ["REPRO_BREAKER_ACTION"] + os.environ["PRISMA_RECONNECT_COOLDOWN_SECONDS"] = "1" + os.environ["LITELLM_LOG"] = "INFO" + os.environ["PRISMA_HEALTH_WATCHDOG_PROBE_TIMEOUT_SECONDS"] = "0.5" + os.environ["PRISMA_WATCHDOG_RECONNECT_TIMEOUT_SECONDS"] = "1" + os.environ["PRISMA_RECONNECT_CIRCUIT_BREAKER_ENABLED"] = "true" + os.environ["PRISMA_RECONNECT_CIRCUIT_BREAKER_ACTION"] = breaker_action + os.environ["PRISMA_RECONNECT_CIRCUIT_BREAKER_WINDOW_SECONDS"] = "60" + os.environ["PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_ATTEMPTS"] = "10" + os.environ["PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_FAILURES"] = "2" + os.environ["PRISMA_RECONNECT_CIRCUIT_BREAKER_MAX_ENGINE_DEATHS"] = ( + "1" if breaker_action == "exit" else "10" + ) + + +def handle_sigterm(_signum, _frame) -> None: + print("Received SIGTERM from Prisma reconnect circuit breaker; exiting cleanly.") + raise SystemExit(0) + + +async def main() -> None: + add_repo_root_to_python_path() + configure_breaker_for_fast_repro() + signal.signal(signal.SIGTERM, handle_sigterm) + + from litellm.proxy.utils import PrismaClient + + proxy_logging = types.SimpleNamespace(failure_handler=noop_failure_handler) + client = PrismaClient( + database_url=os.environ["DATABASE_URL"], + proxy_logging_obj=proxy_logging, + ) + client._db_reconnect_cooldown_seconds = 0 + client._db_health_watchdog_interval_seconds = 0 + client._db_health_watchdog_probe_timeout_seconds = 0.5 + client._db_watchdog_reconnect_timeout_seconds = 1 + + breaker_action = os.environ["REPRO_BREAKER_ACTION"] + print(f"Breaker action: {breaker_action}") + print(f"DATABASE_URL: {os.environ['DATABASE_URL']}") + if breaker_action == "exit": + print("Expected: engine death opens breaker and sends SIGTERM.") + await client.attempt_db_reconnect( + reason="engine_process_death", + force=True, + timeout_seconds=1, + engine_pid=1234, + ) + else: + print("Expected: breaker opens in log mode and process remains alive.") + for attempt in range(1, 3): + result = await client.attempt_db_reconnect( + reason="db_health_watchdog_connection_error", + force=True, + timeout_seconds=1, + ) + print(f"Reconnect attempt {attempt} result: {result}") + print("Log action complete: breaker opened and process is still alive.") + return + + raise RuntimeError("repro failed: circuit breaker did not terminate process") + + +if __name__ == "__main__": + asyncio.run(main()) +""" + + +pytestmark = pytest.mark.skipif( + not os.getenv("DATABASE_URL"), + reason="Requires DATABASE_URL for Prisma DB e2e test", +) + + +def _run_repro(action: str) -> subprocess.CompletedProcess[str]: + env = { + **os.environ, + "DATABASE_URL": os.environ["DATABASE_URL"], + "LITELLM_REPO_ROOT": str(REPO_ROOT), + "REPRO_BREAKER_ACTION": action, + } + return subprocess.run( + [ + sys.executable, + "-c", + REPRO_CODE, + ], + cwd=REPO_ROOT, + env=env, + check=False, + capture_output=True, + text=True, + timeout=60, + ) + + +def test_prisma_reconnect_circuit_breaker_log_action_does_not_exit(): + result = _run_repro("log") + output = result.stdout + result.stderr + + assert result.returncode == 0, output + assert "Breaker action: log" in output + assert "Prisma DB reconnect circuit breaker opened" in output + assert "Log action complete: breaker opened and process is still alive." in output + assert "Received SIGTERM from Prisma reconnect circuit breaker" not in output + + +def test_prisma_reconnect_circuit_breaker_exit_action_exits_cleanly(): + result = _run_repro("exit") + output = result.stdout + result.stderr + + assert result.returncode == 0, output + assert "Breaker action: exit" in output + assert "Prisma DB reconnect circuit breaker opened" in output + assert "Received SIGTERM from Prisma reconnect circuit breaker; exiting cleanly." in output