mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
d10ad78509
commit
8c01888e74
3 changed files with 218 additions and 35 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue