diff --git a/litellm/proxy/roi_calculator/sync.py b/litellm/proxy/roi_calculator/sync.py index d447d1d1942..d551fe1d0b5 100644 --- a/litellm/proxy/roi_calculator/sync.py +++ b/litellm/proxy/roi_calculator/sync.py @@ -25,6 +25,7 @@ from litellm.types.roi_calculator import ( PR_CONCURRENCY: Final = 3 _ESTIMATE_ADAPTER: Final = TypeAdapter(ROIEstimate) _REPORT_ADAPTER: Final = TypeAdapter(ROIReport) +_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object]) class _ConfigParam(Protocol): @@ -42,10 +43,10 @@ class _DailySpendTable(Protocol): async def group_by( self, *, - by: Sequence[Literal["user_id", "date"]], - sum: Mapping[str, object], - where: Mapping[str, object], - order: Mapping[str, object], + by: list[Literal["user_id", "date"]], + sum: dict[str, object], + where: dict[str, object], + order: dict[str, object], ) -> Sequence[Mapping[str, object]]: ... @@ -53,8 +54,7 @@ class _UserTable(Protocol): async def find_many( self, *, - where: Mapping[str, object], - select: Mapping[str, bool], + where: dict[str, object], ) -> Sequence[Mapping[str, object]]: ... @@ -87,7 +87,7 @@ class _DailySpendGroup(BaseModel): model_config = ConfigDict(from_attributes=True) user_id: str | None - date: datetime + date: str sums: _DailySpendSums = Field(alias="_sum") @@ -111,45 +111,51 @@ async def read_spend( database: Final = prisma_client.db daily_table: Final = database.litellm_dailyuserspend + group_by: Final[list[Literal["user_id", "date"]]] = [ + "user_id", + "date", + ] # mutable-ok: prisma client serializer only accepts builtin dict/list + sums: Final[dict[str, object]] = { + "spend": True, + "api_requests": True, + } # mutable-ok: prisma client serializer only accepts builtin dict/list + date_filter: Final[dict[str, object]] = { + "date": { + "gte": start.isoformat(), + "lte": end.isoformat(), + } + } # mutable-ok: prisma client serializer only accepts builtin dict/list + order: Final[dict[str, object]] = { + "date": "asc" + } # mutable-ok: prisma client serializer only accepts builtin dict/list groups: Final = _DAILY_SPEND_GROUPS.validate_python( await daily_table.group_by( - by=("user_id", "date"), - sum=MappingProxyType({"spend": True, "api_requests": True}), - where=MappingProxyType( - { - "date": MappingProxyType( - {"gte": start.isoformat(), "lte": end.isoformat()} - ) - } - ), - order=MappingProxyType({"date": "asc"}), + by=group_by, + sum=sums, + where=date_filter, + order=order, ) ) - user_ids: Final = tuple( - sorted(frozenset(group.user_id for group in groups if group.user_id)) - ) + user_ids: Final = tuple(sorted(frozenset(group.user_id for group in groups if group.user_id))) user_table: Final = database.litellm_usertable + user_filter: Final[dict[str, object]] = { + "user_id": {"in": list(user_ids)} + } # mutable-ok: prisma client serializer only accepts builtin dict/list users: Final = _USER_EMAILS.validate_python( await user_table.find_many( - where=MappingProxyType({"user_id": MappingProxyType({"in": user_ids})}), - select=MappingProxyType({"user_id": True, "user_email": True}), + where=user_filter, ) if user_ids else () ) emails: Final[Mapping[str, str]] = MappingProxyType( - { - user.user_id: normalize_email(user.user_email) - for user in users - if normalize_email(user.user_email) - } + {user.user_id: normalize_email(user.user_email) for user in users if normalize_email(user.user_email)} ) return tuple( ROISpendRecord( - date=group.date.date().isoformat(), + date=group.date, user_id=group.user_id or "", - email=emails.get(group.user_id or "", "") - or normalize_email(group.user_id), + email=emails.get(group.user_id or "", "") or normalize_email(group.user_id), spend=group.sums.spend, requests=group.sums.api_requests, ) @@ -179,9 +185,7 @@ class SyncClock(Protocol): class _StatusUpdate(TypedDict, total=False): running: ReadOnly[bool] - phase: ReadOnly[ - Literal["idle", "spend", "repositories", "estimates", "complete", "cancelled", "error"] - ] + phase: ReadOnly[Literal["idle", "spend", "repositories", "estimates", "complete", "cancelled", "error"]] stage: ReadOnly[str] done: ReadOnly[int] total: ReadOnly[int] @@ -256,9 +260,7 @@ class SyncManager: needs_attention=0, error=None, ) - self._task = asyncio.create_task( - self._run(settings, repository, spend_reader, complete, github_transport) - ) + self._task = asyncio.create_task(self._run(settings, repository, spend_reader, complete, github_transport)) return True async def cancel(self) -> bool: @@ -287,13 +289,10 @@ class SyncManager: start: Final = end - timedelta(days=settings.backfill_days - 1) spend: Final = await spend_reader(start, end) self._update_status(phase="repositories", stage="Reading configured repositories") - pull_groups: Final = await asyncio.gather( - *(github.pulls(repo, start, end) for repo in settings.repos) - ) + pull_groups: Final = await asyncio.gather(*(github.pulls(repo, start, end) for repo in settings.repos)) queue: Final = tuple( chain.from_iterable( - ((repo, pull) for pull in pulls) - for repo, pulls in zip(settings.repos, pull_groups, strict=True) + ((repo, pull) for pull in pulls) for repo, pulls in zip(settings.repos, pull_groups, strict=True) ) ) context: Final = cache_context(settings) @@ -319,11 +318,7 @@ class SyncManager: cached_by_index: Final[Mapping[int, ROIPullRecord]] = MappingProxyType( {index: self._cached_record(pull) for index, pull in cached} ) - pending: Final = tuple( - item - for item in indexed_queue - if item[0] not in cached_by_index - ) + pending: Final = tuple(item for item in indexed_queue if item[0] not in cached_by_index) reused_count: Final = len(cached) self._update_status( phase="estimates", @@ -376,7 +371,10 @@ class SyncManager: warnings=(), ) await github.close() - await repository.set_param("roi_calculator_report", report) + report_json: Final[dict[str, object]] = _JSON_OBJECT_ADAPTER.validate_python( + _REPORT_ADAPTER.dump_python(report, mode="json") + ) + await repository.set_param("roi_calculator_report", report_json) self._update_status(phase="complete", stage="Up to date") except asyncio.CancelledError: self._update_status(phase="cancelled", stage="Sync cancelled") @@ -400,9 +398,7 @@ class SyncManager: self._update_status(running=False) def _update_status(self, **update: Unpack[_StatusUpdate]) -> None: - status: Final = ROISyncStatus.model_validate( - MappingProxyType({**self._status.model_dump(), **update}) - ) + status: Final = ROISyncStatus.model_validate(MappingProxyType({**self._status.model_dump(), **update})) self._status = status async def _previous_report(self, repository: _ReportRepository) -> ROIReport | None: @@ -415,9 +411,7 @@ class SyncManager: return None def _cached_record(self, pull: ROIPullRecord) -> ROIPullRecord: - estimate: Final = _ESTIMATE_ADAPTER.validate_python( - MappingProxyType({**pull["estimate"], "cached": True}) - ) + estimate: Final = _ESTIMATE_ADAPTER.validate_python(MappingProxyType({**pull["estimate"], "cached": True})) return ROIPullRecord( repo=pull["repo"], number=pull["number"], diff --git a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py index 45bf9341921..7e021b9da7f 100644 --- a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py @@ -1,6 +1,7 @@ +import json from collections.abc import Mapping from types import MappingProxyType -from typing import Final +from typing import Final, cast import pytest from fastapi import FastAPI @@ -18,6 +19,12 @@ from litellm.types.roi_calculator import ROISettings _JSON_HEADERS: Final = MappingProxyType({"content-type": "application/json"}) +def _assert_json_round_trip(value: object) -> None: + serialized: Final = json.dumps(value) + decoded: Final[object] = cast(object, json.loads(serialized)) + assert decoded == value + + class _Parameter: def __init__(self, param_value: object) -> None: self.param_value: Final = param_value @@ -32,6 +39,7 @@ class _ConfigRepository: return _Parameter(value) if value is not None else None async def set_param(self, param_name: str, param_value: object) -> object: + _assert_json_round_trip(param_value) self.values = MappingProxyType({**self.values, param_name: param_value}) return self.values[param_name] diff --git a/tests/unit/proxy/roi_calculator/test_sync.py b/tests/unit/proxy/roi_calculator/test_sync.py index 01fe93c8e5b..aeaaa72eb91 100644 --- a/tests/unit/proxy/roi_calculator/test_sync.py +++ b/tests/unit/proxy/roi_calculator/test_sync.py @@ -1,8 +1,9 @@ import asyncio +import json from collections.abc import Mapping, Sequence from datetime import date, datetime, timezone from types import MappingProxyType -from typing import Final, Literal +from typing import Final, Literal, cast import httpx import pytest @@ -55,6 +56,14 @@ _COMMITS_JSON: Final = """[ } } ]""" + + +def _assert_json_round_trip(value: object) -> None: + serialized: Final = json.dumps(value) + decoded: Final[object] = cast(object, json.loads(serialized)) + assert decoded == value + + class _Parameter: def __init__(self, param_value: object) -> None: self.param_value: Final = param_value @@ -69,6 +78,7 @@ class _ReportRepository: return _Parameter(value) if value is not None else None async def set_param(self, param_name: str, param_value: object) -> object: + _assert_json_round_trip(param_value) self.values = MappingProxyType({**self.values, param_name: param_value}) return self.values[param_name] @@ -82,34 +92,27 @@ class _DailySpendTable: where: Mapping[str, object], order: Mapping[str, object], ) -> Sequence[Mapping[str, object]]: - assert by == ("user_id", "date") - assert sum == MappingProxyType({"spend": True, "api_requests": True}) - assert where == MappingProxyType( - {"date": MappingProxyType({"gte": "2026-09-01", "lte": "2026-09-30"})} - ) - assert order == MappingProxyType({"date": "asc"}) + _assert_json_round_trip({"by": by, "sum": sum, "where": where, "order": order}) + assert by == ["user_id", "date"] + assert sum == {"spend": True, "api_requests": True} + assert where == {"date": {"gte": "2026-09-01", "lte": "2026-09-30"}} + assert order == {"date": "asc"} return ( - MappingProxyType( - { - "user_id": "u1", - "date": datetime(2026, 9, 12, tzinfo=timezone.utc), - "_sum": MappingProxyType({"spend": 12.5, "api_requests": 2}), - } - ), - MappingProxyType( - { - "user_id": "team@example.com", - "date": datetime(2026, 9, 13, tzinfo=timezone.utc), - "_sum": MappingProxyType({"spend": 3.0, "api_requests": 1}), - } - ), - MappingProxyType( - { - "user_id": "missing", - "date": datetime(2026, 9, 14, tzinfo=timezone.utc), - "_sum": MappingProxyType({"spend": 1.0, "api_requests": 1}), - } - ), + { + "user_id": "u1", + "date": "2026-09-12", + "_sum": {"spend": 12.5, "api_requests": 2}, + }, + { + "user_id": "team@example.com", + "date": "2026-09-13", + "_sum": {"spend": 3.0, "api_requests": 1}, + }, + { + "user_id": "missing", + "date": "2026-09-14", + "_sum": {"spend": 1.0, "api_requests": 1}, + }, ) @@ -118,12 +121,9 @@ class _UserTable: self, *, where: Mapping[str, object], - select: Mapping[str, bool], ) -> Sequence[Mapping[str, str | None]]: - assert where == MappingProxyType( - {"user_id": MappingProxyType({"in": ("missing", "team@example.com", "u1")})} - ) - assert select == MappingProxyType({"user_id": True, "user_email": True}) + _assert_json_round_trip({"where": where}) + assert where == {"user_id": {"in": ["missing", "team@example.com", "u1"]}} return (MappingProxyType({"user_id": "u1", "user_email": " Alice@Example.com "}),)