perf(guardrails): aggregate usage units in one sorted pass

The flush and the usage endpoints summed units with a scan per distinct key,
quadratic in rows times keys; group sorted rows instead. Skip payloads without
a request_id like the metrics path, type the flush key as a NamedTuple, and drop
the (guardrail_id, date) index that the primary key already covers
This commit is contained in:
mateo-berri 2026-08-17 17:19:05 -07:00
parent 5437139b94
commit 8ba2263d4c
9 changed files with 58 additions and 34 deletions

View file

@ -14,7 +14,3 @@ CREATE TABLE "LiteLLM_DailyGuardrailUsageUnits" (
-- CreateIndex
CREATE INDEX "LiteLLM_DailyGuardrailUsageUnits_date_idx" ON "LiteLLM_DailyGuardrailUsageUnits"("date");
-- CreateIndex
CREATE INDEX "LiteLLM_DailyGuardrailUsageUnits_guardrail_id_date_idx" ON "LiteLLM_DailyGuardrailUsageUnits"("guardrail_id", "date");

View file

@ -1082,7 +1082,6 @@ model LiteLLM_DailyGuardrailUsageUnits {
@@id([guardrail_id, date, team_id, api_key, usage_unit])
@@index([date])
@@index([guardrail_id, date])
}
// Daily policy metrics for usage dashboard (one row per policy per day)

View file

@ -6,6 +6,7 @@ GET /guardrails/usage/overview, /guardrails/usage/detail/:id, /guardrails/usage/
import json
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import datetime, timedelta, timezone
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, overload
@ -113,11 +114,14 @@ async def _find_daily_guardrail_usage_units(
return await _daily_guardrail_usage_units_table(prisma_client).find_many(where=where)
def _counter_name(row: "prisma_models.LiteLLM_DailyGuardrailUsageUnits") -> str:
return row.usage_unit
def _sum_counter_units(rows: "Iterable[prisma_models.LiteLLM_DailyGuardrailUsageUnits]") -> Mapping[str, int]:
materialized: Final = tuple(rows)
counter_names: Final = frozenset(r.usage_unit for r in materialized)
ordered: Final = sorted(rows, key=_counter_name)
return MappingProxyType(
{name: sum(int(r.units) for r in materialized if r.usage_unit == name) for name in counter_names}
{name: sum(int(r.units) for r in group) for name, group in groupby(ordered, key=_counter_name)}
)
@ -125,8 +129,8 @@ def _units_by(
rows: "Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits]",
key_of: "Callable[[prisma_models.LiteLLM_DailyGuardrailUsageUnits], str]",
) -> Mapping[str, Mapping[str, int]]:
keys: Final = frozenset(key_of(r) for r in rows)
return MappingProxyType({key: _sum_counter_units(r for r in rows if key_of(r) == key) for key in keys})
ordered: Final = sorted(rows, key=key_of)
return MappingProxyType({key: _sum_counter_units(group) for key, group in groupby(ordered, key=key_of)})
# --- Response models ---

View file

@ -7,8 +7,10 @@ import json
from collections import defaultdict
from collections.abc import Iterator, Mapping, Sequence
from datetime import datetime, timezone
from itertools import groupby
from operator import itemgetter
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, NamedTuple
from litellm._logging import verbose_proxy_logger
from litellm.proxy.utils import PrismaClient
@ -21,8 +23,13 @@ from litellm.repositories.table_repositories import (
if TYPE_CHECKING:
from prisma import types as prisma_types
_UsageUnitKey = tuple[str, str, str, str, str]
"""(guardrail_id, date, team_id, api_key, usage_unit)"""
class _UsageUnitKey(NamedTuple):
guardrail_id: str
date: str
team_id: str
api_key: str
usage_unit: str
def _guardrail_status_to_action(status: str | None) -> str:
@ -77,44 +84,44 @@ def _parse_payload_start_time(payload: Mapping[str, Any]) -> datetime | None:
def _iter_usage_unit_increments(logs_to_process: Sequence[Mapping[str, Any]]) -> Iterator[tuple[_UsageUnitKey, int]]:
for payload in logs_to_process:
start_time = _parse_payload_start_time(payload)
if start_time is None:
if not payload.get("request_id") or start_time is None:
continue
date_key = _date_str(start_time)
team_id = str(payload.get("team_id") or "")
api_key = str(payload.get("api_key") or "")
for entry in _parse_guardrail_info_from_payload(payload):
guardrail_id = entry.get("guardrail_id") or entry.get("guardrail_name") or ""
guardrail_id = str(entry.get("guardrail_id") or entry.get("guardrail_name") or "")
usage = entry.get("guardrail_usage")
if not guardrail_id or not isinstance(usage, dict):
continue
for unit_name, units in usage.items():
if isinstance(units, int) and not isinstance(units, bool) and units > 0:
yield (guardrail_id, date_key, team_id, api_key, unit_name), units
yield _UsageUnitKey(guardrail_id, date_key, team_id, api_key, str(unit_name)), units
def _sum_usage_unit_increments(logs_to_process: Sequence[Mapping[str, Any]]) -> Mapping[_UsageUnitKey, int]:
increments: Final = tuple(_iter_usage_unit_increments(logs_to_process))
keys: Final = frozenset(k for k, _ in increments)
return MappingProxyType({key: sum(u for k, u in increments if k == key) for key in keys})
ordered: Final = sorted(_iter_usage_unit_increments(logs_to_process), key=itemgetter(0))
return MappingProxyType(
{key: sum(units for _, units in group) for key, group in groupby(ordered, key=itemgetter(0))}
)
async def _upsert_usage_unit_row(prisma_client: PrismaClient, key: _UsageUnitKey, units: int) -> None:
guardrail_id, date_key, team_id, api_key, usage_unit = key
row: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsCreateInput] = {
"guardrail_id": guardrail_id,
"date": date_key,
"team_id": team_id,
"api_key": api_key,
"usage_unit": usage_unit,
"guardrail_id": key.guardrail_id,
"date": key.date,
"team_id": key.team_id,
"api_key": key.api_key,
"usage_unit": key.usage_unit,
"units": units,
}
where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereUniqueInput] = {
"guardrail_id_date_team_id_api_key_usage_unit": {
"guardrail_id": guardrail_id,
"date": date_key,
"team_id": team_id,
"api_key": api_key,
"usage_unit": usage_unit,
"guardrail_id": key.guardrail_id,
"date": key.date,
"team_id": key.team_id,
"api_key": key.api_key,
"usage_unit": key.usage_unit,
}
}
data: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsUpsertInput] = {

View file

@ -1082,7 +1082,6 @@ model LiteLLM_DailyGuardrailUsageUnits {
@@id([guardrail_id, date, team_id, api_key, usage_unit])
@@index([date])
@@index([guardrail_id, date])
}
// Daily policy metrics for usage dashboard (one row per policy per day)

View file

@ -1082,7 +1082,6 @@ model LiteLLM_DailyGuardrailUsageUnits {
@@id([guardrail_id, date, team_id, api_key, usage_unit])
@@index([date])
@@index([guardrail_id, date])
}
// Daily policy metrics for usage dashboard (one row per policy per day)

View file

@ -256,6 +256,8 @@ async def test_overview_reports_usage_units_per_row_and_total():
row = next(r for r in resp.rows if r.id == "yaml-uuid")
assert row.usageUnits == {"topicPolicyUnits": 4, "contentPolicyUnits": 5}
assert resp.totalUsageUnits == {"topicPolicyUnits": 11, "contentPolicyUnits": 5}
units_where = prisma.db.litellm_dailyguardrailusageunits.find_many.call_args.kwargs["where"]
assert units_where == {"date": {"gte": START, "lte": END}}
@pytest.mark.asyncio
@ -289,6 +291,8 @@ async def test_detail_breaks_units_down_by_day_team_and_key():
"hash-1": {"contentPolicyUnits": 2, "topicPolicyUnits": 1},
"hash-2": {"contentPolicyUnits": 1},
}
units_where = prisma.db.litellm_dailyguardrailusageunits.find_many.call_args.kwargs["where"]
assert units_where == {"guardrail_id": {"in": ["yaml-pii", "yaml-1"]}, "date": {"gte": START, "lte": END}}
# ---- logs -------------------------------------------------------------------

View file

@ -122,3 +122,19 @@ async def test_zero_and_non_int_usage_counters_are_skipped():
assert _units_upserts(prisma) == {
("bedrock-guard", "2026-08-17", "team-a", "hashed-key-1", "topicPolicyUnits"): 1,
}
@pytest.mark.asyncio
async def test_payload_without_request_id_is_skipped_like_the_metrics_path():
prisma = _prisma()
logs = [
{**_payload("ignored", usage={"topicPolicyUnits": 5}), "request_id": None},
_payload("r2", usage={"topicPolicyUnits": 1}),
]
await process_spend_logs_guardrail_usage(prisma, logs)
assert _units_upserts(prisma) == {
("bedrock-guard", "2026-08-17", "team-a", "hashed-key-1", "topicPolicyUnits"): 1,
}
assert prisma.db.litellm_dailyguardrailmetrics.upsert.call_args.kwargs["data"]["create"]["requests_evaluated"] == 1

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22906
"limit": 22903
},
"LIT002": {
"limit": 26896
"limit": 26894
},
"LIT003": {
"limit": 269