feat(proxy): share database connections across workers with an in-container pgbouncer

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-04 00:03:10 +00:00
parent cf3af0f486
commit dafa02daf5
9 changed files with 763 additions and 3 deletions

View file

@ -110,7 +110,7 @@ USER root
RUN echo "https://packages.wolfi.dev/os" >> /etc/apk/repositories
# node (without npm) is required by the prisma CLI at runtime
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile pgbouncer
WORKDIR /app
ENV PATH="/app/.venv/bin:${PATH}" \

View file

@ -101,7 +101,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
# node (without npm) is required by the prisma CLI at runtime
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile pgbouncer
WORKDIR /app
ENV PATH="/app/.venv/bin:${PATH}" \

View file

@ -128,7 +128,7 @@ RUN for i in 1 2 3; do \
apk upgrade --no-cache && break || sleep 5; \
done && \
for i in 1 2 3; do \
apk add --no-cache python-3.13 bash openssl tzdata libsndfile nodejs && break || sleep 5; \
apk add --no-cache python-3.13 bash openssl tzdata libsndfile nodejs pgbouncer && break || sleep 5; \
done
# Copy only what runtime needs. The application is installed inside the venv;

View file

@ -117,6 +117,14 @@ spec:
- name: DATABASE_URL_READ_REPLICA
value: {{ .Values.db.readReplicaUrl | quote }}
{{- end }}
{{- if .Values.db.connectionPool.enabled }}
- name: LITELLM_PGBOUNCER_ENABLED
value: "true"
- name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS
value: {{ .Values.db.connectionPool.maxDbConnections | quote }}
- name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN
value: {{ .Values.db.connectionPool.maxClientConn | quote }}
{{- end }}
- name: PROXY_MASTER_KEY
valueFrom:
secretKeyRef:

View file

@ -0,0 +1,61 @@
suite: test in-container connection pool
templates:
- deployment.yaml
- configmap-litellm.yaml
tests:
- it: should not emit pgbouncer env vars by default
template: deployment.yaml
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_PGBOUNCER_ENABLED
value: "true"
- notContains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS
value: "20"
- it: should enable the pool with the default sizing when connectionPool.enabled is set
template: deployment.yaml
set:
db.connectionPool.enabled: true
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_PGBOUNCER_ENABLED
value: "true"
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS
value: "20"
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN
value: "1000"
- it: should pass custom sizing through as strings next to the worker count
template: deployment.yaml
set:
numWorkers: 4
db.connectionPool.enabled: true
db.connectionPool.maxDbConnections: 8
db.connectionPool.maxClientConn: 400
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS
value: "8"
- contains:
path: spec.template.spec.containers[0].env
content:
name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN
value: "400"
- contains:
path: spec.template.spec.containers[0].args
content: "4"

View file

@ -309,6 +309,20 @@ db:
# only (e.g. when IAM_TOKEN_DB_AUTH supplies the token at runtime).
readReplicaUrl: ""
# In-container connection pool (PgBouncer, transaction mode) shared by every
# worker in the pod. Without it each --num_workers worker opens its own
# connection_limit connections to Postgres, so a pod's footprint against the
# database's connection ceiling is workers x connection_limit and grows with
# every replica. With it, the pod holds at most maxDbConnections upstream
# connections no matter how many workers run; the workers connect to the pool
# over loopback, with no extra network hop. Migrations still go straight to
# Postgres. Starting profile for numWorkers: 4 is maxDbConnections: 20, so
# a database with a 5000-connection ceiling fits roughly 200 replicas.
connectionPool:
enabled: false
maxDbConnections: 20
maxClientConn: 1000
# Use the Stackgres Helm chart to deploy an instance of a Stackgres cluster.
# The Stackgres Operator must already be installed within the target
# Kubernetes cluster.

View file

