From 16a5d4df8d7e28f8ac1480504cd161a2141c3154 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 10 Sep 2026 07:48:53 +0000 Subject: [PATCH 1/2] fix(proxy): bound tool and guardrail index create_many by the spend-log statement budgets One flush drains up to MAX_LOGS_PER_INTERVAL source transactions or logs, but a transaction fans out to one LiteLLM_SpendLogToolIndex row per tool and a log to one LiteLLM_SpendLogGuardrailIndex row per guardrail, so the index create_many payload was unbounded. Both index writes now go through spend_log_write_batches(SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS). The tool index write moves out of the rollup batch_() so the split reduces the query-engine payload; replayed index rows are no-ops under skip_duplicates, and the daily rollup upserts stay in one transaction Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/spend_log_tool_index.py | 36 ++++--- litellm/proxy/guardrails/usage_tracking.py | 13 +-- .../proxy/db/test_spend_log_tool_index.py | 93 ++++++++++++------- .../proxy/guardrails/test_usage_tracking.py | 53 +++++++++++ 4 files changed, 147 insertions(+), 48 deletions(-) diff --git a/litellm/proxy/db/spend_log_tool_index.py b/litellm/proxy/db/spend_log_tool_index.py index 93bbc567430..c01fe15fc09 100644 --- a/litellm/proxy/db/spend_log_tool_index.py +++ b/litellm/proxy/db/spend_log_tool_index.py @@ -5,8 +5,10 @@ At request time the spend writer builds one ToolUsageTransaction per request tha invoked tools (MCP namespaced tool name plus response tool_calls; declared-but-not- invoked tools are excluded) and queues it on the prisma client. The spend-log flush job drains the queue into LiteLLM_SpendLogToolIndex (per-request drill-down) and -LiteLLM_DailyToolSpend (the per-day rollup the Cost Optimization card reads) in a -single transaction, so a failed flush never leaves a partial rollup increment. +LiteLLM_DailyToolSpend (the per-day rollup the Cost Optimization card reads). The +index rows are keyed on (request_id, tool_name) and written with skip_duplicates, +so they go out as bounded standalone statements; every rollup upsert stays in one +transaction, so a failed flush never leaves a partial rollup increment. """ from __future__ import annotations @@ -19,7 +21,10 @@ from datetime import datetime, timezone from itertools import groupby from typing import TYPE_CHECKING, Any, Final +from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES +from litellm.proxy.db.spend_log_batching import spend_log_write_batches +from litellm.repositories.table_repositories import SpendLogToolIndexRepository if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient @@ -98,14 +103,19 @@ async def flush_tool_usage_transactions( transactions: Sequence[ToolUsageTransaction], n_retry_times: int = 3, ) -> None: - """Write index rows and rollup upserts for a drained queue batch in one - transaction. Retries only ConnectError, the one failure that proves the - statements never reached the database. Post-send failures (Read timeouts - and errors) are ambiguous and are NOT retried: the engine can abandon the - transaction open on the pooled connection, so a retry's statements stack - into the same transaction and one commit applies both increment sets. - Ambiguous failures drop the batch; the caller logs it at error. Callers - must not add their own retry around this function.""" + """Write the index rows as bounded standalone statements, then every rollup + upsert for the drained queue batch in one transaction. One flush fans out to + transactions x tools index rows, so the index write is split by the spend-log + statement budgets; a split inside ``batch_()`` would not help, since the + batcher ships every queued statement to the query engine as one payload. + Retries only ConnectError, the one failure that proves the statements never + reached the database; replayed index rows are no-ops under skip_duplicates. + Post-send failures (Read timeouts and errors) are ambiguous and are NOT + retried: the engine can abandon the transaction open on the pooled + connection, so a retry's statements stack into the same transaction and one + commit applies both increment sets. Ambiguous failures drop the batch; the + caller logs it at error. Callers must not add their own retry around this + function.""" if not transactions: return @@ -119,10 +129,14 @@ async def flush_tool_usage_transactions( key=lambda entry: (entry[0], entry[1]), ) + index_table: Final = SpendLogToolIndexRepository(prisma_client).table for attempt in range(n_retry_times + 1): try: + for statement_rows in spend_log_write_batches( + index_rows, SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS + ): + await index_table.create_many(data=statement_rows, skip_duplicates=True) async with prisma_client.db.batch_() as batcher: - batcher.litellm_spendlogtoolindex.create_many(data=index_rows, skip_duplicates=True) for (date_key, tool_name), grouped in groupby(per_tool_day, key=lambda entry: (entry[0], entry[1])): entries = tuple(grouped) spend = sum(entry[2] for entry in entries) diff --git a/litellm/proxy/guardrails/usage_tracking.py b/litellm/proxy/guardrails/usage_tracking.py index df967058cf0..797323794d2 100644 --- a/litellm/proxy/guardrails/usage_tracking.py +++ b/litellm/proxy/guardrails/usage_tracking.py @@ -17,8 +17,10 @@ from typing import TYPE_CHECKING, Any, Final, NamedTuple, TypeVar from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger +from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import billed_guardrail_cost_by_unit from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES +from litellm.proxy.db.spend_log_batching import spend_log_write_batches from litellm.proxy.utils import PrismaClient from litellm.repositories.table_repositories import ( DailyGuardrailMetricsRepository, @@ -401,13 +403,12 @@ async def process_spend_logs_guardrail_usage( return try: - # Insert index rows (skip duplicates by request_id + guardrail_id) - if index_rows: + index_table: Final = SpendLogGuardrailIndexRepository(prisma_client).table + for statement_rows in spend_log_write_batches( + index_rows, SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS + ): try: - await SpendLogGuardrailIndexRepository(prisma_client).table.create_many( - data=index_rows, - skip_duplicates=True, - ) + await index_table.create_many(data=statement_rows, skip_duplicates=True) except Exception as e: verbose_proxy_logger.debug("Guardrail usage tracking: index create_many skipped: %s", e) diff --git a/tests/test_litellm/proxy/db/test_spend_log_tool_index.py b/tests/test_litellm/proxy/db/test_spend_log_tool_index.py index 9c4fbbf41aa..ff2be799717 100644 --- a/tests/test_litellm/proxy/db/test_spend_log_tool_index.py +++ b/tests/test_litellm/proxy/db/test_spend_log_tool_index.py @@ -1,15 +1,17 @@ """ Tests for the tool usage writer: ToolUsageTransaction construction (invoked tools -only) and the flush that writes LiteLLM_SpendLogToolIndex plus the -LiteLLM_DailyToolSpend rollup in one transaction. +only) and the flush that writes LiteLLM_SpendLogToolIndex in bounded statements +plus the LiteLLM_DailyToolSpend rollup in one transaction. """ from types import SimpleNamespace from typing import Any from unittest.mock import AsyncMock, MagicMock +import httpx import pytest +from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_ROWS from litellm.proxy.db.spend_log_tool_index import ( ToolUsageTransaction, build_tool_usage_transaction, @@ -35,11 +37,24 @@ class _FakeBatcher: return None +def _prisma(batch_: MagicMock) -> MagicMock: + prisma = MagicMock() + prisma.db.batch_ = batch_ + prisma.db.litellm_spendlogtoolindex.create_many = AsyncMock() + return prisma + + def _prisma_with_batcher() -> tuple[MagicMock, _FakeBatcher]: batcher = _FakeBatcher() - prisma = MagicMock() - prisma.db.batch_ = MagicMock(return_value=batcher) - return prisma, batcher + return _prisma(MagicMock(return_value=batcher)), batcher + + +def _index_rows_written(prisma: MagicMock) -> list[tuple[str, str]]: + return [ + (row["request_id"], row["tool_name"]) + for call in prisma.db.litellm_spendlogtoolindex.create_many.call_args_list + for row in call.kwargs["data"] + ] class TestBuildToolUsageTransaction: @@ -228,9 +243,8 @@ class TestFlushToolUsageTransactions: prisma_client=prisma, transactions=[_transaction("r1", tool_names=("tool_a", "tool_b"), spend=0.10, total_tokens=100)], ) - index_rows = batcher.litellm_spendlogtoolindex.create_many.call_args.kwargs["data"] - assert [(r["request_id"], r["tool_name"]) for r in index_rows] == [("r1", "tool_a"), ("r1", "tool_b")] - assert batcher.litellm_spendlogtoolindex.create_many.call_args.kwargs["skip_duplicates"] is True + assert _index_rows_written(prisma) == [("r1", "tool_a"), ("r1", "tool_b")] + assert prisma.db.litellm_spendlogtoolindex.create_many.call_args.kwargs["skip_duplicates"] is True upserts = { c.kwargs["where"]["date_tool_name"]["tool_name"]: c.kwargs["data"] @@ -267,19 +281,46 @@ class TestFlushToolUsageTransactions: assert data["update"]["request_count"] == {"increment": 2} @pytest.mark.asyncio - async def test_index_rows_and_rollup_share_one_transaction(self): - # Both writes go through the same batch_() so a failed flush cannot leave - # index rows without their rollup increments (or vice versa); increments - # are not idempotent, so partial states must be unreachable. + async def test_index_rows_are_written_in_bounded_statements_outside_the_rollup_transaction(self): prisma, batcher = _prisma_with_batcher() - await flush_tool_usage_transactions( - prisma_client=prisma, - transactions=[_transaction("r1")], - ) + tool_names = tuple(f"tool_{i}" for i in range(50)) + transactions = [_transaction(f"r{i}", tool_names=tool_names) for i in range(5)] + await flush_tool_usage_transactions(prisma_client=prisma, transactions=transactions) + + statements = prisma.db.litellm_spendlogtoolindex.create_many.call_args_list + assert [len(call.kwargs["data"]) for call in statements] == [100, 100, 50] + assert all(len(call.kwargs["data"]) <= SPEND_LOG_WRITE_BATCH_MAX_ROWS for call in statements) + assert all(call.kwargs["skip_duplicates"] is True for call in statements) + assert _index_rows_written(prisma) == [ + (txn.request_id, tool_name) for txn in transactions for tool_name in tool_names + ] + batcher.litellm_spendlogtoolindex.create_many.assert_not_called() + prisma.db.batch_.assert_called_once() + assert batcher.litellm_dailytoolspend.upsert.call_count == len(tool_names) + + @pytest.mark.asyncio + async def test_index_connection_error_is_retried_before_the_rollup_is_attempted(self, monkeypatch): + prisma, batcher = _prisma_with_batcher() + prisma.db.litellm_spendlogtoolindex.create_many = AsyncMock(side_effect=[httpx.ConnectError("down"), None]) + + async def fake_sleep(seconds: float) -> None: + return None + + monkeypatch.setattr("litellm.proxy.db.spend_log_tool_index.asyncio.sleep", fake_sleep) + await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) + assert prisma.db.litellm_spendlogtoolindex.create_many.await_count == 2 prisma.db.batch_.assert_called_once() - batcher.litellm_spendlogtoolindex.create_many.assert_called_once() batcher.litellm_dailytoolspend.upsert.assert_called_once() + @pytest.mark.asyncio + async def test_ambiguous_index_error_drops_the_batch_without_touching_the_rollup(self): + prisma, _ = _prisma_with_batcher() + prisma.db.litellm_spendlogtoolindex.create_many = AsyncMock(side_effect=httpx.ReadTimeout("ambiguous")) + with pytest.raises(httpx.ReadTimeout): + await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) + prisma.db.litellm_spendlogtoolindex.create_many.assert_awaited_once() + prisma.db.batch_.assert_not_called() + @pytest.mark.asyncio async def test_empty_batch_touches_nothing(self): prisma, _ = _prisma_with_batcher() @@ -290,11 +331,8 @@ class TestFlushToolUsageTransactions: async def test_connection_errors_retry_and_succeed(self, monkeypatch): # A failed batch commits nothing, so retrying a connection error cannot # double-count; the flush must retry rather than drop the batch. - import httpx - batcher = _FakeBatcher() - prisma = MagicMock() - prisma.db.batch_ = MagicMock(side_effect=[httpx.ConnectError("down"), batcher]) + prisma = _prisma(MagicMock(side_effect=[httpx.ConnectError("down"), batcher])) sleeps: list[float] = [] async def fake_sleep(seconds: float) -> None: @@ -308,10 +346,7 @@ class TestFlushToolUsageTransactions: @pytest.mark.asyncio async def test_connection_errors_exhaust_retries_then_raise(self, monkeypatch): - import httpx - - prisma = MagicMock() - prisma.db.batch_ = MagicMock(side_effect=httpx.ConnectError("down")) + prisma = _prisma(MagicMock(side_effect=httpx.ConnectError("down"))) async def fake_sleep(seconds: float) -> None: return None @@ -325,8 +360,7 @@ class TestFlushToolUsageTransactions: @pytest.mark.asyncio async def test_non_connection_errors_do_not_retry(self): - prisma = MagicMock() - prisma.db.batch_ = MagicMock(side_effect=ValueError("bad data")) + prisma = _prisma(MagicMock(side_effect=ValueError("bad data"))) with pytest.raises(ValueError, match="bad data"): await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) prisma.db.batch_.assert_called_once() @@ -338,11 +372,8 @@ class TestFlushToolUsageTransactions: # unknown; the engine can leave the transaction open on the pooled # connection, so a retry's statements would stack into it and one # commit would apply both increment sets. These must never retry. - import httpx - error = getattr(httpx, ambiguous_error)("ambiguous") - prisma = MagicMock() - prisma.db.batch_ = MagicMock(side_effect=error) + prisma = _prisma(MagicMock(side_effect=error)) with pytest.raises((httpx.ReadTimeout, httpx.ReadError)): await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) prisma.db.batch_.assert_called_once() diff --git a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py index 110de7dbe70..3acb8317c73 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py @@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, MagicMock import httpx import pytest +from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_ROWS from litellm.proxy.guardrails.usage_tracking import ( _MAX_PENDING_ROWS, PendingRollups, @@ -511,3 +512,55 @@ async def test_requeued_cost_is_added_to_the_next_flush(): costs = _cost_upserts(recovered) assert costs["contentPolicyUnits"] == (pytest.approx(0.45), 0) assert costs["someFutureCounter"] == (0.0, 7) + + +def _fan_out_payload(request_id: str, guardrail_ids: tuple[str, ...]) -> dict[str, Any]: + return { + "request_id": request_id, + "startTime": datetime(2026, 8, 17, 12, 0, tzinfo=timezone.utc), + "team_id": "team-a", + "api_key": "hashed-key-1", + "metadata": json.dumps( + {"guardrail_information": [{"guardrail_id": gid, "guardrail_status": "success"} for gid in guardrail_ids]} + ), + } + + +def _index_rows_written(prisma: MagicMock) -> list[tuple[str, str]]: + return [ + (row["request_id"], row["guardrail_id"]) + for call in prisma.db.litellm_spendlogguardrailindex.create_many.call_args_list + for row in call.kwargs["data"] + ] + + +@pytest.mark.asyncio +async def test_index_rows_are_written_in_row_bounded_statements(): + """ + LIT-5931: the drain caps logs, not logs x guardrails, so a fan-out must be + split into statements the query engine can afford instead of one create_many. + """ + prisma = _prisma() + guardrail_ids = tuple(f"guard-{i}" for i in range(50)) + logs = [_fan_out_payload(f"r{i}", guardrail_ids) for i in range(5)] + + await process_spend_logs_guardrail_usage(prisma, logs, pending=PendingRollups()) + + statements = prisma.db.litellm_spendlogguardrailindex.create_many.call_args_list + assert [len(call.kwargs["data"]) for call in statements] == [100, 100, 50] + assert all(len(call.kwargs["data"]) <= SPEND_LOG_WRITE_BATCH_MAX_ROWS for call in statements) + assert all(call.kwargs["skip_duplicates"] is True for call in statements) + assert _index_rows_written(prisma) == [(f"r{i}", gid) for i in range(5) for gid in guardrail_ids] + + +@pytest.mark.asyncio +async def test_one_failing_index_statement_does_not_drop_the_others_or_the_rollup(): + prisma = _prisma() + prisma.db.litellm_spendlogguardrailindex.create_many.side_effect = [None, httpx.ReadTimeout("ambiguous"), None] + guardrail_ids = tuple(f"guard-{i}" for i in range(50)) + logs = [_fan_out_payload(f"r{i}", guardrail_ids) for i in range(5)] + + await process_spend_logs_guardrail_usage(prisma, logs, pending=PendingRollups()) + + assert prisma.db.litellm_spendlogguardrailindex.create_many.await_count == 3 + assert prisma.db.litellm_dailyguardrailmetrics.upsert.await_count == len(guardrail_ids) From e0eb761e2ab2572d8517668473d2ffba9d9897af Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 10 Sep 2026 08:06:00 +0000 Subject: [PATCH 2/2] test(proxy): pin the row budget in the index fan-out tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/db/test_spend_log_tool_index.py | 6 +++--- .../test_litellm/proxy/guardrails/test_usage_tracking.py | 9 +++++---- 2 files changed, 8 insertions(+), 7 deletions(-) diff --git a/tests/test_litellm/proxy/db/test_spend_log_tool_index.py b/tests/test_litellm/proxy/db/test_spend_log_tool_index.py index ff2be799717..282c2a7cfaa 100644 --- a/tests/test_litellm/proxy/db/test_spend_log_tool_index.py +++ b/tests/test_litellm/proxy/db/test_spend_log_tool_index.py @@ -11,7 +11,7 @@ from unittest.mock import AsyncMock, MagicMock import httpx import pytest -from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_ROWS +from litellm.proxy.db import spend_log_tool_index from litellm.proxy.db.spend_log_tool_index import ( ToolUsageTransaction, build_tool_usage_transaction, @@ -281,7 +281,8 @@ class TestFlushToolUsageTransactions: assert data["update"]["request_count"] == {"increment": 2} @pytest.mark.asyncio - async def test_index_rows_are_written_in_bounded_statements_outside_the_rollup_transaction(self): + async def test_index_rows_are_written_in_bounded_statements_outside_the_rollup_transaction(self, monkeypatch): + monkeypatch.setattr(spend_log_tool_index, "SPEND_LOG_WRITE_BATCH_MAX_ROWS", 100) prisma, batcher = _prisma_with_batcher() tool_names = tuple(f"tool_{i}" for i in range(50)) transactions = [_transaction(f"r{i}", tool_names=tool_names) for i in range(5)] @@ -289,7 +290,6 @@ class TestFlushToolUsageTransactions: statements = prisma.db.litellm_spendlogtoolindex.create_many.call_args_list assert [len(call.kwargs["data"]) for call in statements] == [100, 100, 50] - assert all(len(call.kwargs["data"]) <= SPEND_LOG_WRITE_BATCH_MAX_ROWS for call in statements) assert all(call.kwargs["skip_duplicates"] is True for call in statements) assert _index_rows_written(prisma) == [ (txn.request_id, tool_name) for txn in transactions for tool_name in tool_names diff --git a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py index 3acb8317c73..85f22e1f307 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py @@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock import httpx import pytest -from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_ROWS +from litellm.proxy.guardrails import usage_tracking from litellm.proxy.guardrails.usage_tracking import ( _MAX_PENDING_ROWS, PendingRollups, @@ -535,11 +535,12 @@ def _index_rows_written(prisma: MagicMock) -> list[tuple[str, str]]: @pytest.mark.asyncio -async def test_index_rows_are_written_in_row_bounded_statements(): +async def test_index_rows_are_written_in_row_bounded_statements(monkeypatch): """ LIT-5931: the drain caps logs, not logs x guardrails, so a fan-out must be split into statements the query engine can afford instead of one create_many. """ + monkeypatch.setattr(usage_tracking, "SPEND_LOG_WRITE_BATCH_MAX_ROWS", 100) prisma = _prisma() guardrail_ids = tuple(f"guard-{i}" for i in range(50)) logs = [_fan_out_payload(f"r{i}", guardrail_ids) for i in range(5)] @@ -548,13 +549,13 @@ async def test_index_rows_are_written_in_row_bounded_statements(): statements = prisma.db.litellm_spendlogguardrailindex.create_many.call_args_list assert [len(call.kwargs["data"]) for call in statements] == [100, 100, 50] - assert all(len(call.kwargs["data"]) <= SPEND_LOG_WRITE_BATCH_MAX_ROWS for call in statements) assert all(call.kwargs["skip_duplicates"] is True for call in statements) assert _index_rows_written(prisma) == [(f"r{i}", gid) for i in range(5) for gid in guardrail_ids] @pytest.mark.asyncio -async def test_one_failing_index_statement_does_not_drop_the_others_or_the_rollup(): +async def test_one_failing_index_statement_does_not_drop_the_others_or_the_rollup(monkeypatch): + monkeypatch.setattr(usage_tracking, "SPEND_LOG_WRITE_BATCH_MAX_ROWS", 100) prisma = _prisma() prisma.db.litellm_spendlogguardrailindex.create_many.side_effect = [None, httpx.ReadTimeout("ambiguous"), None] guardrail_ids = tuple(f"guard-{i}" for i in range(50))