mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: cache _has_batch_processed_column; guard cleanup from aborting poll; narrow fallback except
This commit is contained in:
parent
a14b951248
commit
679b3b247e
3 changed files with 86 additions and 43 deletions
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue