fix(db): apply the configured connection params to the read replica URL (#37691)

The read replica never received the operator's DB pool settings, so its
Prisma pool fell back to `num_physical_cpus * 2 + 1` and the configured cap
was not enforced. Both startup paths now pass the same params to the reader:
the CLI, and the componentized entrypoints that go through
`DatabaseURLSettings.apply_to_env`.

Only pool and timeout params are inherited, through a single allowlist both
paths share. Anything that decides which tables a query resolves against
stays on the writer, including entries smuggled in through
`database_extra_connection_params`, so a writer `search_path` cannot repoint
reader queries. Params the operator pinned on the replica URL still win.

Co-authored-by: Yassin Kortam <yassin.kortam@gmail.com>
Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-08-20 17:23:06 -07:00 • committed by GitHub
parent fb3dd0fb98
commit 43995bcb75
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 371 additions and 46 deletions

View file

@ -29,12 +29,16 @@ can run alongside a password-auth reader (or a precomputed reader URL). Reader
token auth is gated on the same global toggle as the writer: the chart only
emits the reader token env vars when the writer also uses token auth.
Reader-side fields fall back to the writer's user / name / schema / port /
password when their ``*_READ_REPLICA`` counterpart is unset.
password when their ``*_READ_REPLICA`` counterpart is unset, and to the
writer's connection params (pool size, timeouts, pgbouncer mode) for the
ones the reader URL does not pin itself.
"""
import os
import urllib.parse
from collections.abc import Mapping
from functools import partial
from types import MappingProxyType
from typing import Annotated, Final, cast
from pydantic import AliasChoices, BeforeValidator, Field
@ -62,6 +66,51 @@ SUPPORTED_DB_SCHEMES: Final[frozenset[str]] = frozenset({"postgresql", "postgres
_MISSING_SCHEME: Final = "<missing scheme>"
# An allowlist, deliberately not a denylist: only these pool and timeout params
# follow the writer to the read replica, so nothing that decides which tables a
# query resolves against (``schema``, or a ``search_path`` inside ``options``)
# can ever repoint the reader. Without them the reader pool silently falls back
# to Prisma's default size.
CONNECTION_PARAM_KEYS: Final[frozenset[str]] = frozenset(
{
"connection_limit",
"pool_timeout",
"connect_timeout",
"socket_timeout",
"pgbouncer",
}
)
def add_missing_query_params(url: str, params: Mapping[str, str | int | float]) -> str:
"""Return ``url`` with the ``params`` it does not already carry appended.
Params the operator pinned on the URL win, so a hand-tuned replica URL keeps
its values. Returns the URL untouched when there is nothing to add, leaving
its existing encoding alone.
"""
parsed: Final = urllib.parse.urlsplit(url)
existing: Final = tuple(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True))
pinned: Final = frozenset(key for key, _ in existing)
additions: Final = tuple((key, str(value)) for key, value in params.items() if key not in pinned)
if not additions:
return url
query: Final = urllib.parse.urlencode(existing + additions)
return urllib.parse.urlunsplit(parsed._replace(query=query))
def reader_shareable_params(params: Mapping[str, str | int | float]) -> Mapping[str, str | int | float]:
"""Return the subset of ``params`` the read replica is allowed to inherit."""
return MappingProxyType({key: value for key, value in params.items() if key in CONNECTION_PARAM_KEYS})
def connection_params_from_url(url: str) -> Mapping[str, str | int | float]:
"""Return the connection params on ``url`` that the read replica shares."""
return reader_shareable_params(
MappingProxyType({key: value for key, value in urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query)})
)
def unsupported_db_scheme(database_url: str) -> str | None:
"""Return the connection URL scheme when it is not PostgreSQL, else None.
@ -326,8 +375,14 @@ class DatabaseURLSettings(BaseSettings):
self._raise_for_unsupported_scheme()
wrote_writer: Final = self.apply_writer_url_to_env()
reader_url: Final = self.build_reader_url()
# The reader inherits the writer's connection params (pool size, timeouts,
# pgbouncer mode). Without this the reader pool ignores the configured cap
# and falls back to Prisma's `num_physical_cpus * 2 + 1` default.
reader_url: Final = self.build_reader_url() or self.database_url_read_replica
if reader_url is not None:
os.environ["DATABASE_URL_READ_REPLICA"] = reader_url
os.environ["DATABASE_URL_READ_REPLICA"] = add_missing_query_params(
reader_url,
connection_params_from_url(os.environ.get("DATABASE_URL", "")),
)
return wrote_writer

View file

@ -1224,6 +1224,8 @@ def run_server(
if os.getenv("DATABASE_URL", None) is not None or os.getenv("DIRECT_URL", None) is not None:
from litellm.proxy.db.db_url_settings import (
add_missing_query_params,
reader_shareable_params,
unsupported_db_scheme,
unsupported_db_scheme_message,
)
@ -1273,6 +1275,24 @@ def run_server(
database_url = os.getenv("DIRECT_URL")
modified_url = append_query_params(database_url, connection_url_params)
os.environ["DIRECT_URL"] = modified_url
# The reader pool is a real pool against the same configured cap, so it
# gets the allowlisted pool params. Schema-affecting ones, including any
# the operator smuggled in through database_extra_connection_params, stay
# on the writer. Anything pinned on the replica URL wins, unlike the
# writer where the config is applied on top.
read_replica_url: Final[str | None] = os.getenv("DATABASE_URL_READ_REPLICA")
if read_replica_url:
reader_options: Final[str] = _pg_options_with_timeouts(
_url_query_value(read_replica_url, "options"),
db_statement_timeout,
db_lock_timeout,
)
os.environ["DATABASE_URL_READ_REPLICA"] = add_missing_query_params(
_with_query_value(read_replica_url, "options", reader_options)
if reader_options
else read_replica_url,
reader_shareable_params(connection_url_params),
)
subprocess.run(["prisma"], capture_output=True)
is_prisma_runnable = True
except FileNotFoundError:

View file

@ -12,6 +12,7 @@ clobber a pre-existing ``DATABASE_URL_READ_REPLICA``. A pre-existing
"""
import os
import urllib.parse
from unittest.mock import patch
import pytest
@ -524,6 +525,137 @@ def test_apply_to_env_accepts_pinned_postgres(monkeypatch):
assert _apply() is False
# ---------------------------------------------------------------------------
# Connection params on the read replica
# ---------------------------------------------------------------------------
def test_reader_inherits_writer_connection_params(monkeypatch):
"""The reader is a second pool: without the writer's params it sizes itself
from Prisma's default and the operator's cap is not enforced."""
monkeypatch.setenv(
"DATABASE_URL",
"postgresql://u:p@writer.example.com:5432/db?connection_limit=3&pool_timeout=20&pgbouncer=true",
)
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA", "postgresql://u:p@reader.example.com:5432/db"
)
_apply()
query = urllib.parse.parse_qs(
urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query
)
assert query["connection_limit"] == ["3"]
assert query["pool_timeout"] == ["20"]
assert query["pgbouncer"] == ["true"]
def test_reader_keeps_its_own_pinned_connection_params(monkeypatch):
monkeypatch.setenv(
"DATABASE_URL",
"postgresql://u:p@writer.example.com:5432/db?connection_limit=3&pool_timeout=20",
)
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA",
"postgresql://u:p@reader.example.com:5432/db?connection_limit=50",
)
_apply()
query = urllib.parse.parse_qs(
urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query
)
assert query["connection_limit"] == ["50"]
assert query["pool_timeout"] == ["20"]
def test_assembled_reader_url_inherits_writer_connection_params(monkeypatch):
"""A reader assembled from the discrete DATABASE_*_READ_REPLICA vars must
carry the params too, and must not inherit the writer's schema."""
monkeypatch.setenv(
"DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db?connection_limit=3&schema=writer_schema"
)
monkeypatch.setenv("DATABASE_USER", "litellm")
monkeypatch.setenv("DATABASE_NAME", "litellm_db")
monkeypatch.setenv("DATABASE_PASSWORD", "s3cr3t")
monkeypatch.setenv("DATABASE_HOST_READ_REPLICA", "reader.example.com")
_apply()
reader_url = os.environ["DATABASE_URL_READ_REPLICA"]
assert reader_url.startswith("postgresql://litellm:s3cr3t@reader.example.com:5432/litellm_db?")
query = urllib.parse.parse_qs(urllib.parse.urlsplit(reader_url).query)
assert query["connection_limit"] == ["3"]
assert "schema" not in query
def test_reader_does_not_inherit_writer_options(monkeypatch):
"""A writer search_path must not follow the reader, or reader queries resolve
against the wrong schema."""
monkeypatch.setenv(
"DATABASE_URL",
"postgresql://u:p@writer.example.com:5432/db?connection_limit=3&options=-c%20search_path%3Dwriter_schema",
)
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "postgresql://u:p@reader.example.com:5432/db")
_apply()
query = urllib.parse.parse_qs(urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)
assert query["connection_limit"] == ["3"]
assert "options" not in query
def test_reader_does_not_inherit_an_unvetted_writer_param(monkeypatch):
"""Inheritance is an allowlist, so a param nobody vetted for the reader stays
on the writer. Flipping this to a denylist would let the next schema-affecting
param leak through by default."""
monkeypatch.setenv(
"DATABASE_URL",
"postgresql://u:p@writer.example.com:5432/db?connection_limit=3&application_name=writer&novel_param=x",
)
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "postgresql://u:p@reader.example.com:5432/db")
_apply()
query = urllib.parse.parse_qs(urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)
assert query["connection_limit"] == ["3"]
assert "application_name" not in query
assert "novel_param" not in query
def test_reader_keeps_its_own_options_when_writer_params_are_appended(monkeypatch):
"""Appending the writer's pool params must leave the reader's own search_path
intact, since that is what decides which tables its queries resolve against."""
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db?connection_limit=3")
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA",
"postgresql://u:p@reader.example.com:5432/db?options=-c%20search_path%3Dreader_schema",
)
_apply()
query = urllib.parse.parse_qs(urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)
assert query["options"] == ["-c search_path=reader_schema"]
assert query["connection_limit"] == ["3"]
def test_reader_url_left_alone_when_writer_has_no_params(monkeypatch):
"""No params to inherit must mean the reader URL is not rewritten at all."""
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA",
"postgresql://u:p@reader.example.com:5432/db?options=-c%20search_path%3Dapp",
)
_apply()
assert (
os.environ["DATABASE_URL_READ_REPLICA"]
== "postgresql://u:p@reader.example.com:5432/db?options=-c%20search_path%3Dapp"
)
def test_unsupported_db_scheme_message_names_var_and_scheme():
msg = unsupported_db_scheme_message("DIRECT_URL", "sqlite")
assert "DIRECT_URL" in msg

