mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
49871cee15
commit
1e6854a2e2
3 changed files with 89 additions and 87 deletions
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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 "}),)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue