From 4c00a6e189d95f34a72036802937b38769f561b5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 16:35:32 -0700 Subject: [PATCH 1/7] fix(batches): account a batch's cost once, from the first retrieve that sees it final Every retrieve of a batch through the proxy shares one spend row, the batch id plus the batch cost suffix, and spend log inserts skip duplicates. A poll that landed while the batch was still validating or in progress wrote that row at $0 and no later retrieve could overwrite it, and every completed retrieve after the first added the cost to the key, team, and user counters again with no new row to show for it. The cost callback now writes nothing for a batch retrieve until the batch is final, releasing the poll's budget reservation instead, and once it is final it charges only when no spend row for that batch is queued for flush or already stored. Batch cost rows are flushed to the database right away so a second instance sees them, and the logger prices a batch only once it is final, which also covers a failed batch that never produced an output file. --- litellm/batches/batch_utils.py | 20 ++ litellm/litellm_core_utils/litellm_logging.py | 13 +- litellm/proxy/db/db_spend_update_writer.py | 3 +- .../proxy/hooks/proxy_track_cost_callback.py | 56 ++++- .../openai_files_endpoints/common_utils.py | 8 +- litellm/proxy/utils.py | 10 +- .../test_litellm/batches/test_batch_utils.py | 55 +++- .../test_litellm_logging.py | 82 ++++++ .../proxy/db/test_db_spend_update_writer.py | 10 +- .../hooks/test_proxy_track_cost_callback.py | 238 +++++++++++++----- .../prisma_and_spend/test_spend_functions.py | 15 ++ 11 files changed, 428 insertions(+), 82 deletions(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 97be5f77d79..eaac3bf0e9f 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -25,6 +25,26 @@ class BatchCostUsageResult: failed_requests: int +_TERMINAL_BATCH_STATUSES: Final = frozenset({"completed", "failed", "cancelled", "expired"}) + + +def batch_cost_is_final(batch: Batch) -> bool: + """Whether this retrieve of the batch is the one to account its cost from. + + A batch still in flight has nothing to price, and a "completed" batch can report + no output_file_id for a moment before the output populates; pricing either records + $0 under the batch's single spend row and pins it there. Final means a completed + batch whose output file has arrived or whose counts prove no line succeeded, or + any other terminal status (failed, cancelled, expired). + """ + if batch.status not in _TERMINAL_BATCH_STATUSES: + return False + if batch.status != "completed" or batch.output_file_id is not None: + return True + request_counts: Final = batch.request_counts + return request_counts is not None and request_counts.total > 0 and request_counts.completed == 0 + + async def calculate_batch_cost_and_usage( file_content_dictionary: list[dict], custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"], diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index c31c4323157..09ddd1b9720 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -36,7 +36,7 @@ from litellm._logging import ( verbose_logger, ) from litellm._uuid import uuid -from litellm.batches.batch_utils import _handle_completed_batch +from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final from litellm.caching.caching import DualCache, InMemoryCache from litellm.caching.caching_handler import LLMCachingHandler from litellm.constants import ( @@ -2899,13 +2899,6 @@ class Logging(LiteLLMLoggingBaseClass): ): # polling job will query these frequently, don't spam db logs return - from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, - ) - - # check if file id is a unified file id - is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(result.id) - batch_cost: Final = kwargs.get("batch_cost", None) batch_usage = kwargs.get("batch_usage", None) batch_models = kwargs.get("batch_models", None) @@ -2913,9 +2906,7 @@ class Logging(LiteLLMLoggingBaseClass): batch_failed_requests: Final = kwargs.get("batch_failed_requests", None) has_explicit_batch_data: Final = all(x is not None for x in (batch_cost, batch_usage, batch_models)) - should_compute_batch_data: Final = ( - not is_base64_unified_file_id or not has_explicit_batch_data and result.status == "completed" - ) + should_compute_batch_data: Final = not has_explicit_batch_data and batch_cost_is_final(result) if has_explicit_batch_data: result._hidden_params["response_cost"] = batch_cost result._hidden_params["batch_models"] = batch_models diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e6880d521f1..ff48b00dc70 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -82,6 +82,7 @@ else: RESPONSES_SESSION_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value}) +IMMEDIATE_FLUSH_CALL_TYPES: Final = RESPONSES_SESSION_CALL_TYPES | frozenset({CallTypes.aretrieve_batch.value}) class _SpendBatch(Protocol): @@ -939,7 +940,7 @@ class DBSpendUpdateWriter: from litellm.proxy.utils import enqueue_spend_logs, request_spend_log_flush await enqueue_spend_logs(prisma_client, (payload,)) - if payload.get("call_type") in RESPONSES_SESSION_CALL_TYPES: + if payload.get("call_type") in IMMEDIATE_FLUSH_CALL_TYPES: request_spend_log_flush() else: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 7254b05db2e..95e61fdd98c 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm._logging import verbose_proxy_logger +from litellm.batches.batch_utils import batch_cost_is_final from litellm.constants import BACKGROUND_INTERACTION_COST_POLLING_ENABLED from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import ( @@ -33,17 +34,21 @@ from litellm.proxy.spend_tracking.spend_log_error_logger import ( from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_error_information_for_spend_logs, get_request_model_access_groups, + get_spend_logs_id, ) from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( CallTypes, + LiteLLMBatch, StandardLoggingPayload, StandardLoggingPayloadErrorInformation, ) from litellm.utils import get_end_user_id_for_cost_tracking if TYPE_CHECKING: - from litellm.proxy.utils import ProxyLogging + from prisma.types import LiteLLM_SpendLogsWhereUniqueInput + + from litellm.proxy.utils import PrismaClient, ProxyLogging _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset( { @@ -224,6 +229,7 @@ class _ProxyDBLogger(CustomLogger): ): from litellm.proxy.proxy_server import ( increment_spend_counters, + prisma_client, proxy_logging_obj, update_cache, ) @@ -248,6 +254,18 @@ class _ProxyDBLogger(CustomLogger): ) _write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata) budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata) + if ( + isinstance(completion_response, LiteLLMBatch) + and kwargs.get("call_type") == CallTypes.aretrieve_batch.value + ): + batch_spend_log_id: Final = get_spend_logs_id( + CallTypes.aretrieve_batch.value, completion_response.model_dump(), kwargs + ) + if not await _batch_cost_is_trackable_now( + batch=completion_response, spend_log_id=batch_spend_log_id, prisma_client=prisma_client + ): + await _release_budget_reservation(budget_reservation=budget_reservation) + return user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None)) team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None)) org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None)) @@ -491,6 +509,42 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: bucket[key] = value +async def _batch_cost_is_trackable_now( + batch: LiteLLMBatch, spend_log_id: str | None, prisma_client: "PrismaClient | None" +) -> bool: + """A batch is billed exactly once, from the first retrieve that sees it final. + + Every retrieve of one batch shares a single spend row (its id plus the batch cost + suffix), so a poll that lands before the output exists would write that row at $0 + and pin it there, and every retrieve after the first would add the cost to the + key, team, and user counters again. + """ + if not batch_cost_is_final(batch): + verbose_proxy_logger.debug("Cost tracking deferred for batch %s still in status %s", batch.id, batch.status) + return False + if prisma_client is None or spend_log_id is None: + return True + if not await _spend_log_already_recorded(prisma_client=prisma_client, request_id=spend_log_id): + return True + verbose_proxy_logger.debug( + "Cost tracking skipped for batch %s: spend row %s already recorded", batch.id, spend_log_id + ) + return False + + +async def _spend_log_already_recorded(prisma_client: "PrismaClient", request_id: str) -> bool: + from litellm.proxy.utils import spend_log_is_queued + + if await spend_log_is_queued(prisma_client, request_id): + return True + spend_log_row: Final[LiteLLM_SpendLogsWhereUniqueInput] = {"request_id": request_id} + try: + return await prisma_client.db.litellm_spendlogs.find_unique(where=spend_log_row) is not None + except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; an unreadable DB must not drop the batch's only spend row + verbose_proxy_logger.warning("Could not check for an existing spend row %s, tracking anyway: %s", request_id, e) + return False + + def _is_unbilled_interaction_response(completion_response: object) -> bool: from litellm.interactions.background_cost_polling import missing_usage_is_expected from litellm.types.interactions import InteractionsAPIResponse diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 15eeddbc489..b1f282a0978 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -15,6 +15,7 @@ from typing import ( runtime_checkable, ) +from litellm.batches.batch_utils import batch_cost_is_final from litellm.proxy._types import ProxyException from litellm.repositories.table_repositories import ( ManagedFileRepository, @@ -1357,12 +1358,7 @@ def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool: enumerated the batch and none succeeded. A zero or unknown total means counts are unreported, so stay eligible and let the next poller pass revisit it. (#37713) """ - if response.output_file_id is not None: - return True - request_counts = response.request_counts - if request_counts is None: - return False - return request_counts.total > 0 and request_counts.completed == 0 + return batch_cost_is_final(response) async def update_batch_in_database( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index accf7b720fb..f64b51bc6c6 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -6251,7 +6251,9 @@ def request_spend_log_flush() -> None: The Responses API hands the client an id it can chain from straight away, and that lookup reads the DB, so the row cannot sit in this worker's queue for a poll interval. - Repeated requests coalesce into the monitor's next pass, so the batching holds. + A batch's cost row is what every other worker checks before charging the same batch + again, so it cannot wait either. Repeated requests coalesce into the monitor's next + pass, so the batching holds. """ PrismaClient.spend_log_flush_requested.set() @@ -6266,6 +6268,12 @@ async def _wait_for_spend_log_flush_request(interval: float) -> bool: return True +async def spend_log_is_queued(prisma_client: PrismaClient, request_id: str) -> bool: + """Whether a spend log with ``request_id`` is still waiting for the next flush.""" + async with prisma_client._spend_log_transactions_lock: + return any(row.get("request_id") == request_id for row in prisma_client.spend_log_transactions) + + async def dequeue_spend_logs(prisma_client: PrismaClient, limit: int) -> list[dict[str, object]]: """Take up to ``limit`` of the oldest queued spend logs off the queue. diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index c86c7c4df03..8d4f68164b4 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -21,11 +21,12 @@ from types import MappingProxyType import httpx import pytest import respx +from openai.types.batch import BatchRequestCounts import litellm import litellm.batches.batch_utils as bu -from litellm.types.utils import Usage +from litellm.types.utils import LiteLLMBatch, Usage # --------------------------------------------------------------------------- # # Builders for batch OUTPUT file rows. @@ -1718,3 +1719,55 @@ def test_unparsable_bedrock_batch_usage_warns(caplog): assert usage.total_tokens == 0 assert "does not understand" in caplog.text assert "inputTextTokenCount" in caplog.text + + +# --------------------------------------------------------------------------- # +# batch_cost_is_final +# --------------------------------------------------------------------------- # + +def _retrieved_batch( + status: str, output_file_id: str | None = None, counts: BatchRequestCounts | None = None +) -> LiteLLMBatch: + return LiteLLMBatch( + id="batch_abc", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="file-in", + object="batch", + status=status, + output_file_id=output_file_id, + request_counts=counts, + ) + + +class TestBatchCostIsFinal: + """Every retrieve of one batch writes the same spend row, so the first retrieve + that prices it decides the row for good. A poll before the output exists must + therefore not count as final: pricing it recorded $0 and pinned it (LIT-7048).""" + + @pytest.mark.parametrize("status", ["validating", "in_progress", "finalizing", "cancelling"]) + def test_in_flight_batch_is_not_final(self, status): + assert bu.batch_cost_is_final(_retrieved_batch(status)) is False + + def test_completed_with_output_is_final(self): + assert bu.batch_cost_is_final(_retrieved_batch("completed", output_file_id="file-out")) is True + + def test_completed_without_output_and_unknown_counts_is_not_final(self): + assert bu.batch_cost_is_final(_retrieved_batch("completed")) is False + + def test_completed_without_output_and_zero_counts_is_not_final(self): + counts = BatchRequestCounts(total=0, completed=0, failed=0) + assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False + + def test_completed_without_output_but_successful_lines_is_not_final(self): + counts = BatchRequestCounts(total=2, completed=2, failed=0) + assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False + + def test_completed_without_output_and_every_line_failed_is_final(self): + counts = BatchRequestCounts(total=2, completed=0, failed=2) + assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is True + + @pytest.mark.parametrize("status", ["failed", "expired", "cancelled"]) + def test_other_terminal_statuses_are_final(self, status): + assert bu.batch_cost_is_final(_retrieved_batch(status)) is True diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 16a99713a06..4583429bd11 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -632,6 +632,88 @@ class TestRetrieveBatchCostPassesModelIdentity: assert captured["model_info"]["input_cost_per_token"] == 0.0 +class TestRetrieveBatchPricesOnlyFinalBatches: + """Regression (LIT-7048): retrieving a provider-id batch priced it on every poll. + + Every retrieve of one batch logs under the same spend row, so pricing a poll + that landed before the output existed wrote that row at $0 and pinned it there. + Only a final batch gets priced; an in-flight poll carries no cost at all. + """ + + @staticmethod + def _logging_obj() -> LitellmLogging: + obj = LitellmLogging( + model="gpt-5.6-luna", + messages=[{"role": "user", "content": "Hey"}], + stream=False, + call_type="aretrieve_batch", + start_time=time.time(), + litellm_call_id="batch-call-2", + function_id="f", + ) + obj.custom_llm_provider = "openai" + return obj + + @staticmethod + def _batch(status: str, output_file_id: str | None): + from litellm.types.utils import LiteLLMBatch + + return LiteLLMBatch( + id="batch_6a9c99e185588190877d391f8b9d7f8a", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="file-in", + object="batch", + status=status, + output_file_id=output_file_id, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("status", "output_file_id"), + [("validating", None), ("in_progress", None), ("finalizing", None), ("completed", None)], + ) + async def test_non_final_batch_is_not_priced(self, monkeypatch, status, output_file_id) -> None: + from litellm.litellm_core_utils import litellm_logging as logging_module + + handle_completed_batch = AsyncMock() + monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) + batch = self._batch(status, output_file_id) + + with contextlib.suppress(Exception): + await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) + + handle_completed_batch.assert_not_awaited() + assert "response_cost" not in batch._hidden_params + + @pytest.mark.asyncio + async def test_completed_batch_with_output_is_priced(self, monkeypatch) -> None: + from litellm.batches.batch_utils import BatchCostUsageResult + from litellm.litellm_core_utils import litellm_logging as logging_module + from litellm.types.utils import Usage + + handle_completed_batch = AsyncMock( + return_value=BatchCostUsageResult( + cost=8e-06, + usage=Usage(prompt_tokens=26, completion_tokens=9, total_tokens=35), + models=["gpt-5.6-luna"], + successful_requests=2, + failed_requests=0, + ) + ) + monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) + batch = self._batch("completed", "file-out") + + with contextlib.suppress(Exception): + await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) + + handle_completed_batch.assert_awaited_once() + assert batch._hidden_params["response_cost"] == 8e-06 + assert batch.usage is not None + assert batch.usage.total_tokens == 35 + + class TestAnthropicPassthroughCustomPricing: """Verify the Anthropic pass-through handler forwards custom pricing.""" 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 11ef911de3e..e1b2d151c6d 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 @@ -2934,12 +2934,16 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey @pytest.mark.asyncio @pytest.mark.parametrize( "call_type, expects_flush", - [("aresponses", True), ("responses", True), ("acompletion", False)], + [("aresponses", True), ("responses", True), ("aretrieve_batch", True), ("acompletion", False)], ) -async def test_insert_spend_log_asks_for_an_immediate_flush_on_responses_calls(call_type: str, expects_flush: bool): +async def test_insert_spend_log_asks_for_an_immediate_flush_on_rows_other_workers_read_back( + call_type: str, expects_flush: bool +): """ A `previous_response_id` chained straight off the previous turn reads the DB, so a - Responses row cannot sit in this worker's queue until the monitor's next poll. + Responses row cannot sit in this worker's queue until the monitor's next poll. A + batch's cost row is what another worker checks before charging the same batch again + (LIT-7048), so it cannot wait either. """ from litellm.proxy.utils import PrismaClient diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 8043a1aca3f..b7037e8d621 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1,4 +1,3 @@ - from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch @@ -7,6 +6,7 @@ import pytest from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.proxy_track_cost_callback import ( + _batch_cost_is_trackable_now, _get_budget_reservation_from_metadata, _ProxyDBLogger, _should_track_cost_callback, @@ -70,9 +70,7 @@ async def test_async_post_call_failure_hook(): # Check that metadata was properly updated assert "litellm_params" in call_args["kwargs"] - assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == { - "request_id": "test_request_id" - } + assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == {"request_id": "test_request_id"} metadata = call_args["kwargs"]["litellm_params"]["metadata"] assert metadata["user_api_key"] == "test_api_key" assert metadata["status"] == "failure" @@ -336,9 +334,7 @@ async def test_should_continue_failure_tracking_when_budget_release_fails(): ) assert mock_invalidate_budget_reservation_counters.await_count == 1 assert ( - mock_invalidate_budget_reservation_counters.await_args.kwargs[ - "budget_reservation" - ] + mock_invalidate_budget_reservation_counters.await_args.kwargs["budget_reservation"] is user_api_key_dict.budget_reservation ) assert user_api_key_dict.budget_reservation["finalized"] is True @@ -433,36 +429,21 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object(): "entries": [{"counter_key": "spend:key:test_api_key"}], } + assert _get_budget_reservation_from_metadata(metadata={"user_api_key_auth": dict(UserAPIKeyAuth())}) is None assert ( _get_budget_reservation_from_metadata( - metadata={"user_api_key_auth": dict(UserAPIKeyAuth())} - ) - is None - ) - assert ( - _get_budget_reservation_from_metadata( - metadata={ - "user_api_key_auth": UserAPIKeyAuth( - budget_reservation=budget_reservation - ) - } + metadata={"user_api_key_auth": UserAPIKeyAuth(budget_reservation=budget_reservation)} ) == budget_reservation ) assert ( _get_budget_reservation_from_metadata( - metadata={ - "user_api_key_auth": dict( - UserAPIKeyAuth(budget_reservation=budget_reservation) - ) - } + metadata={"user_api_key_auth": dict(UserAPIKeyAuth(budget_reservation=budget_reservation))} ) == budget_reservation ) assert ( - _get_budget_reservation_from_metadata( - metadata={"user_api_key_budget_reservation": budget_reservation} - ) + _get_budget_reservation_from_metadata(metadata={"user_api_key_budget_reservation": budget_reservation}) is budget_reservation ) @@ -470,9 +451,7 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object(): @pytest.mark.asyncio async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails(): proxy_logging_obj = MagicMock() - proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( - side_effect=Exception("db unavailable") - ) + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=Exception("db unavailable")) increment_spend_counters = AsyncMock() budget_reservation = {"reserved_cost": 0.5, "entries": []} @@ -508,9 +487,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_db_u async def test_update_database_and_spend_counters_preserves_db_exception_when_release_fails(): proxy_logging_obj = MagicMock() db_exception = RuntimeError("db unavailable") - proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( - side_effect=db_exception - ) + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=db_exception) increment_spend_counters = AsyncMock() budget_reservation = {"reserved_cost": 0.5, "entries": []} @@ -554,12 +531,8 @@ async def test_update_database_and_spend_counters_preserves_db_exception_when_re budget_reservation=budget_reservation, ) assert mock_log_exception.call_count == 2 - mock_log_exception.assert_any_call( - "Failed to release budget reservation after database update failed" - ) - mock_log_exception.assert_any_call( - "Failed to invalidate budget reservation counters after release failed" - ) + mock_log_exception.assert_any_call("Failed to release budget reservation after database update failed") + mock_log_exception.assert_any_call("Failed to invalidate budget reservation counters after release failed") increment_spend_counters.assert_not_awaited() @@ -778,6 +751,169 @@ async def test_track_cost_callback_defers_in_progress_background_interaction(): mock_proxy_logging.failed_tracking_alert.assert_not_called() +def _batch_retrieve_kwargs(call_type: str, reservation: dict | None = None) -> dict: + metadata = { + "user_api_key": "hashed_key", + "user_api_key_user_id": "user-1", + "user_api_key_team_id": "team-1", + **({"user_api_key_budget_reservation": reservation} if reservation is not None else {}), + } + return { + "call_type": call_type, + "model": "gpt-5.6-luna", + "litellm_call_id": "test-call-id", + "litellm_params": {"metadata": metadata}, + "standard_logging_object": {"response_cost": 0.0, "request_tags": None}, + "stream": False, + } + + +def _retrieved_batch(status: str, output_file_id: str | None): + from litellm.types.utils import LiteLLMBatch + + return LiteLLMBatch( + id="batch_abc", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="file-in", + object="batch", + status=status, + output_file_id=output_file_id, + ) + + +def _prisma_client_with(queued_request_ids: tuple[str, ...], stored_row: object) -> MagicMock: + import asyncio + + prisma_client = MagicMock() + prisma_client._spend_log_transactions_lock = asyncio.Lock() + prisma_client.spend_log_transactions = [{"request_id": request_id} for request_id in queued_request_ids] + prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(return_value=stored_row) + return prisma_client + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "output_file_id", "spend_log_id", "prisma_client", "trackable"), + [ + ("in_progress", None, "batch_abc_batch_cost", None, False), + ("in_progress", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), + ("completed", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), + ("completed", "file-out", "batch_abc_batch_cost", None, True), + ("completed", "file-out", None, _prisma_client_with((), None), True), + ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with(("batch_abc_batch_cost",), None), False), + ( + "completed", + "file-out", + "batch_abc_batch_cost", + _prisma_client_with((), {"request_id": "batch_abc_batch_cost"}), + False, + ), + ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with((), None), True), + ("failed", None, "batch_abc_batch_cost", _prisma_client_with((), None), True), + ], + ids=[ + "in_progress_without_db", + "in_progress_never_consults_db", + "completed_without_output_yet", + "final_without_db", + "final_without_spend_log_id", + "final_row_queued_for_flush", + "final_row_already_stored", + "final_first_sighting", + "failed_first_sighting", + ], +) +async def test_batch_cost_is_trackable_now(status, output_file_id, spend_log_id, prisma_client, trackable): + """ + A batch is billed from the first retrieve that sees it final and never again: + a poll before that wrote the shared spend row at $0 and pinned it there, and + every completed retrieve after the first charged the key again (LIT-7048). + """ + assert ( + await _batch_cost_is_trackable_now( + batch=_retrieved_batch(status, output_file_id), spend_log_id=spend_log_id, prisma_client=prisma_client + ) + is trackable + ) + + +@pytest.mark.asyncio +async def test_batch_cost_is_trackable_now_when_the_spend_row_lookup_fails(): + """An unreadable spend log table must not drop the batch's only spend row.""" + prisma_client = _prisma_client_with((), None) + prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(side_effect=RuntimeError("db unreachable")) + + assert ( + await _batch_cost_is_trackable_now( + batch=_retrieved_batch("completed", "file-out"), + spend_log_id="batch_abc_batch_cost", + prisma_client=prisma_client, + ) + is True + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("call_type", "status", "output_file_id", "stored_row", "charged"), + [ + ("aretrieve_batch", "in_progress", None, None, False), + ("aretrieve_batch", "completed", "file-out", {"request_id": "batch_abc_batch_cost"}, False), + ("aretrieve_batch", "completed", "file-out", None, True), + ("acreate_batch", "validating", None, None, True), + ], + ids=["retrieve_before_final", "retrieve_already_recorded", "retrieve_first_final", "create_before_final"], +) +async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # test-quality-ok: whether the spend writer runs and whether the poll's reservation is handed back is the whole observable contract of the gate + call_type, status, output_file_id, stored_row, charged +): + """ + Only retrieves are gated, since creating a batch is its own billable request. + A retrieve that writes nothing hands its budget reservation back instead. + """ + logger = _ProxyDBLogger() + budget_reservation = None if charged else {"reserved_cost": 0.5, "entries": []} + kwargs = _batch_retrieve_kwargs(call_type, reservation=budget_reservation) + + with ( + patch( # test-quality-ok: prisma_client is a proxy_server global the callback reads lazily, no seam + "litellm.proxy.proxy_server.prisma_client", _prisma_client_with((), stored_row) + ), + patch( # test-quality-ok: increment_spend_counters is a proxy_server global the callback reads lazily, no seam + "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock + ), + patch( # test-quality-ok: update_cache is a proxy_server global the callback reads lazily, no seam + "litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock + ), + patch( # test-quality-ok: callback imports proxy_logging_obj off proxy_server in its body, no seam + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_proxy_logging, + patch( # test-quality-ok: the release is imported inside the callback's helper, no seam + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", new_callable=AsyncMock + ) as mock_release_budget_reservation, + ): + mock_proxy_logging.failed_tracking_alert = AsyncMock() + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() + mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=_retrieved_batch(status, output_file_id), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + mock_proxy_logging.failed_tracking_alert.assert_not_called() + if charged: + mock_proxy_logging.db_spend_update_writer.update_database.assert_awaited_once() + mock_release_budget_reservation.assert_not_awaited() + else: + mock_proxy_logging.db_spend_update_writer.update_database.assert_not_called() + mock_release_budget_reservation.assert_awaited_once_with(budget_reservation=budget_reservation) + + def _in_progress_interaction_kwargs(reservation: dict) -> dict: return { "call_type": "acreate_interaction", @@ -1101,10 +1237,7 @@ async def test_async_post_call_failure_hook_propagates_trace_id_from_logging_obj # standard_logging_object should have been propagated from logging obj assert call_kwargs.get("standard_logging_object") is not None - assert ( - call_kwargs["standard_logging_object"]["trace_id"] - == "trace-id-from-logging-obj" - ) + assert call_kwargs["standard_logging_object"]["trace_id"] == "trace-id-from-logging-obj" # litellm_trace_id should also be propagated as a fallback assert call_kwargs.get("litellm_trace_id") == "trace-id-from-logging-obj" @@ -1691,9 +1824,7 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend(): "metadata": {}, "proxy_server_request": {"request_id": "rid"}, "response_cost": 3.5e-05, - "combined_usage_object": Usage( - prompt_tokens=30, completion_tokens=1, total_tokens=31 - ), + "combined_usage_object": Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31), } with patch( @@ -1772,15 +1903,10 @@ async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata(): assert mock_increment.call_args.kwargs["team_id"] == "team-123" assert mock_increment.call_args.kwargs["org_id"] == "org-456" - update_kwargs = ( - mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs - ) + update_kwargs = mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs assert update_kwargs["user_id"] == "mcp-user@example.com" assert update_kwargs["team_id"] == "team-123" - assert ( - kwargs["litellm_params"]["metadata"]["user_api_key_user_id"] - == "mcp-user@example.com" - ) + assert kwargs["litellm_params"]["metadata"]["user_api_key_user_id"] == "mcp-user@example.com" @pytest.mark.parametrize( @@ -1828,9 +1954,7 @@ def test_should_track_cost_callback_pass_through_without_owner(call_type, expect ], ) @pytest.mark.asyncio -async def test_track_cost_callback_logs_unauthenticated_pass_through_request( - call_type, expect_spend_log -): +async def test_track_cost_callback_logs_unauthenticated_pass_through_request(call_type, expect_spend_log): """Regression for LIT-3782: a pass-through request with auth=false reaches the cost callback with no key/user/team/end-user. Before the fix the spend-log write was skipped and the request never appeared in request/usage logs. It @@ -1876,9 +2000,7 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request( end_time=datetime.now(), ) - assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == ( - 1 if expect_spend_log else 0 - ) + assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (1 if expect_spend_log else 0) class _FakeDeploymentLookup: diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py index a1eb88a7834..fc97d760226 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -6,6 +6,7 @@ Symbols pinned here: - ``update_spend_logs_job`` - ``_monitor_spend_logs_queue`` - ``_raise_failed_update_spend_exception`` + - ``spend_log_is_queued`` """ from __future__ import annotations @@ -22,6 +23,7 @@ from litellm.proxy.utils import ( _monitor_spend_logs_queue, _raise_failed_update_spend_exception, drain_spend_logs_queue, + spend_log_is_queued, update_daily_tag_spend, update_spend, update_spend_logs_job, @@ -629,3 +631,16 @@ def test_raise_failed_update_spend_exception_raises_original_error() -> None: with pytest.raises(ValueError, match="specific"): asyncio.run(_runner()) + + +@pytest.mark.asyncio +async def test_spend_log_is_queued_matches_only_rows_awaiting_flush( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="batch_abc_batch_cost")] + + assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is True + assert await spend_log_is_queued(mock_prisma_client, "batch_abc") is False + + mock_prisma_client.spend_log_transactions = [] + assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is False From b067e836f8df15bb3dbefe345bf503b7d3811b9a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 18:58:11 -0700 Subject: [PATCH 2/7] fix(batches): claim the batch cost spend row in the database before charging The cost callback used to look for an existing `_batch_cost` row before charging a completed batch, which left a window where concurrent retrieves on any instance all charged the key, and it would honor a row any request had written under that id. The spend update writer now inserts the batch cost row itself with `create_many(skip_duplicates=True)` and only the retrieve whose insert lands charges the key, team, and user. An existing row only takes the charge when it is a successful `aretrieve_batch` row, so a client-chosen `x-litellm-call-id` on another endpoint cannot suppress billing. Batch cost rows no longer get their own immediate flush path `batch_cost_is_final` now treats the proxy's normalized `complete` status like `completed`, which the enterprise batch cost poller relies on when it decides whether a completed batch is safe to retire. Tests build that status with `model_copy` since the OpenAI `Batch` model rejects it The `test-quality-ok` markers sit on the `patch(` lines the gate keys on, and the logging tests no longer wrap the priced retrieve in `contextlib.suppress` --- litellm/batches/batch_utils.py | 5 +- litellm/proxy/db/db_spend_update_writer.py | 73 ++++++++-- .../proxy/hooks/proxy_track_cost_callback.py | 68 +++------- litellm/proxy/utils.py | 10 +- .../test_litellm/batches/test_batch_utils.py | 14 +- .../test_litellm_logging.py | 12 +- .../proxy/db/test_db_spend_update_writer.py | 128 +++++++++++++++++- .../hooks/test_proxy_track_cost_callback.py | 116 ++++------------ .../prisma_and_spend/test_spend_functions.py | 15 -- 9 files changed, 249 insertions(+), 192 deletions(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index eaac3bf0e9f..959c7498479 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -25,7 +25,8 @@ class BatchCostUsageResult: failed_requests: int -_TERMINAL_BATCH_STATUSES: Final = frozenset({"completed", "failed", "cancelled", "expired"}) +_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"}) +_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"}) def batch_cost_is_final(batch: Batch) -> bool: @@ -39,7 +40,7 @@ def batch_cost_is_final(batch: Batch) -> bool: """ if batch.status not in _TERMINAL_BATCH_STATUSES: return False - if batch.status != "completed" or batch.output_file_id is not None: + if batch.status not in _COMPLETED_BATCH_STATUSES or batch.output_file_id is not None: return True request_counts: Final = batch.request_counts return request_counts is not None and request_counts.total > 0 and request_counts.completed == 0 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ff48b00dc70..3fad351224b 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -82,7 +82,10 @@ else: RESPONSES_SESSION_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value}) -IMMEDIATE_FLUSH_CALL_TYPES: Final = RESPONSES_SESSION_CALL_TYPES | frozenset({CallTypes.aretrieve_batch.value}) + + +def _is_batch_cost_row(payload: SpendLogsPayload) -> bool: + return payload.get("call_type") == CallTypes.aretrieve_batch.value and payload.get("status") == "success" class _SpendBatch(Protocol): @@ -216,7 +219,12 @@ class DBSpendUpdateWriter: start_time: datetime | None, end_time: datetime | None, response_cost: float | None, - ) -> None: + ) -> bool: + """Record the request's spend, answering whether its cost still needs charging. + + False only for a batch retrieve whose cost row another retrieve already wrote, + so the caller leaves the key, team, and user counters alone (LIT-7048). + """ from litellm.proxy.proxy_server import ( disable_spend_logs, litellm_proxy_budget_name, @@ -233,7 +241,7 @@ class DBSpendUpdateWriter: team_id, ) if ProxyUpdateSpend.disable_spend_updates() is True: - return + return True if token is not None and isinstance(token, str) and token.startswith("sk-"): hashed_token = hash_token(token=token) else: @@ -264,10 +272,8 @@ class DBSpendUpdateWriter: payload["team_id"] = team_id if disable_spend_logs is False: - await self._insert_spend_log_to_db( - payload=payload, - prisma_client=prisma_client, - ) + 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, @@ -307,6 +313,7 @@ class DBSpendUpdateWriter: ) verbose_proxy_logger.debug("Runs spend update on all tables") + return True except Exception: spend_log_error( "Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue " @@ -319,7 +326,55 @@ class DBSpendUpdateWriter: org_id, end_user_id, ) - return + 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): + 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) + + 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. + """ + from litellm.repositories.table_repositories import SpendLogsRepository + + request_id: Final = payload["request_id"] + spend_logs: Final = SpendLogsRepository(prisma_client).table + try: + claimed: Final = await spend_logs.create_many( + data=[prisma_client.jsonify_object(payload)], # mutable-ok: prisma create_many takes a list + skip_duplicates=True, + ) + if claimed == 1: + return True + existing: Final = await spend_logs.find_unique( + where={"request_id": request_id} # mutable-ok: prisma where clause + ) + except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; an unreachable DB queues the row like any other spend log + verbose_proxy_logger.warning( + "Could not claim spend row %s for a batch's cost, queueing it: %s", request_id, e + ) + 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" + ): + 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), + ) + return True async def _enqueue_tool_usage_transaction( self, @@ -940,7 +995,7 @@ class DBSpendUpdateWriter: from litellm.proxy.utils import enqueue_spend_logs, request_spend_log_flush await enqueue_spend_logs(prisma_client, (payload,)) - if payload.get("call_type") in IMMEDIATE_FLUSH_CALL_TYPES: + if payload.get("call_type") in RESPONSES_SESSION_CALL_TYPES: request_spend_log_flush() else: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 95e61fdd98c..f0c889a8cb4 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -34,7 +34,6 @@ from litellm.proxy.spend_tracking.spend_log_error_logger import ( from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_error_information_for_spend_logs, get_request_model_access_groups, - get_spend_logs_id, ) from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( @@ -46,9 +45,7 @@ from litellm.types.utils import ( from litellm.utils import get_end_user_id_for_cost_tracking if TYPE_CHECKING: - from prisma.types import LiteLLM_SpendLogsWhereUniqueInput - - from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.proxy.utils import ProxyLogging _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset( { @@ -229,7 +226,6 @@ class _ProxyDBLogger(CustomLogger): ): from litellm.proxy.proxy_server import ( increment_spend_counters, - prisma_client, proxy_logging_obj, update_cache, ) @@ -257,15 +253,15 @@ class _ProxyDBLogger(CustomLogger): if ( isinstance(completion_response, LiteLLMBatch) and kwargs.get("call_type") == CallTypes.aretrieve_batch.value + and not batch_cost_is_final(completion_response) ): - batch_spend_log_id: Final = get_spend_logs_id( - CallTypes.aretrieve_batch.value, completion_response.model_dump(), kwargs + verbose_proxy_logger.debug( + "Cost tracking deferred for batch %s still in status %s", + completion_response.id, + completion_response.status, ) - if not await _batch_cost_is_trackable_now( - batch=completion_response, spend_log_id=batch_spend_log_id, prisma_client=prisma_client - ): - await _release_budget_reservation(budget_reservation=budget_reservation) - return + await _release_budget_reservation(budget_reservation=budget_reservation) + return user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None)) team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None)) org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None)) @@ -307,7 +303,7 @@ class _ProxyDBLogger(CustomLogger): call_type=call_type, ): ## UPDATE DATABASE - await _update_database_and_spend_counters( + charged: Final = await _update_database_and_spend_counters( proxy_logging_obj=proxy_logging_obj, increment_spend_counters=increment_spend_counters, user_api_key=user_api_key, @@ -324,6 +320,8 @@ class _ProxyDBLogger(CustomLogger): request_tags=tags, model_access_groups=model_access_groups, ) + if not charged: + return # update cache (fire-and-forget for backward compat: # cached object fields, soft budget alerts, etc.) @@ -509,42 +507,6 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: bucket[key] = value -async def _batch_cost_is_trackable_now( - batch: LiteLLMBatch, spend_log_id: str | None, prisma_client: "PrismaClient | None" -) -> bool: - """A batch is billed exactly once, from the first retrieve that sees it final. - - Every retrieve of one batch shares a single spend row (its id plus the batch cost - suffix), so a poll that lands before the output exists would write that row at $0 - and pin it there, and every retrieve after the first would add the cost to the - key, team, and user counters again. - """ - if not batch_cost_is_final(batch): - verbose_proxy_logger.debug("Cost tracking deferred for batch %s still in status %s", batch.id, batch.status) - return False - if prisma_client is None or spend_log_id is None: - return True - if not await _spend_log_already_recorded(prisma_client=prisma_client, request_id=spend_log_id): - return True - verbose_proxy_logger.debug( - "Cost tracking skipped for batch %s: spend row %s already recorded", batch.id, spend_log_id - ) - return False - - -async def _spend_log_already_recorded(prisma_client: "PrismaClient", request_id: str) -> bool: - from litellm.proxy.utils import spend_log_is_queued - - if await spend_log_is_queued(prisma_client, request_id): - return True - spend_log_row: Final[LiteLLM_SpendLogsWhereUniqueInput] = {"request_id": request_id} - try: - return await prisma_client.db.litellm_spendlogs.find_unique(where=spend_log_row) is not None - except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; an unreadable DB must not drop the batch's only spend row - verbose_proxy_logger.warning("Could not check for an existing spend row %s, tracking anyway: %s", request_id, e) - return False - - def _is_unbilled_interaction_response(completion_response: object) -> bool: from litellm.interactions.background_cost_polling import missing_usage_is_expected from litellm.types.interactions import InteractionsAPIResponse @@ -636,9 +598,9 @@ async def _update_database_and_spend_counters( budget_reservation: dict | None, request_tags: list[str] | None = None, model_access_groups: Sequence[str] | None = None, -) -> None: +) -> bool: try: - await proxy_logging_obj.db_spend_update_writer.update_database( + charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database( token=user_api_key, response_cost=response_cost, user_id=user_id, @@ -663,6 +625,9 @@ async def _update_database_and_spend_counters( "Failed to invalidate budget reservation counters after release failed" ) raise + if not charged: + await _release_budget_reservation(budget_reservation=budget_reservation) + return False try: await increment_spend_counters( @@ -688,6 +653,7 @@ async def _update_database_and_spend_counters( finally: budget_reservation["finalized"] = True raise + return True async def _release_budget_reservation(budget_reservation: dict | None) -> None: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f64b51bc6c6..accf7b720fb 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -6251,9 +6251,7 @@ def request_spend_log_flush() -> None: The Responses API hands the client an id it can chain from straight away, and that lookup reads the DB, so the row cannot sit in this worker's queue for a poll interval. - A batch's cost row is what every other worker checks before charging the same batch - again, so it cannot wait either. Repeated requests coalesce into the monitor's next - pass, so the batching holds. + Repeated requests coalesce into the monitor's next pass, so the batching holds. """ PrismaClient.spend_log_flush_requested.set() @@ -6268,12 +6266,6 @@ async def _wait_for_spend_log_flush_request(interval: float) -> bool: return True -async def spend_log_is_queued(prisma_client: PrismaClient, request_id: str) -> bool: - """Whether a spend log with ``request_id`` is still waiting for the next flush.""" - async with prisma_client._spend_log_transactions_lock: - return any(row.get("request_id") == request_id for row in prisma_client.spend_log_transactions) - - async def dequeue_spend_logs(prisma_client: PrismaClient, limit: int) -> list[dict[str, object]]: """Take up to ``limit`` of the oldest queued spend logs off the queue. diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 8d4f68164b4..976a96f2db1 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -1735,10 +1735,10 @@ def _retrieved_batch( endpoint="/v1/chat/completions", input_file_id="file-in", object="batch", - status=status, + status="validating", output_file_id=output_file_id, request_counts=counts, - ) + ).model_copy(update={"status": status}) class TestBatchCostIsFinal: @@ -1750,8 +1750,9 @@ class TestBatchCostIsFinal: def test_in_flight_batch_is_not_final(self, status): assert bu.batch_cost_is_final(_retrieved_batch(status)) is False - def test_completed_with_output_is_final(self): - assert bu.batch_cost_is_final(_retrieved_batch("completed", output_file_id="file-out")) is True + @pytest.mark.parametrize("status", ["completed", "complete"]) + def test_completed_with_output_is_final(self, status): + assert bu.batch_cost_is_final(_retrieved_batch(status, output_file_id="file-out")) is True def test_completed_without_output_and_unknown_counts_is_not_final(self): assert bu.batch_cost_is_final(_retrieved_batch("completed")) is False @@ -1764,9 +1765,10 @@ class TestBatchCostIsFinal: counts = BatchRequestCounts(total=2, completed=2, failed=0) assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False - def test_completed_without_output_and_every_line_failed_is_final(self): + @pytest.mark.parametrize("status", ["completed", "complete"]) + def test_completed_without_output_and_every_line_failed_is_final(self, status): counts = BatchRequestCounts(total=2, completed=0, failed=2) - assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is True + assert bu.batch_cost_is_final(_retrieved_batch(status, counts=counts)) is True @pytest.mark.parametrize("status", ["failed", "expired", "cancelled"]) def test_other_terminal_statuses_are_final(self, status): diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 4583429bd11..174efaa4679 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -665,14 +665,14 @@ class TestRetrieveBatchPricesOnlyFinalBatches: endpoint="/v1/chat/completions", input_file_id="file-in", object="batch", - status=status, + status="validating", output_file_id=output_file_id, - ) + ).model_copy(update={"status": status}) @pytest.mark.asyncio @pytest.mark.parametrize( ("status", "output_file_id"), - [("validating", None), ("in_progress", None), ("finalizing", None), ("completed", None)], + [("validating", None), ("in_progress", None), ("finalizing", None), ("completed", None), ("complete", None)], ) async def test_non_final_batch_is_not_priced(self, monkeypatch, status, output_file_id) -> None: from litellm.litellm_core_utils import litellm_logging as logging_module @@ -681,8 +681,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches: monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) batch = self._batch(status, output_file_id) - with contextlib.suppress(Exception): - await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) + await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) handle_completed_batch.assert_not_awaited() assert "response_cost" not in batch._hidden_params @@ -705,8 +704,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches: monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) batch = self._batch("completed", "file-out") - with contextlib.suppress(Exception): - await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) + await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) handle_completed_batch.assert_awaited_once() assert batch._hidden_params["response_cost"] == 8e-06 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 e1b2d151c6d..500a0e7bb06 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 @@ -7,6 +7,7 @@ import re from collections.abc import Callable from contextlib import asynccontextmanager from datetime import datetime, timezone +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, call, patch import pytest @@ -2934,16 +2935,14 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey @pytest.mark.asyncio @pytest.mark.parametrize( "call_type, expects_flush", - [("aresponses", True), ("responses", True), ("aretrieve_batch", True), ("acompletion", False)], + [("aresponses", True), ("responses", True), ("acompletion", False)], ) async def test_insert_spend_log_asks_for_an_immediate_flush_on_rows_other_workers_read_back( call_type: str, expects_flush: bool ): """ A `previous_response_id` chained straight off the previous turn reads the DB, so a - Responses row cannot sit in this worker's queue until the monitor's next poll. A - batch's cost row is what another worker checks before charging the same batch again - (LIT-7048), so it cannot wait either. + Responses row cannot sit in this worker's queue until the monitor's next poll. """ from litellm.proxy.utils import PrismaClient @@ -2961,6 +2960,127 @@ async def test_insert_spend_log_asks_for_an_immediate_flush_on_rows_other_worker PrismaClient.spend_log_flush_requested.clear() +def _batch_cost_payload() -> dict: + return { + **_minimal_spend_payload(), + "request_id": "batch_abc_batch_cost", + "call_type": "aretrieve_batch", + "status": "success", + } + + +def _spend_logs_prisma(inserted: int, existing: object) -> MagicMock: + prisma = _tool_usage_prisma() + prisma.jsonify_object = lambda data: dict(data) + prisma.db.litellm_spendlogs.create_many = AsyncMock(return_value=inserted) + prisma.db.litellm_spendlogs.find_unique = AsyncMock(return_value=existing) + return prisma + + +async def _update_database_with(db_writer: DBSpendUpdateWriter, prisma: MagicMock, payload: dict) -> 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 + ), + patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam + "litellm.proxy.proxy_server.prisma_client", prisma + ), + patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam + "litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget" + ), + patch( # test-quality-ok: update_database imports the payload builder inside its body, no seam + "litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload", + return_value=payload, + ), + ): + charged = await db_writer.update_database( + token="test-token", + user_id="test-user", + end_user_id=None, + team_id=None, + org_id=None, + kwargs={"model": "gpt-5.6-luna", "call_type": "aretrieve_batch"}, + completion_response=None, + start_time=datetime.now(timezone.utc), + end_time=datetime.now(timezone.utc), + response_cost=0.25, + ) + await asyncio.sleep(0) + return charged + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("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, None, True), + ], + ids=[ + "first_retrieve_owns_the_row", + "another_retrieve_already_charged", + "failed_retrieve_holds_the_row", + "client_chosen_call_id_holds_the_row", + "row_gone_between_insert_and_lookup", + ], +) +async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote_its_row( + inserted: int, existing: object, charged: bool +): + """ + 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. + """ + 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()) is charged + + claimed_rows = prisma.db.litellm_spendlogs.create_many.await_args.kwargs + assert claimed_rows["skip_duplicates"] is True + assert [(row["request_id"], row["spend"]) for row in claimed_rows["data"]] == [("batch_abc_batch_cost", 0.25)] + assert prisma.spend_log_transactions == [] + assert db_writer._batch_database_updates.await_count == (1 if charged else 0) + + +@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.""" + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(0, None) + prisma.db.litellm_spendlogs.create_many = AsyncMock(side_effect=RuntimeError("db unreachable")) + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True + + assert [row["request_id"] for row in prisma.spend_log_transactions] == ["batch_abc_batch_cost"] + assert db_writer._batch_database_updates.await_count == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "payload", + [{**_batch_cost_payload(), "call_type": "acompletion"}, {**_batch_cost_payload(), "status": "failure"}], + ids=["not_a_batch_retrieve", "failed_batch_retrieve"], +) +async def test_update_database_queues_every_other_spend_row_for_the_next_flush(payload: dict): + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(1, None) + + assert await _update_database_with(db_writer, prisma, payload) is True + + prisma.db.litellm_spendlogs.create_many.assert_not_called() + assert prisma.spend_log_transactions == [payload] + assert db_writer._batch_database_updates.await_count == 1 + + @pytest.mark.asyncio @pytest.mark.parametrize( "injected_deployment, attributed", diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index b7037e8d621..2965f8b4006 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1,3 +1,4 @@ +import asyncio from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch @@ -6,7 +7,6 @@ import pytest from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.proxy_track_cost_callback import ( - _batch_cost_is_trackable_now, _get_budget_reservation_from_metadata, _ProxyDBLogger, _should_track_cost_callback, @@ -783,110 +783,46 @@ def _retrieved_batch(status: str, output_file_id: str | None): ) -def _prisma_client_with(queued_request_ids: tuple[str, ...], stored_row: object) -> MagicMock: - import asyncio - - prisma_client = MagicMock() - prisma_client._spend_log_transactions_lock = asyncio.Lock() - prisma_client.spend_log_transactions = [{"request_id": request_id} for request_id in queued_request_ids] - prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(return_value=stored_row) - return prisma_client - - @pytest.mark.asyncio @pytest.mark.parametrize( - ("status", "output_file_id", "spend_log_id", "prisma_client", "trackable"), + ("call_type", "status", "output_file_id", "row_claimed", "spend_written", "charged"), [ - ("in_progress", None, "batch_abc_batch_cost", None, False), - ("in_progress", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), - ("completed", None, "batch_abc_batch_cost", _prisma_client_with((), None), False), - ("completed", "file-out", "batch_abc_batch_cost", None, True), - ("completed", "file-out", None, _prisma_client_with((), None), True), - ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with(("batch_abc_batch_cost",), None), False), - ( - "completed", - "file-out", - "batch_abc_batch_cost", - _prisma_client_with((), {"request_id": "batch_abc_batch_cost"}), - False, - ), - ("completed", "file-out", "batch_abc_batch_cost", _prisma_client_with((), None), True), - ("failed", None, "batch_abc_batch_cost", _prisma_client_with((), None), True), + ("aretrieve_batch", "in_progress", None, True, False, False), + ("aretrieve_batch", "completed", None, True, False, False), + ("aretrieve_batch", "completed", "file-out", False, True, False), + ("aretrieve_batch", "completed", "file-out", True, True, True), + ("aretrieve_batch", "failed", None, True, True, True), + ("acreate_batch", "validating", None, True, True, True), ], ids=[ - "in_progress_without_db", - "in_progress_never_consults_db", - "completed_without_output_yet", - "final_without_db", - "final_without_spend_log_id", - "final_row_queued_for_flush", - "final_row_already_stored", - "final_first_sighting", - "failed_first_sighting", + "retrieve_before_final", + "retrieve_completed_without_output_yet", + "retrieve_after_another_retrieve_charged", + "retrieve_first_final", + "retrieve_failed_batch", + "create_before_final", ], ) -async def test_batch_cost_is_trackable_now(status, output_file_id, spend_log_id, prisma_client, trackable): - """ - A batch is billed from the first retrieve that sees it final and never again: - a poll before that wrote the shared spend row at $0 and pinned it there, and - every completed retrieve after the first charged the key again (LIT-7048). - """ - assert ( - await _batch_cost_is_trackable_now( - batch=_retrieved_batch(status, output_file_id), spend_log_id=spend_log_id, prisma_client=prisma_client - ) - is trackable - ) - - -@pytest.mark.asyncio -async def test_batch_cost_is_trackable_now_when_the_spend_row_lookup_fails(): - """An unreadable spend log table must not drop the batch's only spend row.""" - prisma_client = _prisma_client_with((), None) - prisma_client.db.litellm_spendlogs.find_unique = AsyncMock(side_effect=RuntimeError("db unreachable")) - - assert ( - await _batch_cost_is_trackable_now( - batch=_retrieved_batch("completed", "file-out"), - spend_log_id="batch_abc_batch_cost", - prisma_client=prisma_client, - ) - is True - ) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("call_type", "status", "output_file_id", "stored_row", "charged"), - [ - ("aretrieve_batch", "in_progress", None, None, False), - ("aretrieve_batch", "completed", "file-out", {"request_id": "batch_abc_batch_cost"}, False), - ("aretrieve_batch", "completed", "file-out", None, True), - ("acreate_batch", "validating", None, None, True), - ], - ids=["retrieve_before_final", "retrieve_already_recorded", "retrieve_first_final", "create_before_final"], -) -async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # test-quality-ok: whether the spend writer runs and whether the poll's reservation is handed back is the whole observable contract of the gate - call_type, status, output_file_id, stored_row, charged +async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # test-quality-ok: whether the spend writer runs, whether the counters move, and whether the poll's reservation is handed back is the whole observable contract of the gate + call_type, status, output_file_id, row_claimed, spend_written, charged ): """ - Only retrieves are gated, since creating a batch is its own billable request. - A retrieve that writes nothing hands its budget reservation back instead. + A poll before the batch is final used to pin its shared spend row at $0, and every + completed retrieve after the first charged the key again (LIT-7048). Only retrieves + are gated, since creating a batch is its own billable request, and a retrieve that + charges nothing hands its budget reservation back instead. """ logger = _ProxyDBLogger() budget_reservation = None if charged else {"reserved_cost": 0.5, "entries": []} kwargs = _batch_retrieve_kwargs(call_type, reservation=budget_reservation) with ( - patch( # test-quality-ok: prisma_client is a proxy_server global the callback reads lazily, no seam - "litellm.proxy.proxy_server.prisma_client", _prisma_client_with((), stored_row) - ), patch( # test-quality-ok: increment_spend_counters is a proxy_server global the callback reads lazily, no seam "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock - ), + ) as mock_increment_spend_counters, patch( # test-quality-ok: update_cache is a proxy_server global the callback reads lazily, no seam "litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock - ), + ) as mock_update_cache, patch( # test-quality-ok: callback imports proxy_logging_obj off proxy_server in its body, no seam "litellm.proxy.proxy_server.proxy_logging_obj" ) as mock_proxy_logging, @@ -895,7 +831,7 @@ async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # ) as mock_release_budget_reservation, ): mock_proxy_logging.failed_tracking_alert = AsyncMock() - mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock(return_value=row_claimed) mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock() await logger._PROXY_track_cost_callback( @@ -904,13 +840,15 @@ async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # start_time=datetime.now(), end_time=datetime.now(), ) + await asyncio.sleep(0) mock_proxy_logging.failed_tracking_alert.assert_not_called() + assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (1 if spend_written else 0) + assert mock_increment_spend_counters.await_count == (1 if charged else 0) + assert mock_update_cache.await_count == (1 if charged else 0) if charged: - mock_proxy_logging.db_spend_update_writer.update_database.assert_awaited_once() mock_release_budget_reservation.assert_not_awaited() else: - mock_proxy_logging.db_spend_update_writer.update_database.assert_not_called() mock_release_budget_reservation.assert_awaited_once_with(budget_reservation=budget_reservation) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py index fc97d760226..a1eb88a7834 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -6,7 +6,6 @@ Symbols pinned here: - ``update_spend_logs_job`` - ``_monitor_spend_logs_queue`` - ``_raise_failed_update_spend_exception`` - - ``spend_log_is_queued`` """ from __future__ import annotations @@ -23,7 +22,6 @@ from litellm.proxy.utils import ( _monitor_spend_logs_queue, _raise_failed_update_spend_exception, drain_spend_logs_queue, - spend_log_is_queued, update_daily_tag_spend, update_spend, update_spend_logs_job, @@ -631,16 +629,3 @@ def test_raise_failed_update_spend_exception_raises_original_error() -> None: with pytest.raises(ValueError, match="specific"): asyncio.run(_runner()) - - -@pytest.mark.asyncio -async def test_spend_log_is_queued_matches_only_rows_awaiting_flush( - mock_prisma_client: Any, make_spend_log_row: Any -) -> None: - mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="batch_abc_batch_cost")] - - assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is True - assert await spend_log_is_queued(mock_prisma_client, "batch_abc") is False - - mock_prisma_client.spend_log_transactions = [] - assert await spend_log_is_queued(mock_prisma_client, "batch_abc_batch_cost") is False 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 3/7] 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.""" From 0fb3951b2cadca5d9091b7586d0a5a15e4600b42 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 22:32:06 -0700 Subject: [PATCH 4/7] fix(spend): charge a batch once when an older proxy left its cost row at $0 A proxy without this fix wrote the batch's cost row on every poll while the batch was still running, so that row reads $0 and the insert that claims the charge has nowhere to land. The retrieve that charges the batch now writes its own payload over that row under a where clause that still names spend 0.0, so exactly one retrieve takes it over and every later one reads the charge and charges nothing --- litellm/proxy/db/db_spend_update_writer.py | 42 +++++++++- .../proxy/db/test_db_spend_update_writer.py | 76 ++++++++++++++++++- 2 files changed, 112 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 48312c025dd..ae3bc5663eb 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -14,6 +14,7 @@ import time import traceback from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload import litellm @@ -379,9 +380,44 @@ class DBSpendUpdateWriter: 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.debug( - "Spend row %s charged nothing for this batch, so this retrieve charges it", request_id - ) + return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client) + + async def _take_over_uncharged_batch_cost_row( + self, payload: SpendLogsPayload, prisma_client: "PrismaClient" + ) -> bool: + """Take the batch's cost row over from the poll that left it charging nothing. + + A pre-upgrade proxy wrote that row every time it polled the batch while it was still + running, so the charge is still to be made and the row still has to end up carrying + it. The row stops matching the moment it carries a charge, so it is one retrieve that + takes it over and charges, and every later one reads the charge and charges nothing. + """ + from litellm.repositories.table_repositories import SpendLogsRepository + + request_id: Final = payload["request_id"] + if payload["spend"] <= 0: + verbose_proxy_logger.debug( + "Cost tracking skipped: this batch costs nothing and spend row %s says so", request_id + ) + return False + try: + taken_over: Final = await SpendLogsRepository(prisma_client).table.update_many( + data=prisma_client.jsonify_object( + MappingProxyType({field: value for field, value in payload.items() if field != "request_id"}) + ), + where={ # mutable-ok: prisma where clause + "request_id": request_id, + "call_type": CallTypes.aretrieve_batch.value, + "status": "success", + "spend": 0.0, + }, + ) + except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; a row it cannot take over charges the batch + verbose_proxy_logger.warning("Could not take over spend row %s for a batch's cost: %s", request_id, e) + return True + if taken_over == 0: + verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id) + return False return True async def _enqueue_tool_usage_transaction( 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 41f2be08545..33b7af06e1a 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 @@ -2969,16 +2969,21 @@ def _batch_cost_payload() -> dict: } -def _spend_logs_prisma(inserted: int, existing: object) -> MagicMock: +def _spend_logs_prisma(inserted: int, existing: object, taken_over: int = 1) -> MagicMock: prisma = _tool_usage_prisma() prisma.jsonify_object = lambda data: dict(data) prisma.db.litellm_spendlogs.create_many = AsyncMock(return_value=inserted) prisma.db.litellm_spendlogs.find_unique = AsyncMock(return_value=existing) + prisma.db.litellm_spendlogs.update_many = AsyncMock(return_value=taken_over) return prisma async def _update_database_with( - db_writer: DBSpendUpdateWriter, prisma: MagicMock, payload: dict, disable_spend_logs: bool = False + db_writer: DBSpendUpdateWriter, + prisma: MagicMock, + payload: dict, + disable_spend_logs: bool = False, + response_cost: float = 0.25, ) -> bool: with ( patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam @@ -3005,7 +3010,7 @@ async def _update_database_with( completion_response=None, start_time=datetime.now(timezone.utc), end_time=datetime.now(timezone.utc), - response_cost=0.25, + response_cost=response_cost, ) await asyncio.sleep(0) return charged @@ -3054,6 +3059,71 @@ 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( + ("taken_over", "charged"), + [(1, True), (0, False)], + ids=["this_retrieve_takes_it_over", "another_one_got_there_first"], +) +async def test_update_database_charges_a_batch_whose_row_a_pre_upgrade_poll_left_at_zero( + taken_over: int, charged: bool +): + """ + A proxy without this fix wrote the batch's row at $0 on every poll of a running batch, + and the row outlives the upgrade, so the charge has to land on the row itself. Charging + without writing it there would charge again on every later retrieve (LIT-7048). + """ + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0) + prisma = _spend_logs_prisma(0, existing, taken_over) + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is charged + + taken = prisma.db.litellm_spendlogs.update_many.await_args.kwargs + assert taken["where"] == { + "request_id": "batch_abc_batch_cost", + "call_type": "aretrieve_batch", + "status": "success", + "spend": 0.0, + } + assert taken["data"]["spend"] == 0.25 + assert "request_id" not in taken["data"] + assert db_writer._batch_database_updates.await_count == (1 if charged else 0) + + +@pytest.mark.asyncio +async def test_update_database_charges_a_batch_whose_zero_row_it_could_not_take_over(): + """A DB that refuses the takeover must not swallow the batch's cost.""" + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0) + prisma = _spend_logs_prisma(0, existing) + prisma.db.litellm_spendlogs.update_many = AsyncMock(side_effect=RuntimeError("db unreachable")) + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True + + assert db_writer._batch_database_updates.await_count == 1 + + +@pytest.mark.asyncio +async def test_update_database_leaves_a_batch_that_cost_nothing_to_the_retrieve_that_wrote_its_row(): + """ + A batch every line of which failed costs $0, so its row reads $0 for the honest reason + and the retrieve that wrote it is still the one that accounted it. Taking that row over + on every later retrieve would count one batch as many requests. + """ + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0) + prisma = _spend_logs_prisma(0, existing) + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload(), response_cost=0.0) is False + + prisma.db.litellm_spendlogs.update_many.assert_not_called() + assert db_writer._batch_database_updates.await_count == 0 + + @pytest.mark.asyncio @pytest.mark.parametrize( ("inserted", "existing", "charged"), From 24f0be80219cd3403ec28cddbed9be946fad1cb0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 22:47:41 -0700 Subject: [PATCH 5/7] fix(spend): leave a batch uncharged when the database refuses the takeover The takeover of a $0 row an older proxy left behind used to charge the batch when the update could not reach the database. That leaves the row still reading $0, so every later retrieve finds the same row and charges the batch again, which is the repeat charging this PR exists to stop. The retrieve that does take the row over is the one that charges, and a batch nobody retrieves again after that failure is never charged, the same as one whose proxy died inside the write window. --- litellm/proxy/db/db_spend_update_writer.py | 8 +++++--- .../proxy/db/test_db_spend_update_writer.py | 12 ++++++++---- 2 files changed, 13 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ae3bc5663eb..ee7802a45d3 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -412,9 +412,11 @@ class DBSpendUpdateWriter: "spend": 0.0, }, ) - except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; a row it cannot take over charges the batch - verbose_proxy_logger.warning("Could not take over spend row %s for a batch's cost: %s", request_id, e) - return True + except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; the next retrieve takes the row over + verbose_proxy_logger.warning( + "Could not take over spend row %s, leaving this batch's cost to the next retrieve: %s", request_id, e + ) + return False if taken_over == 0: verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id) return False 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 33b7af06e1a..4efb94b60aa 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 @@ -3093,17 +3093,21 @@ async def test_update_database_charges_a_batch_whose_row_a_pre_upgrade_poll_left @pytest.mark.asyncio -async def test_update_database_charges_a_batch_whose_zero_row_it_could_not_take_over(): - """A DB that refuses the takeover must not swallow the batch's cost.""" +async def test_update_database_leaves_a_batch_whose_zero_row_it_could_not_take_over_to_the_next_retrieve(): + """ + A DB that refuses the takeover leaves the row reading $0, so charging here would charge + the batch again on every later retrieve. The retrieve that does take the row over is the + one that charges. + """ db_writer = DBSpendUpdateWriter() db_writer._batch_database_updates = AsyncMock() existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0) prisma = _spend_logs_prisma(0, existing) prisma.db.litellm_spendlogs.update_many = AsyncMock(side_effect=RuntimeError("db unreachable")) - assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False - assert db_writer._batch_database_updates.await_count == 1 + assert db_writer._batch_database_updates.await_count == 0 @pytest.mark.asyncio From fcb6d2267c09389ae1aa9e80e04a35caa5b8b470 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 23:51:44 -0700 Subject: [PATCH 6/7] fix(spend): keep a batch's claim row out of the logs a proxy was told not to write disable_spend_logs has to keep meaning that no request gets logged, and the row that makes a batch chargeable exactly once is the one row it cannot drop, so with logging off that row now carries only what tells the retrieves apart. SPEND_LOGS_URL deployments get their copy back too: the claim writes straight to this table, so the row is queued as well when an external writer is the one that takes the spend logs. --- litellm/proxy/db/db_spend_update_writer.py | 49 +++++++-- .../proxy/db/test_db_spend_update_writer.py | 99 +++++++++++++++++++ 2 files changed, 141 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ee7802a45d3..bdc014d7f13 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -89,6 +89,21 @@ def _is_batch_cost_row(payload: SpendLogsPayload) -> bool: return payload.get("call_type") == CallTypes.aretrieve_batch.value and payload.get("status") == "success" +_BATCH_COST_CLAIM_FIELDS: Final = frozenset({"request_id", "call_type", "spend", "startTime", "endTime", "status"}) + + +def _batch_cost_row_to_write(payload: SpendLogsPayload, disable_spend_logs: bool) -> Mapping[str, object]: + """Reduce a batch's cost row to what tells the retrieves apart when logging is off. + + A proxy run with spend logs disabled still needs one row per batch to charge it once, + so the row is written either way, but it carries no request of its own: no metadata, + no requester IP, no key, model, or token counts (LIT-7048). + """ + if disable_spend_logs is False: + return payload + return MappingProxyType({field: value for field, value in payload.items() if field in _BATCH_COST_CLAIM_FIELDS}) + + class _SpendBatch(Protocol): litellm_usertable: BatchTable litellm_verificationtoken: BatchTable @@ -336,12 +351,16 @@ class DBSpendUpdateWriter: 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) + return await self._claim_batch_cost_spend_log( + payload=payload, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs + ) if disable_spend_logs is False: await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client) return True - async def _claim_batch_cost_spend_log(self, payload: SpendLogsPayload, prisma_client: "PrismaClient") -> bool: + async def _claim_batch_cost_spend_log( + self, payload: SpendLogsPayload, prisma_client: "PrismaClient", disable_spend_logs: bool + ) -> 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 @@ -353,13 +372,17 @@ class DBSpendUpdateWriter: from litellm.repositories.table_repositories import SpendLogsRepository request_id: Final = payload["request_id"] + row: Final = _batch_cost_row_to_write(payload, disable_spend_logs) spend_logs: Final = SpendLogsRepository(prisma_client).table try: claimed: Final = await spend_logs.create_many( - data=[prisma_client.jsonify_object(payload)], # mutable-ok: prisma create_many takes a list + data=[prisma_client.jsonify_object(row)], # mutable-ok: prisma create_many takes a list skip_duplicates=True, ) if claimed == 1: + await self._forward_batch_cost_row( + row=row, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs + ) return True existing: Final = await spend_logs.find_unique( where={"request_id": request_id} # mutable-ok: prisma where clause @@ -368,7 +391,7 @@ class DBSpendUpdateWriter: verbose_proxy_logger.warning( "Could not claim spend row %s for a batch's cost, queueing it: %s", request_id, e ) - await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client) + await self._insert_spend_log_to_db(payload=prisma_client.jsonify_object(row), prisma_client=prisma_client) return True if existing is None or existing.call_type != CallTypes.aretrieve_batch.value or existing.status != "success": verbose_proxy_logger.warning( @@ -380,10 +403,22 @@ class DBSpendUpdateWriter: if existing.spend > 0: verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id) return False - return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client) + return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client, row=row) + + async def _forward_batch_cost_row( + self, row: Mapping[str, object], prisma_client: "PrismaClient", disable_spend_logs: bool + ) -> None: + """Queue the claimed row for an external spend log writer, which the claim went around. + + With ``SPEND_LOGS_URL`` set the queue posts every spend log to that writer instead of + inserting it, so a batch's cost row reaches it only by being queued here as well. + """ + if disable_spend_logs is True or os.getenv("SPEND_LOGS_URL") is None: + return + await self._insert_spend_log_to_db(payload=prisma_client.jsonify_object(row), prisma_client=prisma_client) async def _take_over_uncharged_batch_cost_row( - self, payload: SpendLogsPayload, prisma_client: "PrismaClient" + self, payload: SpendLogsPayload, prisma_client: "PrismaClient", row: Mapping[str, object] ) -> bool: """Take the batch's cost row over from the poll that left it charging nothing. @@ -403,7 +438,7 @@ class DBSpendUpdateWriter: try: taken_over: Final = await SpendLogsRepository(prisma_client).table.update_many( data=prisma_client.jsonify_object( - MappingProxyType({field: value for field, value in payload.items() if field != "request_id"}) + MappingProxyType({field: value for field, value in row.items() if field != "request_id"}) ), where={ # mutable-ok: prisma where clause "request_id": request_id, 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 4efb94b60aa..a8aeaced55f 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 @@ -1,6 +1,7 @@ import asyncio import copy import json +import os import re @@ -3169,6 +3170,104 @@ async def test_update_database_writes_no_ordinary_spend_row_with_spend_logs_disa assert db_writer._batch_database_updates.await_count == 1 +_BATCH_CLAIM_FIELDS = {"request_id", "call_type", "status", "spend", "startTime", "endTime"} + + +def _logged_batch_cost_payload() -> dict: + return { + **_batch_cost_payload(), + "api_key": "0e5b0e9e5f", + "model": "gpt-5.6-luna", + "user": "test-user", + "metadata": '{"batch_models": ["gpt-5.6-luna"]}', + "requester_ip_address": "127.0.0.1", + "proxy_server_request": '{"headers": {"user-agent": "litellm-batch-cost-check"}}', + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("disable_spend_logs", "logs_the_request"), + [(False, True), (True, False)], + ids=["spend_logs_on", "spend_logs_off"], +) +async def test_update_database_claims_a_batch_without_logging_the_request_that_polled_it( + disable_spend_logs: bool, logs_the_request: bool +): + """ + disable_spend_logs has to keep meaning that no request gets logged, and the batch's cost + row is the one row it cannot drop, so with logging off that row carries only what tells + the retrieves apart: no metadata, no requester IP, no key, model, or token counts. + """ + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(1, None) + payload = _logged_batch_cost_payload() + + assert await _update_database_with(db_writer, prisma, payload, disable_spend_logs) is True + + claimed = prisma.db.litellm_spendlogs.create_many.await_args.kwargs["data"][0] + assert set(claimed) == (set(payload) if logs_the_request else _BATCH_CLAIM_FIELDS) + assert claimed["spend"] == 0.25 + assert db_writer._batch_database_updates.await_count == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("spend_logs_url", "forwarded"), + [("http://spend-logs.internal", True), (None, False)], + ids=["an_external_writer_takes_the_rows", "rows_are_written_to_this_db"], +) +async def test_update_database_sends_a_claimed_batch_cost_row_on_to_an_external_spend_log_writer( + monkeypatch, spend_logs_url: str | None, forwarded: bool +): + """ + SPEND_LOGS_URL makes the flush post spend logs to that writer instead of inserting them, + and the claim writes straight to this table, so the batch's row reaches the writer only + by being queued as well. Queueing it with no writer configured would insert it twice. + """ + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(1, None) + if spend_logs_url is None: + monkeypatch.delenv("SPEND_LOGS_URL", raising=False) + else: + monkeypatch.setenv("SPEND_LOGS_URL", spend_logs_url) + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True + + queued = [row["request_id"] for row in prisma.spend_log_transactions] + assert queued == (["batch_abc_batch_cost"] if forwarded else []) + + +@pytest.mark.asyncio +async def test_update_database_forwards_no_batch_cost_row_a_later_retrieve_had_already_claimed(monkeypatch): + """The retrieve that lost the claim charges nothing, so it must not post a row either.""" + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25) + prisma = _spend_logs_prisma(0, existing) + monkeypatch.setenv("SPEND_LOGS_URL", "http://spend-logs.internal") + + assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False + + assert prisma.spend_log_transactions == [] + + +@pytest.mark.asyncio +async def test_update_database_queues_only_the_claim_for_a_batch_it_could_not_write_with_logs_disabled(): + """A refused claim is retried through the queue, so what it queues has to stay unlogged too.""" + db_writer = DBSpendUpdateWriter() + db_writer._batch_database_updates = AsyncMock() + prisma = _spend_logs_prisma(0, None) + prisma.db.litellm_spendlogs.create_many = AsyncMock(side_effect=RuntimeError("db unreachable")) + + assert await _update_database_with(db_writer, prisma, _logged_batch_cost_payload(), True) is True + + assert [set(row) for row in prisma.spend_log_transactions] == [_BATCH_CLAIM_FIELDS] + 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.""" From defd8661f4e359994bf57a7f9e51ed8c479f17ff Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sun, 6 Sep 2026 00:04:24 -0700 Subject: [PATCH 7/7] refactor(spend): stop queueing a batch's claim row for a writer the proxy never builds SPEND_LOGS_URL only diverts spend logs when db_writer_client is set, and nothing in the proxy ever assigns that global, so the queued copy was only ever skipped as a duplicate by the local insert. --- litellm/proxy/db/db_spend_update_writer.py | 15 ------- .../proxy/db/test_db_spend_update_writer.py | 43 ------------------- 2 files changed, 58 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index bdc014d7f13..9230be8055e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -380,9 +380,6 @@ class DBSpendUpdateWriter: skip_duplicates=True, ) if claimed == 1: - await self._forward_batch_cost_row( - row=row, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs - ) return True existing: Final = await spend_logs.find_unique( where={"request_id": request_id} # mutable-ok: prisma where clause @@ -405,18 +402,6 @@ class DBSpendUpdateWriter: return False return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client, row=row) - async def _forward_batch_cost_row( - self, row: Mapping[str, object], prisma_client: "PrismaClient", disable_spend_logs: bool - ) -> None: - """Queue the claimed row for an external spend log writer, which the claim went around. - - With ``SPEND_LOGS_URL`` set the queue posts every spend log to that writer instead of - inserting it, so a batch's cost row reaches it only by being queued here as well. - """ - if disable_spend_logs is True or os.getenv("SPEND_LOGS_URL") is None: - return - await self._insert_spend_log_to_db(payload=prisma_client.jsonify_object(row), prisma_client=prisma_client) - async def _take_over_uncharged_batch_cost_row( self, payload: SpendLogsPayload, prisma_client: "PrismaClient", row: Mapping[str, object] ) -> bool: 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 a8aeaced55f..0bca7c9492c 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 @@ -1,7 +1,6 @@ import asyncio import copy import json -import os import re @@ -3212,48 +3211,6 @@ async def test_update_database_claims_a_batch_without_logging_the_request_that_p assert db_writer._batch_database_updates.await_count == 1 -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("spend_logs_url", "forwarded"), - [("http://spend-logs.internal", True), (None, False)], - ids=["an_external_writer_takes_the_rows", "rows_are_written_to_this_db"], -) -async def test_update_database_sends_a_claimed_batch_cost_row_on_to_an_external_spend_log_writer( - monkeypatch, spend_logs_url: str | None, forwarded: bool -): - """ - SPEND_LOGS_URL makes the flush post spend logs to that writer instead of inserting them, - and the claim writes straight to this table, so the batch's row reaches the writer only - by being queued as well. Queueing it with no writer configured would insert it twice. - """ - db_writer = DBSpendUpdateWriter() - db_writer._batch_database_updates = AsyncMock() - prisma = _spend_logs_prisma(1, None) - if spend_logs_url is None: - monkeypatch.delenv("SPEND_LOGS_URL", raising=False) - else: - monkeypatch.setenv("SPEND_LOGS_URL", spend_logs_url) - - assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True - - queued = [row["request_id"] for row in prisma.spend_log_transactions] - assert queued == (["batch_abc_batch_cost"] if forwarded else []) - - -@pytest.mark.asyncio -async def test_update_database_forwards_no_batch_cost_row_a_later_retrieve_had_already_claimed(monkeypatch): - """The retrieve that lost the claim charges nothing, so it must not post a row either.""" - db_writer = DBSpendUpdateWriter() - db_writer._batch_database_updates = AsyncMock() - existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25) - prisma = _spend_logs_prisma(0, existing) - monkeypatch.setenv("SPEND_LOGS_URL", "http://spend-logs.internal") - - assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False - - assert prisma.spend_log_transactions == [] - - @pytest.mark.asyncio async def test_update_database_queues_only_the_claim_for_a_batch_it_could_not_write_with_logs_disabled(): """A refused claim is retried through the queue, so what it queues has to stay unlogged too."""