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.
This commit is contained in:
Marty Sullivan 2026-08-14 02:18:40 -04:00
parent ec52858865
commit c9e9c279fe
3 changed files with 83 additions and 2 deletions

View file

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

View file

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

View file

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