fix: cache _has_batch_processed_column; guard cleanup from aborting poll; narrow fallback except

This commit is contained in:
Ishaan Jaffer 2026-03-12 18:28:10 -07:00
parent a14b951248
commit 679b3b247e
3 changed files with 86 additions and 43 deletions

View file

@ -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},

View file

@ -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={

View file

@ -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(