diff --git a/litellm/proxy/db/pgbouncer.py b/litellm/proxy/db/pgbouncer.py index f2f8053cf80..668aca0bf6e 100644 --- a/litellm/proxy/db/pgbouncer.py +++ b/litellm/proxy/db/pgbouncer.py @@ -17,6 +17,11 @@ how many workers run. Migrations and the schema diff run in the supervisor before the pooler is started, so they always go straight to Postgres. ``DATABASE_URL_READ_REPLICA`` is left untouched. + +The pooler holds the database password from startup, so it cannot be combined +with ``IAM_TOKEN_DB_AUTH`` or ``AZURE_POSTGRESQL_AUTH``: those rotate the +password inside every worker on their own schedule, and PgBouncer would keep +authenticating upstream with the expired token. """ from __future__ import annotations @@ -41,6 +46,7 @@ from pydantic import Field from pydantic_settings import BaseSettings, SettingsConfigDict from litellm._logging import verbose_proxy_logger +from litellm.proxy.db.token_auth import AZURE_POSTGRESQL_AUTH_ENV_VAR, IAM_TOKEN_DB_AUTH_ENV_VAR PGBOUNCER_ENV_PREFIX: Final = "LITELLM_PGBOUNCER_" PGBOUNCER_LISTEN_ADDR: Final = "127.0.0.1" @@ -50,6 +56,11 @@ PGBOUNCER_RESTART_DELAY_SECONDS: Final = 1.0 PGBOUNCER_READY_TIMEOUT_SECONDS: Final = 15.0 PGBOUNCER_STOP_GRACE_SECONDS: Final = 10.0 PGBOUNCER_UNPRIVILEGED_USER: Final = "nobody" +PGBOUNCER_TOKEN_AUTH_CONFLICT: Final = ( + f"the in-container pgbouncer cannot be combined with {IAM_TOKEN_DB_AUTH_ENV_VAR} or " + f"{AZURE_POSTGRESQL_AUTH_ENV_VAR}: each worker rotates the database password on its own schedule and the pooler " + "would keep using the expired token upstream. Disable the pooler or use a static database password" +) # Prisma's client-side TLS params describe the hop to Postgres, which becomes # PgBouncer's server side. They move into ``server_tls_*`` and must not stay on @@ -251,8 +262,10 @@ class PgBouncerProcess: """Runs ``argv`` as a foreground child and restarts it whenever it exits on its own. Prisma reconnects by itself after a failed query, so a PgBouncer crash - costs the requests in flight and nothing else once the replacement is - listening again. + costs the requests in flight plus one failed query per idle pooled + connection the crash severed, and nothing else once the replacement is + listening again. A replacement that cannot be spawned or exits again is + retried every ``restart_delay_seconds`` until ``stop`` is called. """ def __init__( @@ -319,10 +332,21 @@ class PgBouncerProcess: status, self.restart_delay_seconds, ) + self._restart_after_delay() + + def _restart_after_delay(self) -> None: time.sleep(self.restart_delay_seconds) if self._stopping.is_set(): return - self._watch(self._spawn()) + try: + self._watch(self._spawn()) + except OSError as spawn_error: + verbose_proxy_logger.error( + "In-container pgbouncer could not be restarted (%s); retrying in %.1fs.", + spawn_error, + self.restart_delay_seconds, + ) + threading.Thread(target=self._restart_after_delay, daemon=True, name="litellm-pgbouncer-supervisor").start() def stop(self) -> None: self._stopping.set() @@ -338,13 +362,17 @@ class PgBouncerProcess: process.wait() -def start_in_container_pgbouncer(settings: PgBouncerSettings, upstream_url: str) -> str | PgBouncerError: +def start_in_container_pgbouncer( + settings: PgBouncerSettings, upstream_url: str, token_auth_enabled: bool = False +) -> str | PgBouncerError: """Start the pooler for ``upstream_url`` and return the loopback URL the workers must use. The pooler lives as long as this process: it is stopped from ``atexit`` once the worker manager has returned. PgBouncer refuses to run as root, so a root proxy (the default image) has it drop to ``nobody``. """ + if token_auth_enabled: + return PgBouncerError(PGBOUNCER_TOKEN_AUTH_CONFLICT) runtime_dir: Final = Path(tempfile.mkdtemp(prefix="litellm-pgbouncer-")) atexit.register(shutil.rmtree, runtime_dir, ignore_errors=True) run_as_user: Final = PGBOUNCER_UNPRIVILEGED_USER if os.geteuid() == 0 else None diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 0a07fd95d59..1302aa80b8f 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1366,7 +1366,9 @@ def run_server( pgbouncer_settings: Final = PgBouncerSettings() upstream_database_url: Final = os.getenv("DATABASE_URL") if pgbouncer_settings.enabled and upstream_database_url is not None: - pooled_database_url: Final = start_in_container_pgbouncer(pgbouncer_settings, upstream_database_url) + pooled_database_url: Final = start_in_container_pgbouncer( + pgbouncer_settings, upstream_database_url, token_auth_enabled=wants_rds_iam or wants_azure_entra + ) if isinstance(pooled_database_url, PgBouncerError): print( f"\033[1;31mLiteLLM Proxy: LITELLM_PGBOUNCER_ENABLED is set but the in-container pgbouncer " diff --git a/tests/test_litellm/proxy/db/test_pgbouncer.py b/tests/test_litellm/proxy/db/test_pgbouncer.py index 0b3b438ef09..77e9f3c8827 100644 --- a/tests/test_litellm/proxy/db/test_pgbouncer.py +++ b/tests/test_litellm/proxy/db/test_pgbouncer.py @@ -225,6 +225,25 @@ class TestPgBouncerProcess: pooler.stop() assert _wait_until(lambda: not _listening(port)) + def test_a_failed_restart_is_retried_until_the_pooler_is_back( + self, tmp_path: Path, caplog: pytest.LogCaptureFixture + ): + port: Final = _free_port() + script: Final = _fake_pooler(tmp_path, port) + pooler: Final = PgBouncerProcess(argv=(str(script),), port=port, restart_delay_seconds=0.1) + assert pooler.start() is None + first_pid: Final = pooler.pid + assert first_pid is not None + hidden: Final = script.rename(tmp_path / "hidden") + with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name): + os.kill(first_pid, signal.SIGKILL) + assert _wait_until(lambda: any("could not be restarted" in record.message for record in caplog.records)) + assert not _listening(port) + hidden.rename(script) + assert _wait_until(lambda: pooler.pid not in (None, first_pid) and _listening(port)) + pooler.stop() + assert _wait_until(lambda: not _listening(port)) + def test_a_stopped_pooler_is_not_restarted(self, tmp_path: Path, caplog: pytest.LogCaptureFixture): port: Final = _free_port() pooler: Final = PgBouncerProcess( @@ -276,6 +295,17 @@ class TestStartInContainerPgBouncer: assert isinstance(outcome, PgBouncerError) assert not _listening(port) + def test_token_auth_is_refused_without_starting_anything(self, tmp_path: Path): + port: Final = _free_port() + settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(_fake_pooler(tmp_path, port))) + outcome: Final = start_in_container_pgbouncer( + settings, "postgresql://app:pw@db/litellm", token_auth_enabled=True + ) + assert isinstance(outcome, PgBouncerError) + assert "IAM_TOKEN_DB_AUTH" in outcome.reason + assert "AZURE_POSTGRESQL_AUTH" in outcome.reason + assert not _listening(port) + class TestPgBouncerSettings: def test_reads_the_litellm_pgbouncer_env_vars(self, monkeypatch: pytest.MonkeyPatch):