@ -0,0 +1,373 @@
"""In-container PgBouncer shared by every proxy worker.
Each uvicorn worker owns a Prisma query engine with its own pool of
``connection_limit`` server connections, so the connections a pod holds open
against Postgres scale as ``workers * connection_limit`` and a database with a
fixed connection ceiling runs out of room as pods and workers are added.
When ``LITELLM_PGBOUNCER_ENABLED`` is set, the supervisor process starts one
PgBouncer next to the workers (no extra network hop: it listens on loopback
inside the pod) in transaction pooling mode, points ``DATABASE_URL`` at it
with ``pgbouncer=true`` so Prisma stops using server-side prepared statements,
and keeps it running for the life of the proxy. Every worker's pool then
becomes cheap client connections to PgBouncer while the upstream connection
count is capped at ``LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS`` per pod, no matter
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.
"""
from __future__ import annotations
import atexit
import os
import shlex
import shutil
import socket
import subprocess
import tempfile
import threading
import time
import urllib.parse
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
from pydantic import Field
from pydantic_settings import BaseSettings, SettingsConfigDict
from litellm._logging import verbose_proxy_logger
PGBOUNCER_ENV_PREFIX: Final = "LITELLM_PGBOUNCER_"
PGBOUNCER_LISTEN_ADDR: Final = "127.0.0.1"
PGBOUNCER_INI_NAME: Final = "pgbouncer.ini"
PGBOUNCER_USERLIST_NAME: Final = "userlist.txt"
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"
# 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
# the loopback URL: the listener speaks plain TCP and Prisma would refuse it
# under ``sslmode=require``.
PRISMA_TLS_PARAM_KEYS: Final[frozenset[str]] = frozenset(
{"sslmode", "sslcert", "sslaccept", "sslidentity", "sslpassword"}
)
POOLED_URL_DROPPED_KEYS: Final[frozenset[str]] = PRISMA_TLS_PARAM_KEYS | frozenset(("options", "pgbouncer"))
PGBOUNCER_SSLMODES: Final[frozenset[str]] = frozenset(
{"disable", "allow", "prefer", "require", "verify-ca", "verify-full"}
)
class PgBouncerSettings(BaseSettings):
"""``LITELLM_PGBOUNCER_*`` env vars, read once in the supervisor."""
model_config = SettingsConfigDict(
env_prefix=PGBOUNCER_ENV_PREFIX, case_sensitive=False, extra="ignore", frozen=True
)
enabled: bool = False
port: int = Field(default=6432, ge=1, le=65535)
max_db_connections: int = Field(default=20, ge=1)
max_client_conn: int = Field(default=1000, ge=1)
binary: str = "pgbouncer"
@dataclass(frozen=True, slots=True)
class PgBouncerPlan:
ini: str
userlist: str
pooled_url: str
@dataclass(frozen=True, slots=True)
class PgBouncerError:
reason: str
def _single_quoted(value: str) -> str:
"""Quote for SQL and for PgBouncer's ``[databases]`` connection string: both double a literal ``'``."""
return "'" + value.replace("'", "''") + "'"
def _userlist_quote(value: str) -> str:
return '"' + value.replace('"', '""') + '"'
def _split_option(tokens: Sequence[str]) -> tuple[str, Sequence[str]] | None:
"""Split the first ``-c name=value`` / ``-cname=value`` / ``--name=value`` off ``tokens``."""
head: Final = tokens[0]
if head == "-c":
return (tokens[1], tokens[2:]) if len(tokens) > 1 else None
if head.startswith(("-c", "--")):
return head[2:], tokens[1:]
return None
def _option_settings(tokens: Sequence[str]) -> tuple[str, ...] | None:
"""The ``name=value`` settings in a libpq ``options`` string, or None if it holds anything else."""
if not tokens:
return ()
split: Final = _split_option(tokens)
if split is None or "=" not in split[0]:
return None
tail: Final = _option_settings(split[1])
return None if tail is None else (split[0], *tail)
def _connect_query(options: str) -> str | PgBouncerError:
"""Turn Prisma's ``options=-c name=value ...`` startup param into ``SET`` statements.
PgBouncer rejects any ``-c`` setting in ``options`` that is not one of the
handful it tracks (``statement_timeout`` and ``lock_timeout`` are not), so
the settings are applied to each new server connection instead. Every
client shares them, which is what the single ``DATABASE_URL`` gave anyway.
"""
settings: Final = _option_settings(tuple(shlex.split(options)))
if settings is None:
return PgBouncerError(f"cannot translate the DATABASE_URL options {options!r} into PgBouncer settings")
return "; ".join(
f"SET {name.strip()} TO {_single_quoted(value.strip())}"
for name, value in (setting.split("=", 1) for setting in settings)
)
def _server_tls_settings(sslmode: str, sslcert: str, sslaccept: str) -> tuple[str, ...] | PgBouncerError:
if sslmode not in PGBOUNCER_SSLMODES:
return PgBouncerError(f"unsupported sslmode {sslmode!r} on DATABASE_URL")
verify: Final = sslmode in ("verify-ca", "verify-full") or (sslmode == "require" and sslaccept == "strict")
if verify and not sslcert:
return PgBouncerError(
"DATABASE_URL asks for a verified TLS connection but names no CA bundle; "
"add sslcert=<ca.pem> (or sslrootcert=) so the in-container PgBouncer can verify Postgres"
)
mode: Final = "verify-full" if verify else sslmode
return (f"server_tls_sslmode = {mode}", *((f"server_tls_ca_file = {sslcert}",) if sslcert else ()))
def plan_pgbouncer(
upstream_url: str,
settings: PgBouncerSettings,
runtime_dir: Path,
run_as_user: str | None,
) -> PgBouncerPlan | PgBouncerError:
"""Render the PgBouncer config for ``upstream_url`` and the loopback URL Prisma uses instead.
Params describing Prisma's own pool (``connection_limit``, ``pool_timeout``,
...) stay on the pooled URL; the TLS params and ``options`` describe the hop
to Postgres and move into the PgBouncer config. ``run_as_user`` is the
unprivileged user PgBouncer drops to when the proxy runs as root, which
PgBouncer itself refuses to do.
"""
parsed: Final = urllib.parse.urlsplit(upstream_url)
params: Final[Mapping[str, str]] = MappingProxyType(
dict(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True))
)
dbname: Final = urllib.parse.unquote(parsed.path.lstrip("/"))
username: Final = urllib.parse.unquote(parsed.username or "")
password: Final = None if parsed.password is None else urllib.parse.unquote(parsed.password)
if not parsed.hostname or not username or password is None or not dbname:
return PgBouncerError(
"DATABASE_URL must carry a host, user, password and database name for the in-container PgBouncer"
)
if "sslidentity" in params:
return PgBouncerError("client certificates (sslidentity) are not supported with the in-container PgBouncer")
tls: Final = _server_tls_settings(
params.get("sslmode", "prefer"), params.get("sslcert", ""), params.get("sslaccept", "")
)
if isinstance(tls, PgBouncerError):
return tls
connect_query: Final = _connect_query(params["options"]) if params.get("options") else ""
if isinstance(connect_query, PgBouncerError):
return connect_query
upstream: Final = " ".join(
(
f"host={_single_quoted(parsed.hostname)}",
f"port={parsed.port or 5432}",
f"dbname={_single_quoted(dbname)}",
f"user={_single_quoted(username)}",
f"password={_single_quoted(password)}",
*((f"connect_query={_single_quoted(connect_query)}",) if connect_query else ()),
)
)
ini: Final = "\n".join(
(
"[databases]",
f"{dbname} = {upstream}",
"",
"[pgbouncer]",
f"listen_addr = {PGBOUNCER_LISTEN_ADDR}",
f"listen_port = {settings.port}",
f"unix_socket_dir = {runtime_dir}",
f"auth_file = {runtime_dir / PGBOUNCER_USERLIST_NAME}",
"auth_type = scram-sha-256",
"pool_mode = transaction",
f"max_client_conn = {settings.max_client_conn}",
f"default_pool_size = {settings.max_db_connections}",
f"max_db_connections = {settings.max_db_connections}",
"ignore_startup_parameters = extra_float_digits",
*tls,
*((f"user = {run_as_user}",) if run_as_user else ()),
"",
)
)
userlist: Final = f"{_userlist_quote(username)} {_userlist_quote(password)}\n"
pooled_query: Final = urllib.parse.urlencode(
(*((key, value) for key, value in params.items() if key not in POOLED_URL_DROPPED_KEYS), ("pgbouncer", "true"))
)
credentials: Final = f"{urllib.parse.quote(username, safe='')}:{urllib.parse.quote(password, safe='')}"
pooled_url: Final = urllib.parse.urlunsplit(
parsed._replace(netloc=f"{credentials}@{PGBOUNCER_LISTEN_ADDR}:{settings.port}", query=pooled_query)
)
return PgBouncerPlan(ini=ini, userlist=userlist, pooled_url=pooled_url)
def write_pgbouncer_files(plan: PgBouncerPlan, runtime_dir: Path, run_as_user: str | None) -> Path:
"""Write the ini and userlist (both hold the password, so mode 0600) and return the ini path.
``run_as_user`` is the user PgBouncer drops to when started as root; it has
to own the files it re-reads on reload and the socket directory.
"""
ini_path: Final = runtime_dir / PGBOUNCER_INI_NAME
userlist_path: Final = runtime_dir / PGBOUNCER_USERLIST_NAME
for path, content in ((userlist_path, plan.userlist), (ini_path, plan.ini)):
path.touch(mode=0o600)
path.write_text(content, encoding="utf-8")
if run_as_user is not None:
runtime_dir.chmod(0o700)
for path in (runtime_dir, ini_path, userlist_path):
shutil.chown(path, user=run_as_user)
return ini_path
def _port_open(port: int) -> bool:
try:
with socket.create_connection((PGBOUNCER_LISTEN_ADDR, port), timeout=0.5):
return True
except OSError:
return False
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.
"""
def __init__(
self,
argv: Sequence[str],
port: int,
restart_delay_seconds: float = PGBOUNCER_RESTART_DELAY_SECONDS,
) -> None:
self.argv: Final = tuple(argv)
self.port: Final = port
self.restart_delay_seconds: Final = restart_delay_seconds
self._stopping: Final = threading.Event()
self._lock: Final = threading.Lock()
self._process: subprocess.Popen[bytes] | None = None
@property
def pid(self) -> int | None:
with self._lock:
return None if self._process is None else self._process.pid
def _spawn(self) -> subprocess.Popen[bytes]:
process: Final = subprocess.Popen(self.argv)
with self._lock:
self._process = process
return process
def _wait_ready(self, process: subprocess.Popen[bytes], timeout_seconds: float) -> PgBouncerError | None:
deadline: Final = time.monotonic() + timeout_seconds
while time.monotonic() < deadline:
if process.poll() is not None:
return PgBouncerError(f"pgbouncer exited with status {process.returncode} during startup")
if _port_open(self.port):
return None
time.sleep(0.1)
return PgBouncerError(
f"pgbouncer did not start listening on {PGBOUNCER_LISTEN_ADDR}:{self.port} within {timeout_seconds:.0f}s"
)
def start(self, ready_timeout_seconds: float = PGBOUNCER_READY_TIMEOUT_SECONDS) -> PgBouncerError | None:
"""Spawn PgBouncer, wait until it accepts connections, then supervise it from a daemon thread."""
try:
process: Final = self._spawn()
except OSError as spawn_error:
return PgBouncerError(f"could not start {self.argv[0]!r}: {spawn_error}")
not_ready: Final = self._wait_ready(process, ready_timeout_seconds)
if not_ready is not None:
self.stop()
return not_ready
self._watch(process)
return None
def _watch(self, process: subprocess.Popen[bytes]) -> None:
threading.Thread(
target=self._supervise, args=(process,), daemon=True, name="litellm-pgbouncer-supervisor"
).start()
def _supervise(self, process: subprocess.Popen[bytes]) -> None:
status: Final = process.wait()
if self._stopping.is_set():
return
verbose_proxy_logger.error(
"In-container pgbouncer (pid %s) exited with status %s; restarting in %.1fs.",
process.pid,
status,
self.restart_delay_seconds,
)
time.sleep(self.restart_delay_seconds)
if self._stopping.is_set():
return
self._watch(self._spawn())
def stop(self) -> None:
self._stopping.set()
with self._lock:
process: Final = self._process
if process is None or process.poll() is not None:
return
process.terminate()
try:
process.wait(timeout=PGBOUNCER_STOP_GRACE_SECONDS)
except subprocess.TimeoutExpired:
process.kill()
process.wait()
def start_in_container_pgbouncer(settings: PgBouncerSettings, upstream_url: str) -> 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``.
"""
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
plan: Final = plan_pgbouncer(upstream_url, settings, runtime_dir, run_as_user)
if isinstance(plan, PgBouncerError):
return plan
ini_path: Final = write_pgbouncer_files(plan, runtime_dir, run_as_user)
pooler: Final = PgBouncerProcess(argv=(settings.binary, str(ini_path)), port=settings.port)
failed: Final = pooler.start()
if failed is not None:
return failed
atexit.register(pooler.stop)
verbose_proxy_logger.info(
"In-container pgbouncer (pid %s) listening on %s:%s; capping this pod at %s upstream database connections.",
pooler.pid,
PGBOUNCER_LISTEN_ADDR,
settings.port,
settings.max_db_connections,
)
return plan.pooled_url

