From 679b3b247ebff9222d3855f96630a9bd120323fa Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Thu, 12 Mar 2026 18:28:10 -0700 Subject: [PATCH] fix: cache _has_batch_processed_column; guard cleanup from aborting poll; narrow fallback except --- .../proxy/common_utils/check_batch_cost.py | 92 +++++++++++-------- .../common_utils/check_responses_cost.py | 7 +- .../proxy_unit_tests/test_check_batch_cost.py | 30 +++++- 3 files changed, 86 insertions(+), 43 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 66761e84e57..5609225f9f1 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -33,6 +33,9 @@ class CheckBatchCost: self.proxy_logging_obj: ProxyLogging = proxy_logging_obj self.prisma_client: PrismaClient = prisma_client self.llm_router: Router = llm_router + # Cached after the first poll cycle. Once we know the column is absent we skip + # the guaranteed-failing primary query on every subsequent cycle. + self._has_batch_processed_column: bool = True async def _get_user_info(self, batch_id, user_id) -> dict: """ @@ -74,6 +77,26 @@ class CheckBatchCost: f"(older than {MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days) as stale_expired" ) + async def _fallback_find_jobs(self) -> list: + """Query batch jobs without the batch_processed filter (for older schemas).""" + return await self.prisma_client.db.litellm_managedobjecttable.find_many( + where={ + "file_purpose": "batch", + "status": { + "not_in": [ + "failed", + "expired", + "cancelled", + "complete", + "completed", + "stale_expired", + ] + }, + }, + take=MAX_OBJECTS_PER_POLL_CYCLE, + order={"created_at": "asc"}, + ) + async def check_batch_cost(self): """ Check if the batch JOB has been tracked. @@ -95,46 +118,39 @@ class CheckBatchCost: get_model_id_from_unified_batch_id, ) - await self._cleanup_stale_managed_objects() + try: + await self._cleanup_stale_managed_objects() + except Exception as cleanup_err: + verbose_proxy_logger.warning( + f"CheckBatchCost: stale cleanup failed (poll will continue): {cleanup_err}" + ) # 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={ - "file_purpose": "batch", - "batch_processed": False, - "status": {"not_in": ["failed", "expired", "cancelled", "stale_expired"]}, - }, - take=MAX_OBJECTS_PER_POLL_CYCLE, - order={"created_at": "asc"}, - ) - 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" - ) - jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( - where={ - "file_purpose": "batch", - "status": { - "not_in": [ - "failed", - "expired", - "cancelled", - "complete", - "completed", - "stale_expired", - ] + # self._has_batch_processed_column is cached after the first probe so that + # older schemas don't pay a guaranteed-failing primary query + warning on + # every subsequent poll cycle. + if self._has_batch_processed_column: + try: + jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( + where={ + "file_purpose": "batch", + "batch_processed": False, + "status": {"not_in": ["failed", "expired", "cancelled", "stale_expired"]}, }, - }, - take=MAX_OBJECTS_PER_POLL_CYCLE, - order={"created_at": "asc"}, - ) + take=MAX_OBJECTS_PER_POLL_CYCLE, + order={"created_at": "asc"}, + ) + except Exception as query_err: + if "batch_processed" not in str(query_err).lower() and "unknown column" not in str(query_err).lower() and "does not exist" not in str(query_err).lower(): + raise + # Permanent schema gap — cache the result so future cycles skip straight to fallback + self._has_batch_processed_column = False + verbose_proxy_logger.warning( + "CheckBatchCost: batch_processed column not found, querying without it" + ) + jobs = await self._fallback_find_jobs() + else: + jobs = await self._fallback_find_jobs() for job in jobs: # get the model from the job unified_object_id = job.unified_object_id @@ -297,7 +313,7 @@ class CheckBatchCost: "status": "complete", "file_object": response.model_dump_json(), } - if _has_batch_processed_column: + if self._has_batch_processed_column: update_data["batch_processed"] = True await self.prisma_client.db.litellm_managedobjecttable.update( where={"id": job.id}, diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index 6ee2bce272a..54fbc7abcc5 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -61,7 +61,12 @@ class CheckResponsesCost: - Cost is automatically tracked by litellm.aget_responses() - Mark completed/failed/cancelled responses as complete in the database """ - await self._cleanup_stale_managed_objects() + try: + await self._cleanup_stale_managed_objects() + except Exception as cleanup_err: + verbose_proxy_logger.warning( + f"CheckResponsesCost: stale cleanup failed (poll will continue): {cleanup_err}" + ) jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( where={ diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 94c33669e8b..7606ef49c2c 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -94,7 +94,7 @@ class TestCheckBatchCost: mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( return_value=0 ) - # First find_many (primary query) raises; second (fallback) returns empty list + # First find_many (primary query) raises with a schema error; second (fallback) returns empty mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( side_effect=[Exception("column batch_processed does not exist"), []] ) @@ -104,12 +104,34 @@ class TestCheckBatchCost: calls = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list assert len(calls) == 2 fallback_where = calls[1][1]["where"] - # Fallback must not reference batch_processed assert "batch_processed" not in fallback_where - # Fallback must exclude stale_expired assert "stale_expired" in fallback_where["status"]["not_in"] - # Fallback must still paginate assert calls[1][1]["take"] == MAX_OBJECTS_PER_POLL_CYCLE + # Column absence is now cached — next call should go straight to fallback + assert check_batch_cost_instance._has_batch_processed_column is False + + @pytest.mark.asyncio + async def test_column_absence_cached_across_cycles( + self, check_batch_cost_instance, mock_prisma_client + ): + """After column absence is discovered, subsequent cycles skip the primary query entirely.""" + from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=0 + ) + # Simulate column already known absent from a previous cycle + check_batch_cost_instance._has_batch_processed_column = False + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[] + ) + + await check_batch_cost_instance.check_batch_cost() + + # Only one find_many call — the fallback directly, no primary query attempt + assert mock_prisma_client.db.litellm_managedobjecttable.find_many.call_count == 1 + fallback_where = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args[1]["where"] + assert "batch_processed" not in fallback_where @pytest.mark.asyncio async def test_fallback_completion_update_omits_batch_processed(