mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #20394 from BerriAI/litellm_spend_fix_2
[Fix] Unique Constraint on Daily Tables + Logging When Updates Fail
This commit is contained in:
parent
bad72ac140
commit
fae5865f87
2 changed files with 232 additions and 99 deletions
|
|
@ -1187,119 +1187,130 @@ class DBSpendUpdateWriter:
|
|||
)
|
||||
break
|
||||
|
||||
async with prisma_client.db.batch_() as batcher:
|
||||
for _, transaction in transactions_to_process.items():
|
||||
entity_id = transaction.get(entity_id_field)
|
||||
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: {
|
||||
# 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"),
|
||||
"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
|
||||
)
|
||||
}
|
||||
|
||||
# Add endpoint to update_data so existing rows get their endpoint field updated
|
||||
update_data["endpoint"] = transaction.get("endpoint") or ""
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
update_data["request_id"] = transaction.get("request_id")
|
||||
|
||||
table.upsert(
|
||||
where=where_clause,
|
||||
data={
|
||||
"create": common_data,
|
||||
"update": update_data,
|
||||
},
|
||||
)
|
||||
# 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
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s"
|
||||
|
|
|
|||
|
|
@ -132,7 +132,7 @@ async def test_update_daily_spend_with_null_entity_id():
|
|||
assert create_data["model"] == "gpt-4"
|
||||
assert create_data["custom_llm_provider"] == "openai"
|
||||
assert create_data["mcp_namespaced_tool_name"] == ""
|
||||
assert create_data["endpoint"] is None
|
||||
assert create_data["endpoint"] == ""
|
||||
assert create_data["prompt_tokens"] == 10
|
||||
assert create_data["completion_tokens"] == 20
|
||||
assert create_data["spend"] == 0.1
|
||||
|
|
@ -194,7 +194,7 @@ async def test_update_daily_spend_sorting():
|
|||
"model_group": None,
|
||||
"mcp_namespaced_tool_name": "",
|
||||
"custom_llm_provider": "openai",
|
||||
"endpoint": None,
|
||||
"endpoint": "",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"spend": 0.1,
|
||||
|
|
@ -838,4 +838,126 @@ async def test_endpoint_field_is_correctly_mapped_from_call_type():
|
|||
assert transaction["date"] == "2024-01-01"
|
||||
assert transaction["api_key"] == "test-key"
|
||||
assert transaction["model"] == "gpt-4"
|
||||
assert transaction["custom_llm_provider"] == "openai"
|
||||
assert transaction["custom_llm_provider"] == "openai"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure():
|
||||
"""
|
||||
Test that when batch upsert fails, detailed error information is logged.
|
||||
This ensures proper debugging information is available for issues like unique constraint violations.
|
||||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
# Setup
|
||||
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
|
||||
|
||||
# Create a transaction
|
||||
daily_spend_transactions = {
|
||||
"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,
|
||||
}
|
||||
}
|
||||
|
||||
# Create a mock proxy_logging_obj with failure_handler as AsyncMock
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.failure_handler = AsyncMock()
|
||||
|
||||
# Mock the logger to capture exception calls
|
||||
with patch.object(verbose_proxy_logger, 'exception') as mock_exception_logger:
|
||||
# Call the method and expect it to raise the exception
|
||||
with pytest.raises(Exception, match="Unique constraint violation"):
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0, # No retries to make test faster
|
||||
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",
|
||||
)
|
||||
|
||||
# Verify that exception was logged with detailed information
|
||||
assert mock_exception_logger.called
|
||||
call_args = mock_exception_logger.call_args[0][0]
|
||||
assert "Daily user spend batch upsert failed" in call_args
|
||||
assert "Table: litellm_dailyuserspend" in call_args
|
||||
assert "Constraint: user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint" in call_args
|
||||
assert "Batch size: 1" in call_args
|
||||
assert "Unique constraint violation" in call_args
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_daily_spend_re_raises_exception_after_logging():
|
||||
"""
|
||||
Test that when batch upsert fails, the exception is properly re-raised after logging.
|
||||
This ensures that error handling continues to work correctly upstream.
|
||||
"""
|
||||
# Setup
|
||||
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
|
||||
daily_spend_transactions = {
|
||||
"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,
|
||||
}
|
||||
}
|
||||
|
||||
# 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
|
||||
|
||||
# Create a mock proxy_logging_obj with failure_handler as AsyncMock
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.failure_handler = AsyncMock()
|
||||
|
||||
# Verify the exception is re-raised
|
||||
with pytest.raises(ValueError, match="Database connection lost"):
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0, # No retries to make test faster
|
||||
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",
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue