fix(batches): stop uncostable batches from starving the cost poll page

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-08-12 23:45:36 +00:00
parent 964f0755ee
commit c11ebbed27
2 changed files with 265 additions and 4 deletions

View file

@ -3,7 +3,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
"""
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
from typing import TYPE_CHECKING, Any, Dict, Final, List, Optional, Tuple
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -23,6 +23,15 @@ if TYPE_CHECKING:
CHECK_BATCH_COST_USER_AGENT = "LiteLLM Proxy/CheckBatchCost"
TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
"completed",
"complete",
"failed",
"expired",
"cancelled",
"stale_expired",
)
class CheckBatchCost:
def __init__(
@ -132,11 +141,11 @@ class CheckBatchCost:
in non-terminal states as 'stale_expired'. These will never complete and
should not be polled.
"""
cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
result = await self.prisma_client.db.litellm_managedobjecttable.update_many(
cutoff: Final = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
result: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
where={
"file_purpose": "batch",
"status": {"not_in": ["completed", "complete", "failed", "expired", "cancelled", "stale_expired"]},
"status": {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)},
"created_at": {"lt": cutoff},
},
data={"status": "stale_expired"},
@ -147,6 +156,26 @@ class CheckBatchCost:
f"(older than {MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days) as stale_expired"
)
if not self._has_batch_processed_column:
return
# A row already in a terminal status is never rewritten by the sweep above, so
# without this it keeps a poll-page slot forever and starves newer batches.
retired: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
where={
"file_purpose": "batch",
"batch_processed": False,
"status": {"in": ["complete", "completed"]},
"created_at": {"lt": cutoff},
},
data={"batch_processed": True},
)
if retired > 0:
verbose_proxy_logger.warning(
f"CheckBatchCost: gave up on {retired} completed managed objects older than "
f"{MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days that were never costed"
)
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(
@ -167,6 +196,54 @@ class CheckBatchCost:
order={"created_at": "asc"},
)
async def _retire_job(self, job: "LiteLLM_ManagedObjectTable", reason: str) -> None:
"""
Take a row that can never be costed out of the poll page. Leaving it selectable
would burn one of the MAX_OBJECTS_PER_POLL_CYCLE slots on every future cycle, and
once enough such rows accumulate no newer batch is ever reached. Older schemas
without batch_processed can only be excluded through the status filter.
"""
data: Final = (
{"batch_processed": True}
if self._has_batch_processed_column
else {"status": "stale_expired"}
)
try:
await self.prisma_client.db.litellm_managedobjecttable.update(
where={"id": job.id},
data=data,
)
except Exception as db_err:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to retire uncostable job {job.id} ({reason}): {db_err}"
)
return
verbose_proxy_logger.warning(
f"CheckBatchCost: job {job.id} can never be costed ({reason}), "
"so it will no longer be polled"
)
@staticmethod
def _has_unified_id_without_model(job: "LiteLLM_ManagedObjectTable") -> bool:
"""A unified id that decodes but carries no model_id can never be routed."""
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
get_model_id_from_unified_batch_id,
)
decoded: Final = _is_base64_encoded_unified_file_id(job.unified_object_id)
return bool(decoded) and get_model_id_from_unified_batch_id(decoded) is None
@staticmethod
def _is_batch_gone_at_provider(error: Exception) -> bool:
"""A 404 from the provider means it dropped its record of the batch, so no later
retrieve can ever succeed."""
import openai
from litellm.exceptions import NotFoundError
return isinstance(error, (NotFoundError, openai.NotFoundError))
@staticmethod
def _record_error(
prom_logger: Optional["PrometheusLogger"], error_type: str
@ -645,6 +722,8 @@ class CheckBatchCost:
for job in jobs:
routing = self._resolve_job_routing(job, prom_logger)
if routing is None:
if self._has_unified_id_without_model(job):
await self._retire_job(job, "unified object id has no model id")
continue
model_id, batch_id = routing
@ -667,6 +746,8 @@ class CheckBatchCost:
)
if prom_logger:
prom_logger.record_check_batch_cost_error("provider_retrieval_error")
if self._is_batch_gone_at_provider(e):
await self._retire_job(job, f"batch {batch_id} no longer exists at the provider")
continue
## RETRIEVE THE BATCH JOB OUTPUT FILE

View file

