mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
a7ff709f3d
commit
488594e03f
7 changed files with 210 additions and 94 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue