refactor(proxy): update reconnect circuit breaker logic and enhance testing

- Changed default action for the reconnect circuit breaker from 'exit' to 'log'.
- Added a new method to determine if termination is required based on the action.
- Updated the logic for opening the reconnect circuit breaker to include termination conditions.
- Enhanced unit tests to validate the new behavior of the reconnect circuit breaker, ensuring it does not terminate on log action.
- Introduced end-to-end tests to verify the behavior of the reconnect circuit breaker under different action configurations.
This commit is contained in:
harish-berri 2026-05-16 22:50:44 +00:00
parent d10ad78509
commit 8c01888e74
3 changed files with 218 additions and 35 deletions

View file

@ -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

View file

@ -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"

View file

@ -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