diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index b2a5b5c63ee..adcc881ca24 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -56,6 +56,10 @@ def append_query_params(url: Optional[str], params: dict) -> str: class ProxyInitializationHelpers: + # Auto-recycle workers in multi-worker mode to cap process RSS growth from + # allocator fragmentation / long-lived object accumulation over time. + DEFAULT_MULTI_WORKER_MAX_REQUESTS_BEFORE_RESTART = 1000 + @staticmethod def _echo_litellm_version(): pkg_version = importlib.metadata.version("litellm") # type: ignore @@ -546,7 +550,7 @@ class ProxyInitializationHelpers: "--max_requests_before_restart", default=None, type=int, - help="Restart worker after this many requests (uvicorn: limit_max_requests, gunicorn: max_requests)", + help="Restart worker after this many requests (uvicorn: limit_max_requests, gunicorn: max_requests). Auto-enabled in multi-worker mode.", envvar="MAX_REQUESTS_BEFORE_RESTART", ) def run_server( # noqa: PLR0915 @@ -881,6 +885,25 @@ def run_server( # noqa: PLR0915 ) return + effective_max_requests_before_restart = max_requests_before_restart + if effective_max_requests_before_restart is None and num_workers > 1: + effective_max_requests_before_restart = ( + ProxyInitializationHelpers.DEFAULT_MULTI_WORKER_MAX_REQUESTS_BEFORE_RESTART + ) + print( # noqa + "LiteLLM Proxy: Auto-enabling worker recycling in multi-worker mode: " + f"max_requests_before_restart={effective_max_requests_before_restart}. " + "Override with --max_requests_before_restart / MAX_REQUESTS_BEFORE_RESTART." + ) + elif ( + effective_max_requests_before_restart is not None + and effective_max_requests_before_restart <= 0 + ): + print( # noqa + "LiteLLM Proxy: max_requests_before_restart must be > 0; disabling worker recycling." + ) + effective_max_requests_before_restart = None + uvicorn_args = ProxyInitializationHelpers._get_default_unvicorn_init_args( host=host, port=port, @@ -888,8 +911,8 @@ def run_server( # noqa: PLR0915 keepalive_timeout=keepalive_timeout, ) # Optional: recycle uvicorn workers after N requests - if max_requests_before_restart is not None: - uvicorn_args["limit_max_requests"] = max_requests_before_restart + if effective_max_requests_before_restart is not None: + uvicorn_args["limit_max_requests"] = effective_max_requests_before_restart if run_gunicorn is False and run_hypercorn is False: if ssl_certfile_path is not None and ssl_keyfile_path is not None: print( # noqa @@ -915,7 +938,7 @@ def run_server( # noqa: PLR0915 ssl_certfile_path=ssl_certfile_path, ssl_keyfile_path=ssl_keyfile_path, keepalive_timeout=keepalive_timeout, - max_requests_before_restart=max_requests_before_restart, + max_requests_before_restart=effective_max_requests_before_restart, ) elif run_hypercorn is True: ProxyInitializationHelpers._init_hypercorn_server( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f6613b5548f..542d7a776a7 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3567,23 +3567,28 @@ class PrismaClient: except (PermissionError, OSError): return True - @staticmethod - def _reap_all_zombies() -> set: - """Reap ALL zombie child processes via waitpid(-1, WNOHANG). + def _reap_all_zombies(self, target_pid: Optional[int] = None) -> set: + """Reap a tracked engine child process with waitpid(pid, WNOHANG). - Returns a set of reaped PIDs. As PID 1 in Docker (or any - process that spawns children), we must reap ALL terminated - children to prevent zombie accumulation. + Historically this used waitpid(-1, WNOHANG), which can reap unrelated + child processes owned by the current worker process. In multi-worker + server setups, that may interfere with other child lifecycle handlers. + Restricting reaping to the Prisma engine PID avoids that cross-process + interference while still cleaning up the expected zombie child. """ + pid_to_reap = target_pid if target_pid is not None else self._engine_pid reaped: set = set() - while True: - try: - pid, _ = os.waitpid(-1, os.WNOHANG) - if pid == 0: - break + if pid_to_reap <= 0: + return reaped + + try: + pid, _ = os.waitpid(pid_to_reap, os.WNOHANG) + if pid == pid_to_reap: reaped.add(pid) - except ChildProcessError: - break + except ChildProcessError: + # Child already reaped by another handler. + pass + return reaped def _try_waitpid_watch(self, pid: int) -> bool: @@ -3609,7 +3614,7 @@ class PrismaClient: "prisma-query-engine PID %s already dead at watch start.", pid, ) self._engine_confirmed_dead = True - self._reap_all_zombies() + self._reap_all_zombies(pid) self._cleanup_engine_watcher() asyncio.create_task( self.attempt_db_reconnect( @@ -3663,7 +3668,7 @@ class PrismaClient: dead_pid, ) self._engine_confirmed_dead = True - self._reap_all_zombies() + self._reap_all_zombies(dead_pid) self._cleanup_engine_watcher() asyncio.create_task( self.attempt_db_reconnect( @@ -3717,7 +3722,7 @@ class PrismaClient: dead_pid, ) self._engine_confirmed_dead = True - self._reap_all_zombies() + self._reap_all_zombies(dead_pid) self._cleanup_engine_watcher() asyncio.create_task( self.attempt_db_reconnect( @@ -3740,7 +3745,7 @@ class PrismaClient: self._engine_pid, ) self._engine_confirmed_dead = True - self._reap_all_zombies() + self._reap_all_zombies(self._engine_pid) self._cleanup_engine_watcher() await self.attempt_db_reconnect( reason="engine_process_death", @@ -3840,7 +3845,7 @@ class PrismaClient: "prisma-query-engine PID %s is dead; reconnecting.", dead_pid, ) - self._reap_all_zombies() + self._reap_all_zombies(dead_pid) self._cleanup_engine_watcher() self._engine_confirmed_dead = False diff --git a/tests/litellm/proxy/test_prisma_engine_watchdog.py b/tests/litellm/proxy/test_prisma_engine_watchdog.py index 011b8002db2..6d0374febcc 100644 --- a/tests/litellm/proxy/test_prisma_engine_watchdog.py +++ b/tests/litellm/proxy/test_prisma_engine_watchdog.py @@ -385,9 +385,12 @@ async def test_try_waitpid_watch_handles_already_dead_engine(engine_client) -> N waitpid_calls = iter([(1234, 0)]) def mock_waitpid(pid, flags): - if pid == -1: - raise ChildProcessError - return next(waitpid_calls) + if pid == 1234: + try: + return next(waitpid_calls) + except StopIteration: + raise ChildProcessError + raise ChildProcessError with ( patch("os.waitpid", side_effect=mock_waitpid), @@ -401,6 +404,26 @@ async def test_try_waitpid_watch_handles_already_dead_engine(engine_client) -> N created_coros[0].close() +def test_reap_all_zombies_reaps_only_target_engine_pid(engine_client) -> None: + """_reap_all_zombies should only wait on the tracked engine PID.""" + engine_client._engine_pid = 4321 + with patch("os.waitpid", return_value=(4321, 0)) as mock_waitpid: + reaped = engine_client._reap_all_zombies() + + assert reaped == {4321} + mock_waitpid.assert_called_once_with(4321, os.WNOHANG) + + +def test_reap_all_zombies_with_explicit_target_pid(engine_client) -> None: + """Explicit target PID should be used instead of _engine_pid.""" + engine_client._engine_pid = 1111 + with patch("os.waitpid", return_value=(2222, 0)) as mock_waitpid: + reaped = engine_client._reap_all_zombies(target_pid=2222) + + assert reaped == {2222} + mock_waitpid.assert_called_once_with(2222, os.WNOHANG) + + @pytest.mark.asyncio async def test_on_engine_death_from_thread_triggers_reconnect(engine_client) -> None: """waitpid thread callback schedules attempt_db_reconnect.""" diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 039d0f21f5c..85913ec8f7e 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -428,6 +428,103 @@ class TestProxyInitializationHelpers: call_args = mock_uvicorn_run.call_args assert call_args[1]["limit_max_requests"] == 123 + @patch("uvicorn.run") + @patch("builtins.print") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + def test_max_requests_before_restart_auto_enabled_for_multi_worker( + self, mock_setup_db, mock_print, mock_uvicorn_run + ): + """When not explicitly set, multi-worker mode auto-enables worker recycling.""" + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import ProxyInitializationHelpers, run_server + + runner = CliRunner() + + mock_app = MagicMock() + mock_proxy_config = MagicMock() + mock_key_mgmt = MagicMock() + mock_save_worker_config = MagicMock() + + clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")} + with patch.dict( + os.environ, clean_env, clear=True, + ), patch.dict( + "sys.modules", + { + "proxy_server": MagicMock( + app=mock_app, + ProxyConfig=mock_proxy_config, + KeyManagementSettings=mock_key_mgmt, + save_worker_config=mock_save_worker_config, + ) + }, + ), patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args: + mock_get_args.return_value = { + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, + } + + result = runner.invoke(run_server, ["--local", "--num_workers", "4"]) + + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + mock_uvicorn_run.assert_called_once() + call_args = mock_uvicorn_run.call_args + assert ( + call_args[1]["limit_max_requests"] + == ProxyInitializationHelpers.DEFAULT_MULTI_WORKER_MAX_REQUESTS_BEFORE_RESTART + ) + + @patch("uvicorn.run") + @patch("builtins.print") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + def test_max_requests_before_restart_not_set_for_single_worker_default( + self, mock_setup_db, mock_print, mock_uvicorn_run + ): + """Single-worker mode should not auto-set limit_max_requests.""" + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + runner = CliRunner() + + mock_app = MagicMock() + mock_proxy_config = MagicMock() + mock_key_mgmt = MagicMock() + mock_save_worker_config = MagicMock() + + clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")} + with patch.dict( + os.environ, clean_env, clear=True, + ), patch.dict( + "sys.modules", + { + "proxy_server": MagicMock( + app=mock_app, + ProxyConfig=mock_proxy_config, + KeyManagementSettings=mock_key_mgmt, + save_worker_config=mock_save_worker_config, + ) + }, + ), patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args: + mock_get_args.return_value = { + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, + } + + result = runner.invoke(run_server, ["--local", "--num_workers", "1"]) + + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + mock_uvicorn_run.assert_called_once() + call_args = mock_uvicorn_run.call_args + assert "limit_max_requests" not in call_args[1] + @patch.dict(os.environ, {}, clear=True) def test_construct_database_url_from_env_vars(self): """Test the construct_database_url_from_env_vars function with various scenarios"""