View file

@ -18,6 +18,7 @@ from pydantic import BaseModel, ConfigDict
import litellm
from litellm.constants import DEFAULT_NUM_WORKERS_LITELLM_PROXY
from litellm.proxy.db.pgbouncer import PgBouncerError, PgBouncerSettings, start_in_container_pgbouncer
from litellm.proxy.db.query_engine_reaper import start_query_engine_reaper
if TYPE_CHECKING:
@ -1362,6 +1363,19 @@ def run_server(
print(
f"Unable to connect to DB. DATABASE_URL found in environment, but prisma package not found." # noqa: F541
)
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)
if isinstance(pooled_database_url, PgBouncerError):
print(
f"\033[1;31mLiteLLM Proxy: LITELLM_PGBOUNCER_ENABLED is set but the in-container pgbouncer "
f"could not start: {pooled_database_url.reason}\033[0m",
file=sys.stderr,
flush=True,
)
sys.exit(1)
os.environ["DATABASE_URL"] = pooled_database_url
if port == 4000 and ProxyInitializationHelpers._is_port_in_use(port):
port = random.randint(1024, 49152)

View file

@ -0,0 +1,290 @@
import configparser
import logging
import os
import signal
import socket
import stat
import sys
import textwrap
import time
import urllib.parse
from collections.abc import Callable
from pathlib import Path
from typing import Final
import pytest
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.pgbouncer import (
PgBouncerError,
PgBouncerPlan,
PgBouncerProcess,
PgBouncerSettings,
plan_pgbouncer,
start_in_container_pgbouncer,
write_pgbouncer_files,
)
UPSTREAM: Final = (
"postgresql://app:p%40ss%27w@db.internal:5433/litellm"
"?schema=public&connection_limit=10&pool_timeout=20"
"&sslmode=require&sslaccept=strict&sslcert=/certs/ca.pem"
"&options=-c%20statement_timeout%3D7000%20-c%20lock_timeout%3D3000"
)
SETTINGS: Final = PgBouncerSettings(enabled=True, port=6543, max_db_connections=8, max_client_conn=400)
def _plan(url: str = UPSTREAM, run_as_user: str | None = None) -> PgBouncerPlan:
plan: Final = plan_pgbouncer(url, SETTINGS, Path("/run/pgb"), run_as_user)
assert isinstance(plan, PgBouncerPlan), plan
return plan
def _ini(plan: PgBouncerPlan) -> configparser.ConfigParser:
parser: Final = configparser.ConfigParser(interpolation=None)
parser.read_string(plan.ini)
return parser
def _query(url: str) -> dict[str, str]:
return dict(urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query, keep_blank_values=True))
class TestPlanPgBouncer:
def test_upstream_credentials_and_timeouts_move_into_the_pgbouncer_config(self):
ini: Final = _ini(_plan())
assert ini["databases"]["litellm"] == (
"host='db.internal' port=5433 dbname='litellm' user='app' password='p@ss''w' "
"connect_query='SET statement_timeout TO ''7000''; SET lock_timeout TO ''3000'''"
)
assert _plan().userlist == '"app" "p@ss\'w"\n'
def test_an_upstream_without_a_port_is_reached_on_the_postgres_default(self):
ini: Final = _ini(_plan("postgresql://app:pw@db/litellm"))
assert ini["databases"]["litellm"] == "host='db' port=5432 dbname='litellm' user='app' password='pw'"
def test_pool_is_sized_from_settings_in_transaction_mode(self):
pgb: Final = _ini(_plan())["pgbouncer"]
assert pgb["pool_mode"] == "transaction"
assert pgb["max_db_connections"] == "8"
assert pgb["default_pool_size"] == "8"
assert pgb["max_client_conn"] == "400"
assert pgb["auth_type"] == "scram-sha-256"
assert pgb["listen_addr"] == "127.0.0.1"
assert pgb["listen_port"] == "6543"
assert pgb["auth_file"] == "/run/pgb/userlist.txt"
assert pgb["unix_socket_dir"] == "/run/pgb"
def test_pooled_url_points_prisma_at_loopback_without_prepared_statements(self):
pooled: Final = urllib.parse.urlsplit(_plan().pooled_url)
assert (pooled.hostname, pooled.port, pooled.path) == ("127.0.0.1", 6543, "/litellm")
assert (pooled.username, pooled.password) == ("app", "p%40ss%27w")
assert _query(_plan().pooled_url) == {
"schema": "public",
"connection_limit": "10",
"pool_timeout": "20",
"pgbouncer": "true",
}
def test_verified_tls_becomes_server_side_verify_full_with_the_ca_bundle(self):
pgb: Final = _ini(_plan())["pgbouncer"]
assert pgb["server_tls_sslmode"] == "verify-full"
assert pgb["server_tls_ca_file"] == "/certs/ca.pem"
def test_unverified_require_stays_require_without_a_ca_file(self):
pgb: Final = _ini(_plan("postgresql://app:pw@db/litellm?sslmode=require"))["pgbouncer"]
assert pgb["server_tls_sslmode"] == "require"
assert "server_tls_ca_file" not in pgb
def test_no_tls_params_default_to_prefer(self):
assert _ini(_plan("postgresql://app:pw@db/litellm"))["pgbouncer"]["server_tls_sslmode"] == "prefer"
def test_verification_without_a_ca_bundle_is_refused(self):
outcome: Final = plan_pgbouncer(
"postgresql://app:pw@db/litellm?sslmode=require&sslaccept=strict", SETTINGS, Path("/run/pgb"), None
)
assert isinstance(outcome, PgBouncerError)
assert "sslcert" in outcome.reason
def test_client_certificates_are_refused(self):
outcome: Final = plan_pgbouncer(
"postgresql://app:pw@db/litellm?sslidentity=/certs/client.p12", SETTINGS, Path("/run/pgb"), None
)
assert isinstance(outcome, PgBouncerError)
assert "sslidentity" in outcome.reason
@pytest.mark.parametrize(
"url",
[
"postgresql://app@db/litellm",
"postgresql://app:pw@db",
"postgresql://:pw@db/litellm",
],
)
def test_urls_missing_forwardable_credentials_are_refused(self, url: str):
outcome: Final = plan_pgbouncer(url, SETTINGS, Path("/run/pgb"), None)
assert isinstance(outcome, PgBouncerError)
def test_every_options_spelling_becomes_a_set_statement(self):
options: Final = urllib.parse.quote("-c a=1 -cb=2 --c=3")
ini: Final = _ini(_plan(f"postgresql://app:pw@db/litellm?options={options}"))
assert ini["databases"]["litellm"].endswith("connect_query='SET a TO ''1''; SET b TO ''2''; SET c TO ''3'''")
def test_options_that_are_not_settings_are_refused(self):
outcome: Final = plan_pgbouncer(
"postgresql://app:pw@db/litellm?options=-c%20search_path", SETTINGS, Path("/run/pgb"), None
)
assert isinstance(outcome, PgBouncerError)
assert "options" in outcome.reason
def test_run_as_user_is_only_written_when_given(self):
assert _ini(_plan(run_as_user="nobody"))["pgbouncer"]["user"] == "nobody"
assert "user" not in _ini(_plan())["pgbouncer"]
class TestWritePgBouncerFiles:
def test_files_hold_the_plan_and_are_private_to_the_owner(self, tmp_path: Path):
ini_path: Final = write_pgbouncer_files(_plan(), tmp_path, None)
assert ini_path == tmp_path / "pgbouncer.ini"
assert ini_path.read_text() == _plan().ini
assert (tmp_path / "userlist.txt").read_text() == _plan().userlist
for path in (ini_path, tmp_path / "userlist.txt"):
assert stat.S_IMODE(path.stat().st_mode) == 0o600
def _free_port() -> int:
with socket.socket() as probe:
probe.bind(("127.0.0.1", 0))
return probe.getsockname()[1]
def _fake_pooler(tmp_path: Path, port: int, exit_immediately: bool = False) -> Path:
"""An executable that listens on ``port`` like PgBouncer would (or exits at once), ignoring its ini argument."""
script: Final = tmp_path / "fake-pgbouncer"
script.write_text(
textwrap.dedent(
f"""\
#!{sys.executable}
import socket, sys, time
if {exit_immediately!r}:
sys.exit(3)
listener = socket.socket()
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
listener.bind(("127.0.0.1", {port}))
listener.listen()
while True:
conn, _ = listener.accept()
conn.close()
"""
)
)
script.chmod(0o700)
return script
def _listening(port: int) -> bool:
try:
with socket.create_connection(("127.0.0.1", port), timeout=0.5):
return True
except OSError:
return False
def _wait_until(condition: Callable[[], bool], timeout_seconds: float = 5.0) -> bool:
deadline: Final = time.monotonic() + timeout_seconds
while time.monotonic() < deadline:
if condition():
return True
time.sleep(0.05)
return False
class TestPgBouncerProcess:
def test_start_waits_for_the_listener_and_stop_ends_it(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(argv=(str(_fake_pooler(tmp_path, port)),), port=port)
assert pooler.start() is None
assert _listening(port)
pid: Final = pooler.pid
assert pid is not None
pooler.stop()
assert _wait_until(lambda: not _listening(port))
with pytest.raises(ProcessLookupError):
os.kill(pid, 0)
def test_a_crashed_pooler_is_restarted_with_a_new_pid(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, restart_delay_seconds=0.1
)
assert pooler.start() is None
first_pid: Final = pooler.pid
assert first_pid is not None
os.kill(first_pid, signal.SIGKILL)
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(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, restart_delay_seconds=0.1
)
assert pooler.start() is None
with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name):
pooler.stop()
time.sleep(0.5)
assert not _listening(port)
assert caplog.records == []
def test_a_pooler_that_exits_during_startup_is_reported(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(argv=(str(_fake_pooler(tmp_path, port, exit_immediately=True)),), port=port)
outcome: Final = pooler.start()
assert isinstance(outcome, PgBouncerError)
assert "status 3" in outcome.reason
def test_a_missing_binary_is_reported(self):
outcome: Final = PgBouncerProcess(argv=("/nonexistent/pgbouncer",), port=_free_port()).start()
assert isinstance(outcome, PgBouncerError)
assert "/nonexistent/pgbouncer" in outcome.reason
def test_a_pooler_that_never_listens_times_out(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(argv=(str(_fake_pooler(tmp_path, _free_port())),), port=port)
outcome: Final = pooler.start(ready_timeout_seconds=0.5)
assert isinstance(outcome, PgBouncerError)
assert "did not start listening" in outcome.reason
pid: Final = pooler.pid
assert pid is not None
with pytest.raises(ProcessLookupError):
os.kill(pid, 0)
class TestStartInContainerPgBouncer:
def test_returns_the_loopback_url_once_the_pooler_listens(self, tmp_path: Path):
port: Final = _free_port()
settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(_fake_pooler(tmp_path, port)))
pooled: Final = start_in_container_pgbouncer(settings, "postgresql://app:pw@db/litellm?connection_limit=5")
assert pooled == f"postgresql://app:pw@127.0.0.1:{port}/litellm?connection_limit=5&pgbouncer=true"
assert _listening(port)
def test_a_bad_upstream_url_is_reported_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@db/litellm")
assert isinstance(outcome, PgBouncerError)
assert not _listening(port)
class TestPgBouncerSettings:
def test_reads_the_litellm_pgbouncer_env_vars(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_PGBOUNCER_ENABLED", "true")
monkeypatch.setenv("LITELLM_PGBOUNCER_PORT", "7000")
monkeypatch.setenv("LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS", "12")
settings: Final = PgBouncerSettings()
assert (settings.enabled, settings.port, settings.max_db_connections) == (True, 7000, 12)
def test_defaults_are_off(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("LITELLM_PGBOUNCER_ENABLED", raising=False)
assert PgBouncerSettings().enabled is False