mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge a1ffd6f41a into f285229b51
This commit is contained in:
commit
9defe463fa
7 changed files with 1020 additions and 220 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue