From 488594e03ff5a914dfc5d47950db7e1dfd36eb9f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 9 Oct 2026 00:15:38 -0700 Subject: [PATCH] refactor(proxy): inject one UTC clock read per operation into gateway tracking, PTU rollup and Mavvrik export (#45520) * fix(proxy): read the UTC clock once per operation in gateway tracking, PTU rollup and Mavvrik export * refactor(proxy): default the injected clocks to get_utc_datetime instead of three private copies --------- Co-authored-by: yuneng --- .../mavvrik_focus/mavvrik_focus_logger.py | 7 +- litellm/proxy/db/gateway_request_tracking.py | 14 ++- .../spend_tracking/ptu_flat_cost_rollup.py | 26 +++-- .../focus/test_mavvrik_destination.py | 93 +++++++++++++----- .../proxy/db/test_gateway_request_tracking.py | 97 ++++++++++++------- ...est_billable_request_metrics_middleware.py | 16 ++- .../test_ptu_flat_cost_rollup.py | 51 +++++++--- 7 files changed, 210 insertions(+), 94 deletions(-) diff --git a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py index 6e73387326a..f9e54c30e89 100644 --- a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py +++ b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py @@ -20,6 +20,7 @@ overwrite each other within the same day, producing incomplete data. from __future__ import annotations import os +from collections.abc import Callable from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, Final @@ -28,6 +29,7 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import MAVVRIK_FOCUS_EXPORT_JOB_NAME from litellm.integrations.focus.destinations.base import FocusTimeWindow from litellm.integrations.focus.focus_logger import FocusLogger +from litellm.utils import get_utc_datetime if TYPE_CHECKING: from apscheduler.schedulers.asyncio import AsyncIOScheduler @@ -83,7 +85,8 @@ def _is_empty_metrics_marker(marker: object | None) -> bool: class MavvrikFocusLogger(FocusLogger): """FOCUS-based export logger that routes to the Mavvrik destination.""" - def __init__(self, **kwargs: Any) -> None: + def __init__(self, *, clock: Callable[[], datetime] = get_utc_datetime, **kwargs: Any) -> None: + self._clock: Final = clock frequency: Final = os.getenv("MAVVRIK_FOCUS_FREQUENCY", "daily").lower() if frequency != "daily": raise ValueError( @@ -174,7 +177,7 @@ class MavvrikFocusLogger(FocusLogger): # metricsMarker may be a Unix timestamp (int/float) or an ISO date string. marker: Final = await destination.get_metrics_marker() - now: Final = datetime.now(timezone.utc) + now: Final = self._clock() yesterday: Final = now.replace(hour=0, minute=0, second=0, microsecond=0) - timedelta(days=1) last_ingested: Final = _parse_metrics_marker(marker) diff --git a/litellm/proxy/db/gateway_request_tracking.py b/litellm/proxy/db/gateway_request_tracking.py index e3483ca3215..62f2b45a7c0 100644 --- a/litellm/proxy/db/gateway_request_tracking.py +++ b/litellm/proxy/db/gateway_request_tracking.py @@ -20,8 +20,8 @@ the deployment as a whole costs the primary one statement per interval. """ import json -from collections.abc import AsyncIterator, Iterable -from datetime import datetime, timezone +from collections.abc import AsyncIterator, Callable, Iterable +from datetime import datetime from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Final, TypeAlias @@ -40,6 +40,7 @@ from litellm.types.proxy.gateway_requests import ( GatewayRequestKey, GatewayRequestSnapshot, ) +from litellm.utils import get_utc_datetime _GATEWAY_REQUEST_QUEUE_TARGET: Final = "gateway_request_queue" @@ -58,18 +59,15 @@ _BUFFERED_ENTRIES: Final = TypeAdapter(tuple[str | bytes, ...]) _NO_COUNTS: Final[GatewayRequestSnapshot] = MappingProxyType({}) -def _utc_date() -> str: - return datetime.now(timezone.utc).strftime("%Y-%m-%d") - - class GatewayRequestAccumulator: """Sink for the request-metrics middleware. ``record`` is sync and never awaits.""" - def __init__(self) -> None: + def __init__(self, *, clock: Callable[[], datetime] = get_utc_datetime) -> None: + self._clock: Final = clock self._counts: dict[GatewayRequestKey, GatewayRequestCounts] = {} # mutable-ok: bounded fold, drained per flush def record(self, *, category: BillableCategory, route: str, status_code: int) -> None: - key: Final = GatewayRequestKey(date=_utc_date(), category=category.value, route=route) + key: Final = GatewayRequestKey(date=self._clock().strftime("%Y-%m-%d"), category=category.value, route=route) self._counts[key] = self._counts.get(key, _EMPTY).plus(succeeded=200 <= status_code < 300) def drain(self) -> GatewayRequestSnapshot: diff --git a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py index b2bf6f46b3a..e4370fe304d 100644 --- a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py @@ -36,6 +36,7 @@ from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attributio from litellm.repositories.model_repository import ModelRepository from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import PrismaTableRepository +from litellm.utils import get_utc_datetime if TYPE_CHECKING: from prisma import models as prisma_models @@ -607,6 +608,8 @@ async def run_scheduled_ptu_rollup( target_date: date | None = None, alert: Callable[[str], Awaitable[None]] | None = None, router: object | None = None, + *, + clock: Callable[[], datetime] = get_utc_datetime, ) -> RollupResult | None: """Run the daily rollup under a cross-pod lock so only one proxy reconciles a day. @@ -630,8 +633,12 @@ async def run_scheduled_ptu_rollup( if not is_ptu_cost_attribution_enabled(): return None + today: Final = clock().date() + if pod_lock_manager is None or pod_lock_manager.redis_cache is None: - return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router) + return await _run_and_alert( + prisma_client, target_date=target_date, today=today, alert=alert, may_prune=False, router=router + ) if not await pod_lock_manager.acquire_lock(cronjob_id=PTU_ROLLUP_JOB_ID, ttl=PTU_ROLLUP_LOCK_TTL_SECONDS): if await _lock_is_held(pod_lock_manager): @@ -645,10 +652,14 @@ async def run_scheduled_ptu_rollup( "PTU rollup: could not take the rollup lock and no other pod holds it, " "running unguarded rather than skipping the day" ) - return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router) + return await _run_and_alert( + prisma_client, target_date=target_date, today=today, alert=alert, may_prune=False, router=router + ) try: - return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True, router=router) + return await _run_and_alert( + prisma_client, target_date=target_date, today=today, alert=alert, may_prune=True, router=router + ) finally: await pod_lock_manager.release_lock(cronjob_id=PTU_ROLLUP_JOB_ID) @@ -672,6 +683,7 @@ async def _run_and_alert( prisma_client: "PrismaClient", *, target_date: date | None, + today: date, alert: "Callable[[str], Awaitable[None]] | None", may_prune: bool = True, router: object | None = None, @@ -687,8 +699,9 @@ async def _run_and_alert( explicit date means reconcile exactly that day, so it stays a single-day operation. Its failure is contained: the day's own result is returned either way. """ + rollup_date: Final = target_date or today - timedelta(days=1) result: Final = await run_ptu_flat_cost_rollup( - prisma_client, target_date=target_date, may_prune=may_prune, router=router + prisma_client, target_date=rollup_date, may_prune=may_prune, router=router ) if result.rows_failed: await _deliver_alert( @@ -706,13 +719,14 @@ async def _run_and_alert( "by the provider with nothing attributing it here. Extend the window, or retire the deployment.", ) if target_date is None: - await _backfill_and_alert(prisma_client, alert=alert, router=router) + await _backfill_and_alert(prisma_client, today=today, alert=alert, router=router) return result async def _backfill_and_alert( prisma_client: "PrismaClient", *, + today: date, alert: "Callable[[str], Awaitable[None]] | None", router: object | None = None, ) -> None: @@ -722,7 +736,7 @@ async def _backfill_and_alert( caller whatever the catch-up pass does. """ try: - backfill: Final = await run_ptu_flat_cost_backfill(prisma_client, router=router) + backfill: Final = await run_ptu_flat_cost_backfill(prisma_client, today=today, router=router) except Exception as exc: # noqa: BLE001 # the catch-up pass must not fail the day's rollup verbose_proxy_logger.error("PTU backfill: catch-up pass failed, the day's rollup still stands: %s", exc) return diff --git a/tests/unit/integrations/focus/test_mavvrik_destination.py b/tests/unit/integrations/focus/test_mavvrik_destination.py index 1be72b4409b..13f462c4ef3 100644 --- a/tests/unit/integrations/focus/test_mavvrik_destination.py +++ b/tests/unit/integrations/focus/test_mavvrik_destination.py @@ -3,6 +3,7 @@ from __future__ import annotations from datetime import datetime, timedelta, timezone +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -14,6 +15,7 @@ from litellm.integrations.focus.destinations.mavvrik_destination import ( ) VALID_ENDPOINT = "https://api.mavvrik.ai/tenant123" +_FOCUS_NOW: Final = datetime(2026, 2, 1, 0, 0, tzinfo=timezone.utc) def _make_window() -> FocusTimeWindow: @@ -504,7 +506,6 @@ async def test_export_window_passes_max_rows_as_limit(monkeypatch): async def test_run_scheduled_export_catches_up_missed_dates(): """If metricsMarker is 2 days behind, _run_scheduled_export exports missed dates first.""" import polars as pl - from datetime import datetime, timedelta, timezone from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( MavvrikFocusLogger, ) @@ -512,15 +513,15 @@ async def test_run_scheduled_export_catches_up_missed_dates(): FocusMavvrikDestination, ) - logger = MavvrikFocusLogger() + logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW) # metricsMarker = 3 days ago → 2 missed dates (day-2 and day-1) + today's run - now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) - yesterday = now - timedelta(days=1) - two_days_ago = now - timedelta(days=2) - three_days_ago = now - timedelta(days=3) + now: Final = _FOCUS_NOW + yesterday: Final = now - timedelta(days=1) + two_days_ago: Final = now - timedelta(days=2) + three_days_ago: Final = now - timedelta(days=3) - marker_ts = int(three_days_ago.timestamp()) + marker_ts: Final = int(three_days_ago.timestamp()) # Mock destination dest_mock = MagicMock(spec=FocusMavvrikDestination) @@ -551,7 +552,6 @@ async def test_run_scheduled_export_catches_up_missed_dates(): async def test_run_scheduled_export_no_catchup_when_marker_is_current(): """If metricsMarker = yesterday, no catch-up needed — just export yesterday.""" import polars as pl - from datetime import datetime, timedelta, timezone from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( MavvrikFocusLogger, ) @@ -559,11 +559,11 @@ async def test_run_scheduled_export_no_catchup_when_marker_is_current(): FocusMavvrikDestination, ) - logger = MavvrikFocusLogger() + logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW) - now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) - yesterday = now - timedelta(days=1) - marker_ts = int(yesterday.timestamp()) + now: Final = _FOCUS_NOW + yesterday: Final = now - timedelta(days=1) + marker_ts: Final = int(yesterday.timestamp()) dest_mock = MagicMock(spec=FocusMavvrikDestination) dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts) @@ -592,9 +592,9 @@ async def test_run_scheduled_export_skips_catchup_when_marker_is_unparseable(): FocusMavvrikDestination, ) - logger = MavvrikFocusLogger() - now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) - yesterday = now - timedelta(days=1) + logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW) + now: Final = _FOCUS_NOW + yesterday: Final = now - timedelta(days=1) dest_mock = MagicMock(spec=FocusMavvrikDestination) dest_mock.get_metrics_marker = AsyncMock(return_value="not-a-date") @@ -706,7 +706,6 @@ def test_parse_metrics_marker_returns_none_for_garbage(): async def test_catchup_capped_at_max_catchup_days(): """Catch-up must not go further back than _MAX_CATCHUP_DAYS.""" import polars as pl - from datetime import datetime, timedelta, timezone from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( MavvrikFocusLogger, ) @@ -714,14 +713,14 @@ async def test_catchup_capped_at_max_catchup_days(): FocusMavvrikDestination, ) - logger = MavvrikFocusLogger() - max_days = MavvrikFocusLogger._MAX_CATCHUP_DAYS + logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW) + max_days: Final = MavvrikFocusLogger._MAX_CATCHUP_DAYS - now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) - yesterday = now - timedelta(days=1) + now: Final = _FOCUS_NOW + yesterday: Final = now - timedelta(days=1) # Marker is 30 days ago — well beyond the cap - thirty_days_ago = now - timedelta(days=30) - marker_ts = int(thirty_days_ago.timestamp()) + thirty_days_ago: Final = now - timedelta(days=30) + marker_ts: Final = int(thirty_days_ago.timestamp()) dest_mock = MagicMock(spec=FocusMavvrikDestination) dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts) @@ -740,11 +739,57 @@ async def test_catchup_capped_at_max_catchup_days(): assert db_mock.get_usage_data.call_count <= max_days # First catch-up date must not be earlier than (yesterday - max_days + 1) - earliest_allowed = yesterday - timedelta(days=max_days - 1) - first_call_start = db_mock.get_usage_data.call_args_list[0].kwargs["start_time_utc"] + earliest_allowed: Final = yesterday - timedelta(days=max_days - 1) + first_call_start: Final = db_mock.get_usage_data.call_args_list[0].kwargs["start_time_utc"] assert first_call_start.date() >= earliest_allowed.date() +@pytest.mark.asyncio +async def test_run_scheduled_export_uses_one_clock_read_across_midnight(): + import polars as pl + from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import ( + MavvrikFocusLogger, + ) + from litellm.integrations.focus.destinations.mavvrik_destination import ( + FocusMavvrikDestination, + ) + + logger: Final = MavvrikFocusLogger( + clock=iter( + ( + datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc), + datetime(2026, 2, 1, 0, 0, 0, 1, tzinfo=timezone.utc), + ) + ).__next__ + ) + marker_ts: Final = int(datetime(2026, 1, 28, 0, 0, tzinfo=timezone.utc).timestamp()) + dest_mock: Final = MagicMock(spec=FocusMavvrikDestination) + dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts) + db_mock: Final = MagicMock() + db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) + engine_mock: Final = MagicMock() + engine_mock.database = db_mock + engine_mock.destination = dest_mock + logger._engine = engine_mock + + await logger._run_scheduled_export() + + windows: Final = tuple( + (call.kwargs["start_time_utc"], call.kwargs["end_time_utc"]) + for call in db_mock.get_usage_data.call_args_list + ) + assert windows == ( + ( + datetime(2026, 1, 29, 0, 0, tzinfo=timezone.utc), + datetime(2026, 1, 30, 0, 0, tzinfo=timezone.utc), + ), + ( + datetime(2026, 1, 30, 0, 0, tzinfo=timezone.utc), + datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc), + ), + ) + + @pytest.mark.asyncio async def test_register_resets_on_410(): """_registered flag must be False after a 410 so next run re-registers.""" diff --git a/tests/unit/proxy/db/test_gateway_request_tracking.py b/tests/unit/proxy/db/test_gateway_request_tracking.py index a6689b38039..5566a5cb9bc 100644 --- a/tests/unit/proxy/db/test_gateway_request_tracking.py +++ b/tests/unit/proxy/db/test_gateway_request_tracking.py @@ -6,6 +6,7 @@ LiteLLM_DailyGatewayRequests. import asyncio from collections.abc import Awaitable, Callable from datetime import datetime, timezone +from typing import Final import pytest @@ -22,8 +23,12 @@ from litellm.types.proxy.gateway_requests import GatewayRequestCounts, GatewayRe from litellm.proxy.db.log_db_metrics import record_db_io -def _today() -> str: - return datetime.now(timezone.utc).strftime("%Y-%m-%d") +_NOW: Final = datetime(2026, 3, 14, 12, 0, tzinfo=timezone.utc) +_DAY: Final = _NOW.strftime("%Y-%m-%d") + + +def _accumulator() -> GatewayRequestAccumulator: + return GatewayRequestAccumulator(clock=lambda: _NOW) def _record(accumulator: GatewayRequestAccumulator, status_code: int, **overrides) -> None: @@ -38,32 +43,54 @@ def _record(accumulator: GatewayRequestAccumulator, status_code: int, **override def test_folds_repeated_requests_into_one_key(): - acc = GatewayRequestAccumulator() + acc = _accumulator() for _ in range(3): _record(acc, 200) _record(acc, 500) snapshot = acc.drain() assert snapshot == { - GatewayRequestKey(date=_today(), category="llm", route="/chat/completions"): ( + GatewayRequestKey(date=_DAY, category="llm", route="/chat/completions"): ( GatewayRequestCounts(successful_requests=3, failed_requests=1) ) } +def test_records_each_request_under_the_date_of_its_clock_read(): + accumulator: Final = GatewayRequestAccumulator( + clock=iter( + ( + datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc), + datetime(2026, 2, 1, 0, 0, 0, 1, tzinfo=timezone.utc), + ) + ).__next__ + ) + _record(accumulator, 200) + _record(accumulator, 200) + + assert accumulator.drain() == { + GatewayRequestKey(date="2026-01-31", category="llm", route="/chat/completions"): GatewayRequestCounts( + successful_requests=1, failed_requests=0 + ), + GatewayRequestKey(date="2026-02-01", category="llm", route="/chat/completions"): GatewayRequestCounts( + successful_requests=1, failed_requests=0 + ), + } + + @pytest.mark.parametrize( "status_code, expected_successful, expected_failed", [(200, 1, 0), (201, 1, 0), (204, 1, 0), (299, 1, 0), (300, 0, 1), (400, 0, 1), (500, 0, 1)], ) def test_success_boundary_is_2xx(status_code: int, expected_successful: int, expected_failed: int): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, status_code) counts = next(iter(acc.drain().values())) assert (counts.successful_requests, counts.failed_requests) == (expected_successful, expected_failed) def test_distinct_dimensions_do_not_merge(): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200, route="/chat/completions") _record(acc, 200, route="/embeddings") _record(acc, 200, category=BillableCategory.MCP, route="/mcp") @@ -71,14 +98,14 @@ def test_distinct_dimensions_do_not_merge(): def test_drain_empties_the_fold(): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) assert len(acc.drain()) == 1 assert acc.drain() == {} def test_drain_snapshot_is_not_mutated_by_later_records(): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) snapshot = acc.drain() _record(acc, 200) @@ -198,7 +225,7 @@ def test_commit_skips_the_database_entirely_when_nothing_accumulated(): def test_flush_drains_and_commits(): client = FakePrismaClient() - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) asyncio.run(flush_gateway_requests(client, acc)) @@ -217,7 +244,7 @@ class ExplodingClient: def test_flush_swallows_commit_failure_so_the_scheduler_survives(): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) asyncio.run(flush_gateway_requests(ExplodingClient(), acc)) @@ -225,7 +252,7 @@ def test_flush_swallows_commit_failure_so_the_scheduler_survives(): def test_failed_flush_keeps_counts_for_the_next_attempt(): """A dropped flush would silently undercount the SGR source of truth.""" - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) _record(acc, 500) @@ -234,11 +261,11 @@ def test_failed_flush_keeps_counts_for_the_next_attempt(): client = FakePrismaClient() asyncio.run(flush_gateway_requests(client, acc)) - assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)] + assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 1, 1)] def test_restored_counts_merge_with_requests_recorded_meanwhile(): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) asyncio.run(flush_gateway_requests(ExplodingClient(), acc)) @@ -246,7 +273,7 @@ def test_restored_counts_merge_with_requests_recorded_meanwhile(): client = FakePrismaClient() asyncio.run(flush_gateway_requests(client, acc)) - assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)] + assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 2, 0)] class ExplodingDBWithInFlightRequest: @@ -266,14 +293,14 @@ class ExplodingClientWithInFlightRequest: def test_restore_keeps_requests_recorded_while_the_failed_write_was_in_flight(): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) asyncio.run(flush_gateway_requests(ExplodingClientWithInFlightRequest(acc), acc)) client = FakePrismaClient() asyncio.run(flush_gateway_requests(client, acc)) - assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)] + assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 1, 1)] # ── redis buffer ────────────────────────────────────────────────────────────── @@ -340,7 +367,7 @@ def test_non_leader_workers_push_to_redis_and_never_touch_the_database(): redis = FakeRedis() client = FakePrismaClient() for _ in range(3): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) buffer, _ = _buffer(redis, leader=False) asyncio.run(flush_gateway_requests(client, acc, buffer)) @@ -354,21 +381,21 @@ def test_leader_folds_every_workers_snapshot_into_one_statement(): redis = FakeRedis() client = FakePrismaClient() for _ in range(50): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) _record(acc, 500, route="/responses") buffer, _ = _buffer(redis, leader=False) asyncio.run(flush_gateway_requests(client, acc, buffer)) - leader_acc = GatewayRequestAccumulator() + leader_acc = _accumulator() _record(leader_acc, 200) leader, lock = _buffer(redis, leader=True) asyncio.run(flush_gateway_requests(client, leader_acc, leader)) assert len(client.db.statements) == 1 assert _rows_written(client) == [ - (_today(), "llm", "/chat/completions", 51, 0), - (_today(), "llm", "/responses", 0, 50), + (_DAY, "llm", "/chat/completions", 51, 0), + (_DAY, "llm", "/responses", 0, 50), ] assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == [] assert lock.held == [GATEWAY_REQUESTS_JOB_NAME] @@ -387,7 +414,7 @@ def test_leader_keeps_the_lease_so_staggered_pods_cost_one_statement_per_interva for _interval in range(3): for pod in pods: - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) asyncio.run(flush_gateway_requests(client, acc, pod)) @@ -403,16 +430,16 @@ def test_leader_drains_a_backlog_deeper_than_one_capped_pop(): client = FakePrismaClient() workers = MAX_REDIS_BUFFER_DEQUEUE_COUNT * 2 + 1 for _ in range(workers): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) buffer, _ = _buffer(redis, leader=False) asyncio.run(flush_gateway_requests(client, acc, buffer)) leader, _ = _buffer(redis, leader=True) - asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader)) + asyncio.run(flush_gateway_requests(client, _accumulator(), leader)) assert len(client.db.statements) == 1 - assert _rows_written(client) == [(_today(), "llm", "/chat/completions", workers, 0)] + assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", workers, 0)] assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == [] @@ -421,7 +448,7 @@ def test_leader_with_nothing_buffered_writes_nothing(): client = FakePrismaClient() leader, lock = _buffer(redis, leader=True) - asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader)) + asyncio.run(flush_gateway_requests(client, _accumulator(), leader)) assert client.db.statements == [] assert lock.released == [] @@ -430,7 +457,7 @@ def test_leader_with_nothing_buffered_writes_nothing(): def test_leader_requeues_to_redis_when_the_database_commit_fails(): """Counts popped from Redis are gone from every worker; a failed commit must put them back.""" redis = FakeRedis() - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) _record(acc, 200) leader, lock = _buffer(redis, leader=True) @@ -443,8 +470,8 @@ def test_leader_requeues_to_redis_when_the_database_commit_fails(): client = FakePrismaClient() retry, _ = _buffer(redis, leader=True) - asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), retry)) - assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)] + asyncio.run(flush_gateway_requests(client, _accumulator(), retry)) + assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 2, 0)] class ExplodingRedis(FakeRedis): @@ -467,7 +494,7 @@ class UnwritableRedis(FakeRedis): def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_requeue_fail(): """The pop removed the only copy; if Redis will not take it back the leader itself must carry it.""" redis = FakeRedis() - worker_acc = GatewayRequestAccumulator() + worker_acc = _accumulator() _record(worker_acc, 200) _record(worker_acc, 200) worker, _ = _buffer(redis, leader=False) @@ -475,7 +502,7 @@ def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_req degraded = UnwritableRedis() degraded.lists = redis.lists - leader_acc = GatewayRequestAccumulator() + leader_acc = _accumulator() leader, _ = _buffer(degraded, leader=True) asyncio.run(flush_gateway_requests(ExplodingClient(), leader_acc, leader)) assert degraded.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == [] @@ -483,13 +510,13 @@ def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_req client = FakePrismaClient() retry, _ = _buffer(redis, leader=True) asyncio.run(flush_gateway_requests(client, leader_acc, retry)) - assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)] + assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 2, 0)] def test_leader_whose_redis_read_fails_leaves_the_pushed_rows_for_the_next_flush(): """The scheduler job must not raise, and nothing is popped so nothing needs restoring anywhere.""" redis = UnreadableRedis() - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) client = FakePrismaClient() leader, _ = _buffer(redis, leader=True) @@ -502,7 +529,7 @@ def test_leader_whose_redis_read_fails_leaves_the_pushed_rows_for_the_next_flush def test_failed_redis_push_keeps_counts_locally_for_the_next_flush(): - acc = GatewayRequestAccumulator() + acc = _accumulator() _record(acc, 200) _record(acc, 500) buffer, lock = _buffer(ExplodingRedis(), leader=True) @@ -511,7 +538,7 @@ def test_failed_redis_push_keeps_counts_locally_for_the_next_flush(): assert lock.held == [] assert acc.drain() == { - GatewayRequestKey(date=_today(), category="llm", route="/chat/completions"): ( + GatewayRequestKey(date=_DAY, category="llm", route="/chat/completions"): ( GatewayRequestCounts(successful_requests=1, failed_requests=1) ) } diff --git a/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py b/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py index 333906884a8..f2515905707 100644 --- a/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py +++ b/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py @@ -8,7 +8,8 @@ middleware is a transparent pass-through when no recorder is injected. import asyncio import threading -from typing import List, Optional, Tuple +from datetime import datetime, timezone +from typing import Final, List, Optional, Tuple import pytest from starlette.applications import Starlette @@ -28,6 +29,7 @@ from litellm.proxy.middleware.billable_request_metrics_middleware import ( from litellm.proxy.middleware.in_flight_requests_middleware import ( InFlightRequestsMiddleware, ) +from litellm.types.proxy.gateway_requests import GatewayRequestCounts, GatewayRequestKey class FakeRecorder: @@ -505,14 +507,18 @@ def test_varying_model_ids_fold_into_a_single_persisted_key(): that is. The SGR key is persisted, so it must not carry that dimension: a caller who could vary it could mint an unbounded number of table rows. """ - accumulator = GatewayRequestAccumulator() + frozen_now: Final = datetime(2026, 3, 14, 12, 0, tzinfo=timezone.utc) + accumulator = GatewayRequestAccumulator(clock=lambda: frozen_now) for model_id in ("deploy-1", "deploy-2", "deploy-3"): client = TestClient(_make_sink_app(None, accumulator, status_code=200, model_id=model_id)) client.post("/v1/chat/completions") - snapshot = accumulator.drain() - assert len(snapshot) == 1 - assert next(iter(snapshot.values())).successful_requests == 3 + snapshot: Final = accumulator.drain() + assert snapshot == { + GatewayRequestKey(date=frozen_now.strftime("%Y-%m-%d"), category="llm", route="/chat/completions"): ( + GatewayRequestCounts(successful_requests=3, failed_requests=0) + ) + } @pytest.mark.parametrize("status_code", [400, 429, 500, 503]) diff --git a/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py b/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py index 8f25cffecf5..decebcaa2bc 100644 --- a/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py +++ b/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py @@ -3,6 +3,7 @@ import json import types from datetime import date, datetime, timedelta, timezone +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -23,6 +24,7 @@ from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import ( DAY = date(2026, 7, 30) TODAY = date(2026, 7, 31) +_SCHEDULED_NOW: Final = datetime(2026, 7, 31, 12, 0, tzinfo=timezone.utc) # The endpoints require ptu_effective_from alongside the count and rate, so a fixture that @@ -1449,10 +1451,10 @@ async def test_scheduled_rollup_backfills_after_pricing_the_day(): """The catch-up pass runs after the day's own rollup, so it sees yesterday already priced and does not write it a second time.""" table = _FakeSentinelTable() - yesterday = datetime.now(timezone.utc).date() - timedelta(days=1) + yesterday: Final = _SCHEDULED_NOW.date() - timedelta(days=1) prisma = _prisma_for([_windowed_row(effective_from=_midnight(yesterday - timedelta(days=2)))], table) - await run_scheduled_ptu_rollup(prisma) + await run_scheduled_ptu_rollup(prisma, clock=lambda: _SCHEDULED_NOW) yesterday_key = ("t", yesterday.isoformat(), PTU_SENTINEL_API_KEY, "m1") assert table.upsert_keys.count(yesterday_key) == 1 @@ -1475,13 +1477,13 @@ async def test_scheduled_rollup_with_an_explicit_target_date_does_not_backfill() async def test_scheduled_rollup_holds_one_lock_across_both_phases(): """Backfill running outside the lock would let another pod's prune race its writes.""" table = _FakeSentinelTable() - yesterday = datetime.now(timezone.utc).date() - timedelta(days=1) + yesterday: Final = _SCHEDULED_NOW.date() - timedelta(days=1) prisma = _prisma_for([_windowed_row(effective_from=_midnight(yesterday - timedelta(days=3)))], table) rows_at_release = [] lock = _pod_lock(acquired=True) lock.release_lock = AsyncMock(side_effect=lambda **kwargs: rows_at_release.append(len(table.rows))) - await run_scheduled_ptu_rollup(prisma, pod_lock_manager=lock) + await run_scheduled_ptu_rollup(prisma, pod_lock_manager=lock, clock=lambda: _SCHEDULED_NOW) lock.acquire_lock.assert_awaited_once() assert rows_at_release == [4] @@ -1507,12 +1509,12 @@ async def test_scheduled_rollup_alerts_when_a_backfill_charge_never_landed(): """An unpriced day that stays unpriced is the silent underbill this work exists to remove, so it has to reach an operator too.""" table = _FakeSentinelTable() - yesterday = datetime.now(timezone.utc).date() - timedelta(days=1) + yesterday: Final = _SCHEDULED_NOW.date() - timedelta(days=1) prisma = _prisma_for([_windowed_row(effective_from=_midnight(yesterday - timedelta(days=1)))], table) prisma.db.litellm_dailyteamspend.upsert = AsyncMock(side_effect=RuntimeError("db down")) alert = AsyncMock() - await run_scheduled_ptu_rollup(prisma, alert=alert) + await run_scheduled_ptu_rollup(prisma, alert=alert, clock=lambda: _SCHEDULED_NOW) messages = [call.args[0] for call in alert.await_args_list] assert any("backfill" in message for message in messages) @@ -1535,23 +1537,41 @@ async def test_a_broken_alert_channel_does_not_fail_the_backfill(): @pytest.mark.asyncio async def test_scheduled_rollup_with_no_target_date_closes_a_backdated_window(): - """The production call shape from proxy_server.py, on the real clock: no target_date, + """The production call shape from proxy_server.py, on a frozen clock: no target_date, a window backdated 30 days, and every elapsed in-window day has to end up priced with no operator alert raised. Every other rollup test pins target_date, which is exactly why this regression shipped.""" table = _FakeSentinelTable() - today = datetime.now(timezone.utc).date() - opened_on = today - timedelta(days=30) + today: Final = _SCHEDULED_NOW.date() + opened_on: Final = today - timedelta(days=30) prisma = _prisma_for([_windowed_row(effective_from=_midnight(opened_on))], table) alert = AsyncMock() - await run_scheduled_ptu_rollup(prisma, alert=alert) + await run_scheduled_ptu_rollup(prisma, alert=alert, clock=lambda: _SCHEDULED_NOW) expected = [(opened_on + timedelta(days=offset)).isoformat() for offset in range(30)] assert _priced_dates(table) == expected alert.assert_not_awaited() +@pytest.mark.asyncio +async def test_scheduled_rollup_uses_one_day_across_the_midnight_boundary(): + table: Final = _FakeSentinelTable() + clock: Final = iter( + ( + datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc), + datetime(2026, 2, 1, 0, 0, 0, 1, tzinfo=timezone.utc), + ) + ).__next__ + prisma: Final = _prisma_for([_windowed_row(effective_from=_midnight(date(2026, 1, 27)))], table) + + result: Final = await run_scheduled_ptu_rollup(prisma, clock=clock) + + assert result is not None + assert result.day == date(2026, 1, 30) + assert _priced_dates(table) == ["2026-01-27", "2026-01-28", "2026-01-29", "2026-01-30"] + + # --- R8: a rename must not re-price history under the new name ---------------- @@ -2066,19 +2086,22 @@ async def test_the_catch_up_pass_reaches_a_config_declared_deployment(): """The catch-up shares the loader, so config deployments join it without being wired in. That is what prices the elapsed days of a reservation declared before today.""" table = _FakeSentinelTable() - now = datetime.now(timezone.utc) - started = (now - timedelta(days=3)).strftime("%Y-%m-%dT00:00:00Z") + now: Final = _SCHEDULED_NOW + started: Final = (now - timedelta(days=3)).strftime("%Y-%m-%dT00:00:00Z") entry = _router_entry( model_id="cfg-back", model_info={"ptu_count": 100, "cost_per_ptu_per_hour": 0.02, "team_id": "t", "ptu_effective_from": started}, ) await run_scheduled_ptu_rollup( - _prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), router=_router_holding(entry) + _prisma_for([], table), + pod_lock_manager=_pod_lock(acquired=True), + router=_router_holding(entry), + clock=lambda: _SCHEDULED_NOW, ) charged = sorted(day for (_, day, _, model) in table.rows if model == "cfg-back") - yesterday = (now.date() - timedelta(days=1)).isoformat() + yesterday: Final = (now.date() - timedelta(days=1)).isoformat() assert len(charged) == 3, charged assert charged[-1] == yesterday assert all(row["ptu_flat_cost"] == pytest.approx(48.0) for row in table.rows.values())