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>
This commit is contained in:
shivam 2026-07-23 01:14:06 +00:00
parent 9f6b3d24f7
commit 250de625af
2 changed files with 137 additions and 16 deletions

View file

@ -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(

View file

@ -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():
"""