mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
ec52858865
commit
c9e9c279fe
3 changed files with 83 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue