From 250de625af92b23b428566371b35ebf197abbe51 Mon Sep 17 00:00:00 2001 From: shivam Date: Thu, 23 Jul 2026 01:14:06 +0000 Subject: [PATCH] fix(spend): bound daily-spend batch upsert in a server-side transaction Resolves LIT-4715 _update_daily_spend ran its 100-row upsert batch via a bare prisma db.batch_(), which is a single implicit transaction with no server-side bound. When the batch exceeded prisma-client-py's default 30s httpx timeout the client raised ReadTimeout but nothing cancelled the server-side transaction, so its row locks stayed held (orphaned) while httpx.ReadTimeout being in DB_CONNECTION_ERROR_TYPES made the retry loop re-issue the same batch, piling sessions up behind the orphaned locks until the database was exhausted, and risking double counting if an orphaned batch and its retry both committed. Wrap the batch in db.tx(timeout=timedelta(seconds=60)) with a batcher, matching every other spend path, so a slow batch is rolled back server side instead of leaving orphaned locks. Stop blind-retrying httpx.ReadTimeout on this path since the increments are non-idempotent and a timed-out attempt may still commit; the pending increments are dropped rather than replayed. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 13 +- .../proxy/db/test_db_spend_update_writer.py | 140 ++++++++++++++++-- 2 files changed, 137 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 2262141f426..f3a243eccca 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -25,6 +25,8 @@ from typing import ( overload, ) +import httpx + import litellm from litellm._logging import verbose_proxy_logger from litellm.caching import RedisCache @@ -1510,7 +1512,10 @@ class DBSpendUpdateWriter: return try: - async with prisma_client.db.batch_() as batcher: + async with ( + prisma_client.db.tx(timeout=timedelta(seconds=60)) as db_transaction, + db_transaction.batch_() as batcher, + ): for _, transaction in transactions_to_process.items(): entity_id = transaction.get(entity_id_field) @@ -1644,6 +1649,12 @@ class DBSpendUpdateWriter: break + except httpx.ReadTimeout as e: + _raise_failed_update_spend_exception( + e=e, + start_time=start_time, + proxy_logging_obj=proxy_logging_obj, + ) except DB_CONNECTION_ERROR_TYPES as e: if i >= n_retry_times: _raise_failed_update_spend_exception( diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 8149cf90e70..aa2fc5936b8 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -9,9 +9,10 @@ sys.path.insert( ) # Adds the parent directory to the system path -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock, call, patch +import httpx import pytest from redis.exceptions import DataError @@ -20,6 +21,29 @@ from litellm.proxy._types import Litellm_EntityType from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter +def _wire_daily_spend_tx(mock_prisma_client, mock_batcher, batch_exit_side_effect=None): + """ + Wire ``mock_prisma_client`` so db.tx(...).batch_() yields ``mock_batcher``. + + Mirrors how _update_daily_spend runs its upserts inside a server-side-bounded + transaction (db.tx(timeout=...)) rather than a bare db.batch_(). + """ + batch_context = MagicMock() + batch_context.__aenter__ = AsyncMock(return_value=mock_batcher) + if batch_exit_side_effect is None: + batch_context.__aexit__ = AsyncMock(return_value=False) + else: + batch_context.__aexit__ = AsyncMock(side_effect=batch_exit_side_effect) + + mock_transaction = MagicMock() + mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction) + mock_transaction.__aexit__ = AsyncMock(return_value=False) + mock_transaction.batch_ = MagicMock(return_value=batch_context) + + mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) + return mock_transaction + + @pytest.mark.asyncio async def test_daily_spend_tracking_with_disabled_spend_logs(): """ @@ -87,8 +111,8 @@ async def test_update_daily_spend_with_null_entity_id(): mock_prisma_client = MagicMock() mock_batcher = MagicMock() mock_table = MagicMock() - mock_prisma_client.db.batch_.return_value.__aenter__.return_value = mock_batcher mock_batcher.litellm_dailyuserspend = mock_table + _wire_daily_spend_tx(mock_prisma_client, mock_batcher) # Create a transaction with null entity_id daily_spend_transactions = { @@ -163,8 +187,8 @@ async def test_update_daily_spend_sorting(): mock_prisma_client = MagicMock() mock_batcher = MagicMock() mock_table = MagicMock() - mock_prisma_client.db.batch_.return_value.__aenter__.return_value = mock_batcher mock_batcher.litellm_dailyuserspend = mock_table + _wire_daily_spend_tx(mock_prisma_client, mock_batcher) # Create a 50 transactions with out-of-order entity_ids # In reality we sort using multiple fields, but entity_id is sufficient to test sorting @@ -254,8 +278,8 @@ async def test_update_daily_spend_drains_all_batches_over_batch_size(): mock_prisma_client = MagicMock() mock_batcher = MagicMock() mock_table = MagicMock() - mock_prisma_client.db.batch_.return_value.__aenter__.return_value = mock_batcher mock_batcher.litellm_dailyuserspend = mock_table + _wire_daily_spend_tx(mock_prisma_client, mock_batcher) num_entities = 250 daily_spend_transactions = { @@ -287,7 +311,7 @@ async def test_update_daily_spend_drains_all_batches_over_batch_size(): ) assert mock_table.upsert.call_count == num_entities - assert mock_prisma_client.db.batch_.call_count == 3 + assert mock_prisma_client.db.tx.call_count == 3 assert daily_spend_transactions == {} @@ -300,8 +324,8 @@ async def test_update_daily_spend_tag_with_request_id(): mock_prisma_client = MagicMock() mock_batcher = MagicMock() mock_table = MagicMock() - mock_prisma_client.db.batch_.return_value.__aenter__.return_value = mock_batcher mock_batcher.litellm_dailytagspend = mock_table + _wire_daily_spend_tx(mock_prisma_client, mock_batcher) # Create a transaction with request_id daily_spend_transactions = { @@ -357,8 +381,8 @@ async def test_update_daily_spend_with_none_values_in_sorting_fields(): mock_prisma_client = MagicMock() mock_batcher = MagicMock() mock_table = MagicMock() - mock_prisma_client.db.batch_.return_value.__aenter__.return_value = mock_batcher mock_batcher.litellm_dailyuserspend = mock_table + _wire_daily_spend_tx(mock_prisma_client, mock_batcher) # Create transactions with None values in various sorting fields daily_spend_transactions = { @@ -1154,15 +1178,12 @@ async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure(): mock_prisma_client = MagicMock() mock_batcher = MagicMock() mock_table = MagicMock() - mock_batch_context = MagicMock() - mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher) mock_batcher.litellm_dailyuserspend = mock_table # Make the batch context manager's exit raise an exception # This simulates a batch commit failure (e.g., unique constraint violation) test_exception = Exception("Unique constraint violation") - mock_batch_context.__aexit__ = AsyncMock(side_effect=test_exception) - mock_prisma_client.db.batch_.return_value = mock_batch_context + _wire_daily_spend_tx(mock_prisma_client, mock_batcher, batch_exit_side_effect=test_exception) # Create a transaction daily_spend_transactions = { @@ -1228,8 +1249,6 @@ async def test_update_daily_spend_re_raises_exception_after_logging(): mock_prisma_client = MagicMock() mock_batcher = MagicMock() mock_table = MagicMock() - mock_batch_context = MagicMock() - mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher) mock_batcher.litellm_dailyuserspend = mock_table # Create a transaction @@ -1251,8 +1270,7 @@ async def test_update_daily_spend_re_raises_exception_after_logging(): # Create a custom exception to verify it's re-raised custom_exception = ValueError("Database connection lost") - mock_batch_context.__aexit__ = AsyncMock(side_effect=custom_exception) - mock_prisma_client.db.batch_.return_value = mock_batch_context + _wire_daily_spend_tx(mock_prisma_client, mock_batcher, batch_exit_side_effect=custom_exception) # Create a mock proxy_logging_obj with failure_handler as AsyncMock mock_proxy_logging = MagicMock() @@ -1272,6 +1290,98 @@ async def test_update_daily_spend_re_raises_exception_after_logging(): ) +def _single_daily_user_transaction(): + return { + "test_key": { + "user_id": "test-user", + "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, + } + } + + +@pytest.mark.asyncio +async def test_update_daily_spend_runs_inside_server_side_bounded_transaction(): + """ + Regression for LIT-4715. The daily-spend batch must run inside a + server-side-bounded transaction (db.tx(timeout=...)) so a slow batch is rolled + back by the query engine instead of leaving an orphaned transaction holding row + locks after the client disconnects. A bare db.batch_() has no server-side bound. + """ + mock_prisma_client = MagicMock() + mock_batcher = MagicMock() + mock_table = MagicMock() + mock_batcher.litellm_dailyuserspend = mock_table + mock_transaction = _wire_daily_spend_tx(mock_prisma_client, mock_batcher) + + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=1, + prisma_client=mock_prisma_client, + proxy_logging_obj=MagicMock(), + daily_spend_transactions=_single_daily_user_transaction(), + entity_type="user", + entity_id_field="user_id", + table_name="litellm_dailyuserspend", + unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint", + ) + + mock_prisma_client.db.tx.assert_called_once() + _, tx_kwargs = mock_prisma_client.db.tx.call_args + assert isinstance(tx_kwargs.get("timeout"), timedelta) + assert tx_kwargs["timeout"].total_seconds() > 0 + mock_transaction.batch_.assert_called_once() + assert mock_prisma_client.db.batch_.called is False + mock_table.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_update_daily_spend_does_not_retry_or_double_count_on_read_timeout(): + """ + Regression for LIT-4715. A client-side httpx.ReadTimeout on the daily-spend + batch must not be blind-retried: the batch is a non-idempotent increment whose + first attempt may still commit server side, so retrying risks double counting + spend. Assert the batch is attempted exactly once (despite n_retry_times>0) and + the pending increments are dropped rather than replayed on a later flush. + """ + mock_prisma_client = MagicMock() + mock_batcher = MagicMock() + mock_table = MagicMock() + mock_batcher.litellm_dailyuserspend = mock_table + _wire_daily_spend_tx( + mock_prisma_client, + mock_batcher, + batch_exit_side_effect=httpx.ReadTimeout("query engine read timed out"), + ) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.failure_handler = AsyncMock() + + daily_spend_transactions = _single_daily_user_transaction() + + with pytest.raises(httpx.ReadTimeout): + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=3, + prisma_client=mock_prisma_client, + proxy_logging_obj=mock_proxy_logging, + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + table_name="litellm_dailyuserspend", + unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint", + ) + + assert mock_prisma_client.db.tx.call_count == 1 + assert daily_spend_transactions == {} + + @pytest.mark.asyncio async def test_commit_key_spend_updates_includes_last_active(): """