mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Merge remote-tracking branch 'origin/litellm_pgbouncer_in_container_pool' into litellm_perf_1k_rps
This commit is contained in:
commit
f2aa379242
9 changed files with 1175 additions and 3 deletions
19
Dockerfile
19
Dockerfile
|
|
@ -8,9 +8,25 @@ ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7
|
|||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
|
||||
# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x)
|
||||
ARG PGBOUNCER_VERSION=1.25.2
|
||||
ARG PGBOUNCER_SHA256=924ad35113fd0a71c8e2dbe85b5d03445532e2b7b37a9f8a48983beea238b332
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
||||
FROM $LITELLM_BUILD_IMAGE AS pgbouncer-builder
|
||||
ARG PGBOUNCER_VERSION
|
||||
ARG PGBOUNCER_SHA256
|
||||
USER root
|
||||
RUN apk add --no-cache build-base pkgconf libevent-dev openssl-dev curl
|
||||
WORKDIR /build
|
||||
RUN curl -fsSL -o pgbouncer.tar.gz "https://www.pgbouncer.org/downloads/files/${PGBOUNCER_VERSION}/pgbouncer-${PGBOUNCER_VERSION}.tar.gz" && \
|
||||
echo "${PGBOUNCER_SHA256} pgbouncer.tar.gz" | sha256sum -c - && \
|
||||
tar xzf pgbouncer.tar.gz --strip-components=1 && \
|
||||
./configure --prefix=/usr/local --with-openssl=/usr && \
|
||||
make -j"$(nproc)" pgbouncer && \
|
||||
install -m 0755 pgbouncer /usr/local/bin/pgbouncer
|
||||
|
||||
# Admin UI builder. Pinned to the build platform so the architecture-independent
|
||||
# Next.js static export compiles once natively even in a multi-arch build,
|
||||
# instead of once per target arch under QEMU.
|
||||
|
|
@ -110,7 +126,8 @@ 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 libevent
|
||||
COPY --from=pgbouncer-builder /usr/local/bin/pgbouncer /usr/local/bin/pgbouncer
|
||||
|
||||
WORKDIR /app
|
||||
ENV PATH="/app/.venv/bin:${PATH}" \
|
||||
|
|
|
|||
|
|
@ -8,9 +8,25 @@ ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7
|
|||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
|
||||
# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x)
|
||||
ARG PGBOUNCER_VERSION=1.25.2
|
||||
ARG PGBOUNCER_SHA256=924ad35113fd0a71c8e2dbe85b5d03445532e2b7b37a9f8a48983beea238b332
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
||||
FROM $LITELLM_BUILD_IMAGE AS pgbouncer-builder
|
||||
ARG PGBOUNCER_VERSION
|
||||
ARG PGBOUNCER_SHA256
|
||||
USER root
|
||||
RUN apk add --no-cache build-base pkgconf libevent-dev openssl-dev curl
|
||||
WORKDIR /build
|
||||
RUN curl -fsSL -o pgbouncer.tar.gz "https://www.pgbouncer.org/downloads/files/${PGBOUNCER_VERSION}/pgbouncer-${PGBOUNCER_VERSION}.tar.gz" && \
|
||||
echo "${PGBOUNCER_SHA256} pgbouncer.tar.gz" | sha256sum -c - && \
|
||||
tar xzf pgbouncer.tar.gz --strip-components=1 && \
|
||||
./configure --prefix=/usr/local --with-openssl=/usr && \
|
||||
make -j"$(nproc)" pgbouncer && \
|
||||
install -m 0755 pgbouncer /usr/local/bin/pgbouncer
|
||||
|
||||
# Admin UI builder. Pinned to the build platform so the architecture-independent
|
||||
# Next.js static export compiles once natively even in a multi-arch build,
|
||||
# instead of once per target arch under QEMU.
|
||||
|
|
@ -101,7 +117,8 @@ 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 libevent
|
||||
COPY --from=pgbouncer-builder /usr/local/bin/pgbouncer /usr/local/bin/pgbouncer
|
||||
|
||||
WORKDIR /app
|
||||
ENV PATH="/app/.venv/bin:${PATH}" \
|
||||
|
|
|
|||
|
|
@ -7,9 +7,25 @@ ARG PROXY_EXTRAS_SOURCE=published
|
|||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
|
||||
# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x)
|
||||
ARG PGBOUNCER_VERSION=1.25.2
|
||||
ARG PGBOUNCER_SHA256=924ad35113fd0a71c8e2dbe85b5d03445532e2b7b37a9f8a48983beea238b332
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
||||
FROM $LITELLM_BUILD_IMAGE AS pgbouncer-builder
|
||||
ARG PGBOUNCER_VERSION
|
||||
ARG PGBOUNCER_SHA256
|
||||
USER root
|
||||
RUN apk add --no-cache build-base pkgconf libevent-dev openssl-dev curl
|
||||
WORKDIR /build
|
||||
RUN curl -fsSL -o pgbouncer.tar.gz "https://www.pgbouncer.org/downloads/files/${PGBOUNCER_VERSION}/pgbouncer-${PGBOUNCER_VERSION}.tar.gz" && \
|
||||
echo "${PGBOUNCER_SHA256} pgbouncer.tar.gz" | sha256sum -c - && \
|
||||
tar xzf pgbouncer.tar.gz --strip-components=1 && \
|
||||
./configure --prefix=/usr/local --with-openssl=/usr && \
|
||||
make -j"$(nproc)" pgbouncer && \
|
||||
install -m 0755 pgbouncer /usr/local/bin/pgbouncer
|
||||
|
||||
# Admin UI builder. Pinned to the build platform so the architecture-independent
|
||||
# Next.js static export compiles once natively even in a multi-arch build,
|
||||
# instead of once per target arch under QEMU.
|
||||
|
|
@ -128,8 +144,9 @@ 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 libevent && break || sleep 5; \
|
||||
done
|
||||
COPY --from=pgbouncer-builder /usr/local/bin/pgbouncer /usr/local/bin/pgbouncer
|
||||
|
||||
# Copy only what runtime needs. The application is installed inside the venv;
|
||||
# the rest of the builder's /app is source and build metadata that must not
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -344,6 +344,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.
|
||||
|
|
|
|||
487
litellm/proxy/db/pgbouncer.py
Normal file
487
litellm/proxy/db/pgbouncer.py
Normal file
|
|
@ -0,0 +1,487 @@
|
|||
"""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.
|
||||
|
||||
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
|
||||
|
||||
import atexit
|
||||
import os
|
||||
import re
|
||||
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
|
||||
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"
|
||||
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"
|
||||
PGBOUNCER_MIN_VERSION: Final = (1, 19)
|
||||
PGBOUNCER_VERSION_PATTERN: Final = re.compile(r"PgBouncer (\d+)\.(\d+)")
|
||||
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
|
||||
# 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 _option_settings(tokens: Sequence[str]) -> tuple[str, ...] | None:
|
||||
"""The ``name=value`` settings in a libpq ``options`` string, or None if it holds anything else.
|
||||
|
||||
Accepts ``-c name=value``, ``-cname=value`` and ``--name=value``; a
|
||||
detached ``-c`` is folded into the token that follows it first.
|
||||
"""
|
||||
folded: Final = tuple(
|
||||
f"-c{tokens[index + 1]}" if token == "-c" and index + 1 < len(tokens) else token
|
||||
for index, token in enumerate(tokens)
|
||||
if index == 0 or tokens[index - 1] != "-c"
|
||||
)
|
||||
settings: Final = tuple(token[2:] for token in folded if token.startswith(("-c", "--")) and "=" in token[2:])
|
||||
return settings if len(settings) == len(folded) else None
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def _unix_socket_open(path: Path) -> bool:
|
||||
with socket.socket(socket.AF_UNIX) as probe:
|
||||
probe.settimeout(0.5)
|
||||
try:
|
||||
probe.connect(str(path))
|
||||
except OSError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def unix_socket_path(runtime_dir: Path, port: int) -> Path:
|
||||
return runtime_dir / f".s.PGSQL.{port}"
|
||||
|
||||
|
||||
def pgbouncer_version(binary: str) -> tuple[int, int] | PgBouncerError:
|
||||
"""``(major, minor)`` from ``<binary> --version``.
|
||||
|
||||
Readiness relies on PgBouncer exiting when it cannot bind its TCP port,
|
||||
which it does from 1.19 on. Older releases log a warning and serve the unix
|
||||
socket alone, so their socket would vouch for a port held by someone else.
|
||||
"""
|
||||
try:
|
||||
output: Final = subprocess.run(
|
||||
(binary, "--version"), capture_output=True, text=True, check=False, timeout=10
|
||||
).stdout
|
||||
except (OSError, subprocess.TimeoutExpired) as run_error:
|
||||
return PgBouncerError(f"could not run {binary!r} --version: {run_error}")
|
||||
found: Final = PGBOUNCER_VERSION_PATTERN.search(output)
|
||||
if found is None:
|
||||
return PgBouncerError(f"{binary!r} --version did not report a PgBouncer version: {output.strip()!r}")
|
||||
return int(found[1]), int(found[2])
|
||||
|
||||
|
||||
def _end(process: subprocess.Popen[bytes]) -> None:
|
||||
if process.poll() is not None:
|
||||
return
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=PGBOUNCER_STOP_GRACE_SECONDS)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
|
||||
|
||||
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 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, finds its port
|
||||
taken, exits again or never starts listening is retried every
|
||||
``restart_delay_seconds`` until ``stop`` is called.
|
||||
|
||||
A connect probe of ``port`` cannot tell the child from another process
|
||||
that grabbed the port after the availability check, so readiness also
|
||||
needs ``socket_path``: the unix socket PgBouncer creates in the private
|
||||
runtime directory, which it only does once every TCP listener is bound
|
||||
(PgBouncer 1.19 or newer, see ``pgbouncer_version``).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
argv: Sequence[str],
|
||||
port: int,
|
||||
socket_path: Path,
|
||||
restart_delay_seconds: float = PGBOUNCER_RESTART_DELAY_SECONDS,
|
||||
ready_timeout_seconds: float = PGBOUNCER_READY_TIMEOUT_SECONDS,
|
||||
) -> None:
|
||||
self.argv: Final = tuple(argv)
|
||||
self.port: Final = port
|
||||
self.socket_path: Final = socket_path
|
||||
self.restart_delay_seconds: Final = restart_delay_seconds
|
||||
self.ready_timeout_seconds: Final = ready_timeout_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] | PgBouncerError | None:
|
||||
"""Start a child, or None once ``stop`` ran; both take the lock so no child can slip in after a stop.
|
||||
|
||||
The port has to be free first: a listener that is already there would
|
||||
pass the readiness check while the child fails to bind.
|
||||
"""
|
||||
with self._lock:
|
||||
if self._stopping.is_set():
|
||||
return None
|
||||
if _port_open(self.port):
|
||||
return PgBouncerError(f"{PGBOUNCER_LISTEN_ADDR}:{self.port} is already in use by another process")
|
||||
try:
|
||||
process: Final = subprocess.Popen(self.argv)
|
||||
except OSError as spawn_error:
|
||||
return PgBouncerError(f"could not start {self.argv[0]!r}: {spawn_error}")
|
||||
self._process = process
|
||||
return process
|
||||
|
||||
def _wait_ready(self, process: subprocess.Popen[bytes]) -> PgBouncerError | None:
|
||||
deadline: Final = time.monotonic() + self.ready_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) and _unix_socket_open(self.socket_path):
|
||||
return None
|
||||
time.sleep(0.1)
|
||||
if _port_open(self.port):
|
||||
return PgBouncerError(
|
||||
f"{PGBOUNCER_LISTEN_ADDR}:{self.port} is served by another process, not the pgbouncer that was started"
|
||||
)
|
||||
return PgBouncerError(
|
||||
f"pgbouncer did not start listening on {PGBOUNCER_LISTEN_ADDR}:{self.port} "
|
||||
f"within {self.ready_timeout_seconds:.0f}s"
|
||||
)
|
||||
|
||||
def start(self) -> PgBouncerError | None:
|
||||
"""Spawn PgBouncer, wait until it listens on port and unix socket, then supervise it from a daemon thread."""
|
||||
process: Final = self._spawn()
|
||||
if process is None:
|
||||
return PgBouncerError("pgbouncer was stopped before it started")
|
||||
if isinstance(process, PgBouncerError):
|
||||
return process
|
||||
not_ready: Final = self._wait_ready(process)
|
||||
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,
|
||||
)
|
||||
self._restart_after_delay()
|
||||
|
||||
def _restart_after_delay(self) -> None:
|
||||
time.sleep(self.restart_delay_seconds)
|
||||
process: Final = self._spawn()
|
||||
if process is None:
|
||||
return
|
||||
if isinstance(process, PgBouncerError):
|
||||
self._retry_restart(process.reason)
|
||||
return
|
||||
not_ready: Final = self._wait_ready(process)
|
||||
if not_ready is None:
|
||||
self._watch(process)
|
||||
return
|
||||
_end(process)
|
||||
self._retry_restart(not_ready.reason)
|
||||
|
||||
def _retry_restart(self, reason: str) -> None:
|
||||
if self._stopping.is_set():
|
||||
return
|
||||
verbose_proxy_logger.error(
|
||||
"In-container pgbouncer could not be restarted (%s); retrying in %.1fs.", reason, self.restart_delay_seconds
|
||||
)
|
||||
threading.Thread(target=self._restart_after_delay, daemon=True, name="litellm-pgbouncer-supervisor").start()
|
||||
|
||||
def stop(self) -> None:
|
||||
with self._lock:
|
||||
self._stopping.set()
|
||||
process: Final = self._process
|
||||
if process is not None:
|
||||
_end(process)
|
||||
|
||||
|
||||
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)
|
||||
version: Final = pgbouncer_version(settings.binary)
|
||||
if isinstance(version, PgBouncerError):
|
||||
return version
|
||||
if version < PGBOUNCER_MIN_VERSION:
|
||||
return PgBouncerError(
|
||||
f"PgBouncer {version[0]}.{version[1]} keeps running after failing to bind its TCP port, so the proxy "
|
||||
f"cannot tell it apart from another listener; {PGBOUNCER_MIN_VERSION[0]}.{PGBOUNCER_MIN_VERSION[1]} "
|
||||
"or newer is required"
|
||||
)
|
||||
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,
|
||||
socket_path=unix_socket_path(runtime_dir, 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:
|
||||
|
|
@ -1377,6 +1378,21 @@ 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, 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 "
|
||||
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)
|
||||
if prometheus_metrics_port == port:
|
||||
|
|
|
|||
535
tests/test_litellm/proxy/db/test_pgbouncer.py
Normal file
535
tests/test_litellm/proxy/db/test_pgbouncer.py
Normal file
|
|
@ -0,0 +1,535 @@
|
|||
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 concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
from typing import Final, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.db.pgbouncer import (
|
||||
PgBouncerError,
|
||||
PgBouncerPlan,
|
||||
PgBouncerProcess,
|
||||
PgBouncerSettings,
|
||||
pgbouncer_version,
|
||||
plan_pgbouncer,
|
||||
start_in_container_pgbouncer,
|
||||
unix_socket_path,
|
||||
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 _bound_port(sock: socket.socket) -> int:
|
||||
return cast(tuple[str, int], sock.getsockname())[1]
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
with socket.socket() as probe:
|
||||
probe.bind(("127.0.0.1", 0))
|
||||
return _bound_port(probe)
|
||||
|
||||
|
||||
def _fake_pooler(
|
||||
tmp_path: Path,
|
||||
port: int,
|
||||
exit_immediately: bool = False,
|
||||
port_file: Path | None = None,
|
||||
bind_delay_seconds: float = 0.0,
|
||||
version_banner: str = "PgBouncer 1.25.2\nlibevent 2.1.13-stable",
|
||||
) -> Path:
|
||||
"""An executable that listens like PgBouncer: on the TCP port first, then on ``.s.PGSQL.<port>`` in the socket dir.
|
||||
|
||||
Port and socket dir come from the ini it is given, else from ``port`` and
|
||||
``tmp_path``. With ``port_file`` each start reads the port from that file
|
||||
instead. ``bind_delay_seconds`` holds the bind back, like a slow start.
|
||||
``--version`` prints ``version_banner``.
|
||||
"""
|
||||
script: Final = tmp_path / "fake-pgbouncer"
|
||||
script.write_text(
|
||||
textwrap.dedent(
|
||||
f"""\
|
||||
#!{sys.executable}
|
||||
import configparser, os, pathlib, select, socket, sys, time
|
||||
if sys.argv[1:] == ["--version"]:
|
||||
print({version_banner!r})
|
||||
sys.exit(0)
|
||||
if {exit_immediately!r}:
|
||||
sys.exit(3)
|
||||
ini = configparser.ConfigParser()
|
||||
ini.read(sys.argv[1:2])
|
||||
port = ini.getint("pgbouncer", "listen_port", fallback={port})
|
||||
if not {port_file is None!r}:
|
||||
port = int(pathlib.Path({str(port_file)!r}).read_text())
|
||||
socket_dir = ini.get("pgbouncer", "unix_socket_dir", fallback={str(tmp_path)!r})
|
||||
socket_path = f"{{socket_dir}}/.s.PGSQL.{{port}}"
|
||||
time.sleep({bind_delay_seconds!r})
|
||||
listener = socket.socket()
|
||||
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
listener.bind(("127.0.0.1", port))
|
||||
listener.listen()
|
||||
if os.path.exists(socket_path):
|
||||
os.unlink(socket_path)
|
||||
unix_listener = socket.socket(socket.AF_UNIX)
|
||||
unix_listener.bind(socket_path)
|
||||
unix_listener.listen()
|
||||
while True:
|
||||
for ready in select.select([listener, unix_listener], [], [])[0]:
|
||||
conn, _ = ready.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, socket_path=unix_socket_path(tmp_path, 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,
|
||||
socket_path=unix_socket_path(tmp_path, 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_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, socket_path=unix_socket_path(tmp_path, 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_replacement_that_never_listens_is_replaced_again(self, tmp_path: Path, caplog: pytest.LogCaptureFixture):
|
||||
port: Final = _free_port()
|
||||
port_file: Final = tmp_path / "port"
|
||||
port_file.write_text(str(port))
|
||||
script: Final = _fake_pooler(tmp_path, port, port_file=port_file)
|
||||
pooler: Final = PgBouncerProcess(
|
||||
argv=(str(script),),
|
||||
port=port,
|
||||
socket_path=unix_socket_path(tmp_path, port),
|
||||
restart_delay_seconds=0.1,
|
||||
ready_timeout_seconds=0.3,
|
||||
)
|
||||
assert pooler.start() is None
|
||||
first_pid: Final = pooler.pid
|
||||
assert first_pid is not None
|
||||
wrong_port: Final = _free_port()
|
||||
port_file.write_text(str(wrong_port))
|
||||
with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name):
|
||||
os.kill(first_pid, signal.SIGKILL)
|
||||
assert _wait_until(lambda: _listening(wrong_port))
|
||||
assert _wait_until(lambda: any("did not start listening" in record.message for record in caplog.records))
|
||||
port_file.write_text(str(port))
|
||||
assert _wait_until(lambda: _listening(port))
|
||||
assert _wait_until(lambda: not _listening(wrong_port))
|
||||
pooler.stop()
|
||||
assert _wait_until(lambda: not _listening(port))
|
||||
|
||||
def test_stopping_during_the_restart_delay_leaves_no_pooler_behind(self, tmp_path: Path):
|
||||
port: Final = _free_port()
|
||||
pooler: Final = PgBouncerProcess(
|
||||
argv=(str(_fake_pooler(tmp_path, port)),),
|
||||
port=port,
|
||||
socket_path=unix_socket_path(tmp_path, port),
|
||||
restart_delay_seconds=0.3,
|
||||
)
|
||||
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: not _listening(port))
|
||||
pooler.stop()
|
||||
time.sleep(1.0)
|
||||
assert not _listening(port)
|
||||
assert pooler.pid == first_pid
|
||||
|
||||
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,
|
||||
socket_path=unix_socket_path(tmp_path, 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,
|
||||
socket_path=unix_socket_path(tmp_path, port),
|
||||
)
|
||||
outcome: Final = pooler.start()
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert "status 3" in outcome.reason
|
||||
|
||||
def test_a_missing_binary_is_reported(self, tmp_path: Path):
|
||||
outcome: Final = PgBouncerProcess(
|
||||
argv=("/nonexistent/pgbouncer",), port=_free_port(), socket_path=tmp_path / "sock"
|
||||
).start()
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert "/nonexistent/pgbouncer" in outcome.reason
|
||||
|
||||
def test_a_port_owned_by_someone_else_is_refused_before_spawning(self, tmp_path: Path):
|
||||
with socket.socket() as squatter:
|
||||
squatter.bind(("127.0.0.1", 0))
|
||||
squatter.listen()
|
||||
port: Final = _bound_port(squatter)
|
||||
pooler: Final = PgBouncerProcess(
|
||||
argv=(str(_fake_pooler(tmp_path, port)),), port=port, socket_path=unix_socket_path(tmp_path, port)
|
||||
)
|
||||
outcome: Final = pooler.start()
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert f"127.0.0.1:{port} is already in use" in outcome.reason
|
||||
assert pooler.pid is None
|
||||
|
||||
def test_a_replacement_waits_until_a_squatter_leaves_the_port(
|
||||
self, tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
port: Final = _free_port()
|
||||
pooler: Final = PgBouncerProcess(
|
||||
argv=(str(_fake_pooler(tmp_path, port)),),
|
||||
port=port,
|
||||
socket_path=unix_socket_path(tmp_path, port),
|
||||
restart_delay_seconds=0.5,
|
||||
)
|
||||
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: not _listening(port))
|
||||
with socket.socket() as squatter, caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name):
|
||||
squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
squatter.bind(("127.0.0.1", port))
|
||||
squatter.listen()
|
||||
assert _wait_until(lambda: any("already in use" in record.message for record in caplog.records))
|
||||
assert pooler.pid == first_pid
|
||||
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_listener_that_grabs_the_port_after_the_spawn_is_not_taken_for_the_pooler(self, tmp_path: Path):
|
||||
port: Final = _free_port()
|
||||
pooler: Final = PgBouncerProcess(
|
||||
argv=(str(_fake_pooler(tmp_path, port, bind_delay_seconds=0.5)),),
|
||||
port=port,
|
||||
socket_path=unix_socket_path(tmp_path, port),
|
||||
ready_timeout_seconds=3.0,
|
||||
)
|
||||
with socket.socket() as squatter, ThreadPoolExecutor(max_workers=1) as starter:
|
||||
squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
starting: Final = starter.submit(pooler.start)
|
||||
assert _wait_until(lambda: pooler.pid is not None)
|
||||
squatter.bind(("127.0.0.1", port))
|
||||
squatter.listen()
|
||||
outcome: Final = starting.result()
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert "exited with status 1" in outcome.reason
|
||||
|
||||
def test_a_port_served_by_a_stranger_while_the_pooler_is_still_starting_is_reported(self, tmp_path: Path):
|
||||
port: Final = _free_port()
|
||||
pooler: Final = PgBouncerProcess(
|
||||
argv=(str(_fake_pooler(tmp_path, port, bind_delay_seconds=30.0)),),
|
||||
port=port,
|
||||
socket_path=unix_socket_path(tmp_path, port),
|
||||
ready_timeout_seconds=0.5,
|
||||
)
|
||||
with socket.socket() as squatter, ThreadPoolExecutor(max_workers=1) as starter:
|
||||
squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
starting: Final = starter.submit(pooler.start)
|
||||
assert _wait_until(lambda: pooler.pid is not None)
|
||||
squatter.bind(("127.0.0.1", port))
|
||||
squatter.listen()
|
||||
outcome: Final = starting.result()
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert f"127.0.0.1:{port} is served by another process" in outcome.reason
|
||||
pid: Final = pooler.pid
|
||||
assert pid is not None
|
||||
with pytest.raises(ProcessLookupError):
|
||||
os.kill(pid, 0)
|
||||
|
||||
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,
|
||||
socket_path=unix_socket_path(tmp_path, port),
|
||||
ready_timeout_seconds=0.5,
|
||||
)
|
||||
outcome: Final = pooler.start()
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
def test_a_pgbouncer_that_survives_a_failed_tcp_bind_is_refused_without_starting(self, tmp_path: Path):
|
||||
port: Final = _free_port()
|
||||
binary: Final = _fake_pooler(tmp_path, port, version_banner="PgBouncer 1.18.1\nlibevent 2.1.12-stable")
|
||||
settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(binary))
|
||||
outcome: Final = start_in_container_pgbouncer(settings, "postgresql://app:pw@db/litellm")
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert "PgBouncer 1.18" in outcome.reason
|
||||
assert "1.19" in outcome.reason
|
||||
assert not _listening(port)
|
||||
|
||||
def test_the_first_version_that_dies_on_a_failed_tcp_bind_is_accepted(self, tmp_path: Path):
|
||||
port: Final = _free_port()
|
||||
binary: Final = _fake_pooler(tmp_path, port, version_banner="PgBouncer 1.19.0")
|
||||
settings: Final = PgBouncerSettings(enabled=True, port=port, binary=str(binary))
|
||||
assert start_in_container_pgbouncer(settings, "postgresql://app:pw@db/litellm") == (
|
||||
f"postgresql://app:pw@127.0.0.1:{port}/litellm?pgbouncer=true"
|
||||
)
|
||||
assert _listening(port)
|
||||
|
||||
|
||||
class TestPgBouncerVersion:
|
||||
def test_reads_major_and_minor_from_the_banner(self, tmp_path: Path):
|
||||
assert pgbouncer_version(str(_fake_pooler(tmp_path, _free_port()))) == (1, 25)
|
||||
|
||||
def test_a_binary_that_cannot_run_is_reported(self, tmp_path: Path):
|
||||
outcome: Final = pgbouncer_version(str(tmp_path / "missing-pgbouncer"))
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert "missing-pgbouncer" in outcome.reason
|
||||
|
||||
def test_a_banner_without_a_version_is_reported(self, tmp_path: Path):
|
||||
outcome: Final = pgbouncer_version(str(_fake_pooler(tmp_path, _free_port(), version_banner="something else")))
|
||||
assert isinstance(outcome, PgBouncerError)
|
||||
assert "something else" in outcome.reason
|
||||
|
||||
|
||||
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