diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json index 5bd7ed97a55..1f7dc0992c5 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json @@ -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", diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 35fc3510414..861bcba15ad 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -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, ) diff --git a/litellm/constants.py b/litellm/constants.py index 530d678457d..ff6ad4a742e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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 diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index c7bf291a887..aad62d26409 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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 diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 8f4ad26a4fa..5c1f401d240 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -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 diff --git a/tests/unit/enterprise/proxy/test_managed_files_access_check.py b/tests/unit/enterprise/proxy/test_managed_files_access_check.py index ad46798b788..8049b76b974 100644 --- a/tests/unit/enterprise/proxy/test_managed_files_access_check.py +++ b/tests/unit/enterprise/proxy/test_managed_files_access_check.py @@ -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) diff --git a/tests/unit/proxy/common_utils/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py index 4e7effad3bf..6b98b90707c 100644 --- a/tests/unit/proxy/common_utils/test_check_batch_cost.py +++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py @@ -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()