View file

@ -18,8 +18,9 @@ import types
import urllib.parse as urlparse
import uvicorn
import yaml
from litellm.proxy.proxy_cli import ProxyInitializationHelpers
from litellm.proxy.proxy_cli import ProxyInitializationHelpers, run_server
@pytest.mark.xdist_group("proxy_cli")
@ -2242,7 +2243,7 @@ class TestPostgresStatementTimeoutOptions:
yaml.dump({"model_list": [], "general_settings": {"database_statement_timeout": 60}})
)
captured = self._run_server_and_capture_urls(
captured = _run_server_and_capture_urls(
str(config_path), direct_url="postgresql://t:t@localhost:5432/t"
)
@ -2285,57 +2286,174 @@ class TestPostgresStatementTimeoutOptions:
assert "-c search_path=app" in options
assert "-c statement_timeout=60000" in options
@classmethod
@staticmethod
def _run_server_and_capture_database_url(
cls,
config_path: str,
database_url: str = "postgresql://t:t@localhost:5432/t",
) -> str:
return cls._run_server_and_capture_urls(config_path, database_url=database_url)["DATABASE_URL"]
return _run_server_and_capture_urls(config_path, database_url=database_url)["DATABASE_URL"]
@staticmethod
def _run_server_and_capture_urls(
config_path: str,
database_url: str = "postgresql://t:t@localhost:5432/t",
direct_url: str | None = None,
) -> dict:
from litellm.proxy.proxy_cli import run_server
_CAPTURED_DB_ENV_VARS = ("DATABASE_URL", "DIRECT_URL", "DATABASE_URL_READ_REPLICA")
def _run_server_and_capture_urls(
config_path: str,
database_url: str = "postgresql://t:t@localhost:5432/t",
direct_url: str | None = None,
read_replica_url: str | None = None,
) -> dict:
loaded_config = yaml.safe_load(Path(config_path).read_text())
mock_proxy_config = MagicMock()
mock_proxy_config.return_value.get_config = AsyncMock(return_value=loaded_config)
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=mock_proxy_config,
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
clean_env = {k: v for k, v in os.environ.items() if k not in _CAPTURED_DB_ENV_VARS}
clean_env["DATABASE_URL"] = database_url
if direct_url is not None:
clean_env["DIRECT_URL"] = direct_url
if read_replica_url is not None:
clean_env["DATABASE_URL_READ_REPLICA"] = read_replica_url
with (
patch.dict(os.environ, clean_env, clear=True),
patch.dict(
"sys.modules",
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
),
patch("subprocess.run", return_value=MagicMock(returncode=0)),
patch("atexit.register"),
patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False),
patch("litellm.proxy.db.check_migration.check_prisma_schema_diff"),
):
run_server.main(
["--config", config_path, "--local", "--skip_server_startup"],
standalone_mode=False,
)
return {k: os.environ[k] for k in _CAPTURED_DB_ENV_VARS if k in os.environ}
class TestReadReplicaConnectionParams:
"""The reader is a second Prisma client with its own pool. Without the
configured params on DATABASE_URL_READ_REPLICA it sizes itself from Prisma's
`num_physical_cpus * 2 + 1` default, so an operator's cap is not the cap that
gets enforced.
"""
def test_pool_settings_reach_the_read_replica_url(self, tmp_path):
import yaml
loaded_config = yaml.safe_load(Path(config_path).read_text())
mock_proxy_config = MagicMock()
mock_proxy_config.return_value.get_config = AsyncMock(return_value=loaded_config)
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=mock_proxy_config,
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")}
clean_env["DATABASE_URL"] = database_url
if direct_url is not None:
clean_env["DIRECT_URL"] = direct_url
with (
patch.dict(os.environ, clean_env, clear=True),
patch.dict(
"sys.modules",
config_path = tmp_path / "config.yaml"
config_path.write_text(
yaml.dump(
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
),
patch("subprocess.run", return_value=MagicMock(returncode=0)),
patch("atexit.register"),
patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False),
patch("litellm.proxy.db.check_migration.check_prisma_schema_diff"),
):
run_server.main(
["--config", config_path, "--local", "--skip_server_startup"],
standalone_mode=False,
"model_list": [],
"general_settings": {
"database_connection_pool_limit": 3,
"database_connection_pool_timeout": 20,
"database_connect_timeout": 15,
"database_socket_timeout": 120,
"database_disable_prepared_statements": True,
"database_statement_timeout": 60,
},
}
)
return {k: os.environ[k] for k in ("DATABASE_URL", "DIRECT_URL") if k in os.environ}
)
captured = _run_server_and_capture_urls(
str(config_path),
read_replica_url="postgresql://t:t@reader:5432/t",
)
query = urlparse.parse_qs(urlparse.urlparse(captured["DATABASE_URL_READ_REPLICA"]).query)
assert query["connection_limit"] == ["3"]
assert query["pool_timeout"] == ["20"]
assert query["connect_timeout"] == ["15"]
assert query["socket_timeout"] == ["120"]
assert query["pgbouncer"] == ["true"]
assert "-c statement_timeout=60000" in query["options"][0]
def test_operator_pinned_replica_params_win(self, tmp_path):
"""The documented workaround (params pinned on the replica URL) must keep
working, so an operator who tuned the reader separately is not overridden.
"""
import yaml
config_path = tmp_path / "config.yaml"
config_path.write_text(
yaml.dump(
{
"model_list": [],
"general_settings": {
"database_connection_pool_limit": 3,
"database_connection_pool_timeout": 20,
},
}
)
)
captured = _run_server_and_capture_urls(
str(config_path),
read_replica_url="postgresql://t:t@reader:5432/t?connection_limit=50",
)
query = urlparse.parse_qs(urlparse.urlparse(captured["DATABASE_URL_READ_REPLICA"]).query)
assert query["connection_limit"] == ["50"]
assert query["pool_timeout"] == ["20"]
def test_extra_connection_params_never_carry_a_schema_override_to_the_reader(self, tmp_path):
"""database_extra_connection_params is an untyped passthrough, so it can carry a
search_path. The writer keeps it, the reader must not inherit it, or replica
queries resolve against the writer's schema.
"""
config_path = tmp_path / "config.yaml"
config_path.write_text(
yaml.dump(
{
"model_list": [],
"general_settings": {
"database_connection_pool_limit": 3,
"database_extra_connection_params": {
"options": "-c search_path=writer_schema",
"schema": "writer_schema",
"socket_timeout": 90,
},
},
}
)
)
captured = _run_server_and_capture_urls(
str(config_path),
read_replica_url="postgresql://t:t@reader:5432/t",
)
writer_query = urlparse.parse_qs(urlparse.urlparse(captured["DATABASE_URL"]).query)
assert writer_query["options"] == ["-c search_path=writer_schema"]
assert writer_query["schema"] == ["writer_schema"]
reader_query = urlparse.parse_qs(urlparse.urlparse(captured["DATABASE_URL_READ_REPLICA"]).query)
assert reader_query["connection_limit"] == ["3"]
assert reader_query["socket_timeout"] == ["90"]
assert "options" not in reader_query
assert "schema" not in reader_query
def test_replica_url_untouched_when_unset(self, tmp_path):
import yaml
config_path = tmp_path / "config.yaml"
config_path.write_text(yaml.dump({"model_list": [], "general_settings": {}}))
captured = _run_server_and_capture_urls(str(config_path))
assert "DATABASE_URL_READ_REPLICA" not in captured
class TestTokenAuthCliFlags: