fix(proxy): serialize ROI Prisma inputs with builtin containers

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-29 06:29:07 +00:00
parent 49871cee15
commit 1e6854a2e2
3 changed files with 89 additions and 87 deletions

View file

@ -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"],

View file

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

View file

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