From 061c25b5cac15212ded745c51e3299aa4d43f056 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 21:25:03 -0700 Subject: [PATCH] fix(spend): let a batch's charge survive an older proxy's $0 poll row A proxy running the old code wrote _batch_cost at $0 every time it polled a batch that was still running, so after an upgrade the claim found that row and read it as proof the batch had already been charged. Only a row that recorded a charge counts now, which leaves those $0 rows, and any row a client planted under the batch id, to be charged over disable_spend_logs skipped the claim entirely, so under that setting every retrieve of a finished batch charged again. The claim now runs either way and writes the one row per batch that makes the charge exactly once, while the per-request logs stay off --- litellm/proxy/db/db_spend_update_writer.py | 42 +++++++------ .../proxy/db/test_db_spend_update_writer.py | 60 ++++++++++++++++--- 2 files changed, 78 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 3fad351224b..48312c025dd 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -271,9 +271,12 @@ class DBSpendUpdateWriter: if team_id is not None and team_id != "": payload["team_id"] = team_id + if not await self._record_spend_log( + payload=payload, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs + ): + return False + if disable_spend_logs is False: - if not await self._record_spend_log(payload=payload, prisma_client=prisma_client): - return False await self._enqueue_tool_usage_transaction( payload=payload, completion_response=completion_response, @@ -328,19 +331,23 @@ class DBSpendUpdateWriter: ) return True - async def _record_spend_log(self, payload: SpendLogsPayload, prisma_client: "PrismaClient | None") -> bool: - if prisma_client is None or not _is_batch_cost_row(payload): + async def _record_spend_log( + self, payload: SpendLogsPayload, prisma_client: "PrismaClient | None", disable_spend_logs: bool + ) -> bool: + if prisma_client is not None and _is_batch_cost_row(payload): + return await self._claim_batch_cost_spend_log(payload=payload, prisma_client=prisma_client) + if disable_spend_logs is False: await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client) - return True - return await self._claim_batch_cost_spend_log(payload=payload, prisma_client=prisma_client) + return True async def _claim_batch_cost_spend_log(self, payload: SpendLogsPayload, prisma_client: "PrismaClient") -> bool: """Write the batch's cost row now, or learn that another retrieve already did. Every retrieve of one batch shares this row, so the insert that lands first owns the charge and every later one finds the row and charges nothing (LIT-7048). Only - a row a successful retrieve wrote counts: a failed retrieve, or any request whose - client picked the batch id as its call id, cannot take the charge away. + a row that recorded a charge counts: a failed retrieve, a request whose client + picked the batch id as its call id, and the $0 row an older proxy left behind + while the batch was still running all leave the charge to be made. """ from litellm.repositories.table_repositories import SpendLogsRepository @@ -362,17 +369,18 @@ class DBSpendUpdateWriter: ) await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client) return True - if ( - existing is not None - and existing.call_type == CallTypes.aretrieve_batch.value - and existing.status == "success" - ): + if existing is None or existing.call_type != CallTypes.aretrieve_batch.value or existing.status != "success": + verbose_proxy_logger.warning( + "Spend row %s belongs to a %s request, so this batch's cost is charged without a row of its own", + request_id, + getattr(existing, "call_type", None), + ) + return True + if existing.spend > 0: verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id) return False - verbose_proxy_logger.warning( - "Spend row %s belongs to a %s request, so this batch's cost is charged without a row of its own", - request_id, - getattr(existing, "call_type", None), + verbose_proxy_logger.debug( + "Spend row %s charged nothing for this batch, so this retrieve charges it", request_id ) return True 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 500a0e7bb06..41f2be08545 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 @@ -2977,10 +2977,12 @@ def _spend_logs_prisma(inserted: int, existing: object) -> MagicMock: return prisma -async def _update_database_with(db_writer: DBSpendUpdateWriter, prisma: MagicMock, payload: dict) -> bool: +async def _update_database_with( + db_writer: DBSpendUpdateWriter, prisma: MagicMock, payload: dict, disable_spend_logs: bool = False +) -> bool: with ( patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam - "litellm.proxy.proxy_server.disable_spend_logs", False + "litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs ), patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam "litellm.proxy.proxy_server.prisma_client", prisma @@ -3014,14 +3016,16 @@ async def _update_database_with(db_writer: DBSpendUpdateWriter, prisma: MagicMoc ("inserted", "existing", "charged"), [ (1, None, True), - (0, SimpleNamespace(call_type="aretrieve_batch", status="success"), False), - (0, SimpleNamespace(call_type="aretrieve_batch", status="failure"), True), - (0, SimpleNamespace(call_type="aembedding", status="success"), True), + (0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25), False), + (0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0), True), + (0, SimpleNamespace(call_type="aretrieve_batch", status="failure", spend=0.0), True), + (0, SimpleNamespace(call_type="aembedding", status="success", spend=0.25), True), (0, None, True), ], ids=[ "first_retrieve_owns_the_row", "another_retrieve_already_charged", + "an_older_proxy_left_a_zero_row_while_the_batch_ran", "failed_retrieve_holds_the_row", "client_chosen_call_id_holds_the_row", "row_gone_between_insert_and_lookup", @@ -3033,8 +3037,9 @@ async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote """ Every retrieve of one batch shares one spend row, so the insert that lands first is the charge and every later retrieve must leave the counters alone (LIT-7048). A row - written by anything but a successful retrieve, say a request whose client picked the - batch id as its call id, must not be able to take the charge away. + that recorded no charge must not be able to take the charge away: neither one a + client planted under the batch id, nor the $0 row a pre-upgrade proxy wrote every + time it polled the batch while it was still running. """ db_writer = DBSpendUpdateWriter() db_writer._batch_database_updates = AsyncMock() @@ -3049,6 +3054,47 @@ async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote assert db_writer._batch_database_updates.await_count == (1 if charged else 0) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("inserted", "existing", "charged"), + [ + (1, None, True), + (0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25), False), + ], + ids=["first_retrieve_owns_the_row", "another_retrieve_already_charged"], +) +async def test_update_database_charges_a_batch_once_even_with_spend_logs_disabled( + inserted: int, existing: object, charged: bool +): + """ + disable_spend_logs drops the per-request logs, not the batch's charge, so the one row + that makes a batch chargeable exactly once is still written and still read back. + """ + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(inserted, existing) + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload(), True) is charged + + assert prisma.db.litellm_spendlogs.create_many.await_count == 1 + assert db_writer._batch_database_updates.await_count == (1 if charged else 0) + + +@pytest.mark.asyncio +async def test_update_database_writes_no_ordinary_spend_row_with_spend_logs_disabled(): + """The batch carve-out above stays a carve-out: every other row still goes unwritten.""" + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(1, None) + payload = {**_batch_cost_payload(), "call_type": "acompletion"} + + assert await _update_database_with(db_writer, prisma, payload, True) is True + + prisma.db.litellm_spendlogs.create_many.assert_not_called() + assert prisma.spend_log_transactions == [] + assert db_writer._batch_database_updates.await_count == 1 + + @pytest.mark.asyncio async def test_update_database_queues_a_batch_cost_row_it_could_not_claim(): """An unreachable DB must not drop the batch's only spend row, nor its charge."""