fix(spend): let a batch's charge survive an older proxy's $0 poll row

A proxy running the old code wrote <batch id>_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
This commit is contained in:
mateo-berri 2026-09-05 21:25:03 -07:00
parent b067e836f8
commit 061c25b5ca
2 changed files with 78 additions and 24 deletions

View file

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

View file

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