diff --git a/litellm/constants.py b/litellm/constants.py index 938a3c85c29..233a0e2e4df 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1807,6 +1807,7 @@ SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL", SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)) SPEND_LOG_WRITE_BATCH_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_BYTES", 2_000_000))) SPEND_LOG_WRITE_BATCH_MAX_ROWS: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_ROWS", "100"))) +SPEND_ROLLUP_LOCK_TIMEOUT_MS: Final = max(1, int(os.getenv("SPEND_ROLLUP_LOCK_TIMEOUT_MS", "5000"))) SPEND_LOG_QUEUE_SIZE_THRESHOLD: Final = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100)) SPEND_LOG_QUEUE_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_QUEUE_MAX_BYTES", "64000000"))) SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0)) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 5fdde118fb9..7bbb0e63202 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -75,6 +75,7 @@ from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( ) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.db.model_usage_rollup import build_model_usage_transaction +from litellm.proxy.db.rollup_lock_timeout import apply_rollup_lock_timeout from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING from litellm.proxy.spend_tracking.compression_savings import ( extract_compression_saved_tokens, @@ -292,6 +293,7 @@ async def _spend_update_tx( ) -> AsyncGenerator[_SpendTransaction]: tx: Final[_SpendTransactionManager] = prisma_client.db.tx(timeout=timedelta(seconds=60)) async with db_span(call_type, table), tx as transaction: + await apply_rollup_lock_timeout(transaction) yield transaction @@ -2032,11 +2034,17 @@ class DBSpendUpdateWriter: start_time: float, proxy_logging_obj: ProxyLogging, ) -> None: - """Retry a failed spend-update transaction on connection errors or deadlocks, else re-raise.""" + """Retry a failed spend-update transaction on connection errors, deadlocks or a rollup + ``lock_timeout`` (55P03), else re-raise. All three roll the transaction back before any + increment applied, so re-sending the same batch cannot double-count.""" from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.utils import _raise_failed_update_spend_exception - is_retryable = isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) or PrismaDBExceptionHandler.is_deadlock_error(e) + is_retryable = ( + isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) + or PrismaDBExceptionHandler.is_deadlock_error(e) + or PrismaDBExceptionHandler.is_lock_timeout_error(e) + ) if not is_retryable or attempt >= n_retry_times: _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) verbose_proxy_logger.warning( diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index c3526548942..379e0012ea8 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -31,6 +31,8 @@ _CONNECTION_CAPACITY_PHRASES: Final = ( "remaining connection slots are reserved", ) _PRISMA_POOL_TIMEOUT_CODE: Final = "P2024" +_LOCK_TIMEOUT_SQLSTATE: Final = "55P03" +_LOCK_TIMEOUT_PHRASE: Final = "canceling statement due to lock timeout" def _exception_chain(e: BaseException) -> Iterator[BaseException]: @@ -266,6 +268,21 @@ class PrismaDBExceptionHandler: error_message: Final = str(e).lower() return any(phrase in error_message for phrase in _CONNECTION_CAPACITY_PHRASES) + @staticmethod + def is_lock_timeout_error(e: Exception) -> bool: + """True iff Postgres cancelled a statement for exceeding ``lock_timeout`` + (SQLSTATE 55P03). The statement never acquired its lock, so it never + applied and the transaction rolled back: the rows it carried are safe to + re-send, and the pooled connection is free again rather than pinned + behind the holder.""" + import prisma + + if not isinstance(e, _exception_types(prisma.errors.PrismaError)): + return False + if PrismaDBExceptionHandler.postgres_sqlstate(e) == _LOCK_TIMEOUT_SQLSTATE: + return True + return _LOCK_TIMEOUT_PHRASE in str(e).lower() + @staticmethod def postgres_sqlstate(e: Exception) -> str | None: """The SQLSTATE Postgres attached to a failed statement, as prisma surfaces it, or None.""" diff --git a/litellm/proxy/db/model_usage_rollup.py b/litellm/proxy/db/model_usage_rollup.py index 98d53573f30..d83500d66e8 100644 --- a/litellm/proxy/db/model_usage_rollup.py +++ b/litellm/proxy/db/model_usage_rollup.py @@ -9,6 +9,7 @@ from itertools import groupby from typing import TYPE_CHECKING, Final, Protocol from pydantic import TypeAdapter, ValidationError +from typing_extensions import LiteralString from litellm.constants import ( INTERNAL_CALL_ORIGIN_METADATA_KEY, @@ -17,6 +18,7 @@ from litellm.constants import ( ) from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES, SpendLogsPayload from litellm.proxy.db.model_insights_tasks import load_model_insight_tasks +from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient @@ -32,6 +34,8 @@ class _UpsertTable(Protocol): class _ModelUsageBatch(Protocol): litellm_dailymodelusage: _UpsertTable + def execute_raw(self, query: LiteralString, *args: object) -> None: ... + class _ModelUsageBatchManager(Protocol): async def __aenter__(self) -> _ModelUsageBatch: ... @@ -132,6 +136,7 @@ async def flush_model_usage_transactions( for attempt in range(n_retry_times + 1): try: async with _model_usage_batch(prisma_client) as batcher: + batcher.execute_raw(ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting()) for key, grouped in groupby(ordered, key=lambda transaction: transaction.key): entries = tuple(grouped) spend = sum(entry.spend for entry in entries) diff --git a/litellm/proxy/db/prisma_query_span.py b/litellm/proxy/db/prisma_query_span.py index 70c8e2ec3f9..d99f4dcc768 100644 --- a/litellm/proxy/db/prisma_query_span.py +++ b/litellm/proxy/db/prisma_query_span.py @@ -39,6 +39,7 @@ _ROOT_FIELD: Final = re.compile(r"result:\s*(\w+)") _RAW_SQL: Final = re.compile(r'query:\s*"((?:[^"\\]|\\.)*)"') _LEADING_KEYWORD: Final = re.compile(r"(?:\\[nrt]|\s|\()*(\w+)") _SETTING: Final = re.compile(r"(?:\\[nrt]|\s)*SET\s+(?:LOCAL\s+|SESSION\s+)?([A-Za-z_.]+)", re.IGNORECASE) +_SET_CONFIG: Final = re.compile(r"(?:\\[nrt]|\s)*SELECT\s+set_config\(\s*'([A-Za-z_.]+)'", re.IGNORECASE) _CATALOG: Final = re.compile(r"\bpg_\w+|\bto_regclass\b|\binformation_schema\b|\bcurrent_setting\s*\(|^\s*SHOW\b") _PROBE: Final = re.compile(r"(?:\\[nrt]|\s)*SELECT\s+\d+\s*;?(?:\\[nrt]|\s)*$", re.IGNORECASE) _CTE_WRITE: Final = re.compile(r"\b(UPDATE|INSERT|DELETE)\s+(?:INTO\s+|FROM\s+)?(?:\\?\")", re.IGNORECASE) @@ -90,9 +91,13 @@ def sql_relation(sql: str) -> str | None: def sql_operation(sql: str) -> tuple[str | None, str | None]: """``(verb, target)`` for a raw statement: the SQL verb from its leading keyword and the - relation it names, or for ``SET`` the setting it changes.""" + relation it names, or for ``SET`` (and its parameterizable twin ``SELECT set_config``) + the setting it changes.""" if _PROBE.match(sql): return "ping", None + set_config: Final = _SET_CONFIG.match(sql) + if set_config is not None: + return "set", set_config.group(1).lower() keyword: Final = _LEADING_KEYWORD.match(sql) leading: Final = keyword.group(1).upper() if keyword is not None else "" cte_write: Final = _CTE_WRITE.search(sql) if leading == "WITH" else None diff --git a/litellm/proxy/db/rollup_lock_timeout.py b/litellm/proxy/db/rollup_lock_timeout.py new file mode 100644 index 00000000000..f16847e9c96 --- /dev/null +++ b/litellm/proxy/db/rollup_lock_timeout.py @@ -0,0 +1,34 @@ +"""The row-lock budget every spend rollup transaction runs under. + +A rollup upsert that waits on a row another pod holds keeps its pooled +connection for as long as it waits. ``lock_timeout`` bounds that wait, so a +long holder costs the fleet one aborted statement per waiter instead of a +connection pinned until the chain drains. The aborted statement never applied, +so the caller requeues its rows (SQLSTATE 55P03, see +``PrismaDBExceptionHandler.is_lock_timeout_error``). + +``set_config(..., true)`` is ``SET LOCAL``: it lasts for the enclosing +transaction only, and takes the value as a bind parameter, which ``SET`` +cannot. Issued as the first statement of every rollup transaction. +""" + +from typing import Final, Protocol + +from typing_extensions import LiteralString + +from litellm.constants import SPEND_ROLLUP_LOCK_TIMEOUT_MS + +ROLLUP_LOCK_TIMEOUT_SQL: Final = "SELECT set_config('lock_timeout', $1::text, true)" + + +def rollup_lock_timeout_setting() -> str: + return f"{SPEND_ROLLUP_LOCK_TIMEOUT_MS}ms" + + +class _RawExecutor(Protocol): + async def execute_raw(self, query: LiteralString, *args: object) -> int: ... + + +async def apply_rollup_lock_timeout(transaction: _RawExecutor) -> None: + """Bound the row-lock waits of the open ``transaction`` for the rest of its life.""" + _ = await transaction.execute_raw(ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting()) diff --git a/litellm/proxy/db/spend_log_tool_index.py b/litellm/proxy/db/spend_log_tool_index.py index b0d8bb9aba1..81338988c66 100644 --- a/litellm/proxy/db/spend_log_tool_index.py +++ b/litellm/proxy/db/spend_log_tool_index.py @@ -15,15 +15,18 @@ from __future__ import annotations import asyncio import random -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timezone from itertools import groupby -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Any, Final, Protocol + +from typing_extensions import LiteralString from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES from litellm.proxy.db.db_span import db_span +from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting from litellm.proxy.db.spend_log_batching import spend_log_write_batches from litellm.repositories.table_repositories import SpendLogToolIndexRepository @@ -99,6 +102,27 @@ def build_tool_usage_transaction( ) +class _ToolRollupTable(Protocol): + def upsert(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... + + +class _ToolRollupBatch(Protocol): + litellm_dailytoolspend: _ToolRollupTable + + def execute_raw(self, query: LiteralString, *args: object) -> None: ... + + +class _ToolRollupBatchManager(Protocol): + async def __aenter__(self) -> _ToolRollupBatch: ... + + async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... + + +def _tool_rollup_batch(prisma_client: PrismaClient) -> _ToolRollupBatchManager: + batch: Final[_ToolRollupBatchManager] = prisma_client.db.batch_() + return batch + + async def flush_tool_usage_transactions( prisma_client: PrismaClient, transactions: Sequence[ToolUsageTransaction], @@ -140,8 +164,9 @@ async def flush_tool_usage_transactions( await index_table.create_many(data=statement_rows, skip_duplicates=True) async with ( db_span("commit_daily_tool_spend", "LiteLLM_DailyToolSpend"), - prisma_client.db.batch_() as batcher, + _tool_rollup_batch(prisma_client) as batcher, ): + batcher.execute_raw(ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting()) for (date_key, tool_name), grouped in groupby(per_tool_day, key=lambda entry: (entry[0], entry[1])): entries = tuple(grouped) spend = sum(entry[2] for entry in entries) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d90ff1581e2..a6e6209298b 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -7645,6 +7645,13 @@ async def update_daily_tag_spend( verbose_proxy_logger.error("Error updating daily tag spend: %s", e) +def _rollup_batch_never_applied(e: Exception) -> bool: + """True when a drained rollup batch provably never reached its rows: Postgres had no + connection for it (53300) or cancelled it under ``lock_timeout`` before it took the row + lock (55P03). Either way re-sending the batch cannot double-count, so it is requeued.""" + return PrismaDBExceptionHandler.is_database_capacity_error(e) or PrismaDBExceptionHandler.is_lock_timeout_error(e) + + async def update_spend_logs_job( prisma_client: PrismaClient, db_writer_client: AsyncHTTPHandler | None, @@ -7710,9 +7717,11 @@ async def _run_spend_logs_job( ) # Tool usage tracking: drain the request-time queue into the tool index and the - # 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. + # LiteLLM_DailyToolSpend rollup. A batch Postgres had no connection for, or whose + # rollup statement was cancelled by lock_timeout, never applied, so it is requeued; + # any other failure is dropped because a replay could double-count the rollup. Only + # a full server (53300) stops the job: a timed-out row belongs to this rollup alone, + # the writers below have their own rows and keep flushing. 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) :] @@ -7724,19 +7733,20 @@ async def _run_spend_logs_job( transactions=tool_usage_to_process, ) except Exception as tool_tracking_err: - if PrismaDBExceptionHandler.is_database_capacity_error(tool_tracking_err): + if _rollup_batch_never_applied(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", + "Spend tracking - database out of connections or rollup row lock timed out during tool usage " + "flush, requeued %s tool usage transactions for the next flush: %s", len(tool_usage_to_process), tool_tracking_err, ) - raise + if PrismaDBExceptionHandler.is_database_capacity_error(tool_tracking_err): + raise else: verbose_proxy_logger.error( "Spend tracking - tool usage flush failed; %s tool usage transactions dropped: %s", @@ -7744,6 +7754,7 @@ async def _run_spend_logs_job( tool_tracking_err, ) + # Model usage rollup: same requeue and stop rules as the tool rollup above async with prisma_client._model_usage_transactions_lock: model_usage_to_process: Final = prisma_client.model_usage_transactions prisma_client.model_usage_transactions = [] @@ -7752,11 +7763,28 @@ async def _run_spend_logs_job( await flush_model_usage_transactions(prisma_client=prisma_client, transactions=model_usage_to_process) except Exception as model_usage_err: - verbose_proxy_logger.error( - "Spend tracking - model usage flush failed; %s model usage transactions dropped: %s", - len(model_usage_to_process), - model_usage_err, - ) + if _rollup_batch_never_applied(model_usage_err): + async with ( + prisma_client._model_usage_transactions_lock # pyright: ignore[reportPrivateUsage] # needed to requeue + ): + prisma_client.model_usage_transactions = [ + *model_usage_to_process, + *prisma_client.model_usage_transactions, + ] + verbose_proxy_logger.warning( + "Spend tracking - database out of connections or rollup row lock timed out during model usage " + "flush, requeued %s model usage transactions for the next flush: %s", + len(model_usage_to_process), + model_usage_err, + ) + if PrismaDBExceptionHandler.is_database_capacity_error(model_usage_err): + raise + else: + verbose_proxy_logger.error( + "Spend tracking - model usage flush failed; %s model usage transactions dropped: %s", + len(model_usage_to_process), + model_usage_err, + ) await flush_baseline_accounting(prisma_client) diff --git a/tests/unit/proxy/db/test_daily_spend_bulk_upsert.py b/tests/unit/proxy/db/test_daily_spend_bulk_upsert.py index 7893fb82281..5bce9137666 100644 --- a/tests/unit/proxy/db/test_daily_spend_bulk_upsert.py +++ b/tests/unit/proxy/db/test_daily_spend_bulk_upsert.py @@ -13,6 +13,7 @@ from litellm.proxy.db.daily_spend_bulk_upsert import ( merge_by_conflict_key, ) from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter +from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL TAG_TABLE = DAILY_SPEND_TABLES["tag"] USER_TABLE = DAILY_SPEND_TABLES["user"] @@ -146,8 +147,12 @@ def test_non_tag_tables_carry_no_request_id_column(): class _RecordingDb: def __init__(self) -> None: self.statements: list[tuple[str, tuple[object, ...]]] = [] + self.session_settings: list[tuple[str, tuple[object, ...]]] = [] async def execute_raw(self, query: str, *args: object) -> int: + if query == ROLLUP_LOCK_TIMEOUT_SQL: + self.session_settings.append((query, args)) + return 0 self.statements.append((query, args)) return len(args) diff --git a/tests/unit/proxy/db/test_db_spend_update_writer.py b/tests/unit/proxy/db/test_db_spend_update_writer.py index fe5e31d00b2..8822ecd5bd0 100644 --- a/tests/unit/proxy/db/test_db_spend_update_writer.py +++ b/tests/unit/proxy/db/test_db_spend_update_writer.py @@ -19,6 +19,7 @@ from redis.exceptions import DataError import litellm from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceTypes +from litellm.constants import SPEND_ROLLUP_LOCK_TIMEOUT_MS from litellm.proxy._types import DailyTagSpendTransaction, Litellm_EntityType, SpendUpdateQueueItem from litellm.proxy.db.db_spend_update_writer import ( _TEAM_ADVISORY_LOCK_SQL, @@ -33,6 +34,7 @@ from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdate from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) +from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL from tests.unit.proxy.db.fake_prisma_engine import engine_call @@ -103,11 +105,21 @@ async def test_update_database_attributes_router_rejected_failure_to_model_group ) with ( - patch("litellm.proxy.proxy_server.disable_spend_logs", True), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam - patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam - patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam - patch("litellm.proxy.proxy_server.llm_router", llm_router), # test-quality-ok: get_llm_router reads this proxy_server module global at call time; no injection seam + patch( + "litellm.proxy.proxy_server.disable_spend_logs", True + ), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam + patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam + patch( + "litellm.proxy.proxy_server.user_api_key_cache", MagicMock() + ), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam + patch( + "litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget" + ), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam + patch( + "litellm.proxy.proxy_server.llm_router", llm_router + ), # test-quality-ok: get_llm_router reads this proxy_server module global at call time; no injection seam ): await db_writer.update_database( token="test-token", @@ -320,7 +332,9 @@ async def test_a_routed_request_reaches_the_auto_router_rollup_whether_or_not_sp } with ( - patch("litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam + patch( + "litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs + ), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam patch("litellm.proxy.proxy_server.prisma_client", prisma), patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"), patch( @@ -356,9 +370,13 @@ class _RecordingDb: def __init__(self, execute_raw: Callable[[], int] | None = None) -> None: self.statements: list[Statement] = [] + self.session_settings: list[Statement] = [] self._execute_raw = execute_raw async def execute_raw(self, query: str, *args: object) -> int: + if query == ROLLUP_LOCK_TIMEOUT_SQL: + self.session_settings.append((query, args)) + return 0 self.statements.append((query, args)) if self._execute_raw is not None: return self._execute_raw() @@ -752,7 +770,7 @@ async def test_update_tag_db_with_valid_tags(): """ Test that _update_tag_db correctly processes valid tags and adds them to the spend update queue. """ - from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem + from litellm.proxy._types import Litellm_EntityType writer = DBSpendUpdateWriter() mock_prisma = MagicMock() @@ -1048,7 +1066,7 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster ), ) - lock_call, spend_call = mock_transaction.execute_raw.await_args_list + _lock_timeout_call, lock_call, spend_call = mock_transaction.execute_raw.await_args_list lock_statement, locked_team_id = lock_call.args assert lock_statement is _TEAM_ADVISORY_LOCK_SQL assert locked_team_id == team_id @@ -1092,7 +1110,7 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u ), ) - *lock_calls, spend_call = mock_transaction.execute_raw.await_args_list + _lock_timeout_call, *lock_calls, spend_call = mock_transaction.execute_raw.await_args_list _statement, members = spend_call.args assert [lock_call.args for lock_call in lock_calls] == [ (_TEAM_ADVISORY_LOCK_SQL, "eng"), @@ -1730,7 +1748,7 @@ async def test_endpoint_field_is_correctly_mapped_from_call_type(): for key, transaction in update_dict.items(): # Verify endpoint is included in the key - assert key == f"test-user_2024-01-01_test-key_gpt-4_openai_/chat/completions" + assert key == "test-user_2024-01-01_test-key_gpt-4_openai_/chat/completions" # Verify endpoint is set in the transaction assert transaction["endpoint"] == "/chat/completions" @@ -2912,17 +2930,22 @@ async def test_daily_transaction_carries_compression_saved_tokens(): @pytest.mark.asyncio -@pytest.mark.parametrize("estimate, recorded_savings, expected", [ - pytest.param(None, None, -0.005, id="plain-classifier-cost"), - pytest.param({"version": 1, "status": "unknown"}, None, 0.0, id="unknown"), - pytest.param({"version": 2, "status": "unknown"}, None, 0.0, id="unknown-v2"), - pytest.param({"version": 1, "status": "unknown"}, -0.003, 0.0, id="unknown-stale-value"), - pytest.param({"version": 0, "status": "estimated"}, -0.003, 0.0, id="unsupported-version"), - pytest.param({"version": 1, "status": "estimated"}, -0.003, -0.003, id="estimated"), - pytest.param(None, -0.003, -0.003, id="legacy"), -]) +@pytest.mark.parametrize( + "estimate, recorded_savings, expected", + [ + pytest.param(None, None, -0.005, id="plain-classifier-cost"), + pytest.param({"version": 1, "status": "unknown"}, None, 0.0, id="unknown"), + pytest.param({"version": 2, "status": "unknown"}, None, 0.0, id="unknown-v2"), + pytest.param({"version": 1, "status": "unknown"}, -0.003, 0.0, id="unknown-stale-value"), + pytest.param({"version": 0, "status": "estimated"}, -0.003, 0.0, id="unsupported-version"), + pytest.param({"version": 1, "status": "estimated"}, -0.003, -0.003, id="estimated"), + pytest.param(None, -0.003, -0.003, id="legacy"), + ], +) async def test_daily_transaction_compression_saved_tokens_zero_when_absent( - estimate: dict[str, object] | None, recorded_savings: float | None, expected: float, + estimate: dict[str, object] | None, + recorded_savings: float | None, + expected: float, ) -> None: """Requests without any compression metadata produce a zero count.""" writer = DBSpendUpdateWriter() @@ -2941,12 +2964,14 @@ async def test_daily_transaction_compression_saved_tokens_zero_when_absent( "prompt_tokens": 100, "completion_tokens": 10, "spend": 0.01, - "metadata": json.dumps({ - "usage_object": {"prompt_tokens": 100, "completion_tokens": 10}, - "routing_decision": {"savings_baseline_model": "anthropic/claude-sonnet-5", "classifier_cost": 0.005}, - "autorouter_savings": recorded_savings, - "autorouter_savings_estimate": estimate, - }), + "metadata": json.dumps( + { + "usage_object": {"prompt_tokens": 100, "completion_tokens": 10}, + "routing_decision": {"savings_baseline_model": "anthropic/claude-sonnet-5", "classifier_cost": 0.005}, + "autorouter_savings": recorded_savings, + "autorouter_savings_estimate": estimate, + } + ), } transaction = await writer._common_add_spend_log_transaction_to_daily_transaction( @@ -3417,6 +3442,9 @@ async def test_failed_per_entity_increment_from_redis_restores_only_what_may_sti def batch_(self): return _BatchContext() + async def execute_raw(self, query: str, *args: object) -> int: + return 0 + async def __aenter__(self): return self @@ -3439,9 +3467,7 @@ async def test_failed_per_entity_increment_from_redis_restores_only_what_may_sti ) mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once() - restored = mock_redis_update_buffer.restore_transactions_to_redis.call_args.kwargs[ - "db_spend_update_transactions" - ] + restored = mock_redis_update_buffer.restore_transactions_to_redis.call_args.kwargs["db_spend_update_transactions"] assert restored["user_list_transactions"] is None assert restored["team_list_transactions"] == {"team-1": 1.5} assert restored["key_list_transactions"] == ({"key-1": 1.5} if safe_to_resend else None) @@ -3985,6 +4011,77 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey proxy_logging.failure_handler.assert_not_called() +_SPEND_UPDATE_PATHS: Final = [ + ("key_list_transactions", "sk-abc"), + ("user_list_transactions", "user-1"), + ("team_list_transactions", "team-1"), + ("team_member_list_transactions", "team_id::team-1::user_id::user-1"), + ("org_list_transactions", "org-1"), + ("org_member_list_transactions", "organization_id::org-1::user_id::user-1"), + ("tag_list_transactions", "tag-1"), + ("agent_list_transactions", "agent-1"), +] + + +@pytest.mark.parametrize("transactions_key, sample_key", _SPEND_UPDATE_PATHS) +@pytest.mark.asyncio +async def test_commit_spend_updates_retries_a_rollup_lock_timeout_then_commits( + monkeypatch, transactions_key, sample_key +): + """A spend UPDATE cancelled under ``lock_timeout`` (55P03) never took its row lock, so the + transaction rolled back with nothing applied. Every entity path retries it like a deadlock + and lands the increment exactly once, instead of raising and dropping the drained batch.""" + slept = [] + monkeypatch.setattr( + "litellm.proxy.db.db_spend_update_writer.asyncio.sleep", + AsyncMock(side_effect=lambda s: slept.append(s)), + ) + + mock_batcher = MagicMock() + mock_prisma_client = MagicMock() + mock_prisma_client.db.tx = MagicMock(side_effect=[_failing_tx(_lock_timeout_error()), _good_tx(mock_batcher)]) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + proxy_logging.call_details = {} + + await DBSpendUpdateWriter()._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=3, + proxy_logging_obj=proxy_logging, + db_spend_update_transactions=_empty_spend_transactions(**{transactions_key: {sample_key: 0.5}}), + ) + + assert { + "transactions_opened": mock_prisma_client.db.tx.call_count, + "backoff_sleeps": len(slept), + "failure_handler_calls": proxy_logging.failure_handler.await_count, + } == {"transactions_opened": 2, "backoff_sleeps": 1, "failure_handler_calls": 0} + + +@pytest.mark.asyncio +async def test_commit_spend_updates_raises_after_exhausting_lock_timeout_retries(monkeypatch): + """A row that stays locked past every retry still surfaces as a failure rather than + looping forever or being swallowed; the budget is the same one deadlocks get.""" + monkeypatch.setattr("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", AsyncMock(return_value=None)) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.tx = MagicMock(side_effect=lambda *a, **k: _failing_tx(_lock_timeout_error())) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + with pytest.raises(PrismaDataError, match="lock timeout"): + await DBSpendUpdateWriter()._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=2, + proxy_logging_obj=proxy_logging, + db_spend_update_transactions=_empty_spend_transactions(key_list_transactions={"sk-abc": 0.5}), + ) + + assert mock_prisma_client.db.tx.call_count == 3 + + @pytest.mark.asyncio @pytest.mark.parametrize( "call_type, expects_flush", @@ -4881,3 +4978,92 @@ async def test_shutdown_drain_that_lands_before_the_interrupted_tag_commit_resol assert redis_buffer.restored == [drained], "a tag batch whose COMMIT came back failed must be restored to Redis" (upsert,) = _daily_upserts(final_db, "LiteLLM_DailyTagSpend") assert _row_values(upsert, "api_requests") == [1] + + +def _lock_timeout_error() -> PrismaDataError: + return PrismaDataError( + data={ + "user_facing_error": { + "is_panic": False, + "message": "Error querying the database: canceling statement due to lock timeout", + "meta": {"code": "55P03", "message": "canceling statement due to lock timeout"}, + } + } + ) + + +@pytest.mark.asyncio +async def test_daily_spend_upsert_runs_under_the_rollup_lock_timeout(): + """The bulk daily upsert opens its transaction by bounding how long it may wait on + rows another pod holds; the setting is transaction-local and bound as a parameter.""" + prisma_client = _RecordingPrisma() + daily_spend_transactions = { + "key": { + "user_id": "user-1", + "date": "2024-01-01", + "api_key": "test-api-key", + "model": "gpt-4", + "custom_llm_provider": "openai", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 0.1, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + } + } + + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, + prisma_client=prisma_client, + proxy_logging_obj=MagicMock(), + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + ) + + assert prisma_client.db.session_settings == [(ROLLUP_LOCK_TIMEOUT_SQL, (f"{SPEND_ROLLUP_LOCK_TIMEOUT_MS}ms",))] + assert len(prisma_client.db.statements) == 1 + assert daily_spend_transactions == {} + + +@pytest.mark.asyncio +async def test_daily_spend_rows_survive_a_lock_timeout_for_the_next_flush(): + """A statement cancelled under ``lock_timeout`` (55P03) never applied, so its rows + are not data Postgres refused: they stay queued for the next flush instead of being + dropped the way a constraint violation is.""" + + def cancelled_waiting_for_the_row() -> int: + raise _lock_timeout_error() + + prisma_client = _RecordingPrisma(execute_raw=cancelled_waiting_for_the_row) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + daily_spend_transactions = { + "key": { + "user_id": "user-1", + "date": "2024-01-01", + "api_key": "test-api-key", + "model": "gpt-4", + "custom_llm_provider": "openai", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 0.1, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + } + } + + with pytest.raises(PrismaDataError, match="lock timeout"): + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging, + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + ) + + assert len(prisma_client.db.statements) == 1 + assert list(daily_spend_transactions) == ["key"] diff --git a/tests/unit/proxy/db/test_exception_handler.py b/tests/unit/proxy/db/test_exception_handler.py index af647b2a3d9..75346ca08a6 100644 --- a/tests/unit/proxy/db/test_exception_handler.py +++ b/tests/unit/proxy/db/test_exception_handler.py @@ -1,12 +1,10 @@ import asyncio -import json import sys from typing import Final from unittest.mock import MagicMock, patch import httpx import pytest -from fastapi import HTTPException, Request from prisma import errors as prisma_errors from prisma.engine.errors import BinaryNotFoundError, EngineConnectionError, EngineRequestError from prisma.errors import ( @@ -22,9 +20,7 @@ from prisma.errors import ( UniqueViolationError, ) - import litellm -from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ProxyErrorTypes, ProxyException from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler @@ -55,22 +51,12 @@ def test_is_database_infrastructure_error_prisma_connection_errors(prisma_error) PrismaError(), PrismaError("validation failed on query"), DataError(data={"user_facing_error": {"meta": {"table": "test_table"}}}), - UniqueViolationError( - data={"user_facing_error": {"meta": {"table": "test_table"}}} - ), - ForeignKeyViolationError( - data={"user_facing_error": {"meta": {"table": "test_table"}}} - ), - MissingRequiredValueError( - data={"user_facing_error": {"meta": {"table": "test_table"}}} - ), + UniqueViolationError(data={"user_facing_error": {"meta": {"table": "test_table"}}}), + ForeignKeyViolationError(data={"user_facing_error": {"meta": {"table": "test_table"}}}), + MissingRequiredValueError(data={"user_facing_error": {"meta": {"table": "test_table"}}}), RawQueryError(data={"user_facing_error": {"meta": {"table": "test_table"}}}), - TableNotFoundError( - data={"user_facing_error": {"meta": {"table": "test_table"}}} - ), - RecordNotFoundError( - data={"user_facing_error": {"meta": {"table": "test_table"}}} - ), + TableNotFoundError(data={"user_facing_error": {"meta": {"table": "test_table"}}}), + RecordNotFoundError(data={"user_facing_error": {"meta": {"table": "test_table"}}}), ], ) def test_is_database_transport_error_non_connection_prisma_errors(prisma_error): @@ -82,12 +68,7 @@ def test_is_database_connection_generic_errors(): """ Test non-Prisma error cases for database connection checking """ - assert ( - PrismaDBExceptionHandler.is_database_connection_error( - Exception("Regular error") - ) - == False - ) + assert PrismaDBExceptionHandler.is_database_connection_error(Exception("Regular error")) == False # Test with ProxyException (DB connection) db_proxy_exception = ProxyException( @@ -95,17 +76,11 @@ def test_is_database_connection_generic_errors(): type=ProxyErrorTypes.no_db_connection, param="test-param", ) - assert ( - PrismaDBExceptionHandler.is_database_connection_error(db_proxy_exception) - == True - ) + assert PrismaDBExceptionHandler.is_database_connection_error(db_proxy_exception) == True # Test with non-DB error regular_exception = Exception("Regular error") - assert ( - PrismaDBExceptionHandler.is_database_connection_error(regular_exception) - == False - ) + assert PrismaDBExceptionHandler.is_database_connection_error(regular_exception) == False @pytest.mark.parametrize( @@ -143,12 +118,7 @@ def test_is_database_service_unavailable_error_prisma_p1001_masquerades_as_datae } } ) - assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error( - p1001_as_dataerror - ) - is True - ) + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(p1001_as_dataerror) is True def test_is_prisma_data_error_only_true_for_dataerror(): @@ -196,12 +166,7 @@ def test_is_database_service_unavailable_error_cached_plan_escapes_as_503(): } } ) - assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error( - cached_plan_error - ) - is True - ) + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(cached_plan_error) is True def test_is_database_service_unavailable_error_prisma_engine_malformed_payload(): @@ -229,10 +194,7 @@ def test_is_database_service_unavailable_error_prisma_engine_malformed_payload() prisma_engine_utils.handle_response_errors(None, malformed_payload) assert "no attribute 'get'" in str(exc_info.value) - assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) - is True - ) + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) is True def test_is_prisma_engine_internal_error_excludes_application_attributeerror(): @@ -247,23 +209,15 @@ def test_is_prisma_engine_internal_error_excludes_application_attributeerror(): with pytest.raises(AttributeError) as exc_info: application_bug() - assert ( - PrismaDBExceptionHandler.is_prisma_engine_internal_error(exc_info.value) - is False - ) - assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) - is False - ) + assert PrismaDBExceptionHandler.is_prisma_engine_internal_error(exc_info.value) is False + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) is False def test_is_prisma_engine_internal_error_excludes_data_layer_prisma_error(): """A data-layer ``PrismaError`` (the DB IS reachable and rejected the data) must stay 401. These are always raised from prisma internals, so the check excludes any ``PrismaError`` by type before inspecting the traceback.""" - data_layer_error = UniqueViolationError( - data={"user_facing_error": {"meta": {"table": "t"}}} - ) + data_layer_error = UniqueViolationError(data={"user_facing_error": {"meta": {"table": "t"}}}) with pytest.raises(UniqueViolationError) as exc_info: raise data_layer_error e = exc_info.value @@ -284,9 +238,7 @@ def test_is_database_service_unavailable_error_excludes_non_infra(error): """Data-layer errors (the DB IS reachable and answered) and generic non-DB errors must NOT be classified as service-unavailable, otherwise a genuine 401 would be masked as a transient 503.""" - assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is False - ) + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is False def _wrapped_like_get_user_object(original): @@ -393,22 +345,14 @@ def test_is_database_service_unavailable_error_asyncpg(monkeypatch): monkeypatch.setitem(sys.modules, "asyncpg.exceptions", fake_exceptions) assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error( - PostgresConnectionError("connection reset") - ) + PrismaDBExceptionHandler.is_database_service_unavailable_error(PostgresConnectionError("connection reset")) is True ) assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error( - InterfaceError("connection was closed") - ) - is True + PrismaDBExceptionHandler.is_database_service_unavailable_error(InterfaceError("connection was closed")) is True ) assert ( - PrismaDBExceptionHandler.is_database_service_unavailable_error( - UniqueViolationError("duplicate key") - ) - is False + PrismaDBExceptionHandler.is_database_service_unavailable_error(UniqueViolationError("duplicate key")) is False ) @@ -675,7 +619,10 @@ def test_is_deadlock_error_excludes_non_deadlocks(error): "22021", ), (RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"message": "m"}}}), None), - (RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"code": 42, "message": "m"}}}), None), + ( + RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"code": 42, "message": "m"}}}), + None, + ), (prisma_errors.DataError(data={"user_facing_error": {"meta": None}}), None), ( prisma_errors.DataError( @@ -695,7 +642,9 @@ def test_is_deadlock_error_excludes_non_deadlocks(error): (httpx.ReadTimeout("no reply"), None), ], ) -def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement(error: Exception, sqlstate: str | None): +def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement( + error: Exception, sqlstate: str | None +): """Only a prisma data error carrying Postgres's own error code yields a SQLSTATE, whether in ``meta`` or, for a batched statement, only in the message; a codeless or malformed payload, an engine-level error, and a transport error yield None.""" @@ -905,7 +854,7 @@ def _pool_timeout_error() -> DataError: RawQueryError( data={ "user_facing_error": { - "message": 'Raw query failed. Code: `53300`. Message: `db error: FATAL: sorry, too many clients already`', + "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", } @@ -945,3 +894,60 @@ 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 + + +def _lock_timeout_error( + message: str = "canceling statement due to lock timeout", code: str | None = "55P03" +) -> DataError: + """The shape prisma raises for a statement Postgres cancelled under ``lock_timeout``.""" + user_facing: Final[dict[str, object]] = {"is_panic": False, "message": f"Error querying the database: {message}"} + if code is not None: + user_facing["meta"] = {"code": code, "message": message} + return DataError(data={"user_facing_error": user_facing}) + + +@pytest.mark.parametrize( + "error", + [ + _lock_timeout_error(), + _lock_timeout_error(message="Abbruch der Anweisung wegen Zeitüberschreitung beim Warten auf eine Sperre"), + _lock_timeout_error(code=None), + DataError( + data={ + "user_facing_error": { + "message": 'Error occurred during query execution: ConnectorError(ConnectorError { user_facing_error: None, kind: QueryError(PostgresError { code: "55P03", message: "canceling statement due to lock timeout", severity: "ERROR" }) })' + } + } + ), + ], + ids=["meta_55P03", "meta_55P03_localised_message", "message_only", "batched_statement"], +) +def test_is_lock_timeout_error_recognises_a_statement_cancelled_under_lock_timeout(error: DataError) -> None: + """SQLSTATE 55P03 means the statement never took its lock, so it never applied + and its rows are safe to re-send. It is a wait budget, not a full server, so + it stays out of the connection-capacity classification.""" + assert PrismaDBExceptionHandler.is_lock_timeout_error(error) is True + assert PrismaDBExceptionHandler.is_database_capacity_error(error) is False + + +@pytest.mark.parametrize( + "error", + [ + _lock_timeout_error(message="canceling statement due to statement timeout", code="57014"), + _lock_timeout_error(message="deadlock detected", code="40P01"), + _capacity_error("sorry, too many clients already"), + _pool_timeout_error(), + httpx.ConnectError("connection refused"), + RuntimeError("canceling statement due to lock timeout"), + ], + ids=[ + "statement_timeout_57014", + "deadlock_40P01", + "capacity_53300", + "pool_timeout_P2024", + "transport", + "not_prisma", + ], +) +def test_is_lock_timeout_error_excludes_other_failures(error: Exception) -> None: + assert PrismaDBExceptionHandler.is_lock_timeout_error(error) is False diff --git a/tests/unit/proxy/db/test_model_access_group_spend.py b/tests/unit/proxy/db/test_model_access_group_spend.py index d2d079bb0e4..0ee86e783ea 100644 --- a/tests/unit/proxy/db/test_model_access_group_spend.py +++ b/tests/unit/proxy/db/test_model_access_group_spend.py @@ -64,6 +64,9 @@ class _FakeTransaction: def batch_(self) -> _FakeBatchManager: return _FakeBatchManager(self._batcher) + async def execute_raw(self, query: str, *args: object) -> int: + return 0 + async def __aenter__(self) -> "_FakeTransaction": return self diff --git a/tests/unit/proxy/db/test_model_usage_rollup.py b/tests/unit/proxy/db/test_model_usage_rollup.py index b9b806f4c54..0b7427d71f2 100644 --- a/tests/unit/proxy/db/test_model_usage_rollup.py +++ b/tests/unit/proxy/db/test_model_usage_rollup.py @@ -6,6 +6,7 @@ from unittest.mock import MagicMock import httpx import pytest +from litellm.proxy.db import rollup_lock_timeout from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.model_usage_rollup import ( ModelUsageKey, @@ -14,11 +15,13 @@ from litellm.proxy.db.model_usage_rollup import ( flush_model_usage_transactions, model_usage_task_type, ) +from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL class _FakeBatcher: def __init__(self) -> None: self.litellm_dailymodelusage = MagicMock() + self.execute_raw = MagicMock() async def __aenter__(self) -> "_FakeBatcher": return self @@ -201,3 +204,26 @@ async def test_request_time_path_queues_usage_instead_of_writing_to_the_db() -> assert [transaction.key.model for transaction in prisma.model_usage_transactions] == ["openai/gpt-5.4-mini"] prisma.db.litellm_dailymodelusage.upsert.assert_not_called() + + +@pytest.mark.asyncio +async def test_rollup_transaction_sets_the_lock_timeout_before_its_first_upsert(monkeypatch) -> None: + """The DailyModelUsage rollup shares the DailyToolSpend write shape, so it runs + under the same per-transaction lock budget, issued before any upsert queues.""" + monkeypatch.setattr(rollup_lock_timeout, "SPEND_ROLLUP_LOCK_TIMEOUT_MS", 250) + batcher = _FakeBatcher() + prisma = _prisma(MagicMock(return_value=batcher)) + order: list[str] = [] + batcher.execute_raw.side_effect = lambda *_: order.append("lock_timeout") + batcher.litellm_dailymodelusage.upsert.side_effect = lambda **_: order.append("upsert") + + await flush_model_usage_transactions( + prisma_client=prisma, + transactions=[ + build_model_usage_transaction(_payload()), + build_model_usage_transaction(_payload(model="other")), + ], + ) + + assert batcher.execute_raw.call_args_list == [((ROLLUP_LOCK_TIMEOUT_SQL, "250ms"),)] + assert order == ["lock_timeout", "upsert", "upsert"] diff --git a/tests/unit/proxy/db/test_prisma_query_span.py b/tests/unit/proxy/db/test_prisma_query_span.py index 674679f8a18..abd95dec517 100644 --- a/tests/unit/proxy/db/test_prisma_query_span.py +++ b/tests/unit/proxy/db/test_prisma_query_span.py @@ -106,6 +106,7 @@ def test_a_payload_the_parser_does_not_know_stays_the_legacy_function_named_span ), ('INSERT INTO "LiteLLM_DailyUserSpend" (id) VALUES ($1)', ("insert", "LiteLLM_DailyUserSpend")), ("SET LOCAL lock_timeout = 1000", ("set", "lock_timeout")), + ("SELECT set_config('lock_timeout', $1::text, true)", ("set", "lock_timeout")), ("SELECT COUNT(*) FROM pg_stat_activity", ("select", "pg_catalog")), ("SELECT 1", ("ping", None)), ("SELECT current_setting('transaction_read_only') AS transaction_read_only", ("select", "pg_catalog")), diff --git a/tests/unit/proxy/db/test_spend_log_tool_index.py b/tests/unit/proxy/db/test_spend_log_tool_index.py index 610faebe17c..eaa94c0438f 100644 --- a/tests/unit/proxy/db/test_spend_log_tool_index.py +++ b/tests/unit/proxy/db/test_spend_log_tool_index.py @@ -4,22 +4,23 @@ only) and the flush that writes LiteLLM_SpendLogToolIndex in bounded statements plus the LiteLLM_DailyToolSpend rollup in one transaction. """ -from types import SimpleNamespace from collections.abc import Awaitable, Callable +from types import SimpleNamespace from typing import Any from unittest.mock import AsyncMock, MagicMock import httpx import pytest -from litellm.proxy.db import spend_log_tool_index +from litellm.proxy.db import rollup_lock_timeout, spend_log_tool_index +from litellm.proxy.db.log_db_metrics import record_db_io +from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL from litellm.proxy.db.spend_log_tool_index import ( ToolUsageTransaction, build_tool_usage_transaction, flush_tool_usage_transactions, response_tool_call_names, ) -from litellm.proxy.db.log_db_metrics import record_db_io from tests.unit.proxy.db.fake_prisma_engine import engine_call @@ -32,6 +33,7 @@ class _FakeBatcher: def __init__(self) -> None: self.litellm_spendlogtoolindex = MagicMock() self.litellm_dailytoolspend = MagicMock() + self.execute_raw = MagicMock() async def __aenter__(self) -> "_FakeBatcher": return self @@ -84,7 +86,9 @@ class TestBuildToolUsageTransaction: mcp_namespaced_tool_name=None, spend=0.5, total_tokens=100, - completion_response=SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(tool_calls=None))]), + completion_response=SimpleNamespace( + choices=[SimpleNamespace(message=SimpleNamespace(tool_calls=None))] + ), ) is None ) @@ -395,3 +399,25 @@ async def test_a_tool_usage_flush_renders_one_postgres_span_per_table_written( "postgres.insert LiteLLM_SpendLogToolIndex", "postgres.upsert LiteLLM_DailyToolSpend", ) + + +@pytest.mark.asyncio +async def test_rollup_transaction_sets_the_lock_timeout_before_its_first_upsert(monkeypatch) -> None: + """Every DailyToolSpend rollup transaction opens by bounding how long its upserts + may wait on a row another pod holds, so a waiter costs one cancelled statement + rather than a pooled connection pinned until the holder commits. The budget is the + SPEND_ROLLUP_LOCK_TIMEOUT_MS knob, bound as a parameter, not baked into the SQL.""" + monkeypatch.setattr(rollup_lock_timeout, "SPEND_ROLLUP_LOCK_TIMEOUT_MS", 250) + prisma, batcher = _prisma_with_batcher() + order: list[str] = [] + batcher.execute_raw.side_effect = lambda *_: order.append("lock_timeout") + batcher.litellm_dailytoolspend.upsert.side_effect = lambda **_: order.append("upsert") + + await flush_tool_usage_transactions( + prisma_client=prisma, + transactions=[_transaction("r1", tool_names=("tool_a", "tool_b"))], + ) + + assert batcher.execute_raw.call_args_list == [((ROLLUP_LOCK_TIMEOUT_SQL, "250ms"),)] + assert order == ["lock_timeout", "upsert", "upsert"] + assert "set_config('lock_timeout'" in ROLLUP_LOCK_TIMEOUT_SQL and ", true)" in ROLLUP_LOCK_TIMEOUT_SQL diff --git a/tests/unit/proxy/db/test_update_daily_tag_spend.py b/tests/unit/proxy/db/test_update_daily_tag_spend.py index 530c0d16767..dda41ab4543 100644 --- a/tests/unit/proxy/db/test_update_daily_tag_spend.py +++ b/tests/unit/proxy/db/test_update_daily_tag_spend.py @@ -1,12 +1,12 @@ -from typing import Dict from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest -from litellm.proxy.utils import update_daily_tag_spend from litellm.proxy._types import DailyTagSpendTransaction -import httpx from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter +from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL +from litellm.proxy.utils import update_daily_tag_spend @pytest.mark.asyncio @@ -18,9 +18,7 @@ async def test_update_daily_tag_spend_delegates_to_tag_commit_writer(): proxy_logging_obj.db_spend_update_writer = MagicMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock() - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = ( - AsyncMock() - ) + proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() await update_daily_tag_spend( prisma_client, @@ -43,12 +41,8 @@ async def test_update_daily_tag_spend_logs_error_and_does_not_raise(): redis_update_buffer._should_commit_spend_updates_to_redis.return_value = False proxy_logging_obj.db_spend_update_writer = MagicMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock( - side_effect=ValueError("boom") - ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = ( - AsyncMock() - ) + proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock(side_effect=ValueError("boom")) + proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() with patch("litellm.proxy.utils.verbose_proxy_logger.error") as error_logger: await update_daily_tag_spend( @@ -69,9 +63,7 @@ async def test_update_daily_tag_spend_uses_redis_writer_when_enabled(): proxy_logging_obj.db_spend_update_writer = MagicMock() proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = ( - AsyncMock() - ) + proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() await update_daily_tag_spend( prisma_client, @@ -92,17 +84,20 @@ async def test_daily_tag_spend_retries_then_succeeds(): proxy_logging_obj = MagicMock() # Fail the upsert 3 times with retryable DB errors, then succeed. - prisma_client.db.execute_raw = AsyncMock( - side_effect=[ - httpx.ConnectError("x"), - httpx.ConnectError("x"), - httpx.ConnectError("x"), - 1, - ] - ) + upsert_outcomes = iter([httpx.ConnectError("x"), httpx.ConnectError("x"), httpx.ConnectError("x"), 1]) + + async def execute_raw(query: str, *args: object) -> int: + if query == ROLLUP_LOCK_TIMEOUT_SQL: + return 0 + outcome = next(upsert_outcomes) + if isinstance(outcome, Exception): + raise outcome + return outcome + + prisma_client.db.execute_raw = AsyncMock(side_effect=execute_raw) prisma_client.db.tx.return_value.__aenter__.return_value.execute_raw = prisma_client.db.execute_raw - daily_spend_transactions: Dict[str, DailyTagSpendTransaction] = { + daily_spend_transactions: dict[str, DailyTagSpendTransaction] = { "k": { "tag": "prod-tag", "date": "2026-04-03", @@ -135,7 +130,8 @@ async def test_daily_tag_spend_retries_then_succeeds(): daily_spend_transactions=daily_spend_transactions, ) - assert prisma_client.db.execute_raw.await_count == 4 + upsert_attempts = [c for c in prisma_client.db.execute_raw.await_args_list if c.args[0] != ROLLUP_LOCK_TIMEOUT_SQL] + assert len(upsert_attempts) == 4 assert sleep_mock.await_count == 3 # The batch is one statement, so the successful attempt is a single call carrying # the row rather than one call per key. diff --git a/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py index 7c8d1346944..40b0cbe11c5 100644 --- a/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py @@ -225,6 +225,9 @@ class _InMemoryBatcher: def __init__(self, table: _InMemoryUsageTable) -> None: self.litellm_dailymodelusage = table + def execute_raw(self, query: str, *args: object) -> None: + return None + async def __aenter__(self) -> "_InMemoryBatcher": return self 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 c2921dd877e..7c4f4e582be 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 @@ -15,7 +15,7 @@ import json from collections.abc import Callable from contextlib import suppress from datetime import datetime, timezone -from typing import Any, Dict, Final, List +from typing import Any, Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -111,15 +111,11 @@ async def test_update_daily_tag_spend_redis_path_when_buffered( writer = MagicMock() proxy_logging.db_spend_update_writer = writer writer.redis_update_buffer = MagicMock() - writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( - return_value=True - ) + writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True) writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() writer._commit_daily_tag_spend_to_db = AsyncMock() - await update_daily_tag_spend( - prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging - ) + await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) redis_kwargs = writer._commit_daily_tag_spend_to_db_with_redis.await_args.kwargs pinned = { "redis_calls": writer._commit_daily_tag_spend_to_db_with_redis.await_count, @@ -130,9 +126,7 @@ async def test_update_daily_tag_spend_redis_path_when_buffered( assert pinned == { "redis_calls": 1, "direct_calls": 0, - "redis_kwargs_keys": sorted( - ["prisma_client", "n_retry_times", "proxy_logging_obj"] - ), + "redis_kwargs_keys": sorted(["prisma_client", "n_retry_times", "proxy_logging_obj"]), "redis_n_retries": 3, } @@ -145,15 +139,11 @@ async def test_update_daily_tag_spend_direct_path_when_no_redis( writer = MagicMock() proxy_logging.db_spend_update_writer = writer writer.redis_update_buffer = MagicMock() - writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( - return_value=False - ) + writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=False) writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() writer._commit_daily_tag_spend_to_db = AsyncMock() - await update_daily_tag_spend( - prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging - ) + await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) assert writer._commit_daily_tag_spend_to_db.await_count == 1 assert writer._commit_daily_tag_spend_to_db_with_redis.await_count == 0 @@ -175,9 +165,7 @@ async def test_update_daily_tag_spend_logs_and_swallows_errors( proxy_logging.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock( side_effect=RuntimeError("commit boom") ) - await update_daily_tag_spend( - prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging - ) + await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) @pytest.mark.asyncio @@ -265,15 +253,11 @@ async def test_update_spend_logs_job_processes_and_clears_queue( mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() # Stub auxiliary imports so the test focuses on the spend-logs write path. - import litellm.proxy.guardrails.usage_tracking as guard_mod import litellm.proxy.db.spend_log_tool_index as tool_mod + import litellm.proxy.guardrails.usage_tracking as guard_mod - monkeypatch.setattr( - guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False - ) - monkeypatch.setattr( - tool_mod, "flush_tool_usage_transactions", AsyncMock(), raising=False - ) + monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False) + monkeypatch.setattr(tool_mod, "flush_tool_usage_transactions", AsyncMock(), raising=False) await update_spend_logs_job( prisma_client=mock_prisma_client, @@ -283,12 +267,10 @@ async def test_update_spend_logs_job_processes_and_clears_queue( pinned = { "create_many_calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, "queue_after": mock_prisma_client.spend_log_transactions, - "first_data_request_id": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[ - "data" - ][0]["request_id"], - "skip_duplicates_set": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[ - "skip_duplicates" + "first_data_request_id": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs["data"][0][ + "request_id" ], + "skip_duplicates_set": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs["skip_duplicates"], } assert pinned == { "create_many_calls": 1, @@ -315,9 +297,7 @@ async def test_update_spend_logs_job_requeues_popped_rows_when_write_cancelled( mock_prisma_client.spend_log_transactions.append(row_arriving_mid_flush) raise asyncio.CancelledError() - mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( - side_effect=_cancel_mid_write - ) + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_cancel_mid_write) with pytest.raises(asyncio.CancelledError): await update_spend_logs_job( @@ -326,9 +306,7 @@ async def test_update_spend_logs_job_requeues_popped_rows_when_write_cancelled( proxy_logging_obj=proxy_logging, ) - assert [ - row["request_id"] for row in mock_prisma_client.spend_log_transactions - ] == ["r1", "r2", "r3"] + assert [row["request_id"] for row in mock_prisma_client.spend_log_transactions] == ["r1", "r2", "r3"] @pytest.mark.asyncio @@ -369,12 +347,8 @@ async def test_drain_spend_logs_queue_flushes_rows_queued_while_draining( import litellm.proxy.db.spend_log_tool_index as tool_mod import litellm.proxy.guardrails.usage_tracking as guard_mod - monkeypatch.setattr( - guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False - ) - monkeypatch.setattr( - tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False - ) + monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False) + monkeypatch.setattr(tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False) proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -385,9 +359,7 @@ async def test_drain_spend_logs_queue_flushes_rows_queued_while_draining( async def _write(*args: Any, **kwargs: Any) -> None: written.extend(row["request_id"] for row in kwargs["data"]) if len(written) == 1: - mock_prisma_client.spend_log_transactions.append( - make_spend_log_row(request_id="r2") - ) + mock_prisma_client.spend_log_transactions.append(make_spend_log_row(request_id="r2")) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write) @@ -408,12 +380,8 @@ async def test_drain_spend_logs_queue_stops_monitor_and_keeps_its_popped_rows( import litellm.proxy.db.spend_log_tool_index as tool_mod import litellm.proxy.guardrails.usage_tracking as guard_mod - monkeypatch.setattr( - guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False - ) - monkeypatch.setattr( - tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False - ) + monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False) + monkeypatch.setattr(tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False) proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -460,12 +428,8 @@ async def test_drain_spend_logs_queue_gives_up_after_max_passes( import litellm.proxy.db.spend_log_tool_index as tool_mod import litellm.proxy.guardrails.usage_tracking as guard_mod - monkeypatch.setattr( - guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False - ) - monkeypatch.setattr( - tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False - ) + monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False) + monkeypatch.setattr(tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False) proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -474,9 +438,7 @@ async def test_drain_spend_logs_queue_gives_up_after_max_passes( async def _write_and_refill(*args: Any, **kwargs: Any) -> None: mock_prisma_client.spend_log_transactions.append(make_spend_log_row()) - mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( - side_effect=_write_and_refill - ) + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write_and_refill) await drain_spend_logs_queue( prisma_client=mock_prisma_client, @@ -484,10 +446,7 @@ async def test_drain_spend_logs_queue_gives_up_after_max_passes( proxy_logging_obj=proxy_logging, ) - assert ( - mock_prisma_client.db.litellm_spendlogs.create_many.await_count - == MAX_SPEND_LOG_DRAIN_ITERATIONS - ) + assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == MAX_SPEND_LOG_DRAIN_ITERATIONS @pytest.mark.asyncio @@ -496,8 +455,8 @@ async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty( make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch, ) -> None: - import litellm.proxy.utils as utils_mod import litellm.constants as constants_mod + import litellm.proxy.utils as utils_mod monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False) monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_SIZE_THRESHOLD", 1, raising=False) @@ -530,8 +489,8 @@ async def test_monitor_spend_logs_queue_swallows_errors_and_backs_off( """An exception inside the loop is logged with backoff and the loop continues running rather than crashing the monitor task. """ - import litellm.proxy.utils as utils_mod import litellm.constants as constants_mod + import litellm.proxy.utils as utils_mod monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False) @@ -721,8 +680,7 @@ def test_raise_failed_update_spend_exception_emits_failure_handler() -> None: else None ), "non_blocking_in_traceback": ( - "Non-Blocking" - in proxy_logging.failure_handler.call_args.kwargs["traceback_str"] + "Non-Blocking" in proxy_logging.failure_handler.call_args.kwargs["traceback_str"] if proxy_logging.failure_handler.call_args else False ), @@ -1086,3 +1044,152 @@ async def test_tool_usage_flush_still_drops_ambiguous_failures( ) assert mock_prisma_client.tool_usage_transactions == [] + + +def _postgres_lock_timeout() -> Exception: + """prisma's shape for Postgres SQLSTATE 55P03: the rollup upsert waited past + ``lock_timeout`` for a row another pod held and was cancelled before it took the lock.""" + from prisma.errors import DataError + + return DataError( + data={ + "user_facing_error": { + "is_panic": False, + "message": "Error querying the database: canceling statement due to lock timeout", + "meta": {"code": "55P03", "message": "canceling statement due to lock timeout"}, + } + } + ) + + +@pytest.mark.asyncio +async def test_tool_usage_flush_requeues_when_the_rollup_row_lock_times_out( + mock_prisma_client: MagicMock, +) -> None: + """A rollup transaction cancelled by ``lock_timeout`` never applied: the waiter + gives its connection back instead of pinning it behind the holder, and its + increments go back to the queue head for the next flush rather than being dropped. + Unlike a full server, one contended row says nothing about the other writers, so the + job carries on: the model usage rollup still opens its own transaction and the + auto-router drain still runs.""" + mock_prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock() + rollup_batch = MagicMock() + rollup_batch.__aenter__ = AsyncMock(return_value=MagicMock()) + rollup_batch.__aexit__ = AsyncMock(side_effect=_postgres_lock_timeout()) + mock_prisma_client.db.batch_ = MagicMock(return_value=rollup_batch) + mock_prisma_client.spend_log_transactions = [] + first, second = _tool_usage_transaction("r1"), _tool_usage_transaction("r2") + mock_prisma_client.tool_usage_transactions = [first, second] + mock_prisma_client.model_usage_transactions = [_model_usage_transaction("gpt-4o")] + mock_prisma_client.autorouter_turn_transactions = [MagicMock(name="autorouter_turn")] + + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=MagicMock(), + ) + + assert { + "rollup_transactions": mock_prisma_client.db.batch_.call_count, + "queue_after": mock_prisma_client.tool_usage_transactions, + "autorouter_queue_after": mock_prisma_client.autorouter_turn_transactions, + } == {"rollup_transactions": 2, "queue_after": [first, second], "autorouter_queue_after": []} + + +def _model_usage_transaction(model: str) -> object: + from litellm.proxy.db.model_usage_rollup import ModelUsageKey, ModelUsageTransaction + + return ModelUsageTransaction( + key=ModelUsageKey( + date="2026-10-03", model_group="gpt", model=model, custom_llm_provider="openai", task_type="chat" + ), + spend=0.001, + prompt_tokens=1, + completion_tokens=1, + successful=True, + ) + + +def _model_usage_batch_failing_with(error: Exception, mock_prisma_client: MagicMock) -> MagicMock: + """Point the model usage rollup's ``batch_()`` at a transaction that fails on commit, with + no tool usage queued so the tool rollup never opens one.""" + mock_prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock() + mock_prisma_client.tool_usage_transactions = [] + mock_prisma_client.spend_log_transactions = [] + rollup_batch = MagicMock() + rollup_batch.__aenter__ = AsyncMock(return_value=MagicMock()) + rollup_batch.__aexit__ = AsyncMock(side_effect=error) + mock_prisma_client.db.batch_ = MagicMock(return_value=rollup_batch) + return rollup_batch + + +@pytest.mark.asyncio +async def test_model_usage_flush_requeues_and_carries_on_when_its_row_lock_times_out( + mock_prisma_client: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The DailyModelUsage rollup shares its rows across every pod exactly like the tool rollup. + A batch cancelled under ``lock_timeout`` never applied, so its transactions go back to the + queue head for the next flush instead of being dropped, and because only that rollup's row + was contended the job still drains the auto-router queue behind it.""" + monkeypatch.setattr("litellm.proxy.db.model_usage_rollup.asyncio.sleep", AsyncMock(return_value=None)) + _model_usage_batch_failing_with(_postgres_lock_timeout(), mock_prisma_client) + first, second = _model_usage_transaction("gpt-4o"), _model_usage_transaction("gpt-4o-mini") + mock_prisma_client.model_usage_transactions = [first, second] + mock_prisma_client.autorouter_turn_transactions = [MagicMock(name="autorouter_turn")] + + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=MagicMock(), + ) + + assert { + "queue_after": mock_prisma_client.model_usage_transactions, + "autorouter_queue_after": mock_prisma_client.autorouter_turn_transactions, + } == {"queue_after": [first, second], "autorouter_queue_after": []} + + +@pytest.mark.asyncio +async def test_model_usage_flush_requeues_and_stops_the_job_when_postgres_is_out_of_connections( + mock_prisma_client: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A batch refused for want of a connection (53300) was never sent either, so it is kept, but + a full server will refuse the writers behind it too: the job stops there like the tool + rollup does, so the drain loop backs off instead of re-hitting the server.""" + monkeypatch.setattr("litellm.proxy.db.model_usage_rollup.asyncio.sleep", AsyncMock(return_value=None)) + _model_usage_batch_failing_with(_postgres_out_of_connections(), mock_prisma_client) + first, second = _model_usage_transaction("gpt-4o"), _model_usage_transaction("gpt-4o-mini") + mock_prisma_client.model_usage_transactions = [first, second] + mock_prisma_client.autorouter_turn_transactions = [MagicMock(name="autorouter_turn")] + + 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=MagicMock(), + ) + + assert { + "queue_after": mock_prisma_client.model_usage_transactions, + "autorouter_queue_after": len(mock_prisma_client.autorouter_turn_transactions), + } == {"queue_after": [first, second], "autorouter_queue_after": 1} + + +@pytest.mark.asyncio +async def test_model_usage_flush_still_drops_a_batch_that_may_have_applied( + mock_prisma_client: MagicMock, +) -> None: + """Any other failure is ambiguous about whether the increments landed, so replaying it could + double-count: the batch is logged and dropped, and the job carries on.""" + _model_usage_batch_failing_with(RuntimeError("engine lost mid-commit"), mock_prisma_client) + mock_prisma_client.model_usage_transactions = [_model_usage_transaction("gpt-4o")] + + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=MagicMock(), + ) + + assert mock_prisma_client.model_usage_transactions == []