mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
9f6b3d24f7
commit
250de625af
2 changed files with 137 additions and 16 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue