mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(batches): isolate CheckBatchCost failures per job
Wrap each managed-batch poll iteration in its own error boundary and make Prometheus error recording best-effort so one poisoned batch cannot abort reconciliation for unrelated jobs in the same cycle. Fixes #35357
This commit is contained in:
parent
3c2264cfac
commit
d9b9676d49
2 changed files with 341 additions and 80 deletions
|
|
@ -107,7 +107,12 @@ class CheckBatchCost:
|
|||
prom_logger: Optional["PrometheusLogger"], error_type: str
|
||||
) -> None:
|
||||
if prom_logger is not None:
|
||||
prom_logger.record_check_batch_cost_error(error_type)
|
||||
try:
|
||||
prom_logger.record_check_batch_cost_error(error_type)
|
||||
except Exception as prom_err:
|
||||
verbose_proxy_logger.warning(
|
||||
f"CheckBatchCost: failed to record {error_type} metric: {prom_err}"
|
||||
)
|
||||
|
||||
def _resolve_job_routing(
|
||||
self,
|
||||
|
|
@ -558,94 +563,101 @@ class CheckBatchCost:
|
|||
else:
|
||||
jobs = await self._fallback_find_jobs()
|
||||
for job in jobs:
|
||||
routing = self._resolve_job_routing(job, prom_logger)
|
||||
if routing is None:
|
||||
continue
|
||||
model_id, batch_id = routing
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Querying model ID: {model_id} for cost and usage of batch ID: {batch_id}"
|
||||
)
|
||||
|
||||
try:
|
||||
response = await self.llm_router.aretrieve_batch(
|
||||
model=model_id,
|
||||
batch_id=batch_id,
|
||||
litellm_metadata={
|
||||
"user_api_key_user_id": job.created_by or "default-user-id",
|
||||
"batch_ignore_default_logging": True,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
routing = self._resolve_job_routing(job, prom_logger)
|
||||
if routing is None:
|
||||
continue
|
||||
model_id, batch_id = routing
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Skipping job {job.unified_object_id} because of error querying model ID: {model_id} for cost and usage of batch ID: {batch_id}: {e}"
|
||||
f"Querying model ID: {model_id} for cost and usage of batch ID: {batch_id}"
|
||||
)
|
||||
if prom_logger:
|
||||
prom_logger.record_check_batch_cost_error("provider_retrieval_error")
|
||||
continue
|
||||
|
||||
## RETRIEVE THE BATCH JOB OUTPUT FILE
|
||||
if (
|
||||
response.status == "completed"
|
||||
and response.output_file_id is not None
|
||||
):
|
||||
try:
|
||||
tracked = await self._track_completed_batch_cost(
|
||||
job=job,
|
||||
response=response,
|
||||
model_id=model_id,
|
||||
response = await self.llm_router.aretrieve_batch(
|
||||
model=model_id,
|
||||
batch_id=batch_id,
|
||||
prom_logger=prom_logger,
|
||||
)
|
||||
except Exception as tracking_err:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: failed to track cost for batch {batch_id} "
|
||||
f"(job {job.id}); leaving it unprocessed so the next poll retries: {tracking_err}"
|
||||
)
|
||||
self._record_error(prom_logger, "cost_tracking_error")
|
||||
continue
|
||||
if tracked is None:
|
||||
continue
|
||||
|
||||
# Track this job for the final metrics summary
|
||||
processed_models.append(tracked)
|
||||
|
||||
# mark the job as complete
|
||||
try:
|
||||
update_data: dict = {
|
||||
"status": "complete",
|
||||
"file_object": response.model_dump_json(),
|
||||
}
|
||||
if self._has_batch_processed_column:
|
||||
update_data["batch_processed"] = True
|
||||
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 mark job {job.id} complete in DB: {db_err}"
|
||||
)
|
||||
|
||||
elif response.status in ("failed", "expired", "cancelled"):
|
||||
try:
|
||||
update_data = {
|
||||
"status": response.status,
|
||||
"file_object": response.model_dump_json(),
|
||||
}
|
||||
if self._has_batch_processed_column:
|
||||
update_data["batch_processed"] = True
|
||||
await self.prisma_client.db.litellm_managedobjecttable.update(
|
||||
where={"id": job.id},
|
||||
data=update_data,
|
||||
litellm_metadata={
|
||||
"user_api_key_user_id": job.created_by or "default-user-id",
|
||||
"batch_ignore_default_logging": True,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.info(
|
||||
f"CheckBatchCost: marked job {job.id} as {response.status} in DB"
|
||||
)
|
||||
except Exception as db_err:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: failed to mark job {job.id} as {response.status} in DB: {db_err}"
|
||||
f"Skipping job {job.unified_object_id} because of error querying model ID: {model_id} for cost and usage of batch ID: {batch_id}: {e}"
|
||||
)
|
||||
self._record_error(prom_logger, "provider_retrieval_error")
|
||||
continue
|
||||
|
||||
## RETRIEVE THE BATCH JOB OUTPUT FILE
|
||||
if (
|
||||
response.status == "completed"
|
||||
and response.output_file_id is not None
|
||||
):
|
||||
try:
|
||||
tracked = await self._track_completed_batch_cost(
|
||||
job=job,
|
||||
response=response,
|
||||
model_id=model_id,
|
||||
batch_id=batch_id,
|
||||
prom_logger=prom_logger,
|
||||
)
|
||||
except Exception as tracking_err:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: failed to track cost for batch {batch_id} "
|
||||
f"(job {job.id}); leaving it unprocessed so the next poll retries: {tracking_err}"
|
||||
)
|
||||
self._record_error(prom_logger, "cost_tracking_error")
|
||||
continue
|
||||
if tracked is None:
|
||||
continue
|
||||
|
||||
# Track this job for the final metrics summary
|
||||
processed_models.append(tracked)
|
||||
|
||||
# mark the job as complete
|
||||
try:
|
||||
update_data: dict = {
|
||||
"status": "complete",
|
||||
"file_object": response.model_dump_json(),
|
||||
}
|
||||
if self._has_batch_processed_column:
|
||||
update_data["batch_processed"] = True
|
||||
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 mark job {job.id} complete in DB: {db_err}"
|
||||
)
|
||||
|
||||
elif response.status in ("failed", "expired", "cancelled"):
|
||||
try:
|
||||
update_data = {
|
||||
"status": response.status,
|
||||
"file_object": response.model_dump_json(),
|
||||
}
|
||||
if self._has_batch_processed_column:
|
||||
update_data["batch_processed"] = True
|
||||
await self.prisma_client.db.litellm_managedobjecttable.update(
|
||||
where={"id": job.id},
|
||||
data=update_data,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
f"CheckBatchCost: marked job {job.id} as {response.status} in DB"
|
||||
)
|
||||
except Exception as db_err:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: failed to mark job {job.id} as {response.status} in DB: {db_err}"
|
||||
)
|
||||
except Exception as job_err:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: unhandled error processing job "
|
||||
f"{getattr(job, 'unified_object_id', job.id)}; continuing with next job: {job_err}"
|
||||
)
|
||||
self._record_error(prom_logger, "job_processing_error")
|
||||
continue
|
||||
|
||||
# Record polling run metrics (always, even if nothing was processed)
|
||||
if prom_logger:
|
||||
|
|
|
|||
|
|
@ -420,6 +420,255 @@ class TestCheckBatchCost:
|
|||
), "update() must include batch_processed=True when column is present"
|
||||
assert update_data["status"] == "complete"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prometheus_error_during_failure_handling_does_not_block_siblings(
|
||||
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""Prometheus metric failures while handling a poisoned job must not abort siblings."""
|
||||
from unittest.mock import patch
|
||||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
|
||||
return_value=0
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
failing_job = MagicMock()
|
||||
failing_job.id = "job-failing-prom"
|
||||
failing_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
|
||||
failing_job.created_by = "user-1"
|
||||
|
||||
healthy_job = MagicMock()
|
||||
healthy_job.id = "job-healthy-prom"
|
||||
healthy_job.unified_object_id = "aGVhbHRoeV9iYXRjaF9pZA=="
|
||||
healthy_job.created_by = "user-2"
|
||||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[failing_job, healthy_job]
|
||||
)
|
||||
|
||||
failing_response = MagicMock()
|
||||
failing_response.status = "completed"
|
||||
failing_response.output_file_id = "file-output-fail"
|
||||
|
||||
healthy_response = MagicMock()
|
||||
healthy_response.status = "completed"
|
||||
healthy_response.output_file_id = "file-output-ok"
|
||||
healthy_response.model_dump_json.return_value = (
|
||||
'{"id":"batch-ok","status":"completed"}'
|
||||
)
|
||||
|
||||
mock_llm_router.aretrieve_batch = AsyncMock(
|
||||
side_effect=[failing_response, healthy_response]
|
||||
)
|
||||
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
|
||||
return_value={"api_key": "sk-test"}
|
||||
)
|
||||
|
||||
mock_deployment = MagicMock()
|
||||
mock_deployment.litellm_params.custom_llm_provider = "openai"
|
||||
mock_deployment.litellm_params.model = "gpt-4"
|
||||
mock_deployment.model_info.model_dump.return_value = {}
|
||||
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
|
||||
|
||||
mock_file_content = MagicMock()
|
||||
mock_file_content.content = b'{"id":"req-1"}'
|
||||
|
||||
mock_prom_logger = MagicMock()
|
||||
mock_prom_logger.record_check_batch_cost_error.side_effect = RuntimeError(
|
||||
"metrics backend unavailable"
|
||||
)
|
||||
|
||||
decoded_ids = [
|
||||
"llm_model_id,model-123;llm_batch_id,batch-fail;",
|
||||
None,
|
||||
"llm_model_id,model-123;llm_batch_id,batch-ok;",
|
||||
None,
|
||||
]
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
|
||||
return_value=mock_prom_logger,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
|
||||
side_effect=decoded_ids,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
|
||||
return_value="model-123",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
|
||||
side_effect=["batch-fail", "batch-ok"],
|
||||
),
|
||||
patch(
|
||||
"litellm.files.main.afile_content",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=[
|
||||
ValueError("Failed to get batch output file content"),
|
||||
mock_file_content,
|
||||
],
|
||||
),
|
||||
patch(
|
||||
"litellm.batches.batch_utils._get_file_content_as_dictionary",
|
||||
return_value=[{"id": "req-1"}],
|
||||
),
|
||||
patch(
|
||||
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(
|
||||
0.01,
|
||||
{"prompt_tokens": 10, "completion_tokens": 5},
|
||||
["gpt-4"],
|
||||
),
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
|
||||
return_value=("gpt-4", "openai", None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.litellm_logging.Logging"
|
||||
) as mock_logging_cls,
|
||||
):
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
mock_logging_cls.return_value = mock_logging_obj
|
||||
|
||||
await check_batch_cost_instance.check_batch_cost()
|
||||
|
||||
assert mock_llm_router.aretrieve_batch.await_count == 2
|
||||
assert (
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cost_tracking_failure_does_not_block_sibling_jobs(
|
||||
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""#35357: one poisoned batch must not abort the poll cycle for siblings."""
|
||||
from unittest.mock import patch
|
||||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
|
||||
return_value=0
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
failing_job = MagicMock()
|
||||
failing_job.id = "job-failing-1"
|
||||
failing_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
|
||||
failing_job.created_by = "user-1"
|
||||
|
||||
healthy_job = MagicMock()
|
||||
healthy_job.id = "job-healthy-1"
|
||||
healthy_job.unified_object_id = "aGVhbHRoeV9iYXRjaF9pZA=="
|
||||
healthy_job.created_by = "user-2"
|
||||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[failing_job, healthy_job]
|
||||
)
|
||||
|
||||
failing_response = MagicMock()
|
||||
failing_response.status = "completed"
|
||||
failing_response.output_file_id = "file-output-fail"
|
||||
failing_response.model_dump_json.return_value = (
|
||||
'{"id":"batch-fail","status":"completed"}'
|
||||
)
|
||||
|
||||
healthy_response = MagicMock()
|
||||
healthy_response.status = "completed"
|
||||
healthy_response.output_file_id = "file-output-ok"
|
||||
healthy_response.model_dump_json.return_value = (
|
||||
'{"id":"batch-ok","status":"completed"}'
|
||||
)
|
||||
|
||||
mock_llm_router.aretrieve_batch = AsyncMock(
|
||||
side_effect=[failing_response, healthy_response]
|
||||
)
|
||||
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
|
||||
return_value={"api_key": "sk-test"}
|
||||
)
|
||||
|
||||
mock_deployment = MagicMock()
|
||||
mock_deployment.litellm_params.custom_llm_provider = "openai"
|
||||
mock_deployment.litellm_params.model = "gpt-4"
|
||||
mock_deployment.model_info.model_dump.return_value = {}
|
||||
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
|
||||
|
||||
mock_file_content = MagicMock()
|
||||
mock_file_content.content = b'{"id":"req-1"}'
|
||||
|
||||
decoded_ids = [
|
||||
"llm_model_id,model-123;llm_batch_id,batch-fail;",
|
||||
None,
|
||||
"llm_model_id,model-123;llm_batch_id,batch-ok;",
|
||||
None,
|
||||
]
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
|
||||
side_effect=decoded_ids,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
|
||||
return_value="model-123",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
|
||||
side_effect=["batch-fail", "batch-ok"],
|
||||
),
|
||||
patch(
|
||||
"litellm.files.main.afile_content",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=[
|
||||
ValueError("Failed to get batch output file content"),
|
||||
mock_file_content,
|
||||
],
|
||||
),
|
||||
patch(
|
||||
"litellm.batches.batch_utils._get_file_content_as_dictionary",
|
||||
return_value=[{"id": "req-1"}],
|
||||
),
|
||||
patch(
|
||||
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(
|
||||
0.01,
|
||||
{"prompt_tokens": 10, "completion_tokens": 5},
|
||||
["gpt-4"],
|
||||
),
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
|
||||
return_value=("gpt-4", "openai", None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.litellm_logging.Logging"
|
||||
) as mock_logging_cls,
|
||||
):
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
mock_logging_cls.return_value = mock_logging_obj
|
||||
|
||||
await check_batch_cost_instance.check_batch_cost()
|
||||
|
||||
assert mock_llm_router.aretrieve_batch.await_count == 2
|
||||
assert (
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
|
||||
), "only the healthy sibling should be marked processed"
|
||||
update_call = (
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update.call_args_list[0]
|
||||
)
|
||||
assert update_call[1]["where"] == {"id": "job-healthy-1"}
|
||||
assert update_call[1]["data"]["batch_processed"] is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cost_tracking_failure_leaves_job_unprocessed(
|
||||
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue