mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(proxy): kill orphaned prisma engine subprocess on failed disconnect
This commit is contained in:
parent
81dadb698a
commit
1f04fa2461
4 changed files with 170 additions and 4 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue