fix(check_batch_cost): leave a lagging-output completed batch for the next poll cycle

This commit is contained in:
mateo-berri 2026-08-22 12:48:39 -07:00
parent 9cd6240387
commit af18f77db6
2 changed files with 80 additions and 1 deletions

View file

@ -966,6 +966,16 @@ class CheckBatchCost:
)
elif response.status in PROVIDER_TERMINAL_BATCH_STATUSES:
from litellm.proxy.openai_files_endpoints.common_utils import (
_completed_batch_safe_to_retire,
)
if response.status in ("completed", "complete") and not _completed_batch_safe_to_retire(response):
verbose_proxy_logger.info(
f"CheckBatchCost: batch {batch_id} is completed but its output file id "
f"has not appeared yet; leaving job {job.id} for the next poll cycle"
)
continue
await self._finalize_unbilled_terminal_job(job, response)
# Record polling run metrics (always, even if nothing was processed)

View file

@ -1044,7 +1044,9 @@ class TestCheckBatchCost:
Pre-fix it matched neither the completed-with-output branch nor the
failed/expired/cancelled branch, so batch_processed stayed False and the row
was re-selected on every poll cycle forever. It must now be marked terminal
exactly once, without being billed (no output means nothing to bill).
exactly once, without being billed: request_counts.completed == 0 proves the
missing output file means nothing to bill rather than a lagging output id
(#37713 keeps the lagging case eligible for the next cycle).
"""
import base64
from unittest.mock import patch
@ -1073,6 +1075,7 @@ class TestCheckBatchCost:
mock_response.status = completed_status
mock_response.output_file_id = None
mock_response.error_file_id = "file-error-123"
mock_response.request_counts = MagicMock(completed=0, failed=3, total=3)
mock_response.model_dump_json.return_value = (
f'{{"id":"batch-1","status":"{completed_status}"}}'
)
@ -1107,6 +1110,72 @@ class TestCheckBatchCost:
mock_llm_router.get_deployment_credentials_with_provider.call_count == 0
), "a batch with no output file must not enter the cost-tracking path"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"request_counts",
[MagicMock(completed=7, failed=0, total=7), None],
ids=["lagging_output_id", "unknown_counts"],
)
async def test_completed_with_lagging_output_file_left_for_next_cycle(
self,
check_batch_cost_instance,
mock_prisma_client,
mock_llm_router,
request_counts,
):
"""#37713 regression: a batch can report completed while its output_file_id is
still lagging behind at the provider. Retiring it in that window (or when the
request counts cannot prove there is nothing to bill) permanently loses the
spend record, so the poller must leave the row untouched and revisit it on the
next cycle once the output id has appeared.
"""
import base64
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=1
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-completed-lagging-output-1"
mock_job.unified_object_id = base64.urlsafe_b64encode(
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
mock_job.created_by = "user-1"
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = None
mock_response.error_file_id = None
mock_response.request_counts = request_counts
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
with patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
) as mock_afile_content:
await check_batch_cost_instance.check_batch_cost()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0
), "a completed batch whose output id is still lagging must stay eligible for the next poll"
assert (
mock_afile_content.await_count == 0
), "a batch with no output file must not be billed"
@pytest.mark.asyncio
async def test_non_terminal_status_left_unprocessed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router