mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(batch): sync all status changes to DB in CheckBatchCost polling job
The CheckBatchCost polling job only updated the database when a batch reached "completed" status. Intermediate states (in_progress, finalizing) and terminal failure states (failed, expired, cancelled) were never written back. Since list_user_batches reads exclusively from the DB, enterprise users saw permanently stale "validating" status in the Swagger UI and GET /v1/batches responses. This adds a status synchronization step after aretrieve_batch that updates the DB for all status changes. Terminal failure states also get batch_processed=True to stop futile re-polling.
This commit is contained in:
parent
3dccdde9c8
commit
167b8bb1f5
2 changed files with 172 additions and 0 deletions
|
|
@ -204,6 +204,54 @@ class CheckBatchCost:
|
|||
)
|
||||
continue
|
||||
|
||||
# Sync batch status to DB for all status changes (not just completion).
|
||||
# This ensures list_user_batches returns fresh status from the provider.
|
||||
_current_db_status = job.status
|
||||
_provider_status = response.status
|
||||
_normalized_status = (
|
||||
_provider_status if _provider_status != "completed" else "complete"
|
||||
)
|
||||
|
||||
if _normalized_status != _current_db_status:
|
||||
_is_terminal_failure = _provider_status in (
|
||||
"failed",
|
||||
"expired",
|
||||
"cancelled",
|
||||
)
|
||||
|
||||
if _provider_status != "completed":
|
||||
# For non-completed status changes, update DB now.
|
||||
# Completed batches are handled below with cost tracking.
|
||||
try:
|
||||
update_data: dict = {
|
||||
"status": _normalized_status,
|
||||
"file_object": response.model_dump_json(),
|
||||
}
|
||||
if _is_terminal_failure:
|
||||
if self._has_batch_processed_column:
|
||||
update_data["batch_processed"] = True
|
||||
verbose_proxy_logger.info(
|
||||
f"Batch {batch_id} reached terminal state "
|
||||
f"'{_provider_status}', marking as processed"
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.info(
|
||||
f"Batch {batch_id} status changed: "
|
||||
f"{_current_db_status} -> {_provider_status}"
|
||||
)
|
||||
await self.prisma_client.db.litellm_managedobjecttable.update(
|
||||
where={"id": job.id},
|
||||
data=update_data,
|
||||
)
|
||||
except Exception as db_err:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: failed to sync status for "
|
||||
f"batch {batch_id}: {db_err}"
|
||||
)
|
||||
|
||||
if _is_terminal_failure:
|
||||
continue # No cost tracking needed for failed/expired/cancelled
|
||||
|
||||
## RETRIEVE THE BATCH JOB OUTPUT FILE
|
||||
if (
|
||||
response.status == "completed"
|
||||
|
|
|
|||
|
|
@ -207,3 +207,127 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti
|
|||
assert "file_object" in update_call_kwargs["data"], (
|
||||
"file_object must be written to DB so list_batches reads updated status"
|
||||
)
|
||||
|
||||
|
||||
# --- Status sync tests: polling job updates DB for all status changes ---
|
||||
|
||||
|
||||
def _make_checker_for_status_test(mock_job, batch_response):
|
||||
"""Helper to create a CheckBatchCost instance for status sync tests."""
|
||||
from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[mock_job]
|
||||
)
|
||||
mock_prisma.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0)
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.aretrieve_batch = AsyncMock(return_value=batch_response)
|
||||
|
||||
checker = CheckBatchCost(
|
||||
proxy_logging_obj=MagicMock(),
|
||||
prisma_client=mock_prisma,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
return checker, mock_prisma
|
||||
|
||||
|
||||
def _make_mock_job(status: str) -> MagicMock:
|
||||
"""Create a mock DB job with a specific status."""
|
||||
unified_raw = "litellm_proxy;model_id:model-deploy-xyz;llm_batch_id:batch-123;llm_output_file_id:file-raw-output"
|
||||
mock_job = MagicMock()
|
||||
mock_job.unified_object_id = base64.b64encode(unified_raw.encode()).decode()
|
||||
mock_job.created_by = "user-A"
|
||||
mock_job.id = "job-1"
|
||||
mock_job.status = status
|
||||
return mock_job
|
||||
|
||||
|
||||
def _make_batch_response(status: str):
|
||||
"""Create a LiteLLMBatch with a specific status."""
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
return LiteLLMBatch(
|
||||
id="batch-123",
|
||||
completion_window="24h",
|
||||
created_at=1700000000,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-input",
|
||||
object="batch",
|
||||
status=status,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_batch_cost_updates_db_for_in_progress_status():
|
||||
"""
|
||||
When provider returns in_progress, the polling job should update the DB
|
||||
status and file_object so list_batches shows fresh data.
|
||||
"""
|
||||
mock_job = _make_mock_job(status="validating")
|
||||
batch_response = _make_batch_response(status="in_progress")
|
||||
checker, mock_prisma = _make_checker_for_status_test(mock_job, batch_response)
|
||||
|
||||
await checker.check_batch_cost()
|
||||
|
||||
# DB should be updated with in_progress status
|
||||
mock_prisma.db.litellm_managedobjecttable.update.assert_called_once()
|
||||
update_kwargs = mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs
|
||||
assert update_kwargs["data"]["status"] == "in_progress"
|
||||
assert "file_object" in update_kwargs["data"]
|
||||
# batch_processed should NOT be set for intermediate states
|
||||
assert "batch_processed" not in update_kwargs["data"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_batch_cost_marks_failed_batch_as_processed():
|
||||
"""
|
||||
When provider returns failed, the polling job should update DB status
|
||||
AND set batch_processed=True to stop re-polling.
|
||||
"""
|
||||
mock_job = _make_mock_job(status="validating")
|
||||
batch_response = _make_batch_response(status="failed")
|
||||
checker, mock_prisma = _make_checker_for_status_test(mock_job, batch_response)
|
||||
|
||||
await checker.check_batch_cost()
|
||||
|
||||
mock_prisma.db.litellm_managedobjecttable.update.assert_called_once()
|
||||
update_kwargs = mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs
|
||||
assert update_kwargs["data"]["status"] == "failed"
|
||||
assert update_kwargs["data"]["batch_processed"] is True
|
||||
assert "file_object" in update_kwargs["data"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_batch_cost_marks_expired_batch_as_processed():
|
||||
"""
|
||||
When provider returns expired, the polling job should update DB status
|
||||
AND set batch_processed=True to stop re-polling.
|
||||
"""
|
||||
mock_job = _make_mock_job(status="in_progress")
|
||||
batch_response = _make_batch_response(status="expired")
|
||||
checker, mock_prisma = _make_checker_for_status_test(mock_job, batch_response)
|
||||
|
||||
await checker.check_batch_cost()
|
||||
|
||||
mock_prisma.db.litellm_managedobjecttable.update.assert_called_once()
|
||||
update_kwargs = mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs
|
||||
assert update_kwargs["data"]["status"] == "expired"
|
||||
assert update_kwargs["data"]["batch_processed"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_batch_cost_skips_update_when_status_unchanged():
|
||||
"""
|
||||
When provider returns the same status as DB, no update should be made.
|
||||
"""
|
||||
mock_job = _make_mock_job(status="validating")
|
||||
batch_response = _make_batch_response(status="validating")
|
||||
checker, mock_prisma = _make_checker_for_status_test(mock_job, batch_response)
|
||||
|
||||
await checker.check_batch_cost()
|
||||
|
||||
# No DB update should have been made
|
||||
mock_prisma.db.litellm_managedobjecttable.update.assert_not_called()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue