mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
b067e836f8
commit
061c25b5ca
2 changed files with 78 additions and 24 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue