From 43995bcb75ac29246c387f63bb95c565ec58e6d6 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 20 Aug 2026 17:23:06 -0700 Subject: [PATCH] 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 Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_url_settings.py | 61 +++++- litellm/proxy/proxy_cli.py | 20 ++ .../proxy/db/test_db_url_settings.py | 132 ++++++++++++ tests/test_litellm/proxy/test_proxy_cli.py | 204 ++++++++++++++---- 4 files changed, 371 insertions(+), 46 deletions(-) diff --git a/litellm/proxy/db/db_url_settings.py b/litellm/proxy/db/db_url_settings.py index d393aa1b977..0918b9039da 100644 --- a/litellm/proxy/db/db_url_settings.py +++ b/litellm/proxy/db/db_url_settings.py @@ -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 = "" +# 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 diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 0e3e43accef..0449802abae 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -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: diff --git a/tests/test_litellm/proxy/db/test_db_url_settings.py b/tests/test_litellm/proxy/db/test_db_url_settings.py index e83e8310626..ee4cf7fbb05 100644 --- a/tests/test_litellm/proxy/db/test_db_url_settings.py +++ b/tests/test_litellm/proxy/db/test_db_url_settings.py @@ -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 diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 28f43345350..48c56a41ad5 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -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: