mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Auto-enable worker recycling and scope Prisma waitpid reaping
Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
78c2a39f1c
commit
691e81c9b4
4 changed files with 173 additions and 25 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue