From 6aa84272ec11efc9e1713d0b24545d2987680c3f Mon Sep 17 00:00:00 2001 From: harish-berri Date: Mon, 18 May 2026 23:19:10 +0000 Subject: [PATCH] feat(proxy): enhance PrismaClient with signal handling and exit logging - Added signal handling utilities to format signal names and engine wait statuses. - Implemented logging for Prisma engine exit reasons, capturing detailed exit statuses and signals. - Updated methods to pass wait status information to the event loop for improved diagnostics. - Enhanced tests to validate new signal handling and exit status formatting functionalities. --- litellm/proxy/utils.py | 102 +++++++++++++++--- .../proxy/test_prisma_engine_watchdog.py | 35 +++++- 2 files changed, 119 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 32c887f17b2..74ec2a4438a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4,6 +4,7 @@ import hashlib import inspect import json import os +import signal import smtplib import sys import threading @@ -4202,6 +4203,65 @@ class PrismaClient: except (PermissionError, OSError): return True + @staticmethod + def _format_signal_name(signal_number: int) -> str: + try: + return signal.Signals(signal_number).name + except ValueError: + return f"UNKNOWN_SIGNAL_{signal_number}" + + @staticmethod + def _format_engine_wait_status(wait_status: int) -> str: + if os.WIFEXITED(wait_status): + return f"exit_code={os.WEXITSTATUS(wait_status)}" + elif os.WIFSIGNALED(wait_status): + signal_number = os.WTERMSIG(wait_status) + signal_name = PrismaClient._format_signal_name(signal_number) + core_dumped = ( + os.WCOREDUMP(wait_status) if hasattr(os, "WCOREDUMP") else False + ) + return ( + f"signal={signal_name} signal_number={signal_number} " + f"core_dumped={core_dumped}" + ) + elif os.WIFSTOPPED(wait_status): + signal_number = os.WSTOPSIG(wait_status) + signal_name = PrismaClient._format_signal_name(signal_number) + return f"stopped_by_signal={signal_name} signal_number={signal_number}" + elif hasattr(os, "WIFCONTINUED") and os.WIFCONTINUED(wait_status): + return "continued=True" + else: + return f"raw_wait_status={wait_status}" + + def _format_prisma_engine_exit_reason( + self, + *, + detection_method: str, + wait_status: Optional[int], + ) -> str: + if wait_status is None: + return f"detection_method={detection_method} exit_status=unavailable" + return ( + f"detection_method={detection_method} " + f"{self._format_engine_wait_status(wait_status)}" + ) + + def _log_prisma_engine_exit_reason( + self, + *, + pid: int, + detection_method: str, + wait_status: Optional[int], + ) -> None: + verbose_proxy_logger.error( + "prisma-query-engine PID %s exited; %s; triggering reconnect.", + pid, + self._format_prisma_engine_exit_reason( + detection_method=detection_method, + wait_status=wait_status, + ), + ) + @staticmethod def _reap_all_zombies() -> set: """Reap ALL zombie child processes via waitpid(-1, WNOHANG). @@ -4240,7 +4300,7 @@ class PrismaClient: if sys.platform == "win32": return False try: - probe_pid, _ = os.waitpid(pid, os.WNOHANG) + probe_pid, wait_status = os.waitpid(pid, os.WNOHANG) except ChildProcessError: verbose_proxy_logger.debug( "PID %s is not a child process; skipping waitpid watch.", @@ -4249,8 +4309,13 @@ class PrismaClient: return False if probe_pid == pid: + self._log_prisma_engine_exit_reason( + pid=pid, + detection_method="waitpid watch start", + wait_status=wait_status, + ) verbose_proxy_logger.warning( - "prisma-query-engine PID %s already dead at watch start.", + "prisma-query-engine PID %s already dead at watch start; triggering reconnect.", pid, ) self._engine_confirmed_dead = True @@ -4286,26 +4351,32 @@ class PrismaClient: in its SIGCHLD handler. In that case our waitpid raises ChildProcessError. we still notify the event loop because the engine is dead either way. """ + wait_status: Optional[int] = None try: - os.waitpid(pid, 0) + _, wait_status = os.waitpid(pid, 0) except ChildProcessError: pass except OSError: pass try: - loop.call_soon_threadsafe(self._on_engine_death_from_thread, pid) + loop.call_soon_threadsafe( + self._on_engine_death_from_thread, pid, wait_status + ) except RuntimeError: pass - def _on_engine_death_from_thread(self, dead_pid: int) -> None: + def _on_engine_death_from_thread( + self, dead_pid: int, wait_status: Optional[int] = None + ) -> None: """Called on the event loop thread when the waitpid thread detects engine death.""" if self._engine_confirmed_dead: return if dead_pid != self._engine_pid: return - verbose_proxy_logger.error( - "prisma-query-engine PID %s exited (waitpid thread); triggering reconnect.", - dead_pid, + self._log_prisma_engine_exit_reason( + pid=dead_pid, + detection_method="waitpid thread", + wait_status=wait_status, ) self._engine_confirmed_dead = True self._reap_all_zombies() @@ -4357,9 +4428,10 @@ class PrismaClient: self._engine_pidfd = -1 return dead_pid = self._engine_pid - verbose_proxy_logger.error( - "prisma-query-engine PID %s exited (pidfd event); triggering reconnect.", - dead_pid, + self._log_prisma_engine_exit_reason( + pid=dead_pid, + detection_method="pidfd event", + wait_status=None, ) self._engine_confirmed_dead = True self._reap_all_zombies() @@ -4380,9 +4452,11 @@ class PrismaClient: try: os.kill(self._engine_pid, 0) except ProcessLookupError: - verbose_proxy_logger.error( - "prisma-query-engine PID %s gone; triggering reconnect.", - self._engine_pid, + dead_pid = self._engine_pid + self._log_prisma_engine_exit_reason( + pid=dead_pid, + detection_method="os.kill polling", + wait_status=None, ) self._engine_confirmed_dead = True self._reap_all_zombies() diff --git a/tests/litellm/proxy/test_prisma_engine_watchdog.py b/tests/litellm/proxy/test_prisma_engine_watchdog.py index 0d241f75749..8779c8e61cd 100644 --- a/tests/litellm/proxy/test_prisma_engine_watchdog.py +++ b/tests/litellm/proxy/test_prisma_engine_watchdog.py @@ -16,8 +16,7 @@ Covers: import asyncio import os -import threading -import time +import signal from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -91,6 +90,19 @@ def test_is_engine_alive_returns_true_for_running_process(engine_client): assert engine_client._is_engine_alive() is True +def test_format_engine_wait_status_for_exit_code(engine_client): + wait_status = 7 << 8 + + assert engine_client._format_engine_wait_status(wait_status) == "exit_code=7" + + +def test_format_engine_wait_status_for_signal(engine_client): + assert ( + engine_client._format_engine_wait_status(signal.SIGTERM.value) + == "signal=SIGTERM signal_number=15 core_dumped=False" + ) + + # --------------------------------------------------------------------------- # _poll_engine_proc — calls attempt_db_reconnect on death # --------------------------------------------------------------------------- @@ -399,7 +411,7 @@ def test_try_waitpid_watch_starts_thread_for_child(engine_client): with ( patch("os.waitpid", return_value=(0, 0)), patch("asyncio.get_running_loop", return_value=mock_loop), - patch("threading.Thread", return_value=mock_thread) as mock_thread_cls, + patch("threading.Thread", return_value=mock_thread), ): result = engine_client._try_waitpid_watch(1234) assert result is True @@ -407,6 +419,21 @@ def test_try_waitpid_watch_starts_thread_for_child(engine_client): assert engine_client._engine_wait_thread is mock_thread +def test_waitpid_thread_passes_exit_status_to_event_loop(engine_client): + """waitpid thread forwards the raw wait status so logs can include the exit reason.""" + mock_loop = MagicMock() + wait_status = 9 << 8 + + with patch("os.waitpid", return_value=(1234, wait_status)): + engine_client._waitpid_thread_func(1234, mock_loop) + + mock_loop.call_soon_threadsafe.assert_called_once_with( + engine_client._on_engine_death_from_thread, + 1234, + wait_status, + ) + + @pytest.mark.asyncio async def test_try_waitpid_watch_handles_already_dead_engine(engine_client) -> None: """_try_waitpid_watch detects engine already dead at watch start.""" @@ -451,7 +478,7 @@ async def test_on_engine_death_from_thread_triggers_reconnect(engine_client) -> return MagicMock() with patch("asyncio.create_task", side_effect=capture_task): - engine_client._on_engine_death_from_thread(1234) + engine_client._on_engine_death_from_thread(1234, 7 << 8) assert len(created_coros) == 1 await created_coros[0]