@ -1791,3 +1791,183 @@ class TestBatchCostAttribution:
metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1")
assert metadata["user_api_key_alias"] == "prod-key"
class TestPollPageStarvation:
"""LIT-5462 regression: a row that can never be costed used to keep its slot in the
MAX_OBJECTS_PER_POLL_CYCLE page forever, so once enough of them accumulated no newer
batch was ever polled or costed."""
def _instance(self, prisma, llm_router):
from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost
proxy_logging_obj = MagicMock()
proxy_logging_obj.get_proxy_hook.return_value = None
return CheckBatchCost(
proxy_logging_obj=proxy_logging_obj,
prisma_client=prisma,
llm_router=llm_router,
)
def _prisma(self, jobs):
prisma = MagicMock()
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0)
prisma.db.litellm_managedobjecttable.update = AsyncMock()
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=jobs)
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
return prisma
def _job(self, job_id, unified_object_id):
job = MagicMock()
job.id = job_id
job.unified_object_id = unified_object_id
job.created_by = "user-1"
return job
@staticmethod
def _encode(unified_id: str) -> str:
import base64
return base64.urlsafe_b64encode(unified_id.encode()).decode().rstrip("=")
@pytest.mark.asyncio
async def test_unified_id_without_model_id_is_retired(self):
"""A unified id that decodes but carries no model_id is unroutable no matter what
the config says, so it must leave the poll page instead of being retried forever."""
prisma = self._prisma(
[self._job("job-no-model", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))]
)
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock()
await self._instance(prisma, llm_router).check_batch_cost()
llm_router.aretrieve_batch.assert_not_awaited()
prisma.db.litellm_managedobjecttable.update.assert_awaited_once()
call = prisma.db.litellm_managedobjecttable.update.call_args[1]
assert call["where"] == {"id": "job-no-model"}
assert call["data"] == {"batch_processed": True}
@pytest.mark.asyncio
async def test_provider_404_retires_job(self):
"""The provider dropping its record of the batch is permanent: no later retrieve
can succeed, so the row must stop occupying a slot."""
import litellm
prisma = self._prisma(
[
self._job(
"job-gone",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_deadbeef"),
)
]
)
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock(
side_effect=litellm.NotFoundError(
message="No batch found with id 'batch_deadbeef'.",
model="model-123",
llm_provider="openai",
)
)
await self._instance(prisma, llm_router).check_batch_cost()
prisma.db.litellm_managedobjecttable.update.assert_awaited_once()
assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == {
"batch_processed": True
}
@pytest.mark.asyncio
async def test_transient_provider_error_keeps_job_for_retry(self):
"""A failure that may clear up (timeout, 5xx) must still leave the row unprocessed."""
prisma = self._prisma(
[
self._job(
"job-flaky",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_flaky"),
)
]
)
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock(side_effect=Exception("connection reset"))
await self._instance(prisma, llm_router).check_batch_cost()
prisma.db.litellm_managedobjecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_retirement_falls_back_to_status_without_batch_processed_column(self):
"""Older schemas have no batch_processed column, so the only way to stop selecting
the row is the status filter the poll query already applies."""
prisma = self._prisma(
[self._job("job-legacy", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))]
)
instance = self._instance(prisma, MagicMock())
instance._has_batch_processed_column = False
await instance.check_batch_cost()
assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == {
"status": "stale_expired"
}
@pytest.mark.asyncio
async def test_stale_cleanup_gives_up_on_never_costed_completed_rows(self):
"""A row already in a terminal status is never rewritten by the staleness sweep, so
it needs its own bound or it starves newer batches indefinitely."""
prisma = self._prisma([])
await self._instance(prisma, MagicMock()).check_batch_cost()
calls = prisma.db.litellm_managedobjecttable.update_many.call_args_list
assert len(calls) == 2, "expected the staleness sweep plus the never-costed sweep"
where = calls[1][1]["where"]
assert where["file_purpose"] == "batch"
assert where["batch_processed"] is False
assert where["status"] == {"in": ["complete", "completed"]}
assert "created_at" in where
assert calls[1][1]["data"] == {"batch_processed": True}
@pytest.mark.asyncio
async def test_newer_batch_is_polled_once_dead_rows_are_retired(self):
"""The end state the customer cares about: dead rows retire on the cycle they are
first seen, and the healthy batch behind them keeps getting polled."""
dead_rows = [
self._job("job-no-model", self._encode("litellm_proxy;llm_batch_id:poison-no-model")),
self._job(
"job-gone",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_deadbeef"),
),
]
live_row = self._job(
"job-live",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_live"),
)
prisma = self._prisma(dead_rows + [live_row])
import litellm
in_progress = MagicMock()
in_progress.status = "in_progress"
async def _retrieve(model, batch_id, litellm_metadata):
if batch_id == "batch_deadbeef":
raise litellm.NotFoundError(
message="No batch found", model=model, llm_provider="openai"
)
return in_progress
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock(side_effect=_retrieve)
await self._instance(prisma, llm_router).check_batch_cost()
retired = [
call[1]["where"]["id"]
for call in prisma.db.litellm_managedobjecttable.update.call_args_list
]
assert retired == ["job-no-model", "job-gone"]
assert (
llm_router.aretrieve_batch.await_args_list[-1][1]["batch_id"] == "batch_live"
), "the newer healthy batch must still be polled in the same cycle"