fix(proxy): kill orphaned prisma engine subprocess on failed disconnect

This commit is contained in:
michelligabriele 2026-03-19 19:50:39 +01:00
parent 81dadb698a
commit 1f04fa2461
4 changed files with 170 additions and 4 deletions

View file

@ -5,6 +5,7 @@ This file contains the PrismaWrapper class, which is used to wrap the Prisma cli
import asyncio
import os
import random
import signal
import subprocess
import time
import urllib
@ -45,6 +46,46 @@ class PrismaWrapper:
self._reconnection_lock = asyncio.Lock()
self._last_refresh_time: Optional[datetime] = None
def _get_engine_pid(self) -> int:
"""Get the PID of the current Prisma engine subprocess, or 0 if unavailable."""
try:
engine = self._original_prisma._engine
process = getattr(engine, "process", None) if engine is not None else None
if process is not None:
return process.pid
except (AttributeError, TypeError):
pass
return 0
@staticmethod
def _kill_engine_process(pid: int) -> None:
"""Force-kill an orphaned engine subprocess to prevent DB connection pool leaks.
Called when disconnect() fails and the old engine process may still be
holding open connections. Sends SIGTERM for graceful shutdown, waits
briefly, then SIGKILL as a backstop.
"""
if pid <= 0:
return
try:
os.kill(pid, signal.SIGTERM)
except (ProcessLookupError, PermissionError, OSError):
return # Already dead or inaccessible
verbose_proxy_logger.warning(
"Sent SIGTERM to orphaned prisma-query-engine PID %s after failed disconnect.",
pid,
)
# Brief wait for graceful shutdown, then force-kill
time.sleep(0.5)
try:
os.kill(pid, signal.SIGKILL)
verbose_proxy_logger.warning(
"Sent SIGKILL to prisma-query-engine PID %s (did not exit after SIGTERM).",
pid,
)
except (ProcessLookupError, PermissionError, OSError):
pass # Exited after SIGTERM — expected
def _extract_token_from_db_url(self, db_url: Optional[str]) -> Optional[str]:
"""
Extract the token (password) from the DATABASE_URL.
@ -179,10 +220,13 @@ class PrismaWrapper:
"""Disconnect and reconnect the Prisma client with a new database URL."""
from prisma import Prisma # type: ignore
old_engine_pid = self._get_engine_pid()
try:
await self._original_prisma.disconnect()
except Exception as e:
verbose_proxy_logger.warning(f"Failed to disconnect Prisma client: {e}")
self._kill_engine_process(old_engine_pid)
if http_client is not None:
self._original_prisma = Prisma(http=http_client)

View file

@ -1917,6 +1917,7 @@ class ProxyLogging:
original_exception,
traceback.format_exc(),
),
daemon=True,
).start()
async def post_call_success_hook(
@ -4005,13 +4006,15 @@ class PrismaClient:
)
async def _do_direct_reconnect() -> None:
old_pid = self._get_engine_pid()
try:
await self.db.disconnect()
except Exception as disconnect_err:
verbose_proxy_logger.debug(
"Prisma DB disconnect before reconnect failed (ignored): %s",
verbose_proxy_logger.warning(
"Prisma DB disconnect before reconnect failed: %s",
disconnect_err,
)
PrismaWrapper._kill_engine_process(old_pid)
await self.db.connect()
await self.db.query_raw("SELECT 1")

View file

@ -1,7 +1,8 @@
import json
import os
import signal
import sys
from unittest.mock import AsyncMock, Mock, patch
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import pytest
from fastapi.testclient import TestClient
@ -14,6 +15,14 @@ sys.path.insert(
from litellm.proxy.db.prisma_client import PrismaWrapper, should_update_prisma_schema
@pytest.fixture(autouse=True)
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 mock_module
def test_should_update_prisma_schema(monkeypatch):
# CASE 1: Environment variable behavior
# When DISABLE_SCHEMA_UPDATE is not set -> should update
@ -73,4 +82,79 @@ async def test_recreate_prisma_client_successful_disconnect():
# Verify that the new client replaced the original
assert wrapper._original_prisma != mock_prisma
assert hasattr(wrapper._original_prisma, 'connect')
assert hasattr(wrapper._original_prisma, 'connect')
@pytest.mark.asyncio
async def test_recreate_prisma_client_kills_old_engine_on_disconnect_failure(
mock_prisma_binary,
):
"""When disconnect() fails, recreate_prisma_client must SIGTERM/SIGKILL the old engine PID."""
mock_prisma = AsyncMock()
mock_prisma.disconnect.side_effect = Exception("engine hung")
# Simulate engine subprocess with a known PID
mock_engine = MagicMock()
mock_engine.process.pid = 12345
mock_prisma._engine = mock_engine
wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False)
# Configure the mock Prisma constructor
mock_new_prisma = AsyncMock()
mock_prisma_binary.Prisma.return_value = mock_new_prisma
with (
patch("os.kill") as mock_kill,
patch("time.sleep"),
):
await wrapper.recreate_prisma_client("postgresql://new")
# Verify old engine was killed
mock_kill.assert_any_call(12345, signal.SIGTERM)
# Verify new client was created and connected
mock_new_prisma.connect.assert_awaited_once()
@pytest.mark.asyncio
async def test_recreate_prisma_client_skips_kill_on_successful_disconnect(
mock_prisma_binary,
):
"""When disconnect() succeeds, no kill should be attempted."""
mock_prisma = AsyncMock()
mock_prisma.disconnect.return_value = None
wrapper = PrismaWrapper(original_prisma=mock_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:
await wrapper.recreate_prisma_client("postgresql://new")
mock_kill.assert_not_called()
mock_new_prisma.connect.assert_awaited_once()
@pytest.mark.asyncio
async def test_recreate_prisma_client_handles_missing_engine_pid(
mock_prisma_binary,
):
"""When engine PID is unavailable (no _engine attr), kill is skipped gracefully."""
mock_prisma = AsyncMock()
mock_prisma.disconnect.side_effect = Exception("engine hung")
mock_prisma._engine = None # No engine subprocess
wrapper = PrismaWrapper(original_prisma=mock_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,
patch("time.sleep"),
):
await wrapper.recreate_prisma_client("postgresql://new")
mock_kill.assert_not_called() # PID was 0, kill skipped
mock_new_prisma.connect.assert_awaited_once()

View file

@ -1,5 +1,6 @@
import asyncio
import os
import signal
import sys
import time
from unittest.mock import AsyncMock, MagicMock, patch
@ -279,3 +280,37 @@ async def test_db_health_watchdog_start_stop_lifecycle(mock_proxy_logging):
await client.stop_db_health_watchdog_task()
assert client._db_health_watchdog_task is None
assert dummy_task.cancelled() is True
@pytest.mark.asyncio
async def test_lightweight_reconnect_kills_engine_on_disconnect_failure(mock_proxy_logging):
"""Lightweight reconnect must kill the old engine PID when disconnect() fails."""
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
client.db.disconnect = AsyncMock(side_effect=Exception("disconnect failed"))
client.db.connect = AsyncMock(return_value=None)
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
with (
patch.object(client, "_get_engine_pid", return_value=9999),
patch("os.kill") as mock_kill,
patch("time.sleep"),
):
await client._run_reconnect_cycle(timeout_seconds=5.0)
mock_kill.assert_any_call(9999, signal.SIGTERM)
client.db.connect.assert_awaited_once()
client.db.query_raw.assert_awaited_once_with("SELECT 1")
@pytest.mark.asyncio
async def test_lightweight_reconnect_skips_kill_on_successful_disconnect(mock_proxy_logging):
"""Lightweight reconnect must NOT kill when disconnect() succeeds."""
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
client.db.disconnect = AsyncMock(return_value=None)
client.db.connect = AsyncMock(return_value=None)
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
with patch("os.kill") as mock_kill:
await client._run_reconnect_cycle(timeout_seconds=5.0)
mock_kill.assert_not_called()