mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
964f0755ee
commit
c11ebbed27
2 changed files with 265 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue