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:
yuneng-jiang 2026-07-07 09:11:02 -07:00 • committed by GitHub
commit db60ce9574
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 132 additions and 4 deletions

View file

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

View file

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

View file

@ -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():
"""

View file

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

View file

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

View file

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

View file

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