mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix: process all daily spend batches per flush
_update_daily_spend previously exited after the first 100-row batch, dropping remaining daily-spend transactions accumulated in the same flush window. Process batches until the transaction map is empty and add regression coverage for >BATCH_SIZE updates.
This commit is contained in:
parent
c89496f378
commit
b0304ca589
2 changed files with 214 additions and 165 deletions
|
|
@ -1528,196 +1528,199 @@ class DBSpendUpdateWriter:
|
|||
start_time = time.time()
|
||||
|
||||
try:
|
||||
for i in range(n_retry_times + 1):
|
||||
try:
|
||||
# Sort the transactions to minimize the probability of deadlocks by reducing the chance of concurrent
|
||||
# trasactions locking the same rows/ranges in different orders.
|
||||
transactions_to_process = dict(
|
||||
sorted(
|
||||
daily_spend_transactions.items(),
|
||||
# Normally to avoid deadlocks we would sort by the index, but since we have sprinkled indexes
|
||||
# on our schema like we're discount Salt Bae, we just sort by all fields that have an index,
|
||||
# in an ad-hoc (but hopefully sensible) order of indexes. The actual ordering matters less than
|
||||
# ensuring that all concurrent transactions sort in the same order.
|
||||
# We could in theory use the dict key, as it contains basically the same fields, but this is more
|
||||
# robust to future changes in the key format.
|
||||
# If _update_daily_spend ever gets the ability to write to multiple tables at once, the sorting
|
||||
# should sort by the table first.
|
||||
key=lambda x: (
|
||||
x[1].get("date") or "",
|
||||
x[1].get(entity_id_field) or "",
|
||||
x[1].get("api_key") or "",
|
||||
x[1].get("model") or "",
|
||||
x[1].get("custom_llm_provider") or "",
|
||||
),
|
||||
)[:BATCH_SIZE]
|
||||
)
|
||||
|
||||
if len(transactions_to_process) == 0:
|
||||
verbose_proxy_logger.debug(
|
||||
f"No new transactions to process for daily {entity_type} spend update"
|
||||
)
|
||||
break
|
||||
|
||||
while daily_spend_transactions:
|
||||
batch_succeeded = False
|
||||
for i in range(n_retry_times + 1):
|
||||
try:
|
||||
async with prisma_client.db.batch_() as batcher:
|
||||
for _, transaction in transactions_to_process.items():
|
||||
entity_id = transaction.get(entity_id_field)
|
||||
# Sort the transactions to minimize the probability of deadlocks by reducing the chance of concurrent
|
||||
# trasactions locking the same rows/ranges in different orders.
|
||||
transactions_to_process = dict(
|
||||
sorted(
|
||||
daily_spend_transactions.items(),
|
||||
# Normally to avoid deadlocks we would sort by the index, but since we have sprinkled indexes
|
||||
# on our schema like we're discount Salt Bae, we just sort by all fields that have an index,
|
||||
# in an ad-hoc (but hopefully sensible) order of indexes. The actual ordering matters less than
|
||||
# ensuring that all concurrent transactions sort in the same order.
|
||||
# We could in theory use the dict key, as it contains basically the same fields, but this is more
|
||||
# robust to future changes in the key format.
|
||||
# If _update_daily_spend ever gets the ability to write to multiple tables at once, the sorting
|
||||
# should sort by the table first.
|
||||
key=lambda x: (
|
||||
x[1].get("date") or "",
|
||||
x[1].get(entity_id_field) or "",
|
||||
x[1].get("api_key") or "",
|
||||
x[1].get("model") or "",
|
||||
x[1].get("custom_llm_provider") or "",
|
||||
),
|
||||
)[:BATCH_SIZE]
|
||||
)
|
||||
|
||||
# Construct the where clause dynamically
|
||||
where_clause = {
|
||||
unique_constraint_name: {
|
||||
if len(transactions_to_process) == 0:
|
||||
break
|
||||
|
||||
try:
|
||||
async with prisma_client.db.batch_() as batcher:
|
||||
for _, transaction in transactions_to_process.items():
|
||||
entity_id = transaction.get(entity_id_field)
|
||||
|
||||
# Construct the where clause dynamically
|
||||
where_clause = {
|
||||
unique_constraint_name: {
|
||||
entity_id_field: entity_id,
|
||||
"date": transaction["date"],
|
||||
"api_key": transaction["api_key"],
|
||||
"model": transaction["model"],
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
)
|
||||
or "",
|
||||
"mcp_namespaced_tool_name": transaction.get(
|
||||
"mcp_namespaced_tool_name"
|
||||
)
|
||||
or "",
|
||||
"endpoint": transaction.get("endpoint") or "",
|
||||
}
|
||||
}
|
||||
|
||||
# Get the table dynamically
|
||||
table = getattr(batcher, table_name)
|
||||
|
||||
# Common data structure for both create and update
|
||||
common_data = {
|
||||
entity_id_field: entity_id,
|
||||
"date": transaction["date"],
|
||||
"api_key": transaction["api_key"],
|
||||
"model": transaction["model"],
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
)
|
||||
or "",
|
||||
"model": transaction.get("model"),
|
||||
"model_group": transaction.get("model_group"),
|
||||
"mcp_namespaced_tool_name": transaction.get(
|
||||
"mcp_namespaced_tool_name"
|
||||
)
|
||||
or "",
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
),
|
||||
"endpoint": transaction.get("endpoint") or "",
|
||||
"prompt_tokens": transaction["prompt_tokens"],
|
||||
"completion_tokens": transaction[
|
||||
"completion_tokens"
|
||||
],
|
||||
"spend": transaction["spend"],
|
||||
"api_requests": transaction["api_requests"],
|
||||
"successful_requests": transaction[
|
||||
"successful_requests"
|
||||
],
|
||||
"failed_requests": transaction["failed_requests"],
|
||||
}
|
||||
}
|
||||
|
||||
# Get the table dynamically
|
||||
table = getattr(batcher, table_name)
|
||||
|
||||
# Common data structure for both create and update
|
||||
common_data = {
|
||||
entity_id_field: entity_id,
|
||||
"date": transaction["date"],
|
||||
"api_key": transaction["api_key"],
|
||||
"model": transaction.get("model"),
|
||||
"model_group": transaction.get("model_group"),
|
||||
"mcp_namespaced_tool_name": transaction.get(
|
||||
"mcp_namespaced_tool_name"
|
||||
)
|
||||
or "",
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
),
|
||||
"endpoint": transaction.get("endpoint") or "",
|
||||
"prompt_tokens": transaction["prompt_tokens"],
|
||||
"completion_tokens": transaction[
|
||||
"completion_tokens"
|
||||
],
|
||||
"spend": transaction["spend"],
|
||||
"api_requests": transaction["api_requests"],
|
||||
"successful_requests": transaction[
|
||||
"successful_requests"
|
||||
],
|
||||
"failed_requests": transaction["failed_requests"],
|
||||
}
|
||||
|
||||
# Add cache-related fields if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
common_data[
|
||||
"cache_read_input_tokens"
|
||||
] = transaction.get("cache_read_input_tokens", 0)
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
common_data[
|
||||
"cache_creation_input_tokens"
|
||||
] = transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
common_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
)
|
||||
|
||||
# Create update data structure
|
||||
update_data = {
|
||||
"prompt_tokens": {
|
||||
"increment": transaction["prompt_tokens"]
|
||||
},
|
||||
"completion_tokens": {
|
||||
"increment": transaction["completion_tokens"]
|
||||
},
|
||||
"spend": {"increment": transaction["spend"]},
|
||||
"api_requests": {
|
||||
"increment": transaction["api_requests"]
|
||||
},
|
||||
"successful_requests": {
|
||||
"increment": transaction["successful_requests"]
|
||||
},
|
||||
"failed_requests": {
|
||||
"increment": transaction["failed_requests"]
|
||||
},
|
||||
}
|
||||
|
||||
# Add cache-related fields to update if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
update_data["cache_read_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_read_input_tokens", 0
|
||||
# Add cache-related fields if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
common_data["cache_read_input_tokens"] = (
|
||||
transaction.get("cache_read_input_tokens", 0)
|
||||
)
|
||||
}
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
update_data["cache_creation_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
common_data["cache_creation_input_tokens"] = (
|
||||
transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
)
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
common_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
)
|
||||
|
||||
# Create update data structure
|
||||
update_data = {
|
||||
"prompt_tokens": {
|
||||
"increment": transaction["prompt_tokens"]
|
||||
},
|
||||
"completion_tokens": {
|
||||
"increment": transaction["completion_tokens"]
|
||||
},
|
||||
"spend": {"increment": transaction["spend"]},
|
||||
"api_requests": {
|
||||
"increment": transaction["api_requests"]
|
||||
},
|
||||
"successful_requests": {
|
||||
"increment": transaction["successful_requests"]
|
||||
},
|
||||
"failed_requests": {
|
||||
"increment": transaction["failed_requests"]
|
||||
},
|
||||
}
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
update_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
# Add cache-related fields to update if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
update_data["cache_read_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_read_input_tokens", 0
|
||||
)
|
||||
}
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
update_data["cache_creation_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
}
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
update_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
)
|
||||
|
||||
# Add endpoint to update_data so existing rows get their endpoint field updated
|
||||
update_data["endpoint"] = (
|
||||
transaction.get("endpoint") or ""
|
||||
)
|
||||
|
||||
# Add endpoint to update_data so existing rows get their endpoint field updated
|
||||
update_data["endpoint"] = (
|
||||
transaction.get("endpoint") or ""
|
||||
)
|
||||
table.upsert(
|
||||
where=where_clause,
|
||||
data={
|
||||
"create": common_data,
|
||||
"update": update_data,
|
||||
},
|
||||
)
|
||||
except Exception as batch_error:
|
||||
# Log detailed error information for debugging batch upsert failures
|
||||
# This helps diagnose issues like unique constraint violations
|
||||
verbose_proxy_logger.exception(
|
||||
f"Daily {entity_type} spend batch upsert failed. "
|
||||
f"Table: {table_name}, Constraint: {unique_constraint_name}, "
|
||||
f"Batch size: {len(transactions_to_process)}, "
|
||||
f"Error: {str(batch_error)}"
|
||||
)
|
||||
raise
|
||||
|
||||
table.upsert(
|
||||
where=where_clause,
|
||||
data={
|
||||
"create": common_data,
|
||||
"update": update_data,
|
||||
},
|
||||
)
|
||||
except Exception as batch_error:
|
||||
# Log detailed error information for debugging batch upsert failures
|
||||
# This helps diagnose issues like unique constraint violations
|
||||
verbose_proxy_logger.exception(
|
||||
f"Daily {entity_type} spend batch upsert failed. "
|
||||
f"Table: {table_name}, Constraint: {unique_constraint_name}, "
|
||||
f"Batch size: {len(transactions_to_process)}, "
|
||||
f"Error: {str(batch_error)}"
|
||||
verbose_proxy_logger.debug(
|
||||
f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s"
|
||||
)
|
||||
raise
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s"
|
||||
)
|
||||
# Remove processed transactions
|
||||
for key in transactions_to_process.keys():
|
||||
daily_spend_transactions.pop(key, None)
|
||||
|
||||
# Remove processed transactions
|
||||
for key in transactions_to_process.keys():
|
||||
daily_spend_transactions.pop(key, None)
|
||||
batch_succeeded = True
|
||||
break
|
||||
|
||||
except DB_CONNECTION_ERROR_TYPES as e:
|
||||
if i >= n_retry_times:
|
||||
_raise_failed_update_spend_exception(
|
||||
e=e,
|
||||
start_time=start_time,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await asyncio.sleep(
|
||||
# Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are
|
||||
# cancelled basically at the same time, so if they wait the same time they will also retry at the same time
|
||||
# and thus they are more likely to deadlock again.
|
||||
# Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of
|
||||
# repeated deadlocks, and therefore of exceeding the retry limit.
|
||||
random.uniform(2**i, 2 ** (i + 1))
|
||||
)
|
||||
|
||||
if not batch_succeeded:
|
||||
break
|
||||
|
||||
except DB_CONNECTION_ERROR_TYPES as e:
|
||||
if i >= n_retry_times:
|
||||
_raise_failed_update_spend_exception(
|
||||
e=e,
|
||||
start_time=start_time,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await asyncio.sleep(
|
||||
# Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are
|
||||
# cancelled basically at the same time, so if they wait the same time they will also retry at the same time
|
||||
# and thus they are more likely to deadlock again.
|
||||
# Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of
|
||||
# repeated deadlocks, and therefore of exceeding the retry limit.
|
||||
random.uniform(2**i, 2 ** (i + 1))
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
if "transactions_to_process" in locals():
|
||||
for key in transactions_to_process.keys(): # type: ignore
|
||||
|
|
|
|||
|
|
@ -238,6 +238,52 @@ async def test_update_daily_spend_sorting():
|
|||
mock_table.upsert.assert_has_calls(upsert_calls)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_daily_spend_processes_all_batches():
|
||||
"""
|
||||
Ensure _update_daily_spend processes transactions beyond BATCH_SIZE.
|
||||
|
||||
Regression coverage for a bug where only the first 100 transactions were
|
||||
processed and remaining entries were silently dropped.
|
||||
"""
|
||||
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
|
||||
|
||||
daily_spend_transactions = {}
|
||||
total_transactions = 205 # > BATCH_SIZE (100), should require 3 batches
|
||||
for i in range(total_transactions):
|
||||
daily_spend_transactions[f"test_key_{i}"] = {
|
||||
"user_id": f"user{i}",
|
||||
"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=1,
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
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_table.upsert.call_count == total_transactions
|
||||
assert len(daily_spend_transactions) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_daily_spend_tag_with_request_id():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue