This commit is contained in:
devin-ai-integration[bot] 2026-09-30 16:57:07 -04:00 • committed by GitHub
commit 9defe463fa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 1020 additions and 220 deletions

View file

@ -4796,6 +4796,63 @@
"title": "litellm_check_batch_cost_errors rate",
"type": "timeseries"
},
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"description": "Batches retired after the staleness cutoff without reconciling cost",
"fieldConfig": {
"defaults": {
"color": {
"mode": "palette-classic"
},
"custom": {
"drawStyle": "line",
"fillOpacity": 10,
"lineWidth": 1,
"showPoints": "never",
"spanNulls": false
},
"unit": "short"
},
"overrides": []
},
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 338
},
"id": 112,
"options": {
"legend": {
"calcs": [],
"displayMode": "list",
"placement": "bottom",
"showLegend": true
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "sum(rate(litellm_check_batch_cost_stale_expired_total[$__rate_interval]))",
"legendFormat": "batches with unreconciled cost",
"range": true,
"refId": "A"
}
],
"title": "litellm_check_batch_cost_stale_expired rate",
"type": "timeseries"
},
{
"datasource": {
"type": "prometheus",

View file

@ -3,15 +3,17 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
"""
from collections.abc import Sequence
from dataclasses import dataclass, field
from dataclasses import replace as dataclasses_replace
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Protocol, Tuple, cast
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, Tuple, TypeAlias, cast
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import (
CLI_SESSION_KEY_PREFIX,
MANAGED_OBJECT_STALE_RECONCILE_GRACE_DAYS,
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
MAX_OBJECTS_PER_POLL_CYCLE,
)
@ -43,10 +45,42 @@ TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
)
@dataclass(frozen=True, slots=True)
class _JobBilled:
model: str | None
api_provider: str | None
@dataclass(frozen=True, slots=True)
class _JobSettled:
pass
@dataclass(frozen=True, slots=True)
class _StaleSweepResult:
processed_models: tuple[tuple[str | None, str | None], ...]
polled: int
@dataclass(frozen=True, slots=True)
class _JobUnreconciled:
reason: str
retryable: bool = field(kw_only=True)
_JobOutcome: TypeAlias = _JobBilled | _JobSettled | _JobUnreconciled
class _ManagedObjectRow(Protocol):
@property
def id(self) -> str: ...
@property
def status(self) -> str: ...
@property
def created_at(self) -> datetime: ...
@property
def unified_object_id(self) -> str: ...
@ -57,6 +91,14 @@ class _ManagedObjectRow(Protocol):
def file_object(self) -> object: ...
def _normalize_datetime_to_utc(value: datetime) -> datetime:
return (
value.replace(tzinfo=timezone.utc)
if value.tzinfo is None
else value.astimezone(timezone.utc)
)
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
@ -248,52 +290,130 @@ class CheckBatchCost:
return metadata
async def _cleanup_stale_managed_objects(self) -> None:
"""
Mark managed objects older than MANAGED_OBJECT_STALENESS_CUTOFF_DAYS days
in non-terminal states as 'stale_expired'. These will never complete and
should not be polled.
"""
cutoff: Final = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
result: Final = await _managed_object_table(self.prisma_client).update_many(
async def _cleanup_stale_managed_objects(
self, prom_logger: "PrometheusLogger | None", cutoff: datetime
) -> _StaleSweepResult:
grace_deadline: Final = cutoff - timedelta(
days=MANAGED_OBJECT_STALE_RECONCILE_GRACE_DAYS
)
if self._has_batch_processed_column:
await _managed_object_table(self.prisma_client).update_many(
where={
"file_purpose": "batch",
"batch_processed": True,
"status": {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)},
"created_at": {"lt": cutoff},
},
data={"status": "stale_expired"},
)
batch_processed_filter: Final = (
{"batch_processed": False} if self._has_batch_processed_column else {}
)
status_filter: Final = (
{"not_in": ["failed", "expired", "cancelled", "stale_expired"]}
if self._has_batch_processed_column
else {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)}
)
candidates: Final = await _managed_object_table(self.prisma_client).find_many(
where={
"file_purpose": "batch",
"status": {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)},
**batch_processed_filter,
"status": status_filter,
"created_at": {"lt": cutoff},
},
data={"status": "stale_expired"},
take=MAX_OBJECTS_PER_POLL_CYCLE,
order={"created_at": "asc"},
)
if result > 0:
verbose_proxy_logger.warning(
f"CheckBatchCost: marked {result} stale managed objects "
f"(older than {MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days) as stale_expired"
)
if not self._has_batch_processed_column:
reconciled: Final = [
await self._reconcile_stale_candidate(job, grace_deadline, prom_logger)
for job in candidates
]
return _StaleSweepResult(
processed_models=tuple(model_pair for model_pair in reconciled if model_pair is not None),
polled=len(candidates),
)
async def _reconcile_stale_candidate(
self,
job: "_ManagedObjectRow",
grace_deadline: datetime,
prom_logger: "PrometheusLogger | None",
) -> tuple[str | None, str | None] | None:
try:
match await self._poll_job(job, prom_logger):
case _JobBilled(model=model, api_provider=api_provider):
return model, api_provider
case _JobSettled():
return None
case _JobUnreconciled(retryable=True) if _normalize_datetime_to_utc(
job.created_at
) >= grace_deadline:
return None
case _JobUnreconciled(reason=reason):
await self._expire_unreconciled_job(job, reason, prom_logger)
return None
except Exception as reconciliation_err:
verbose_proxy_logger.warning(
f"CheckBatchCost: stale candidate {job.id} reconciliation failed: "
f"{reconciliation_err}"
)
return None
return None
async def _cleanup_stale_managed_objects_safely(
self, prom_logger: "PrometheusLogger | None", cutoff: datetime
) -> _StaleSweepResult:
try:
return await self._cleanup_stale_managed_objects(prom_logger, cutoff)
except Exception as cleanup_err:
verbose_proxy_logger.warning(
f"CheckBatchCost: stale cleanup failed (poll will continue): {cleanup_err}"
)
return _StaleSweepResult(processed_models=(), polled=0)
async def _expire_unreconciled_job(
self,
job: "_ManagedObjectRow",
reason: str,
prom_logger: "PrometheusLogger | None",
) -> None:
mark_completed_processed: Final = (
self._has_batch_processed_column and job.status in ("complete", "completed")
)
update_where: Final = {
"id": job.id,
**({"batch_processed": False} if self._has_batch_processed_column else {}),
**(
{"status": {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)}}
if not mark_completed_processed
else {}
),
}
update_data: Final = (
{"batch_processed": True}
if mark_completed_processed
else {"status": "stale_expired"}
)
updated: Final = await _managed_object_table(self.prisma_client).update_many(
where=update_where,
data=update_data,
)
if updated == 0:
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 _managed_object_table(self.prisma_client).update_many(
where={
"file_purpose": "batch",
"batch_processed": False,
"status": {"in": ["complete", "completed"]},
"created_at": {"lt": cutoff},
},
data={"batch_processed": True},
verbose_proxy_logger.warning(
f"CheckBatchCost: batch {job.unified_object_id} is older than "
f"{MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days; cost never reconciled ({reason})"
)
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"
)
if prom_logger is not None:
prom_logger.record_check_batch_cost_stale_expired()
async def _fallback_find_jobs(self) -> "Sequence[_ManagedObjectRow]":
async def _fallback_find_jobs(self, cutoff: datetime) -> "Sequence[_ManagedObjectRow]":
"""Query batch jobs without the batch_processed filter (for older schemas)."""
return await _managed_object_table(self.prisma_client).find_many(
where={
"file_purpose": "batch",
"created_at": {"gte": cutoff},
"status": {
"not_in": [
"failed",
@ -552,6 +672,110 @@ class CheckBatchCost:
self._record_error(prom_logger, "invalid_unified_id")
return None
async def _poll_job(
self, job: "_ManagedObjectRow", prom_logger: "PrometheusLogger | None"
) -> _JobOutcome:
routing: Final = 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")
return _JobUnreconciled("job could not be routed", retryable=False)
model_id, batch_id = routing
verbose_proxy_logger.info(
f"Querying model ID: {model_id} for cost and usage of batch ID: {batch_id}"
)
try:
response: Final = await self.llm_router.aretrieve_batch(
model=model_id,
batch_id=batch_id,
litellm_metadata={
"user_api_key_user_id": job.created_by or "default-user-id",
"batch_ignore_default_logging": True,
},
)
except Exception as e:
verbose_proxy_logger.info(
f"Skipping job {job.unified_object_id} because of error querying model ID: {model_id} "
f"for cost and usage of batch ID: {batch_id}: {e}"
)
if prom_logger:
prom_logger.record_check_batch_cost_error("provider_retrieval_error")
batch_gone_at_provider: Final = (
self._is_batch_gone_at_provider(e, batch_id) and self._batch_deployment_exists(model_id)
)
if batch_gone_at_provider:
await self._retire_job(job, f"batch {batch_id} no longer exists at the provider")
return _JobUnreconciled(
f"provider retrieval error: {e}", retryable=not batch_gone_at_provider
)
if response.status in PROVIDER_TERMINAL_BATCH_STATUSES and response.output_file_id is not None:
try:
tracked: Final = await self._track_completed_batch_cost(
job=job,
response=response,
model_id=model_id,
batch_id=batch_id,
prom_logger=prom_logger,
)
except Exception as tracking_err:
if self._is_output_file_gone_at_provider(
tracking_err, response.output_file_id
) and self._batch_deployment_exists(model_id):
verbose_proxy_logger.warning(
f"CheckBatchCost: output file {response.output_file_id} of batch {batch_id} "
f"does not exist at the provider; retiring job {job.id} unbilled"
)
await self._finalize_unbilled_terminal_job(job, response)
return _JobSettled()
verbose_proxy_logger.error(
f"CheckBatchCost: failed to track cost for batch {batch_id} "
f"(job {job.id}); leaving it unprocessed so the next poll retries: {tracking_err}"
)
self._record_error(prom_logger, "cost_tracking_error")
return _JobUnreconciled(f"cost tracking error: {tracking_err}", retryable=True)
if tracked is None:
return _JobUnreconciled(
"cost not tracked: batch claimed by another poller or its deployment is gone",
retryable=False,
)
try:
update_data: Final = {
"status": response.status if response.status != "completed" else "complete",
"file_object": response.model_dump_json(),
**({"batch_processed": True} if self._has_batch_processed_column else {}),
}
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=update_data,
)
except Exception as db_err:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to mark job {job.id} complete in DB: {db_err}"
)
return _JobBilled(model=tracked[0], api_provider=tracked[1])
if response.status in PROVIDER_TERMINAL_BATCH_STATUSES:
from litellm.proxy.openai_files_endpoints.common_utils import (
_completed_batch_safe_to_retire,
)
if response.status in ("completed", "complete") and not _completed_batch_safe_to_retire(response):
verbose_proxy_logger.info(
f"CheckBatchCost: batch {batch_id} is completed but its output file id "
f"has not appeared yet; leaving job {job.id} for the next poll cycle"
)
return _JobUnreconciled(
f"provider status {response.status} with lagging output file", retryable=True
)
await self._finalize_unbilled_terminal_job(job, response)
return _JobSettled()
return _JobUnreconciled(f"provider status {response.status}", retryable=False)
def _resolve_unmanaged_provider_routing(
self,
job: "_ManagedObjectRow",
@ -947,14 +1171,13 @@ class CheckBatchCost:
verbose_proxy_logger.error(f"CheckBatchCost: could not get Prometheus logger: {e}")
prom_logger = None
processed_models: List[Tuple[Optional[str], Optional[str]]] = []
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}"
)
now: Final = datetime.now(timezone.utc)
cutoff: Final = now - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
processed_models: Final[list[tuple[str | None, str | None]]] = []
stale_sweep: Final = await self._cleanup_stale_managed_objects_safely(
prom_logger, cutoff
)
processed_models.extend(stale_sweep.processed_models)
# Look for all batches that have not yet been processed by CheckBatchCost.
# self._has_batch_processed_column is cached after the first probe so that
@ -978,6 +1201,7 @@ class CheckBatchCost:
"stale_expired",
]
},
"created_at": {"gte": cutoff},
},
take=MAX_OBJECTS_PER_POLL_CYCLE,
order={"created_at": "asc"},
@ -991,108 +1215,21 @@ class CheckBatchCost:
verbose_proxy_logger.warning(
"CheckBatchCost: batch_processed column not found, querying without it"
)
jobs = await self._fallback_find_jobs()
jobs = await self._fallback_find_jobs(cutoff)
else:
jobs = await self._fallback_find_jobs()
jobs = await self._fallback_find_jobs(cutoff)
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
verbose_proxy_logger.info(
f"Querying model ID: {model_id} for cost and usage of batch ID: {batch_id}"
)
try:
response = await self.llm_router.aretrieve_batch(
model=model_id,
batch_id=batch_id,
litellm_metadata={
"user_api_key_user_id": job.created_by or "default-user-id",
"batch_ignore_default_logging": True,
},
)
except Exception as e:
verbose_proxy_logger.info(
f"Skipping job {job.unified_object_id} because of error querying model ID: {model_id} for cost and usage of batch ID: {batch_id}: {e}"
)
if prom_logger:
prom_logger.record_check_batch_cost_error("provider_retrieval_error")
if self._is_batch_gone_at_provider(e, batch_id) and self._batch_deployment_exists(model_id):
await self._retire_job(job, f"batch {batch_id} no longer exists at the provider")
continue
## RETRIEVE THE BATCH JOB OUTPUT FILE
if (
response.status in PROVIDER_TERMINAL_BATCH_STATUSES
and response.output_file_id is not None
):
try:
tracked = await self._track_completed_batch_cost(
job=job,
response=response,
model_id=model_id,
batch_id=batch_id,
prom_logger=prom_logger,
)
except Exception as tracking_err:
if self._is_output_file_gone_at_provider(
tracking_err, response.output_file_id
) and self._batch_deployment_exists(model_id):
verbose_proxy_logger.warning(
f"CheckBatchCost: output file {response.output_file_id} of batch {batch_id} "
f"does not exist at the provider; retiring job {job.id} unbilled"
)
await self._finalize_unbilled_terminal_job(job, response)
continue
verbose_proxy_logger.error(
f"CheckBatchCost: failed to track cost for batch {batch_id} "
f"(job {job.id}); leaving it unprocessed so the next poll retries: {tracking_err}"
)
self._record_error(prom_logger, "cost_tracking_error")
continue
if tracked is None:
continue
# Track this job for the final metrics summary
processed_models.append(tracked)
# mark the job as complete
try:
update_data: dict = {
"status": response.status if response.status != "completed" else "complete",
"file_object": response.model_dump_json(),
}
if self._has_batch_processed_column:
update_data["batch_processed"] = True
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=update_data,
)
except Exception as db_err:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to mark job {job.id} complete in DB: {db_err}"
)
elif response.status in PROVIDER_TERMINAL_BATCH_STATUSES:
from litellm.proxy.openai_files_endpoints.common_utils import (
_completed_batch_safe_to_retire,
)
if response.status in ("completed", "complete") and not _completed_batch_safe_to_retire(response):
verbose_proxy_logger.info(
f"CheckBatchCost: batch {batch_id} is completed but its output file id "
f"has not appeared yet; leaving job {job.id} for the next poll cycle"
)
continue
await self._finalize_unbilled_terminal_job(job, response)
match await self._poll_job(job, prom_logger):
case _JobBilled(model=model, api_provider=api_provider):
processed_models.append((model, api_provider))
case _JobSettled():
pass
case _JobUnreconciled():
pass
# Record polling run metrics (always, even if nothing was processed)
if prom_logger:
prom_logger.record_check_batch_cost_run(
jobs_polled=len(jobs),
jobs_polled=len(jobs) + stale_sweep.polled,
processed_models=processed_models if processed_models else None,
)

View file

@ -1806,6 +1806,9 @@ RESET_BUDGET_JOB_LOCK_TTL_SECONDS: Final[int] = 900
PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600))
MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)))
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)))
MANAGED_OBJECT_STALE_RECONCILE_GRACE_DAYS: Final = max(
0, int(os.getenv("MANAGED_OBJECT_STALE_RECONCILE_GRACE_DAYS", "1"))
)
STALE_OBJECT_CLEANUP_BATCH_SIZE: Final = max(1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000)))
# Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and
# CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on

View file

@ -853,6 +853,13 @@ class PrometheusLogger(CustomLogger):
labelnames=["error_type"],
)
self.litellm_check_batch_cost_stale_expired_total = self._counter_factory(
name="litellm_check_batch_cost_stale_expired_total",
documentation="Total number of batches CheckBatchCost gave up on after "
"MANAGED_OBJECT_STALENESS_CUTOFF_DAYS without reconciling cost",
labelnames=[],
)
self.litellm_check_batch_cost_last_run_timestamp = self._gauge_factory(
"litellm_check_batch_cost_last_run_timestamp",
"Unix timestamp of the last CheckBatchCost job run",
@ -3371,6 +3378,12 @@ class PrometheusLogger(CustomLogger):
except Exception as e:
verbose_logger.warning("Error recording check batch cost error metric: %s", e)
def record_check_batch_cost_stale_expired(self):
try:
self.litellm_check_batch_cost_stale_expired_total.inc()
except Exception as e:
verbose_logger.warning("Error recording check batch cost stale expired metric: %s", e)
@staticmethod
def _get_exception_class_name(exception: Exception) -> str:
# Some exception types pin the ``exception_class`` label to a legacy

View file

@ -309,6 +309,7 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_check_batch_cost_jobs_polled",
"litellm_check_batch_cost_jobs_processed_total",
"litellm_check_batch_cost_errors_total",
"litellm_check_batch_cost_stale_expired_total",
"litellm_check_batch_cost_last_run_timestamp",
# MCP tool call metrics
"litellm_mcp_tool_calls_total",
@ -937,6 +938,8 @@ class PrometheusMetricLabels:
litellm_check_batch_cost_errors_total: list[str] = [] # label: error_type (custom)
litellm_check_batch_cost_stale_expired_total: list[str] = []
litellm_check_batch_cost_last_run_timestamp: list[str] = []
# MCP tool call metrics

View file

@ -9,10 +9,14 @@ with deployment credentials, bypassing the managed files access-control hooks.
"""
import base64
import pytest
from collections.abc import Mapping
from datetime import datetime, timezone, tzinfo
from types import SimpleNamespace
from typing import Final, cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.caching.dual_cache import DualCache
@ -20,6 +24,16 @@ from litellm.proxy._types import CallTypes, UserAPIKeyAuth
from litellm.types.utils import LiteLLMBatch
class _FrozenDateTime(datetime):
@classmethod
def now(
cls: type["_FrozenDateTime"],
tz: tzinfo | None = None,
) -> datetime:
fixed_now: Final = datetime(2025, 1, 1, tzinfo=timezone.utc)
return fixed_now.astimezone(tz) if tz is not None else fixed_now.replace(tzinfo=None)
def _make_user_api_key_dict(user_id: str) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-test",
@ -290,12 +304,30 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti
mock_job.created_by = "user-A"
mock_job.id = "job-1"
mock_job.team_id = None
job_created_at: Final = datetime(2025, 1, 1, tzinfo=timezone.utc)
mock_job.created_at = job_created_at
# Mock prisma
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
async def _find_many(
*,
where: dict[str, object],
take: int | None,
order: dict[str, str] | None,
) -> list[MagicMock]:
created_at_filter: Final = where.get("created_at")
if isinstance(created_at_filter, Mapping):
created_at_values: Final = cast(Mapping[str, object], created_at_filter)
cutoff_lt: Final = created_at_values.get("lt")
cutoff_gte: Final = created_at_values.get("gte")
if isinstance(cutoff_lt, datetime) and job_created_at >= cutoff_lt:
return []
if isinstance(cutoff_gte, datetime) and job_created_at < cutoff_gte:
return []
return [mock_job]
mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock(side_effect=_find_many)
mock_prisma.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
@ -348,11 +380,17 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1","response":{"status_code":200,"body":{"id":"cmpl-1","object":"chat.completion","created":1700000000,"model":"gpt-4","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15}}}}\n'
with patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
) as mock_direct_afile_content:
with (
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
) as mock_direct_afile_content,
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
):
await checker.check_batch_cost()
# afile_content should be called directly (not through managed_files_obj)

