fix(proxy): bound tool-usage requeue to safe connection errors

Restrict tool-usage requeuing to safe pre-send connection errors (DB_RETRY_SAFE_ERROR_TYPES) to prevent double-counting of rollups on ambiguous post-send failures.
Enforce MAX_TOOL_USAGE_TRANSACTIONS_IN_MEMORY to prevent unbounded queue growth and memory exhaustion during persistent outages.
Preserve the existing per-statement error handling contract for autorouter session turns.
Test real flush logic using dependency injection on MockPrismaClient.

Signed-off-by: Liang Xu <755674130@qq.com>
This commit is contained in:
Liang Xu 2026-09-27 10:06:33 +08:00
parent f4f8974743
commit f8b51e0c3b
2 changed files with 97 additions and 132 deletions

View file

@ -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

View file

@ -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