mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
cf3af0f486
commit
dafa02daf5
9 changed files with 763 additions and 3 deletions
|
|
@ -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}" \
|
||||
|
|
|
|||
|
|
@ -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}" \
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
61
helm/litellm-helm/tests/connection_pool_tests.yaml
Normal file
61
helm/litellm-helm/tests/connection_pool_tests.yaml
Normal 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"
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
373
litellm/proxy/db/pgbouncer.py
Normal file
373
litellm/proxy/db/pgbouncer.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
290
tests/test_litellm/proxy/db/test_pgbouncer.py
Normal file
290
tests/test_litellm/proxy/db/test_pgbouncer.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue