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.
This commit is contained in:
harish-berri 2026-05-16 19:52:11 +00:00
parent 9ac4092536
commit 3eb6e1d917
2 changed files with 353 additions and 3 deletions

View file

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

View file

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