From fae5865f87c899afc4d4e5fc8c759123114a1724 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 4 Feb 2026 09:42:46 -0800 Subject: [PATCH] Merge pull request #20394 from BerriAI/litellm_spend_fix_2 [Fix] Unique Constraint on Daily Tables + Logging When Updates Fail --- litellm/proxy/db/db_spend_update_writer.py | 203 +++++++++--------- .../proxy/db/test_db_spend_update_writer.py | 128 ++++++++++- 2 files changed, 232 insertions(+), 99 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 429e56c805b..dc928921425 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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" 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 72403b0ba7b..6ccecf59eed 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 @@ -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" \ No newline at end of file + 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", + ) \ No newline at end of file