mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #32323 from BerriAI/litellm_fix_prisma_reconnect_disconnected_client
fix(proxy): recover Prisma DB reconnect loop when client is disconnected
This commit is contained in:
commit
db60ce9574
7 changed files with 132 additions and 4 deletions
|
|
@ -137,8 +137,17 @@ class PrismaWrapper:
|
|||
self.on_engine_replaced: Callable[[], None] | None = None
|
||||
|
||||
def _get_engine_pid(self) -> int:
|
||||
"""Get the PID of the current Prisma engine subprocess, or 0 if unavailable."""
|
||||
"""Get the PID of the current Prisma engine subprocess, or 0 if unavailable.
|
||||
|
||||
Must never raise: it runs inside the reconnect path, where the client
|
||||
may be in any broken state. Prisma's ``_engine`` is a property that
|
||||
raises ``ClientNotConnectedError`` on a disconnected client; if that
|
||||
escaped here, ``recreate_prisma_client`` would fail before it could
|
||||
build a replacement client and the reconnect loop could never recover.
|
||||
"""
|
||||
try:
|
||||
if self._original_prisma.is_connected() is not True:
|
||||
return 0
|
||||
engine = self._original_prisma._engine
|
||||
process = getattr(engine, "process", None) if engine is not None else None
|
||||
if process is not None:
|
||||
|
|
|
|||
|
|
@ -4097,11 +4097,22 @@ class PrismaClient:
|
|||
raise e
|
||||
|
||||
def _get_engine_pid(self) -> int:
|
||||
"""Get the PID of the writer's engine subprocess, or 0 if unavailable.
|
||||
|
||||
Must never raise: prisma's ``_engine`` property raises
|
||||
``ClientNotConnectedError`` on a disconnected client, and an exception
|
||||
escaping from the reconnect path would leave it unable to recover.
|
||||
"""
|
||||
try:
|
||||
engine = self.db._original_prisma._engine # type: ignore[attr-defined]
|
||||
prisma_obj = self.writer_db._original_prisma
|
||||
if prisma_obj.is_connected() is not True:
|
||||
return 0
|
||||
engine = prisma_obj._engine
|
||||
process = getattr(engine, "process", None) if engine is not None else None
|
||||
if process is not None:
|
||||
return process.pid
|
||||
pid = process.pid
|
||||
if isinstance(pid, int):
|
||||
return pid
|
||||
except (AttributeError, TypeError):
|
||||
pass
|
||||
return 0
|
||||
|
|
|
|||
|
|
@ -20,6 +20,31 @@ _PROXY_MODULE_GLOBALS_TO_ISOLATE = (
|
|||
)
|
||||
|
||||
|
||||
class StubClientNotConnectedError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class DisconnectedPrisma:
|
||||
"""Mimics prisma-client-py after disconnect(): ``is_connected()`` is False
|
||||
and the ``_engine`` property raises ``ClientNotConnectedError``."""
|
||||
|
||||
def is_connected(self) -> bool:
|
||||
return False
|
||||
|
||||
@property
|
||||
def _engine(self) -> None:
|
||||
raise StubClientNotConnectedError(
|
||||
"Client is not connected to the query engine, you must call `connect()` "
|
||||
"before attempting to query data."
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def disconnected_prisma() -> DisconnectedPrisma:
|
||||
"""A stand-in for a Prisma client wedged in the disconnected state."""
|
||||
return DisconnectedPrisma()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_proxy_module_globals():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -92,6 +92,7 @@ async def test_recreate_prisma_client_kills_old_engine_on_disconnect_failure(
|
|||
"""When disconnect() fails, recreate_prisma_client must SIGTERM/SIGKILL the old engine PID."""
|
||||
mock_prisma = AsyncMock()
|
||||
mock_prisma.disconnect.side_effect = Exception("engine hung")
|
||||
mock_prisma.is_connected = MagicMock(return_value=True)
|
||||
|
||||
# Simulate engine subprocess with a known PID
|
||||
mock_engine = MagicMock()
|
||||
|
|
@ -122,6 +123,7 @@ async def test_recreate_prisma_client_skips_kill_on_successful_disconnect(
|
|||
):
|
||||
"""When disconnect() succeeds, no kill should be attempted."""
|
||||
mock_prisma = AsyncMock()
|
||||
mock_prisma.is_connected = MagicMock(return_value=True)
|
||||
mock_prisma.disconnect.return_value = None
|
||||
|
||||
wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False)
|
||||
|
|
@ -142,6 +144,7 @@ async def test_recreate_prisma_client_handles_missing_engine_pid(
|
|||
):
|
||||
"""When engine PID is unavailable (no _engine attr), kill is skipped gracefully."""
|
||||
mock_prisma = AsyncMock()
|
||||
mock_prisma.is_connected = MagicMock(return_value=True)
|
||||
mock_prisma.disconnect.side_effect = Exception("engine hung")
|
||||
mock_prisma._engine = None # No engine subprocess
|
||||
|
||||
|
|
@ -158,3 +161,35 @@ async def test_recreate_prisma_client_handles_missing_engine_pid(
|
|||
|
||||
mock_kill.assert_not_called() # PID was 0, kill skipped
|
||||
mock_new_prisma.connect.assert_awaited_once()
|
||||
|
||||
|
||||
def test_get_engine_pid_returns_zero_for_disconnected_client(disconnected_prisma):
|
||||
"""A disconnected client must read as "no engine" instead of raising,
|
||||
otherwise the reconnect path can never recover."""
|
||||
wrapper = PrismaWrapper(
|
||||
original_prisma=disconnected_prisma, iam_token_db_auth=False
|
||||
)
|
||||
|
||||
assert wrapper._get_engine_pid() == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recreate_prisma_client_recovers_from_disconnected_client(
|
||||
mock_prisma_binary, disconnected_prisma
|
||||
):
|
||||
"""recreate_prisma_client must still build a replacement client when the
|
||||
current one is disconnected."""
|
||||
wrapper = PrismaWrapper(
|
||||
original_prisma=disconnected_prisma, iam_token_db_auth=False
|
||||
)
|
||||
|
||||
mock_new_prisma = AsyncMock()
|
||||
mock_prisma_binary.Prisma.return_value = mock_new_prisma
|
||||
|
||||
with patch("os.kill") as mock_kill:
|
||||
result = await wrapper.recreate_prisma_client("postgresql://new")
|
||||
|
||||
assert result is True
|
||||
mock_kill.assert_not_called()
|
||||
assert wrapper._original_prisma is mock_new_prisma
|
||||
mock_new_prisma.connect.assert_awaited_once()
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ def mock_prisma_binary():
|
|||
def _make_wrapper(engine_pid: int = 111, iam: bool = False) -> PrismaWrapper:
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.connect = AsyncMock()
|
||||
mock_prisma.is_connected = MagicMock(return_value=True)
|
||||
mock_prisma._engine = MagicMock()
|
||||
mock_prisma._engine.process.pid = engine_pid
|
||||
return PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=iam)
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ def mock_prisma_binary():
|
|||
"""Mock prisma.Prisma to avoid requiring generated Prisma binaries for unit tests."""
|
||||
mock_module = MagicMock()
|
||||
with patch.dict(sys.modules, {"prisma": mock_module}):
|
||||
yield
|
||||
yield mock_module
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -515,6 +515,43 @@ async def test_engine_confirmed_dead_persists_across_failed_heavy_reconnect(
|
|||
assert client._engine_confirmed_dead is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_heavy_reconnect_recovers_from_disconnected_prisma_client(
|
||||
mock_proxy_logging, mock_prisma_binary, disconnected_prisma
|
||||
):
|
||||
"""Once the active Prisma client is in the disconnected state, every DB
|
||||
call raises ClientNotConnectedError. The heavy reconnect path is the only
|
||||
way out, so it must not re-raise that same error while inspecting the
|
||||
broken client; otherwise `recreate_prisma_client` fails before it can
|
||||
build a replacement and the proxy loops on failed reconnects forever.
|
||||
|
||||
The full real reconnect path (attempt_db_reconnect -> _run_reconnect_cycle
|
||||
-> recreate_prisma_client) must succeed from that wedged state.
|
||||
"""
|
||||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client.db._original_prisma = disconnected_prisma
|
||||
client._engine_confirmed_dead = True
|
||||
client._start_engine_watcher = AsyncMock()
|
||||
|
||||
replacement = MagicMock()
|
||||
replacement.connect = AsyncMock()
|
||||
mock_prisma_binary.Prisma.return_value = replacement
|
||||
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_disconnected_client",
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
assert client.db._original_prisma is replacement
|
||||
replacement.connect.assert_awaited_once()
|
||||
assert client._consecutive_reconnect_failures == 0
|
||||
assert client._engine_confirmed_dead is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_health_watchdog_should_reconnect_degraded_writer(
|
||||
mock_proxy_logging,
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ def test_get_engine_pid_extracts_process_pid(prisma_client: PrismaClient) -> Non
|
|||
fake_engine.process = MagicMock()
|
||||
fake_engine.process.pid = 4242
|
||||
prisma_client.db._original_prisma = MagicMock()
|
||||
prisma_client.db._original_prisma.is_connected = MagicMock(return_value=True)
|
||||
prisma_client.db._original_prisma._engine = fake_engine
|
||||
actual = {
|
||||
"pid": prisma_client._get_engine_pid(),
|
||||
|
|
@ -58,6 +59,15 @@ def test_get_engine_pid_returns_zero_when_engine_attr_missing(
|
|||
assert prisma_client._get_engine_pid() == 0
|
||||
|
||||
|
||||
def test_get_engine_pid_returns_zero_when_client_disconnected(
|
||||
prisma_client: PrismaClient, disconnected_prisma
|
||||
) -> None:
|
||||
"""The reconnect path calls this on an arbitrarily-broken client; it must
|
||||
report "no engine" instead of re-raising ClientNotConnectedError."""
|
||||
prisma_client.db._original_prisma = disconnected_prisma
|
||||
assert prisma_client._get_engine_pid() == 0
|
||||
|
||||
|
||||
def test_is_engine_alive_true_when_pid_zero(prisma_client: PrismaClient) -> None:
|
||||
prisma_client._engine_pid = 0
|
||||
pinned = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue