From 41a3781d4ee747cf5faee2b06204c5c942cec178 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 14:28:52 -0700 Subject: [PATCH] fix(proxy): treat Postgres connection exhaustion as backpressure, not poison rows (#44266) 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 | 54 +++++++- litellm/proxy/db/exception_handler.py | 16 ++- litellm/proxy/proxy_cli.py | 21 ++++ litellm/proxy/utils.py | 29 +++-- tests/unit/proxy/db/test_db_url_settings.py | 50 ++++++++ tests/unit/proxy/db/test_exception_handler.py | 86 ++++++++++++- tests/unit/proxy/test_proxy_cli.py | 73 ++++++++++- .../prisma_and_spend/test_spend_functions.py | 118 ++++++++++++++++++ 8 files changed, 429 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/db/db_url_settings.py b/litellm/proxy/db/db_url_settings.py index e6e97cb1eb3..7d00784182a 100644 --- a/litellm/proxy/db/db_url_settings.py +++ b/litellm/proxy/db/db_url_settings.py @@ -254,6 +254,40 @@ def translate_libpq_ssl_params(url: str, resolve_root_cert: RootCertResolver = p return urllib.parse.urlunsplit(parsed._replace(query=query)) +def postgres_connection_budget_message(writer_limit: str, reader_limit: str | None, num_workers: str) -> str: + """The startup line that states this pod's worst-case Postgres connection demand. + + Prisma's ``connection_limit`` is per query engine, and every uvicorn worker + owns one engine per configured database (writer, plus the reader when + ``DATABASE_URL_READ_REPLICA`` is set). The limits are read off the final URLs, + so a reader that pins its own ``connection_limit`` or an operator override in + ``database_extra_connection_params`` is counted at its real value. The + server-side cap is shared by every pod, so the number an operator has to keep + under ``max_connections`` minus ``superuser_reserved_connections`` is + pods x workers x the per-worker sum, not one engine's limit. + """ + try: + workers: Final = max(1, int(num_workers)) + writer: Final = int(writer_limit) + reader: Final = int(reader_limit) if reader_limit is not None else 0 + except ValueError: + return ( + "LiteLLM Proxy: Postgres connection budget per pod = workers x (writer connection_limit + reader " + f"connection_limit) (workers={num_workers!r}, writer={writer_limit!r}, reader={reader_limit!r}); " + "keep pods x that figure under max_connections minus superuser_reserved_connections" + ) + engines: Final = ( + f"(writer connection_limit {writer} + reader connection_limit {reader})" + if reader_limit is not None + else f"writer connection_limit {writer}" + ) + per_pod: Final = workers * (writer + reader) + return ( + f"LiteLLM Proxy: Postgres connection budget per pod = {workers} worker(s) x {engines} = up to {per_pod} " + "connections; keep pods x that figure under max_connections minus superuser_reserved_connections" + ) + + 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}) @@ -266,16 +300,26 @@ def connection_params_from_url(url: str) -> Mapping[str, str | int | float]: ) +# A re-minted token URL replaces a URL to the same database, so unlike the +# reader allowlist it may also carry ``options``: that is where the server-side +# timeouts (statement, lock, idle-in-transaction) live, and a refresh that +# dropped them would leave the replacement engine's sessions unbounded. +TOKEN_REFRESH_PARAM_KEYS: Final[frozenset[str]] = CONNECTION_PARAM_KEYS | PRISMA_TLS_PARAM_KEYS | frozenset({"options"}) + + def token_refresh_params_from_url(url: str) -> Mapping[str, str | int | float]: """Return the params a re-minted token URL carries over from the URL it replaces. - The pool and timeout params plus Prisma's TLS params (already translated from - libpq spelling), so a refreshed URL keeps verifying the server the way the - first one did. + The pool and timeout params, Prisma's TLS params (already translated from + libpq spelling) and the ``options`` string, so a refreshed URL keeps verifying + the server and bounding its sessions the way the first one did. """ - kept: Final = CONNECTION_PARAM_KEYS | PRISMA_TLS_PARAM_KEYS return MappingProxyType( - {key: value for key, value in urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query) if key in kept} + { + key: value + for key, value in urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query) + if key in TOKEN_REFRESH_PARAM_KEYS + } ) diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index 7eb8f564225..c3526548942 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -30,6 +30,7 @@ _CONNECTION_CAPACITY_PHRASES: Final = ( "too many connections for database", "remaining connection slots are reserved", ) +_PRISMA_POOL_TIMEOUT_CODE: Final = "P2024" def _exception_chain(e: BaseException) -> Iterator[BaseException]: @@ -115,6 +116,8 @@ class PrismaDBExceptionHandler: return True if isinstance(e, _exception_types(prisma.engine.errors.EngineConnectionError)): return True + if PrismaDBExceptionHandler.is_database_capacity_error(e): + return True return isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection @staticmethod @@ -198,6 +201,8 @@ class PrismaDBExceptionHandler: ), ): return True + if PrismaDBExceptionHandler.is_database_capacity_error(e): + return False if isinstance(e, _exception_types(prisma.errors.PrismaError)): error_message: Final = str(e).lower() connection_keywords: Final = ( @@ -245,14 +250,17 @@ class PrismaDBExceptionHandler: @staticmethod def is_database_capacity_error(e: Exception) -> bool: """True iff Postgres refused the pool a new connection (SQLSTATE 53300: - ``too many clients already``, a reserved slot, or a per-role limit). The - server is up but full, so the failure is neither a transport error (which - would tear the engine down and open yet more connections against it) nor - a rejection of the rows being written.""" + ``too many clients already``, a reserved slot, or a per-role limit) or + prisma's own pool timed out handing one over (P2024). The server is up + but full, so the failure is neither a transport error (which would tear + the engine down and open yet more connections against it) nor a + rejection of the rows being written.""" import prisma if not isinstance(e, _exception_types(prisma.errors.PrismaError)): return False + if isinstance(e, prisma.errors.DataError) and e.code == _PRISMA_POOL_TIMEOUT_CODE: + return True if PrismaDBExceptionHandler.postgres_sqlstate(e) == "53300": return True error_message: Final = str(e).lower() diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index a9745565799..9e736757fe2 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -108,12 +108,14 @@ class DatabaseTimeoutSettings(BaseModel): database_statement_timeout: float | None = None database_lock_timeout: float | None = None + database_idle_in_transaction_session_timeout: float | None = None def _pg_options_with_timeouts( existing_options: str, statement_timeout: float | None, lock_timeout: float | None, + idle_in_transaction_session_timeout: float | None = None, ) -> str: """Return the Postgres ``options`` value carrying the configured timeouts. @@ -143,6 +145,7 @@ def _pg_options_with_timeouts( for name, seconds in ( ("statement_timeout", statement_timeout), ("lock_timeout", lock_timeout), + ("idle_in_transaction_session_timeout", idle_in_transaction_session_timeout), ) if seconds is not None and not re.search(rf"(?:-c\s*|--){re.escape(name)}=", existing_options) ) @@ -1160,6 +1163,7 @@ def run_server( db_extra_connection_params: dict | None = None db_statement_timeout: float | None = None db_lock_timeout: float | None = None + db_idle_in_transaction_timeout: float | None = None general_settings = {} ### GET DB TOKEN FOR RDS IAM / AZURE ENTRA AUTH ### @@ -1270,6 +1274,7 @@ def run_server( db_timeouts: Final = DatabaseTimeoutSettings.model_validate(general_settings) db_statement_timeout = db_timeouts.database_statement_timeout db_lock_timeout = db_timeouts.database_lock_timeout + db_idle_in_transaction_timeout = db_timeouts.database_idle_in_transaction_session_timeout if database_url and database_url.startswith("os.environ/"): original_dir: Final = os.getcwd() # set the working directory to where this script is @@ -1303,6 +1308,7 @@ def run_server( DISABLE_PREPARED_STATEMENTS_ENV_VAR, add_missing_query_params, idle_lifetime_params, + postgres_connection_budget_message, reader_shareable_params, translate_libpq_ssl_params, unsupported_db_scheme, @@ -1343,6 +1349,7 @@ def run_server( _url_query_value(resolved_url, "options"), db_statement_timeout, db_lock_timeout, + db_idle_in_transaction_timeout, ) writer_url: Final = ( _with_query_value(resolved_url, "options", pg_options) @@ -1373,6 +1380,7 @@ def run_server( _url_query_value(read_replica_url, "options"), db_statement_timeout, db_lock_timeout, + db_idle_in_transaction_timeout, ) os.environ["DATABASE_URL_READ_REPLICA"] = translate_libpq_ssl_params( add_missing_query_params( @@ -1385,6 +1393,19 @@ def run_server( lifetime_params, ) ) + print( + postgres_connection_budget_message( + writer_limit=_url_query_value(os.getenv("DATABASE_URL"), "connection_limit") + or f"{db_connection_pool_limit}", + reader_limit=( + _url_query_value(os.getenv("DATABASE_URL_READ_REPLICA"), "connection_limit") + or f"{db_connection_pool_limit}" + ) + if read_replica_url + else None, + num_workers=f"{1 if run_hypercorn else num_workers}", + ) + ) from litellm_proxy_extras.prisma_toolchain import prisma_cli_available is_prisma_runnable: Final = prisma_cli_available() diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 4ebff86ac06..8af90e19ecb 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -7661,8 +7661,9 @@ async def _run_spend_logs_job( ) # Tool usage tracking: drain the request-time queue into the tool index and the - # LiteLLM_DailyToolSpend rollup. Never retried; a dropped batch is permanently - # absent from the rollup, so failures log at error. + # LiteLLM_DailyToolSpend rollup. A batch Postgres had no connection for was never + # sent, so it is requeued and the job stops; any other failure is dropped because + # a replay could double-count the rollup. async with prisma_client._tool_usage_transactions_lock: tool_usage_to_process: Final = prisma_client.tool_usage_transactions[:MAX_LOGS_PER_INTERVAL] prisma_client.tool_usage_transactions = prisma_client.tool_usage_transactions[len(tool_usage_to_process) :] @@ -7674,11 +7675,25 @@ async def _run_spend_logs_job( transactions=tool_usage_to_process, ) except Exception as tool_tracking_err: - verbose_proxy_logger.error( - "Spend tracking - tool usage flush failed; %s tool usage transactions dropped: %s", - len(tool_usage_to_process), - tool_tracking_err, - ) + if PrismaDBExceptionHandler.is_database_capacity_error(tool_tracking_err): + async with prisma_client._tool_usage_transactions_lock: + prisma_client.tool_usage_transactions = [ + *tool_usage_to_process, + *prisma_client.tool_usage_transactions, + ] + verbose_proxy_logger.warning( + "Spend tracking - database out of connections during tool usage flush, " + "requeued %s tool usage transactions for the next flush: %s", + len(tool_usage_to_process), + tool_tracking_err, + ) + raise + else: + verbose_proxy_logger.error( + "Spend tracking - tool usage flush failed; %s tool usage transactions dropped: %s", + len(tool_usage_to_process), + tool_tracking_err, + ) async with prisma_client._model_usage_transactions_lock: model_usage_to_process: Final = prisma_client.model_usage_transactions diff --git a/tests/unit/proxy/db/test_db_url_settings.py b/tests/unit/proxy/db/test_db_url_settings.py index f5fb1bda0c1..20b95575965 100644 --- a/tests/unit/proxy/db/test_db_url_settings.py +++ b/tests/unit/proxy/db/test_db_url_settings.py @@ -888,6 +888,17 @@ def test_token_refresh_params_keep_the_prisma_tls_dialect_but_not_the_schema(): } +def test_token_refresh_params_keep_the_options_carrying_the_server_timeouts() -> None: + kept: Final = token_refresh_params_from_url( + "postgresql://u:TOKEN@db.example.com:5432/litellm_db" + "?schema=tenant&connection_limit=5&options=-c%20statement_timeout%3D5000%20-c%20idle_in_transaction_session_timeout%3D60000" + ) + assert dict(kept) == { + "connection_limit": "5", + "options": "-c statement_timeout=5000 -c idle_in_transaction_session_timeout=60000", + } + + def _issue_cert( subject: str, issuer: x509.Certificate | None, issuer_key: ec.EllipticCurvePrivateKey | None, ca: bool ) -> tuple[x509.Certificate, ec.EllipticCurvePrivateKey]: @@ -1199,3 +1210,42 @@ def test_reader_keeps_its_own_pinned_idle_lifetime(monkeypatch): assert os.environ["DATABASE_URL_READ_REPLICA"] == ( "postgresql://u:p@reader.example.com:5432/db?max_idle_connection_lifetime=120" ) + + +@pytest.mark.parametrize( + "writer_limit, reader_limit, num_workers, expected", + [ + ("10", None, "4", "4 worker(s) x writer connection_limit 10 = up to 40 connections"), + ( + "10", + "10", + "4", + "4 worker(s) x (writer connection_limit 10 + reader connection_limit 10) = up to 80 connections", + ), + ( + "10", + "50", + "4", + "4 worker(s) x (writer connection_limit 10 + reader connection_limit 50) = up to 240 connections", + ), + ("10", None, "0", "1 worker(s) x writer connection_limit 10 = up to 10 connections"), + ], + ids=["writer_only", "reader_doubles_engines", "reader_pins_its_own_limit", "worker_floor"], +) +def test_connection_budget_message_sums_the_engines_each_worker_owns( + writer_limit: str, reader_limit: str | None, num_workers: str, expected: str +) -> None: + from litellm.proxy.db.db_url_settings import postgres_connection_budget_message + + message: Final = postgres_connection_budget_message( + writer_limit=writer_limit, reader_limit=reader_limit, num_workers=num_workers + ) + assert expected in message + assert "max_connections minus superuser_reserved_connections" in message + + +def test_connection_budget_message_survives_an_unparseable_limit() -> None: + from litellm.proxy.db.db_url_settings import postgres_connection_budget_message + + message: Final = postgres_connection_budget_message(writer_limit="ten", reader_limit=None, num_workers="2") + assert "writer='ten'" in message diff --git a/tests/unit/proxy/db/test_exception_handler.py b/tests/unit/proxy/db/test_exception_handler.py index c00deea7de5..af647b2a3d9 100644 --- a/tests/unit/proxy/db/test_exception_handler.py +++ b/tests/unit/proxy/db/test_exception_handler.py @@ -8,7 +8,7 @@ import httpx import pytest from fastapi import HTTPException, Request from prisma import errors as prisma_errors -from prisma.engine.errors import BinaryNotFoundError, EngineConnectionError +from prisma.engine.errors import BinaryNotFoundError, EngineConnectionError, EngineRequestError from prisma.errors import ( ClientNotConnectedError, DataError, @@ -861,3 +861,87 @@ def test_postgres_connection_capacity_refusal_is_service_unavailable_not_a_data_ ) def test_is_database_capacity_error_excludes_other_failures(error: Exception) -> None: assert PrismaDBExceptionHandler.is_database_capacity_error(error) is False + + +def _capacity_error(message: str) -> DataError: + """The shape prisma-client-py raises when Postgres refuses the session: the + connector message with no SQLSTATE in ``meta`` and no P-code, so it falls + through ``handle_response_errors`` to the base ``DataError``.""" + return DataError( + data={ + "error": f"Error occurred during query execution:\nConnectorError(ConnectorError {{ user_facing_error: None, kind: QueryError({message}) }})", + "user_facing_error": { + "is_panic": False, + "message": f"Error in connector: Error querying the database: FATAL: {message}", + "backtrace": None, + }, + } + ) + + +def _pool_timeout_error() -> DataError: + return DataError( + data={ + "error": "Error in connector: Error creating a database connection. (Timed out fetching a connection from the pool (connection limit: 10, in use: 10, pool timeout 60))", + "user_facing_error": { + "is_panic": False, + "message": "Timed out fetching a new connection from the connection pool. More info: http://pris.ly/d/connection-pool (Current connection pool timeout: 60, connection limit: 10)", + "meta": {"connection_limit": 10, "timeout": 60}, + "error_code": "P2024", + }, + } + ) + + +@pytest.mark.parametrize( + "error", + [ + _capacity_error("sorry, too many clients already"), + _pool_timeout_error(), + EngineRequestError( + MagicMock(status=500), + '{"is_panic":false,"message":"Error in connector: Error querying the database: FATAL: sorry, too many clients already","backtrace":null}', + ), + RawQueryError( + data={ + "user_facing_error": { + "message": 'Raw query failed. Code: `53300`. Message: `db error: FATAL: sorry, too many clients already`', + "meta": {"code": "53300", "message": "FATAL: sorry, too many clients already"}, + "error_code": "P2010", + } + } + ), + ], + ids=["53300_connector_dataerror", "P2024_pool_timeout", "engine_500", "raw_query_53300"], +) +def test_is_database_capacity_error_recognises_postgres_and_pool_exhaustion(error: Exception) -> None: + """SQLSTATE 53300 and prisma P2024 mean the statement was never sent because + the deployment is over its connection budget. Both are infrastructure + failures (503, never 401), neither is a reason to recreate the engine (that + opens another pool against a server that is already full), and both must + reach the spend-log writer as "requeue", not "bisect".""" + assert PrismaDBExceptionHandler.is_database_capacity_error(error) is True + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is True + assert PrismaDBExceptionHandler.is_database_connection_error(error) is True + assert PrismaDBExceptionHandler.is_database_transport_error(error) is False + assert PrismaDBExceptionHandler.is_permanent_database_fault(error) is False + + +@pytest.mark.parametrize( + "error", + [ + PrismaError("timed out while connecting"), + PrismaError("can't reach database server"), + DataError(data={"user_facing_error": {"message": "Can't reach database server at `127.0.0.1`:`5499`"}}), + httpx.ConnectError("conn refused"), + ], +) +def test_capacity_check_leaves_reachability_failures_on_the_reconnect_path(error: Exception) -> None: + assert PrismaDBExceptionHandler.is_database_capacity_error(error) is False + assert PrismaDBExceptionHandler.is_database_transport_error(error) is True + + +def test_capacity_error_is_seen_through_a_wrapping_exception() -> None: + wrapped: Final = RuntimeError("spend flush failed") + wrapped.__cause__ = _capacity_error("sorry, too many clients already") + assert PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(wrapped) is True diff --git a/tests/unit/proxy/test_proxy_cli.py b/tests/unit/proxy/test_proxy_cli.py index 47071827f2c..51fbaf2ee9f 100644 --- a/tests/unit/proxy/test_proxy_cli.py +++ b/tests/unit/proxy/test_proxy_cli.py @@ -3,6 +3,7 @@ import os from contextlib import nullcontext from pathlib import Path from types import SimpleNamespace +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import click @@ -2964,6 +2965,32 @@ class TestPostgresStatementTimeoutOptions: assert _pg_options_with_timeouts(existing, statement_timeout, lock_timeout) == expected + @pytest.mark.parametrize( + "existing, idle_timeout, expected", + [ + ("", 30, "-c statement_timeout=60000 -c lock_timeout=15000 -c idle_in_transaction_session_timeout=30000"), + ("", None, "-c statement_timeout=60000 -c lock_timeout=15000"), + ( + "-c idle_in_transaction_session_timeout=5000", + 30, + "-c idle_in_transaction_session_timeout=5000 -c statement_timeout=60000 -c lock_timeout=15000", + ), + ], + ids=["idle_set", "idle_unset", "pinned_idle_wins"], + ) + def test_pg_options_with_idle_in_transaction_timeout( + self, + existing: str, + idle_timeout: int | None, + expected: str, + ) -> None: + """A transaction that opened and then stalled holds its connection and its + locks for as long as the client stays silent; ``idle_in_transaction_session_timeout`` + is the only server-side bound on that, so it rides the same ``options`` string.""" + from litellm.proxy.proxy_cli import _pg_options_with_timeouts + + assert _pg_options_with_timeouts(existing, 60, 15, idle_timeout) == expected + def test_timeouts_reach_the_database_url_from_general_settings(self, tmp_path): """The whole point of the setting: it has to land on DATABASE_URL.""" import yaml @@ -2976,6 +3003,7 @@ class TestPostgresStatementTimeoutOptions: "general_settings": { "database_statement_timeout": 60, "database_lock_timeout": 15, + "database_idle_in_transaction_session_timeout": 30, }, } ) @@ -2986,6 +3014,7 @@ class TestPostgresStatementTimeoutOptions: options = urlparse.parse_qs(urlparse.urlparse(modified_url).query)["options"][0] assert "-c statement_timeout=60000" in options assert "-c lock_timeout=15000" in options + assert "-c idle_in_transaction_session_timeout=30000" in options def test_no_options_param_when_unset(self, tmp_path): """Unset must mean today's behavior, not an empty options string.""" @@ -3066,6 +3095,7 @@ def _run_server_and_capture_urls( database_url: str = "postgresql://t:t@localhost:5432/t", direct_url: str | None = None, read_replica_url: str | None = None, + extra_args: tuple[str, ...] = (), ) -> dict: loaded_config = yaml.safe_load(Path(config_path).read_text()) mock_proxy_config = MagicMock() @@ -3098,7 +3128,7 @@ def _run_server_and_capture_urls( patch("litellm.proxy.db.check_migration.check_prisma_schema_diff"), ): run_server.main( - ["--config", config_path, "--local", "--skip_server_startup"], + ["--config", config_path, "--local", "--skip_server_startup", *extra_args], standalone_mode=False, ) return {k: os.environ[k] for k in _CAPTURED_DB_ENV_VARS if k in os.environ} @@ -3172,6 +3202,47 @@ class TestReadReplicaConnectionParams: assert query["connection_limit"] == ["50"] assert query["pool_timeout"] == ["20"] + def test_connection_budget_line_counts_the_limits_the_final_urls_carry( + self, + tmp_path: Path, + capsys: pytest.CaptureFixture[str], + ) -> None: + import yaml + + config_path: Final = tmp_path / "config.yaml" + config_path.write_text( + yaml.dump({"model_list": [], "general_settings": {"database_connection_pool_limit": 3}}) + ) + + _run_server_and_capture_urls( + str(config_path), + read_replica_url="postgresql://t:t@reader:5432/t?connection_limit=50", + ) + + assert ( + "1 worker(s) x (writer connection_limit 3 + reader connection_limit 50) = up to 53 connections" + in capsys.readouterr().out + ) + + def test_connection_budget_line_counts_one_worker_under_hypercorn( + self, + tmp_path: Path, + capsys: pytest.CaptureFixture[str], + ) -> None: + import yaml + + config_path: Final = tmp_path / "config.yaml" + config_path.write_text( + yaml.dump({"model_list": [], "general_settings": {"database_connection_pool_limit": 3}}) + ) + + _run_server_and_capture_urls( + str(config_path), + extra_args=("--run_hypercorn", "--num_workers", "4"), + ) + + assert "1 worker(s) x writer connection_limit 3 = up to 3 connections" in capsys.readouterr().out + 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 diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py index 7aa06bcafae..c2921dd877e 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -12,7 +12,9 @@ from __future__ import annotations import asyncio import json +from collections.abc import Callable from contextlib import suppress +from datetime import datetime, timezone from typing import Any, Dict, Final, List from unittest.mock import AsyncMock, MagicMock @@ -968,3 +970,119 @@ async def test_monitor_spend_logs_queue_pulls_parked_rows_before_each_flush( ) assert seen == [["parked"]] + + +def _postgres_out_of_connections() -> Exception: + """prisma's shape for Postgres SQLSTATE 53300: a base ``DataError`` whose only + hint is the connector message.""" + from prisma.errors import DataError + + return DataError( + data={ + "user_facing_error": { + "is_panic": False, + "message": "Error in connector: Error querying the database: FATAL: sorry, too many clients already", + "backtrace": None, + } + } + ) + + +@pytest.mark.asyncio +async def test_update_spend_logs_job_requeues_whole_batch_when_postgres_is_out_of_connections( + mock_prisma_client: MagicMock, + make_spend_log_row: Callable[..., dict[str, object]], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """53300 is not a poison row: bisecting it would issue one failing statement + per row (each a fresh connection attempt against a full server) and drop + every row. The batch goes back to the queue head untouched, in one attempt, + without the in-job retry loop hammering the server.""" + sleeps: Final[list[float]] = [] + + async def _no_sleep(seconds: float, *_: object, **__: object) -> None: + sleeps.append(seconds) + + monkeypatch.setattr(asyncio, "sleep", _no_sleep) + proxy_logging: Final = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [ + make_spend_log_row(request_id="r1"), + make_spend_log_row(request_id="r2"), + make_spend_log_row(request_id="r3"), + ] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_postgres_out_of_connections()) + + with pytest.raises(Exception, match="too many clients already"): + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + assert { + "create_many_calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, + "queue_after": [row["request_id"] for row in mock_prisma_client.spend_log_transactions], + "backoff_sleeps": sleeps, + } == {"create_many_calls": 1, "queue_after": ["r1", "r2", "r3"], "backoff_sleeps": []} + + +def _tool_usage_transaction(request_id: str) -> object: + from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction + + return ToolUsageTransaction( + request_id=request_id, + date="2026-10-02", + start_time=datetime(2026, 10, 2, tzinfo=timezone.utc), + tool_names=("get_weather",), + spend=0.01, + total_tokens=12, + ) + + +@pytest.mark.asyncio +async def test_tool_usage_flush_requeues_when_postgres_is_out_of_connections( + mock_prisma_client: MagicMock, +) -> None: + """A tool-usage batch Postgres had no connection for was never sent, so it is + safe to keep; dropping it loses the rollup increments for good. The job stops + there, like the spend-log write, so a drain loop does not re-hit the full server.""" + from prisma.errors import DataError + + mock_prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock(side_effect=_postgres_out_of_connections()) + mock_prisma_client.spend_log_transactions = [] + first, second = _tool_usage_transaction("r1"), _tool_usage_transaction("r2") + mock_prisma_client.tool_usage_transactions = [first, second] + + with pytest.raises(DataError, match="too many clients already"): + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=MagicMock(), + ) + + assert { + "index_writes": mock_prisma_client.db.litellm_spendlogtoolindex.create_many.await_count, + "queue_after": mock_prisma_client.tool_usage_transactions, + } == {"index_writes": 1, "queue_after": [first, second]} + + +@pytest.mark.asyncio +async def test_tool_usage_flush_still_drops_ambiguous_failures( + mock_prisma_client: MagicMock, +) -> None: + """Anything other than a connection refusal may have reached the server, and + the rollup increments are not idempotent, so the batch is not replayed.""" + mock_prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock( + side_effect=RuntimeError("engine returned a malformed payload") + ) + mock_prisma_client.spend_log_transactions = [] + mock_prisma_client.tool_usage_transactions = [_tool_usage_transaction("r1")] + + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=MagicMock(), + ) + + assert mock_prisma_client.tool_usage_transactions == []