From d9b9676d49f570bcacf0507ac07fa08caa602e8e Mon Sep 17 00:00:00 2001 From: Souravrajvi0 <144546710+Souravrajvi0@users.noreply.github.com> Date: Fri, 31 Jul 2026 14:47:22 +0000 Subject: [PATCH] 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 --- .../proxy/common_utils/check_batch_cost.py | 172 ++++++------ .../proxy_unit_tests/test_check_batch_cost.py | 249 ++++++++++++++++++ 2 files changed, 341 insertions(+), 80 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index f209ab54f64..ed4a43b6f6d 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -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: diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index b822799fb40..fe0b12225b5 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -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