View file

@ -8,15 +8,22 @@ ARN unified_object_id) batches with no managed unified id.
import asyncio
import json
from collections.abc import Mapping, Sequence
from contextlib import contextmanager
from typing import TYPE_CHECKING
from datetime import datetime, timedelta, timezone, tzinfo
from typing import TYPE_CHECKING, Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.constants import (
MANAGED_OBJECT_STALE_RECONCILE_GRACE_DAYS,
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
)
if TYPE_CHECKING:
from litellm.batches.batch_utils import BatchCostUsageResult
from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost
_IS_B64 = "litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id"
_CLAIM_UNIFIED_BATCH_ID = "dW5pZmllZF9iYXRjaF9pZA=="
@ -137,12 +144,14 @@ class TestCheckBatchCost:
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
stale_call = calls[0]
assert stale_call[1]["data"] == {"status": "stale_expired"}
where = stale_call[1]["where"]
assert where["file_purpose"] == "batch"
assert "stale_expired" in where["status"]["not_in"]
assert "created_at" in where
assert all(call[1]["where"]["file_purpose"] == "batch" for call in calls)
read_calls: Final = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list
assert all(call[1]["where"]["file_purpose"] == "batch" for call in read_calls)
settled_call: Final = calls[0][1]
assert settled_call["where"]["batch_processed"] is True
assert settled_call["where"]["status"]["not_in"]
assert settled_call["where"]["created_at"]["lt"]
assert settled_call["data"] == {"status": "stale_expired"}
@pytest.mark.asyncio
async def test_startup_probe_confirms_batch_processed_support(
@ -224,9 +233,8 @@ class TestCheckBatchCost:
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=1
)
# 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"), []]
side_effect=[[], Exception("column batch_processed does not exist"), []]
)
await check_batch_cost_instance.check_batch_cost()
@ -234,8 +242,9 @@ class TestCheckBatchCost:
calls = (
mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list
)
assert len(calls) == 2
fallback_where = calls[1][1]["where"]
assert len(calls) == 3
assert calls[0][1]["where"]["created_at"]["lt"]
fallback_where: Final = calls[2][1]["where"]
assert "batch_processed" not in fallback_where
assert "stale_expired" in fallback_where["status"]["not_in"]
assert calls[1][1]["take"] == MAX_OBJECTS_PER_POLL_CYCLE
@ -261,9 +270,8 @@ class TestCheckBatchCost:
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
mock_prisma_client.db.litellm_managedobjecttable.find_many.call_count == 2
)
fallback_where = (
mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args[1][
@ -299,7 +307,7 @@ class TestCheckBatchCost:
# Simulate column already known absent (e.g. discovered on a previous cycle)
check_batch_cost_instance._has_batch_processed_column = False
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
# Build a fake batch response whose status triggers the completion branch
@ -404,7 +412,7 @@ class TestCheckBatchCost:
mock_job.id = "job-bedrock-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(side_effect=[[], [mock_job]])
mock_response = MagicMock()
mock_response.status = "completed"
@ -510,7 +518,7 @@ class TestCheckBatchCost:
mock_job.id = "job-poller-rates-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(side_effect=[[], [mock_job]])
mock_response = MagicMock()
mock_response.status = "completed"
@ -610,7 +618,7 @@ class TestCheckBatchCost:
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(side_effect=[[], [mock_job]])
mock_response = MagicMock()
mock_response.status = "completed"
@ -693,7 +701,7 @@ class TestCheckBatchCost:
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
mock_response = MagicMock()
@ -804,8 +812,7 @@ class TestCheckBatchCost:
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = None
mock_job.team_id = None
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(side_effect=[[], [mock_job]])
# A real LiteLLMBatch (not a bare MagicMock): this test runs the real
# litellm_logging.Logging pipeline, which type-checks the result via
@ -930,7 +937,7 @@ class TestCheckBatchCost:
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
mock_response = MagicMock()
@ -1001,7 +1008,7 @@ class TestCheckBatchCost:
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
mock_response = MagicMock()
@ -1087,7 +1094,7 @@ class TestCheckBatchCost:
check_batch_cost_instance._has_batch_processed_column = True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
response = LiteLLMBatch(
@ -1178,7 +1185,7 @@ class TestCheckBatchCost:
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
mock_response = MagicMock()
@ -1263,7 +1270,7 @@ class TestCheckBatchCost:
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
mock_response = MagicMock()
@ -1310,7 +1317,7 @@ class TestCheckBatchCost:
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
mock_response = MagicMock()
@ -1371,7 +1378,7 @@ class TestCheckBatchCost:
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
mock_response = MagicMock()
@ -1484,7 +1491,7 @@ class TestCheckBatchCost:
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(side_effect=[[], [mock_job]])
mock_response = MagicMock()
mock_response.status = "completed"
@ -1597,7 +1604,7 @@ class TestCheckBatchCost:
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
missing_output_file_id = "gs://batch-out/job-1/predictions.jsonl"
@ -1667,7 +1674,7 @@ class TestCheckBatchCost:
check_batch_cost_instance._has_batch_processed_column = True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
side_effect=[[], [mock_job]]
)
raw_output_file_id = "file-batch-output-abc123"
@ -2001,7 +2008,7 @@ class TestUnmanagedVertexRouting:
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
prisma.db.litellm_managedobjecttable.update = AsyncMock()
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[self._job()]
side_effect=[[], [self._job()]]
)
prisma.db.litellm_usertable = MagicMock()
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
@ -2231,7 +2238,7 @@ class TestUnmanagedBedrockRouting:
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
prisma.db.litellm_managedobjecttable.update = AsyncMock()
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[self._job()]
side_effect=[[], [self._job()]]
)
prisma.db.litellm_usertable = MagicMock()
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
@ -2830,7 +2837,7 @@ class TestPollPageStarvation:
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_managedobjecttable.find_many = AsyncMock(side_effect=[[], jobs])
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
return prisma
@ -2958,23 +2965,6 @@ class TestPollPageStarvation:
"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
@ -3061,7 +3051,7 @@ class _FakeManagedObjectRow:
self.team_id = None
self.api_key = None
self.request_tags = None
self.created_at = 1700000000
self.created_at = datetime.max.replace(tzinfo=timezone.utc)
self.file_object = json.dumps(
{"id": "batch-456", "status": "in_progress", "input_file_id": "file-input-1",
"output_file_id": _CLAIM_OUTPUT_FILE_ID}
@ -3073,46 +3063,73 @@ class _FakeManagedObjectTable:
It honours the batch_processed and status filters, so the poller's compare-and-swap
and the managed-files deletion guard both read the same state a shared Postgres row
would give them. Staleness sweeps (the only queries scoped by created_at) never match.
would give them.
"""
def __init__(self, row: _FakeManagedObjectRow, journal: list):
self.row = row
def __init__(self, rows: _FakeManagedObjectRow | list[_FakeManagedObjectRow], journal: list[str]):
self.rows: list[_FakeManagedObjectRow] = rows if isinstance(rows, list) else [rows]
self.row = self.rows[0]
self.journal = journal
self.update_many = AsyncMock(side_effect=self._update_many)
self.update = AsyncMock(side_effect=self._update)
self.find_many = AsyncMock(side_effect=self._find_many)
self.find_first = AsyncMock(return_value=None)
def _matches(self, where: dict) -> bool:
@staticmethod
def _matches(row: _FakeManagedObjectRow, where: dict[str, object]) -> bool:
for key, value in where.items():
if key == "created_at":
return False
if not isinstance(value, Mapping):
return False
cutoff_lt: Final = value.get("lt")
cutoff_gte: Final = value.get("gte")
if isinstance(cutoff_lt, datetime) and row.created_at >= cutoff_lt:
return False
if isinstance(cutoff_gte, datetime) and row.created_at < cutoff_gte:
return False
if key == "status":
if self.row.status in value.get("not_in", []):
if not isinstance(value, Mapping):
return False
if "in" in value and self.row.status not in value["in"]:
excluded: Final = value.get("not_in", ())
included: Final = value.get("in", ())
if isinstance(excluded, Sequence) and row.status in excluded:
return False
elif getattr(self.row, key) != value:
if isinstance(included, Sequence) and "in" in value and row.status not in included:
return False
elif key != "created_at" and getattr(row, key) != value:
return False
return True
async def _update_many(self, *, where: dict, data: dict) -> int:
if not self._matches(where):
return 0
if "batch_processed" in where:
self.journal.append("claim" if data.get("batch_processed") else "release")
for key, value in data.items():
setattr(self.row, key, value)
return 1
async def _update_many(self, *, where: dict[str, object], data: dict[str, object]) -> int:
matched_rows: Final = [row for row in self.rows if self._matches(row, where)]
for row in matched_rows:
if "batch_processed" in where and "batch_processed" in data:
self.journal.append("claim" if data["batch_processed"] else "release")
for key, value in data.items():
setattr(row, key, value)
return len(matched_rows)
async def _update(self, *, where: dict, data: dict) -> None:
async def _update(self, *, where: dict[str, object], data: dict[str, object]) -> None:
self.journal.append("finalize")
for key, value in data.items():
setattr(self.row, key, value)
for row in self.rows:
if self._matches(row, where):
for key, value in data.items():
setattr(row, key, value)
async def _find_many(self, *, where: dict, take=None, order=None) -> list:
return [self.row] if self._matches(where) else []
async def _find_many(
self,
*,
where: dict[str, object],
take: int | None = None,
order: dict[str, str] | None = None,
) -> list[_FakeManagedObjectRow]:
matched_rows: Final = [row for row in self.rows if self._matches(row, where)]
ordered_rows: Final = sorted(
matched_rows,
key=lambda row: row.created_at,
reverse=order == {"created_at": "desc"},
)
return ordered_rows[:take] if take is not None else ordered_rows
class TestMultiPodBatchCostClaim:
@ -3396,3 +3413,535 @@ class TestMultiPodBatchCostClaim:
assert self._claim_calls(prisma) == []
assert journal == ["fetch", "bill", "finalize"]
logging_obj.async_success_handler.assert_awaited_once()
class _FrozenDateTime(datetime):
@classmethod
def now(
cls: type["_FrozenDateTime"],
tz: tzinfo | None = None,
) -> datetime:
fixed_now: Final = datetime(2025, 1, 1, tzinfo=timezone.utc)
return fixed_now.astimezone(tz) if tz is not None else fixed_now.replace(tzinfo=None)
class TestStaleManagedObjectReconciliation:
@staticmethod
def _old_row(
status: str,
batch_processed: bool = False,
job_id: str = "job-claim-1",
created_at: datetime | None = None,
) -> _FakeManagedObjectRow:
row: Final = _FakeManagedObjectRow()
row.id = job_id
row.status = status
row.batch_processed = batch_processed
row.created_at = created_at if created_at is not None else datetime.min.replace(tzinfo=timezone.utc)
return row
@staticmethod
def _created_at_past_reconcile_grace() -> datetime:
now: Final = _FrozenDateTime.now(timezone.utc)
return now - timedelta(
days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS + MANAGED_OBJECT_STALE_RECONCILE_GRACE_DAYS,
hours=1,
)
@staticmethod
def _created_at_within_reconcile_grace() -> datetime:
return _FrozenDateTime.now(timezone.utc) - timedelta(
days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS, hours=4
)
@staticmethod
def _instance(prisma: MagicMock, llm_router: MagicMock) -> "CheckBatchCost":
return TestMultiPodBatchCostClaim._instance(prisma, llm_router)
@staticmethod
def _prisma(
rows: _FakeManagedObjectRow | list[_FakeManagedObjectRow], journal: list[str]
) -> MagicMock:
prisma: Final = TestMultiPodBatchCostClaim._prisma(rows, journal)
return prisma
@staticmethod
def _router(
status: str = "completed", output_file_id: str | None = _CLAIM_OUTPUT_FILE_ID
) -> MagicMock:
router: Final = TestMultiPodBatchCostClaim._router()
router.aretrieve_batch.return_value.status = status
router.aretrieve_batch.return_value.output_file_id = output_file_id
return router
@pytest.mark.asyncio
async def test_stale_completed_batch_is_costed_before_expiration(self):
row: Final = self._old_row("validating")
settled_row: Final = self._old_row(
"in_progress", batch_processed=True, job_id="job-settled"
)
journal: Final[list[str]] = []
prisma: Final = self._prisma([row, settled_row], journal)
prom_logger: Final = MagicMock()
with (
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, self._router()).check_batch_cost()
assert row.status == "complete"
assert row.batch_processed is True
assert settled_row.status == "stale_expired"
assert logging_obj.async_success_handler.await_count == 1
assert logging_obj.async_success_handler.await_args.kwargs["batch_cost"] == 0.01
assert prom_logger.record_check_batch_cost_stale_expired.call_count == 0
prom_logger.record_check_batch_cost_run.assert_called_once_with(
jobs_polled=1,
processed_models=[("gpt-4", "openai")],
)
@pytest.mark.asyncio
async def test_stale_nonterminal_batch_is_expired_and_counted(self):
row: Final = self._old_row("validating")
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, self._router(status="in_progress")).check_batch_cost()
assert row.status == "stale_expired"
logging_obj.async_success_handler.assert_not_awaited()
prom_logger.record_check_batch_cost_stale_expired.assert_called_once_with()
@pytest.mark.asyncio
async def test_stale_nonterminal_expires_without_grace(self) -> None:
row: Final = self._old_row(
"validating",
created_at=self._created_at_within_reconcile_grace(),
)
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, self._router(status="in_progress")).check_batch_cost()
assert row.status == "stale_expired"
assert row.batch_processed is False
logging_obj.async_success_handler.assert_not_awaited()
prom_logger.record_check_batch_cost_stale_expired.assert_called_once_with()
@pytest.mark.parametrize(
"failure", ["provider_retrieval", "cost_tracking", "lagging_output_file"]
)
@pytest.mark.asyncio
async def test_stale_retryable_failure_within_grace_stays_eligible(self, failure: str) -> None:
expected_status: Final = "validating" if failure == "provider_retrieval" else "completed"
row: Final = self._old_row(
expected_status,
created_at=self._created_at_within_reconcile_grace(),
)
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
output_file_id: Final = (
_CLAIM_OUTPUT_FILE_ID if failure == "cost_tracking" else None
)
router: Final = self._router(status="completed", output_file_id=output_file_id)
router.aretrieve_batch.return_value.request_counts = None
async def _raise_during_fetch() -> None:
raise RuntimeError("output file read failed")
during_fetch: Final = _raise_during_fetch if failure == "cost_tracking" else None
if failure == "provider_retrieval":
router.aretrieve_batch = AsyncMock(side_effect=RuntimeError("provider unavailable"))
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(
journal, during_fetch=during_fetch
) as logging_obj,
):
await self._instance(prisma, router).check_batch_cost()
assert row.status == expected_status
assert row.batch_processed is False
logging_obj.async_success_handler.assert_not_awaited()
prom_logger.record_check_batch_cost_stale_expired.assert_not_called()
@pytest.mark.parametrize("condition", ["lagging_output", "tracking_error"])
@pytest.mark.asyncio
async def test_stale_cleanup_gives_up_on_never_costed_completed_rows(self, condition: str):
row: Final = self._old_row(
"completed",
created_at=self._created_at_past_reconcile_grace(),
)
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
async def _raise_during_fetch() -> None:
raise RuntimeError("output file read failed")
during_fetch: Final = _raise_during_fetch if condition == "tracking_error" else None
output_file_id: Final = _CLAIM_OUTPUT_FILE_ID if condition == "tracking_error" else None
router: Final = self._router(status="completed", output_file_id=output_file_id)
router.aretrieve_batch.return_value.request_counts = None
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(
journal, during_fetch=during_fetch
) as logging_obj,
):
await self._instance(prisma, router).check_batch_cost()
assert row.status == "completed"
assert row.batch_processed is True
logging_obj.async_success_handler.assert_not_awaited()
prom_logger.record_check_batch_cost_stale_expired.assert_called_once_with()
@pytest.mark.asyncio
async def test_stale_batch_already_processed_is_not_retrieved(self):
row: Final = self._old_row("in_progress", batch_processed=True)
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
router: Final = self._router(status="in_progress")
with patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
):
await self._instance(prisma, router).check_batch_cost()
router.aretrieve_batch.assert_not_awaited()
assert row.status == "stale_expired"
prom_logger.record_check_batch_cost_stale_expired.assert_not_called()
@pytest.mark.asyncio
async def test_old_schema_stale_completed_batch_is_costed(self):
row: Final = self._old_row("validating")
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
instance: Final = self._instance(prisma, self._router())
instance._has_batch_processed_column = False
with (
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await instance.check_batch_cost()
assert row.status == "complete"
assert logging_obj.async_success_handler.await_count == 1
assert logging_obj.async_success_handler.await_args.kwargs["batch_cost"] == 0.01
prom_logger.record_check_batch_cost_stale_expired.assert_not_called()
@pytest.mark.asyncio
async def test_old_schema_stale_nonterminal_batch_is_expired_and_counted(self):
row: Final = self._old_row("validating")
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
router: Final = self._router(status="in_progress")
instance: Final = self._instance(prisma, router)
instance._has_batch_processed_column = False
with (
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await instance.check_batch_cost()
assert row.status == "stale_expired"
router.aretrieve_batch.assert_awaited_once()
logging_obj.async_success_handler.assert_not_awaited()
prom_logger.record_check_batch_cost_stale_expired.assert_called_once_with()
@pytest.mark.asyncio
async def test_stale_retrieve_error_is_expired_and_counted(self):
row: Final = self._old_row(
"validating",
created_at=self._created_at_past_reconcile_grace(),
)
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
router: Final = self._router()
router.aretrieve_batch = AsyncMock(side_effect=RuntimeError("provider unavailable"))
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, router).check_batch_cost()
assert row.status == "stale_expired"
logging_obj.async_success_handler.assert_not_awaited()
prom_logger.record_check_batch_cost_stale_expired.assert_called_once_with()
@pytest.mark.asyncio
async def test_jobs_polled_includes_stale_sweep(self) -> None:
stale_row: Final = self._old_row("validating")
fresh_row: Final = self._old_row(
"validating",
job_id="job-fresh",
created_at=_FrozenDateTime.now(timezone.utc),
)
journal: Final[list[str]] = []
prisma: Final = self._prisma([stale_row, fresh_row], journal)
prom_logger: Final = MagicMock()
router: Final = self._router()
fresh_response: Final = router.aretrieve_batch.return_value
stale_response: Final = MagicMock()
stale_response.status = "in_progress"
stale_response.output_file_id = None
router.aretrieve_batch = AsyncMock(side_effect=[stale_response, fresh_response])
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, router).check_batch_cost()
assert stale_row.status == "stale_expired"
assert fresh_row.status == "complete"
assert logging_obj.async_success_handler.await_count == 1
prom_logger.record_check_batch_cost_run.assert_called_once_with(
jobs_polled=2,
processed_models=[("gpt-4", "openai")],
)
@pytest.mark.asyncio
async def test_retryable_stale_row_not_polled_twice_in_one_cycle(self) -> None:
stale_row: Final = self._old_row(
"validating",
job_id="job-stale",
created_at=self._created_at_within_reconcile_grace(),
)
fresh_row: Final = self._old_row(
"validating",
job_id="job-fresh",
created_at=_FrozenDateTime.now(timezone.utc),
)
journal: Final[list[str]] = []
prisma: Final = self._prisma([stale_row, fresh_row], journal)
prom_logger: Final = MagicMock()
router: Final = self._router()
fresh_response: Final = router.aretrieve_batch.return_value
router.aretrieve_batch = AsyncMock(
side_effect=[RuntimeError("provider unavailable"), fresh_response]
)
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, router).check_batch_cost()
assert router.aretrieve_batch.await_count == 2
assert stale_row.status == "validating"
assert fresh_row.status == "complete"
logging_obj.async_success_handler.assert_awaited_once()
@pytest.mark.asyncio
async def test_stale_sweep_keeps_results_when_expiration_update_fails(self) -> None:
first_row: Final = self._old_row(
"validating",
job_id="job-first",
created_at=self._created_at_past_reconcile_grace(),
)
second_row: Final = self._old_row(
"validating",
job_id="job-second",
created_at=self._created_at_past_reconcile_grace(),
)
journal: Final[list[str]] = []
prisma: Final = self._prisma([first_row, second_row], journal)
table: Final = prisma.db.litellm_managedobjecttable
table.update_many = AsyncMock(
side_effect=[1, 1, RuntimeError("expiration update failed")]
)
prom_logger: Final = MagicMock()
router: Final = self._router()
billed_response: Final = router.aretrieve_batch.return_value
unreconciled_response: Final = MagicMock()
unreconciled_response.status = "in_progress"
unreconciled_response.output_file_id = None
router.aretrieve_batch = AsyncMock(
side_effect=[billed_response, unreconciled_response]
)
with (
patch(
"litellm_enterprise.proxy.common_utils.check_batch_cost.datetime",
_FrozenDateTime,
),
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, router).check_batch_cost()
assert first_row.status == "complete"
assert first_row.batch_processed is True
assert second_row.status == "validating"
assert logging_obj.async_success_handler.await_count == 1
prom_logger.record_check_batch_cost_run.assert_called_once_with(
jobs_polled=2,
processed_models=[("gpt-4", "openai")],
)
@pytest.mark.asyncio
async def test_concurrent_completion_is_not_overwritten_or_counted(self):
row: Final = self._old_row("validating")
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
response: Final = MagicMock()
response.status = "in_progress"
router: Final = MagicMock()
async def _complete_during_retrieve(**kwargs: object) -> MagicMock:
row.status = "completed"
return response
router.aretrieve_batch = AsyncMock(side_effect=_complete_during_retrieve)
instance: Final = self._instance(prisma, router)
instance._has_batch_processed_column = False
with (
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal),
):
await instance.check_batch_cost()
assert row.status == "completed"
prom_logger.record_check_batch_cost_stale_expired.assert_not_called()
@pytest.mark.asyncio
async def test_stale_row_claimed_by_other_poller_is_not_expired(self) -> None:
row: Final = self._old_row("validating")
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
router: Final = self._router(status="in_progress")
response: Final = router.aretrieve_batch.return_value
async def _other_poller_claims_batch(**kwargs: object) -> MagicMock:
row.batch_processed = True
return response
router.aretrieve_batch = AsyncMock(side_effect=_other_poller_claims_batch)
with (
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal),
):
await self._instance(prisma, router).check_batch_cost()
router.aretrieve_batch.assert_awaited_once()
assert journal == []
assert row.status == "validating"
assert row.batch_processed is True
prom_logger.record_check_batch_cost_stale_expired.assert_not_called()
@pytest.mark.asyncio
async def test_stale_completed_row_with_missing_deployment_is_retired_and_counted(self) -> None:
row: Final = self._old_row("validating")
journal: Final[list[str]] = []
prisma: Final = self._prisma(row, journal)
prom_logger: Final = MagicMock()
router: Final = self._router(status="completed")
router.get_deployment.return_value = None
with (
patch(
"litellm.integrations.prometheus.PrometheusLogger.get_instance",
return_value=prom_logger,
),
TestMultiPodBatchCostClaim._billing_patches(journal) as logging_obj,
):
await self._instance(prisma, router).check_batch_cost()
assert journal == ["fetch"]
assert row.status == "stale_expired"
assert row.batch_processed is False
logging_obj.async_success_handler.assert_not_awaited()
prom_logger.record_check_batch_cost_stale_expired.assert_called_once_with()