diff --git a/gateway/launch.py b/gateway/launch.py index 5f1dafe4f4e..fd241c6c521 100644 --- a/gateway/launch.py +++ b/gateway/launch.py @@ -13,7 +13,7 @@ Run with: import os import sys -from collections.abc import MutableMapping, Sequence +from collections.abc import Mapping, MutableMapping, Sequence from typing import Final from uvicorn.main import main as uvicorn_main @@ -22,6 +22,15 @@ from litellm.proxy.db.db_url_settings import DatabaseURLSettings from litellm.proxy.db.pgbouncer import PgBouncerError, PgBouncerSettings, start_in_container_pgbouncer GATEWAY_APP: Final = "gateway.main:app" +KEEPALIVE_FLAG: Final = "--timeout-keep-alive" + + +def uvicorn_argv(argv: Sequence[str], environ: Mapping[str, str]) -> tuple[str, ...]: + """Honor ``KEEPALIVE_TIMEOUT`` like ``proxy_cli.py`` does, unless the flag was passed explicitly.""" + keepalive: Final = environ.get("KEEPALIVE_TIMEOUT") + if keepalive is None or any(arg == KEEPALIVE_FLAG or arg.startswith(f"{KEEPALIVE_FLAG}=") for arg in argv): + return (GATEWAY_APP, *argv) + return (GATEWAY_APP, *argv, KEEPALIVE_FLAG, keepalive) def pool_database_url( @@ -55,7 +64,7 @@ def main(argv: Sequence[str]) -> None: failed: Final = pool_database_url(settings, PgBouncerSettings(), os.environ) if failed is not None: sys.exit(f"LiteLLM gateway: in-container pgbouncer could not start: {failed.reason}") - uvicorn_main((GATEWAY_APP, *argv), prog_name="uvicorn") + uvicorn_main(uvicorn_argv(argv, os.environ), prog_name="uvicorn") if __name__ == "__main__": diff --git a/tests/test_gateway/test_launch.py b/tests/test_gateway/test_launch.py index 9882e853db1..a13e7dfb554 100644 --- a/tests/test_gateway/test_launch.py +++ b/tests/test_gateway/test_launch.py @@ -7,8 +7,9 @@ from pathlib import Path from typing import Final, cast import pytest +from uvicorn.main import main as uvicorn_main -from gateway.launch import pool_database_url +from gateway.launch import GATEWAY_APP, pool_database_url, uvicorn_argv from litellm.proxy.db.db_url_settings import DatabaseURLSettings from litellm.proxy.db.pgbouncer import PgBouncerError, PgBouncerSettings @@ -70,6 +71,25 @@ def password_env(monkeypatch: pytest.MonkeyPatch) -> dict[str, str]: return dict(DB_ENV) +def _uvicorn_params(argv: tuple[str, ...]) -> dict[str, object]: + return uvicorn_main.make_context("uvicorn", list(argv)).params + + +class TestUvicornArgv: + def test_keepalive_env_reaches_uvicorn(self): + params: Final = _uvicorn_params(uvicorn_argv(("--workers", "4"), {"KEEPALIVE_TIMEOUT": "75"})) + assert params["app"] == GATEWAY_APP + assert params["workers"] == 4 + assert params["timeout_keep_alive"] == 75 + + def test_unset_env_keeps_the_uvicorn_default(self): + assert _uvicorn_params(uvicorn_argv(("--workers", "4"), {}))["timeout_keep_alive"] == 5 + + def test_an_explicit_flag_wins_over_the_env(self): + argv: Final = uvicorn_argv(("--timeout-keep-alive", "30"), {"KEEPALIVE_TIMEOUT": "75"}) + assert _uvicorn_params(argv)["timeout_keep_alive"] == 30 + + class TestPoolDatabaseUrl: def test_disabled_pooler_leaves_the_assembled_url_alone(self, password_env: dict[str, str]): settings: Final = DatabaseURLSettings.from_env()