From c9e9c279fe6270132079db6da41b69c7d8809169 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Fri, 14 Aug 2026 02:18:40 -0400 Subject: [PATCH] fix(batches): decide batch cost ownership once per retrieve The ownership question was asked twice for one retrieve: once before the provider call to decide whether to suppress inline accounting, and again afterwards to decide whether to mark the batch accounted. Between those two points the poller can complete its first successful filtered query and become usable, so the two answers disagree. The retrieve then accounts for the batch inline, having decided the poller was unusable, while the later check sees a usable poller and leaves the marker unset, so the poller accounts for the same batch again and its spend is counted twice. The retrieve now decides once and passes that decision to update_batch_in_database, which prefers it over re-deriving one. Callers that record no cost of their own leave it unset and keep deriving it as before, so the cancel path is unchanged. --- litellm/proxy/batches_endpoints/endpoints.py | 4 +- .../openai_files_endpoints/common_utils.py | 12 +++- .../test_files_common_utils.py | 69 +++++++++++++++++++ 3 files changed, 83 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 563c380d34b..554a286555b 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -497,7 +497,8 @@ async def retrieve_batch( "Batch %s is in non-terminal state %s, syncing with provider", batch_id, response.status ) - if unified_batch_id and batch_cost_poller_is_active(): + poller_owns_accounting: Final = bool(unified_batch_id) and batch_cost_poller_is_active() + if poller_owns_accounting: litellm_metadata = data.get("litellm_metadata") if not isinstance(litellm_metadata, dict): litellm_metadata = {} # mutable-ok: the suppression flag must live inside litellm_metadata for the success handler to read it, and this request carried no mapping to extend @@ -581,6 +582,7 @@ async def retrieve_batch( verbose_proxy_logger=verbose_proxy_logger, db_batch_object=db_batch_object, operation="retrieve", + poller_owns_accounting=poller_owns_accounting, ) ### CALL HOOKS ### - modify outgoing data diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 012c0e6ea5b..c24d7b6e9de 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -1273,6 +1273,7 @@ async def update_batch_in_database( db_batch_object=None, operation: str = "update", user_api_key_dict=None, + poller_owns_accounting: bool | None = None, ): """ Update batch status and object in ManagedObjectTable. @@ -1287,6 +1288,12 @@ async def update_batch_in_database( db_batch_object: Optional existing database object; fetched by unified_object_id when omitted operation: Description of operation ("update", "cancel", etc.) user_api_key_dict: Optional auth context for creating managed file IDs + poller_owns_accounting: Whether the caller already decided that the cost poller + owns this batch's accounting. Callers that suppress their own inline + accounting must pass the same decision they acted on, because re-deciding + here can observe a poller that became usable in between and leave the batch + unmarked after it was already accounted for, billing it twice. Left None by + callers that record no cost themselves. """ import litellm.utils @@ -1336,7 +1343,10 @@ async def update_batch_in_database( "updated_at": litellm.utils.get_utc_datetime(), } - if db_status == "complete" and not batch_cost_poller_is_active(): + poller_owns: Final = ( + batch_cost_poller_is_active() if poller_owns_accounting is None else poller_owns_accounting + ) + if db_status == "complete" and not poller_owns: update_data["batch_processed"] = True try: diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index 20858a2a15f..c16450446a5 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -326,3 +326,72 @@ async def test_update_batch_in_database_is_a_noop_for_unmanaged_batches(monkeypa ) update_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_the_caller_s_accounting_decision_wins_over_a_later_poller_transition(monkeypatch): + """The ownership decision is made before the provider retrieval and acted on there, so + re-deciding afterwards can observe a poller that only just became usable. That split + left the retrieve accounting inline while the row stayed unmarked, so the poller + accounted for the same batch again and billed it twice. Passing the decision through + makes both halves agree even when the poller transitions mid-flight.""" + import litellm.proxy.openai_files_endpoints.common_utils as cu + + # The predicate now reports an active poller, i.e. it flipped during the retrieval. + monkeypatch.setattr(cu, "batch_cost_poller_is_active", lambda: True) + monkeypatch.setattr(cu, "ensure_batch_response_managed_file_ids", AsyncMock()) + + prisma_client = MagicMock() + update_mock = AsyncMock() + prisma_client.db.litellm_managedobjecttable.update = update_mock + db_batch_object = MagicMock() + db_batch_object.status = "in_progress" + + await cu.update_batch_in_database( + batch_id="unified-batch-id", + unified_batch_id="unified-batch-id", + response=_completed_batch(), + managed_files_obj=MagicMock(), + prisma_client=prisma_client, + verbose_proxy_logger=MagicMock(), + db_batch_object=db_batch_object, + operation="retrieve", + poller_owns_accounting=False, + ) + + data = update_mock.await_args.kwargs["data"] + assert data["batch_processed"] is True + assert data["status"] == "complete" + + +@pytest.mark.asyncio +async def test_a_caller_that_handed_off_accounting_still_leaves_the_marker_alone(monkeypatch): + """The mirror case: a caller that suppressed its own accounting must leave the marker + for the poller even if the predicate has since stopped reporting one, otherwise the + batch is retired without anyone having accounted for it.""" + import litellm.proxy.openai_files_endpoints.common_utils as cu + + monkeypatch.setattr(cu, "batch_cost_poller_is_active", lambda: False) + monkeypatch.setattr(cu, "ensure_batch_response_managed_file_ids", AsyncMock()) + + prisma_client = MagicMock() + update_mock = AsyncMock() + prisma_client.db.litellm_managedobjecttable.update = update_mock + db_batch_object = MagicMock() + db_batch_object.status = "in_progress" + + await cu.update_batch_in_database( + batch_id="unified-batch-id", + unified_batch_id="unified-batch-id", + response=_completed_batch(), + managed_files_obj=MagicMock(), + prisma_client=prisma_client, + verbose_proxy_logger=MagicMock(), + db_batch_object=db_batch_object, + operation="retrieve", + poller_owns_accounting=True, + ) + + data = update_mock.await_args.kwargs["data"] + assert "batch_processed" not in data + assert data["status"] == "complete"