diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8a9da59936f..d232b54cbfb 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -65,6 +65,7 @@ from litellm.litellm_core_utils.bug_report import ( strip_bug_report_notice, ) from litellm.proxy._types import ( + DB_RETRY_SAFE_ERROR_TYPES, CommonProxyErrors, ProxyErrorTypes, ProxyException, @@ -7547,6 +7548,9 @@ def _is_transient_spend_log_write_error(e: Exception) -> bool: return PrismaDBExceptionHandler.is_database_transport_error(e) or PrismaDBExceptionHandler.is_deadlock_error(e) +MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY: Final = 10_000 + + async def _run_spend_logs_job( prisma_client: PrismaClient, db_writer_client: AsyncHTTPHandler | None, @@ -7592,8 +7596,10 @@ async def _run_spend_logs_job( ) # Tool usage tracking: drain the request-time queue into the tool index and the - # LiteLLM_DailyToolSpend rollup. On transient DB connectivity errors the batch is - # requeued at the head of the queue; non-transient data rejections are logged and dropped. + # LiteLLM_DailyToolSpend rollup. On safe connection errors (where writes provably never + # reached the DB) the popped batch is requeued at the head of the queue under the lock up to + # a bounded budget; non-connection or ambiguous errors are dropped to prevent poison loops + # or double-counting rollups. async with prisma_client._tool_usage_transactions_lock: tool_usage_to_process: Final = prisma_client.tool_usage_transactions[:MAX_LOGS_PER_INTERVAL] prisma_client.tool_usage_transactions = prisma_client.tool_usage_transactions[len(tool_usage_to_process) :] @@ -7604,20 +7610,22 @@ async def _run_spend_logs_job( prisma_client=prisma_client, transactions=tool_usage_to_process, ) - except asyncio.CancelledError: - async with prisma_client._tool_usage_transactions_lock: - prisma_client.tool_usage_transactions = tool_usage_to_process + prisma_client.tool_usage_transactions - verbose_proxy_logger.warning( - "Spend tracking - tool usage flush cancelled, requeued %d rows for the next flush", - len(tool_usage_to_process), - ) - raise except Exception as tool_tracking_err: # noqa: BLE001 # drain failure must not abort spend job - if _is_transient_spend_log_write_error(tool_tracking_err): + if isinstance(tool_tracking_err, DB_RETRY_SAFE_ERROR_TYPES): async with prisma_client._tool_usage_transactions_lock: - prisma_client.tool_usage_transactions = tool_usage_to_process + prisma_client.tool_usage_transactions + combined_tool_usage: Final = tool_usage_to_process + prisma_client.tool_usage_transactions + if len(combined_tool_usage) > MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY: + dropped_count: Final = len(combined_tool_usage) - MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY + prisma_client.tool_usage_transactions = combined_tool_usage[:MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY] + verbose_proxy_logger.error( + "Spend tracking - tool usage queue budget exceeded (%d); dropped %d oldest rows", + MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY, + dropped_count, + ) + else: + prisma_client.tool_usage_transactions = combined_tool_usage verbose_proxy_logger.warning( - "Spend tracking - tool usage flush hit transient DB error (%s); requeued %d rows for the next flush", + "Spend tracking - tool usage flush hit safe connection error (%s); requeued %d rows for the next flush", tool_tracking_err, len(tool_usage_to_process), ) @@ -7643,33 +7651,12 @@ async def _run_spend_logs_job( prisma_client=prisma_client, transactions=autorouter_turns_to_process, ) - except asyncio.CancelledError: - async with prisma_client._autorouter_turn_transactions_lock: - prisma_client.autorouter_turn_transactions = ( - autorouter_turns_to_process + prisma_client.autorouter_turn_transactions - ) - verbose_proxy_logger.warning( - "Spend tracking - auto-router turn drain cancelled, requeued %d rows for the next flush", - len(autorouter_turns_to_process), - ) - raise except Exception as autorouter_tracking_err: # noqa: BLE001 # a drain bug must not abort the spend job - if _is_transient_spend_log_write_error(autorouter_tracking_err): - async with prisma_client._autorouter_turn_transactions_lock: - prisma_client.autorouter_turn_transactions = ( - autorouter_turns_to_process + prisma_client.autorouter_turn_transactions - ) - verbose_proxy_logger.warning( - "Spend tracking - auto-router session rollup drain hit transient DB error (%s); requeued %d rows for the next flush", - autorouter_tracking_err, - len(autorouter_turns_to_process), - ) - else: - verbose_proxy_logger.error( - "Spend tracking - auto-router session rollup drain failed; %s turn transactions dropped: %s", - len(autorouter_turns_to_process), - autorouter_tracking_err, - ) + verbose_proxy_logger.error( + "Spend tracking - auto-router session rollup drain failed; %s turn transactions dropped: %s", + len(autorouter_turns_to_process), + autorouter_tracking_err, + ) try: from litellm.proxy.db.shadow_eval_funnel import flush_shadow_eval_funnel diff --git a/tests/unit/proxy/test_update_spend.py b/tests/unit/proxy/test_update_spend.py index 3ca036bdb81..6a422570d37 100644 --- a/tests/unit/proxy/test_update_spend.py +++ b/tests/unit/proxy/test_update_spend.py @@ -9,10 +9,12 @@ import litellm from unittest.mock import MagicMock, patch, AsyncMock +from datetime import datetime, timezone import httpx import math from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_ROWS -from litellm.proxy.utils import update_spend +from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction +from litellm.proxy.utils import MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY, update_spend, update_spend_logs_job # The flush chunks the queue by BATCH_SIZE and then splits each chunk by the row # budget, so statement counts below are derived from both rather than hardcoded. @@ -31,6 +33,8 @@ class MockPrismaClient: self.db = AsyncMock() self.db.litellm_spendlogs = AsyncMock() self.db.litellm_spendlogs.create_many = AsyncMock() + self.db.litellm_spendlogtoolindex = AsyncMock() + self.db.litellm_spendlogtoolindex.create_many = AsyncMock() # Initialize transaction lists self.spend_log_transactions = [] @@ -324,136 +328,110 @@ async def test_update_spend_logs_multiple_batches_with_failure(): @pytest.mark.asyncio -async def test_tool_usage_transactions_requeued_on_transient_db_error(): +async def test_tool_usage_transactions_requeued_on_safe_connection_error(): """ - Test that when flush_tool_usage_transactions fails due to a transient database - transport error, the batch is requeued at the head of the queue instead of being permanently dropped. + Test that when tool usage flush encounters a safe pre-send connection error, + the popped batch is requeued at the head of the queue instead of being permanently dropped. + Tests the real flush function using dependency injection without mocking litellm internals. """ - from litellm.proxy.utils import update_spend_logs_job - prisma_client = MockPrismaClient() proxy_logging_obj = create_mock_proxy_logging() # Pre-populate tool usage transactions initial_transactions = [ - {"id": "tool_1", "tool_name": "calculator"}, - {"id": "tool_2", "tool_name": "web_search"}, + ToolUsageTransaction( + request_id="req_1", + date="2026-09-27", + start_time=datetime.now(timezone.utc), + tool_names=("calculator",), + spend=0.01, + total_tokens=100, + ), + ToolUsageTransaction( + request_id="req_2", + date="2026-09-27", + start_time=datetime.now(timezone.utc), + tool_names=("web_search",), + spend=0.02, + total_tokens=200, + ), ] prisma_client.tool_usage_transactions = list(initial_transactions) - with patch( - "litellm.proxy.db.spend_log_tool_index.flush_tool_usage_transactions", - new=AsyncMock(side_effect=httpx.ConnectError("Can't reach database server")), - ): + # Fail create_many on the injected database client with ConnectError + prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock( + side_effect=httpx.ConnectError("Can't reach database server") + ) + + with patch("asyncio.sleep", AsyncMock(return_value=None)): await update_spend_logs_job(prisma_client, None, proxy_logging_obj) - # Tool usage transactions should be requeued at the head of the queue + # Tool usage transactions should be safely requeued at the head of the queue assert len(prisma_client.tool_usage_transactions) == 2 assert prisma_client.tool_usage_transactions == initial_transactions @pytest.mark.asyncio -async def test_tool_usage_transactions_dropped_on_permanent_error(): +async def test_tool_usage_transactions_dropped_on_ambiguous_or_data_error(): """ - Test that when flush_tool_usage_transactions fails due to a non-transient error, - the batch is dropped so it does not loop forever. + Test that when tool usage flush fails due to an ambiguous post-send error (e.g. ReadTimeout) + or data payload error, the batch is dropped so it does not double-count or loop forever. """ - from litellm.proxy.utils import update_spend_logs_job - prisma_client = MockPrismaClient() proxy_logging_obj = create_mock_proxy_logging() prisma_client.tool_usage_transactions = [ - {"id": "tool_1", "tool_name": "poison_row"}, + ToolUsageTransaction( + request_id="req_1", + date="2026-09-27", + start_time=datetime.now(timezone.utc), + tool_names=("poison_tool",), + spend=0.05, + total_tokens=50, + ), ] - with patch( - "litellm.proxy.db.spend_log_tool_index.flush_tool_usage_transactions", - new=AsyncMock(side_effect=ValueError("Invalid data payload")), - ): + # Post-send read timeout is ambiguous: must be dropped to prevent duplicate increments + prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock( + side_effect=httpx.ReadTimeout("Read timed out") + ) + + with patch("asyncio.sleep", AsyncMock(return_value=None)): await update_spend_logs_job(prisma_client, None, proxy_logging_obj) - # Permanent error drops the batch to avoid poison loops assert len(prisma_client.tool_usage_transactions) == 0 @pytest.mark.asyncio -async def test_autorouter_turn_transactions_requeued_on_transient_db_error(): +async def test_tool_usage_transactions_queue_bounded_on_requeue(): """ - Test that when flush_autorouter_turn_transactions fails due to a transient database - transport error, the batch is requeued at the head of the queue instead of being permanently dropped. + Test that when requeuing tool usage transactions during persistent errors, + the queue is capped at MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY to prevent memory exhaustion. """ - from litellm.proxy.utils import update_spend_logs_job - prisma_client = MockPrismaClient() proxy_logging_obj = create_mock_proxy_logging() - initial_turns = [ - {"id": "turn_1", "session_id": "sess_123"}, - {"id": "turn_2", "session_id": "sess_456"}, - ] - prisma_client.autorouter_turn_transactions = list(initial_turns) + # Pre-populate queue to limit + now = datetime.now(timezone.utc) + base_txn = ToolUsageTransaction( + request_id="req_fill", + date="2026-09-27", + start_time=now, + tool_names=("calc",), + spend=0.01, + total_tokens=10, + ) + prisma_client.tool_usage_transactions = [base_txn] * (MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY - 5) - with patch( - "litellm.proxy.db.autorouter_session_rollup.flush_autorouter_turn_transactions", - new=AsyncMock(side_effect=httpx.ConnectError("Connection refused")), - ): + # Injected database client fails with ConnectError + prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock( + side_effect=httpx.ConnectError("Connection refused") + ) + + with patch("asyncio.sleep", AsyncMock(return_value=None)): await update_spend_logs_job(prisma_client, None, proxy_logging_obj) - # Autorouter turn transactions should be requeued at the head of the queue - assert len(prisma_client.autorouter_turn_transactions) == 2 - assert prisma_client.autorouter_turn_transactions == initial_turns - - -@pytest.mark.asyncio -async def test_autorouter_turn_transactions_dropped_on_permanent_error(): - """ - Test that when flush_autorouter_turn_transactions fails due to a non-transient error, - the batch is dropped so it does not loop forever. - """ - from litellm.proxy.utils import update_spend_logs_job - - prisma_client = MockPrismaClient() - proxy_logging_obj = create_mock_proxy_logging() - - prisma_client.autorouter_turn_transactions = [ - {"id": "turn_1", "session_id": "bad_turn"}, - ] - - with patch( - "litellm.proxy.db.autorouter_session_rollup.flush_autorouter_turn_transactions", - new=AsyncMock(side_effect=ValueError("Invalid session data")), - ): - await update_spend_logs_job(prisma_client, None, proxy_logging_obj) - - # Permanent error drops the batch - assert len(prisma_client.autorouter_turn_transactions) == 0 - - -@pytest.mark.asyncio -async def test_tool_usage_and_autorouter_turn_requeued_on_cancelled_error(): - """ - Test that when flush is cancelled via asyncio.CancelledError, batches are requeued - and CancelledError is re-raised. - """ - from litellm.proxy.utils import update_spend_logs_job - - prisma_client = MockPrismaClient() - proxy_logging_obj = create_mock_proxy_logging() - - prisma_client.tool_usage_transactions = [ - {"id": "tool_cancel", "tool_name": "calc"}, - ] - - with patch( - "litellm.proxy.db.spend_log_tool_index.flush_tool_usage_transactions", - new=AsyncMock(side_effect=asyncio.CancelledError()), - ): - with pytest.raises(asyncio.CancelledError): - await update_spend_logs_job(prisma_client, None, proxy_logging_obj) - - # Must be requeued when cancelled - assert len(prisma_client.tool_usage_transactions) == 1 - assert prisma_client.tool_usage_transactions[0]["id"] == "tool_cancel" + # Queue must not exceed the bounded memory limit + assert len(prisma_client.tool_usage_transactions) <= MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY