diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index c9c0cfe8f68..fa00d117c81 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index df527d08af8..e067b4d3175 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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") diff --git a/tests/test_litellm/proxy/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py index 83f07253fc8..9c62c6ffd50 100644 --- a/tests/test_litellm/proxy/db/test_prisma_client.py +++ b/tests/test_litellm/proxy/db/test_prisma_client.py @@ -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') \ No newline at end of file + 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() \ No newline at end of file diff --git a/tests/test_litellm/proxy/db/test_prisma_self_heal.py b/tests/test_litellm/proxy/db/test_prisma_self_heal.py index 03ad95026d8..dbe1f2113b1 100644 --- a/tests/test_litellm/proxy/db/test_prisma_self_heal.py +++ b/tests/test_litellm/proxy/db/test_prisma_self_heal.py @@ -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()