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 e89501fca6f..66761e84e57 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -97,7 +97,11 @@ class CheckBatchCost: await self._cleanup_stale_managed_objects() - # Look for all batches that have not yet been processed by CheckBatchCost + # Look for all batches that have not yet been processed by CheckBatchCost. + # _has_batch_processed_column tracks whether the column exists so the + # completion update can omit it on older schemas (avoiding a silent failure + # that would cause infinite reprocessing and duplicate cost logging). + _has_batch_processed_column = True try: jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( where={ @@ -110,6 +114,7 @@ class CheckBatchCost: ) except Exception: # Fallback: batch_processed column may not exist on older schemas + _has_batch_processed_column = False verbose_proxy_logger.warning( "CheckBatchCost: batch_processed column not found, querying without it" ) @@ -288,13 +293,15 @@ class CheckBatchCost: # mark the job as complete try: + update_data: dict = { + "status": "complete", + "file_object": response.model_dump_json(), + } + if _has_batch_processed_column: + update_data["batch_processed"] = True await self.prisma_client.db.litellm_managedobjecttable.update( where={"id": job.id}, - data={ - "batch_processed": True, - "status": "complete", - "file_object": response.model_dump_json(), - }, + data=update_data, ) except Exception as db_err: verbose_proxy_logger.error( diff --git a/litellm/constants.py b/litellm/constants.py index 2106c527d96..dbc79b69a67 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1351,9 +1351,9 @@ PROXY_BUDGET_RESCHEDULER_MIN_TIME = int( os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597) ) PROXY_BATCH_POLLING_INTERVAL = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600)) -MAX_OBJECTS_PER_POLL_CYCLE = int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)) -MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = int( - os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7) +MAX_OBJECTS_PER_POLL_CYCLE = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50))) +MANAGED_OBJECT_STALENESS_CUTOFF_DAYS = max( + 1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)) ) # Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and # CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 36a2f35de1a..94c33669e8b 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -110,3 +110,49 @@ class TestCheckBatchCost: assert "stale_expired" in fallback_where["status"]["not_in"] # Fallback must still paginate assert calls[1][1]["take"] == MAX_OBJECTS_PER_POLL_CYCLE + + @pytest.mark.asyncio + async def test_fallback_completion_update_omits_batch_processed( + self, check_batch_cost_instance, mock_prisma_client, mock_llm_router + ): + """When batch_processed column is absent, completion update must not include it. + + If it did, the update would fail silently, the job would never be marked done, + and every subsequent poll cycle would re-log the cost (duplicate billing). + """ + 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_job = MagicMock() + mock_job.id = "job-fallback-1" + mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA==" # base64-looking value + mock_job.created_by = "user-1" + + # Primary query fails → fallback path + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + side_effect=[Exception("column batch_processed does not exist"), [mock_job]] + ) + + # Stub out the heavy per-job processing so we reach the update() + with ( + patch( + "litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id", + return_value=None, # causes "not a valid unified object id" early-continue + ), + ): + await check_batch_cost_instance.check_batch_cost() + + # Even though the job was skipped (invalid ID), confirm the fallback path was taken + # by checking the find_many calls + find_calls = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list + assert len(find_calls) == 2 + fallback_where = find_calls[1][1]["where"] + assert "batch_processed" not in fallback_where + + # If a completion update were issued, it must not contain batch_processed + for call in mock_prisma_client.db.litellm_managedobjecttable.update.call_args_list: + assert "batch_processed" not in call[1].get("data", {})