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 <yuneng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-09 00:15:38 -07:00 • committed by GitHub
parent a7ff709f3d
commit 488594e03f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 210 additions and 94 deletions

View file

@ -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)

View file

@ -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:

View file

@ -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

View file

@ -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."""

View file

@ -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)
)
}

View file

@ -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])

View file

@ -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())