Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/usage_tracking.py: 47%
225 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Track guardrail and policy usage for the dashboard: upsert daily metrics and
3insert into SpendLogGuardrailIndex when spend logs are written.
4"""
6import asyncio
7import json
8from collections import defaultdict
9from collections.abc import Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
10from datetime import datetime, timezone
11from functools import partial
12from itertools import groupby
13from operator import itemgetter
14from types import MappingProxyType
15from typing import TYPE_CHECKING, Any, Final, NamedTuple, TypeVar
17from typing_extensions import ReadOnly, TypedDict
19from litellm._logging import verbose_proxy_logger
20from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS
21from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import billed_guardrail_cost_by_unit
22from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES
23from litellm.proxy.db.spend_log_batching import spend_log_write_batches
24from litellm.proxy.utils import PrismaClient
25from litellm.repositories.table_repositories import (
26 DailyGuardrailMetricsRepository,
27 DailyGuardrailUsageUnitsRepository,
28 SpendLogGuardrailIndexRepository,
29)
31if TYPE_CHECKING: 31 ↛ 32line 31 didn't jump to line 32 because the condition on line 31 was never true
32 from prisma import types as prisma_types
35_UPSERT_RETRY_TIMES: Final = 3
36_MAX_PENDING_ROWS: Final = 10_000
38_RowKey = TypeVar("_RowKey")
39_RowValue = TypeVar("_RowValue")
42class _UsageUnitKey(NamedTuple):
43 guardrail_id: str
44 date: str
45 team_id: str
46 api_key: str
47 usage_unit: str
50class _UsageUnitIncrement(NamedTuple):
51 units: int
52 cost: float
53 """USD for the priced share of units."""
54 untracked_units: int
55 """Units recorded with no known price, the share cost leaves out."""
58def _usage_unit_increment(units: int, cost: float | None) -> _UsageUnitIncrement:
59 if cost is None:
60 return _UsageUnitIncrement(units=units, cost=0.0, untracked_units=units)
61 return _UsageUnitIncrement(units=units, cost=cost, untracked_units=0)
64class _MetricsKey(NamedTuple):
65 guardrail_id: str
66 date: str
69class _UsageUnitCompoundKey(TypedDict):
70 guardrail_id: ReadOnly[str]
71 date: ReadOnly[str]
72 team_id: ReadOnly[str]
73 api_key: ReadOnly[str]
74 usage_unit: ReadOnly[str]
77class _UsageUnitWhereUnique(TypedDict):
78 guardrail_id_date_team_id_api_key_usage_unit: ReadOnly[_UsageUnitCompoundKey]
81class PendingRollups:
82 """Rollup rows whose connection-error retries exhausted, held for the next flush."""
84 def __init__(self) -> None:
85 self.lock: Final = asyncio.Lock()
86 self.metrics: Mapping[_MetricsKey, Mapping[str, int]] = MappingProxyType({})
87 self.units: Mapping[_UsageUnitKey, _UsageUnitIncrement] = MappingProxyType({})
90_PENDING_ROLLUPS: Final = PendingRollups()
92_NO_COUNTERS: Final[Mapping[str, int]] = MappingProxyType({})
93_NO_INCREMENT: Final = _UsageUnitIncrement(units=0, cost=0.0, untracked_units=0)
96def _merged_keys(base: Mapping[_RowKey, object], extra: Mapping[_RowKey, object]) -> tuple[_RowKey, ...]:
97 return (*base, *(key for key in extra if key not in base))
100def _summed_increments(increments: Iterable[_UsageUnitIncrement]) -> _UsageUnitIncrement:
101 materialized: Final = tuple(increments)
102 return _UsageUnitIncrement(
103 units=sum(i.units for i in materialized),
104 cost=sum(i.cost for i in materialized),
105 untracked_units=sum(i.untracked_units for i in materialized),
106 )
109def _merged_unit_rows(
110 base: Mapping[_UsageUnitKey, _UsageUnitIncrement], extra: Mapping[_UsageUnitKey, _UsageUnitIncrement]
111) -> Mapping[_UsageUnitKey, _UsageUnitIncrement]:
112 return MappingProxyType(
113 {
114 key: _summed_increments((base.get(key, _NO_INCREMENT), extra.get(key, _NO_INCREMENT)))
115 for key in _merged_keys(base, extra)
116 }
117 )
120def _merged_metric_rows(
121 base: Mapping[_MetricsKey, Mapping[str, int]], extra: Mapping[_MetricsKey, Mapping[str, int]]
122) -> Mapping[_MetricsKey, Mapping[str, int]]:
123 def merged_counters(key: _MetricsKey) -> Mapping[str, int]:
124 base_counters: Final = base.get(key, _NO_COUNTERS)
125 extra_counters: Final = extra.get(key, _NO_COUNTERS)
126 return MappingProxyType(
127 {
128 counter: int(base_counters.get(counter, 0)) + int(extra_counters.get(counter, 0))
129 for counter in _merged_keys(base_counters, extra_counters)
130 }
131 )
133 return MappingProxyType({key: merged_counters(key) for key in _merged_keys(base, extra)})
136def _capped(rows: Mapping[_RowKey, _RowValue], label: str) -> Mapping[_RowKey, _RowValue]:
137 if len(rows) <= _MAX_PENDING_ROWS:
138 return rows
139 verbose_proxy_logger.warning(
140 "Guardrail usage tracking: pending %s requeue exceeds %d rows; dropping the %d oldest (non-fatal)",
141 label,
142 _MAX_PENDING_ROWS,
143 len(rows) - _MAX_PENDING_ROWS,
144 )
145 return MappingProxyType(dict(tuple(rows.items())[len(rows) - _MAX_PENDING_ROWS :]))
148async def _attempt_upsert(
149 upsert_row: Callable[[_RowKey, _RowValue], Awaitable[None]], key: _RowKey, value: _RowValue
150) -> Exception | None:
151 try:
152 await upsert_row(key, value)
153 except Exception as error:
154 return error
155 return None
158async def _upsert_rows_with_retry(
159 rows: Mapping[_RowKey, _RowValue],
160 upsert_row: Callable[[_RowKey, _RowValue], Awaitable[None]],
161 label: str,
162 sleep: Callable[[float], Awaitable[None]],
163 retries_left: int = _UPSERT_RETRY_TIMES,
164) -> Mapping[_RowKey, _RowValue]:
165 """Returns the rows still failing with connection errors once retries exhaust, for requeueing."""
166 outcomes: Final = {key: await _attempt_upsert(upsert_row, key, value) for key, value in rows.items()}
167 for key, error in outcomes.items():
168 if error is not None and not isinstance(error, DB_RETRY_SAFE_ERROR_TYPES):
169 verbose_proxy_logger.warning(
170 "Guardrail usage tracking: %s upsert failed for %s and is not safe to retry (non-fatal): %s",
171 label,
172 key,
173 error,
174 )
175 retryable: Final = MappingProxyType(
176 {key: rows[key] for key, error in outcomes.items() if isinstance(error, DB_RETRY_SAFE_ERROR_TYPES)}
177 )
178 if not retryable:
179 return MappingProxyType({})
180 if retries_left == 0:
181 for key in retryable:
182 verbose_proxy_logger.warning(
183 "Guardrail usage tracking: %s upsert failed for %s after %d retries; requeued for the next flush "
184 "(non-fatal): %s",
185 label,
186 key,
187 _UPSERT_RETRY_TIMES,
188 outcomes[key],
189 )
190 return retryable
191 await sleep(2 ** (_UPSERT_RETRY_TIMES - retries_left))
192 return await _upsert_rows_with_retry(retryable, upsert_row, label, sleep, retries_left - 1)
195def guardrail_status_to_action(status: str | None) -> str:
196 """Map StandardLogging guardrail_status to blocked/passed/flagged/not_run."""
197 if not status:
198 return "passed"
199 s: Final = (status or "").lower()
200 if s == "not_run":
201 return "not_run"
202 if "intervened" in s or "block" in s:
203 return "blocked"
204 if "flagged" in s or "fail" in s or "error" in s:
205 return "flagged"
206 return "passed"
209def _parse_guardrail_info_from_payload(payload: Mapping[str, object]) -> Sequence[Mapping[str, Any]]:
210 """Extract guardrail_information from spend log payload metadata."""
211 meta = payload.get("metadata")
212 if not meta: 212 ↛ 213line 212 didn't jump to line 213 because the condition on line 212 was never true
213 return []
214 if isinstance(meta, str): 214 ↛ 219line 214 didn't jump to line 219 because the condition on line 214 was always true
215 try:
216 meta = json.loads(meta)
217 except (json.JSONDecodeError, TypeError):
218 return []
219 if not isinstance(meta, dict): 219 ↛ 220line 219 didn't jump to line 220 because the condition on line 219 was never true
220 return []
221 info: Final = meta.get("guardrail_information") or meta.get("standard_logging_guardrail_information")
222 if not isinstance(info, list): 222 ↛ 224line 222 didn't jump to line 224 because the condition on line 222 was always true
223 return []
224 return info
227def _date_str(dt: datetime) -> str:
228 """YYYY-MM-DD in UTC."""
229 if dt.tzinfo is None: 229 ↛ 230line 229 didn't jump to line 230 because the condition on line 229 was never true
230 dt = dt.replace(tzinfo=timezone.utc)
231 return dt.astimezone(timezone.utc).strftime("%Y-%m-%d")
234def _parse_payload_start_time(payload: Mapping[str, object]) -> datetime | None:
235 start_time: Final = payload.get("startTime")
236 if isinstance(start_time, datetime): 236 ↛ 237line 236 didn't jump to line 237 because the condition on line 236 was never true
237 return start_time
238 if not isinstance(start_time, str): 238 ↛ 239line 238 didn't jump to line 239 because the condition on line 238 was never true
239 return None
240 try:
241 return datetime.fromisoformat(start_time.replace("Z", "+00:00"))
242 except (ValueError, TypeError):
243 return None
246def _iter_usage_unit_increments(
247 logs_to_process: Sequence[Mapping[str, object]],
248) -> Iterator[tuple[_UsageUnitKey, _UsageUnitIncrement]]:
249 for payload in logs_to_process:
250 start_time = _parse_payload_start_time(payload)
251 if not payload.get("request_id") or start_time is None: 251 ↛ 252line 251 didn't jump to line 252 because the condition on line 251 was never true
252 continue
253 date_key = _date_str(start_time)
254 team_id = str(payload.get("team_id") or "")
255 api_key = str(payload.get("api_key") or "")
256 for entry in _parse_guardrail_info_from_payload(payload): 256 ↛ 257line 256 didn't jump to line 257 because the loop on line 256 never started
257 guardrail_id = str(entry.get("guardrail_id") or entry.get("guardrail_name") or "")
258 usage = entry.get("guardrail_usage")
259 if not guardrail_id or not isinstance(usage, dict):
260 continue
261 cost_by_unit = billed_guardrail_cost_by_unit(entry)
262 for unit_name, units in usage.items():
263 if isinstance(units, int) and not isinstance(units, bool) and units > 0:
264 key = _UsageUnitKey(guardrail_id, date_key, team_id, api_key, str(unit_name))
265 cost = cost_by_unit.get(str(unit_name)) if cost_by_unit is not None else None
266 yield key, _usage_unit_increment(units=units, cost=cost)
269def _sum_usage_unit_increments(
270 logs_to_process: Sequence[Mapping[str, object]],
271) -> Mapping[_UsageUnitKey, _UsageUnitIncrement]:
272 ordered: Final = sorted(_iter_usage_unit_increments(logs_to_process), key=itemgetter(0))
273 return MappingProxyType(
274 {
275 key: _summed_increments(increment for _, increment in group)
276 for key, group in groupby(ordered, key=itemgetter(0))
277 }
278 )
281async def _upsert_usage_unit_row(
282 prisma_client: PrismaClient, key: _UsageUnitKey, increment: _UsageUnitIncrement
283) -> None:
284 row: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsCreateInput] = {
285 "guardrail_id": key.guardrail_id,
286 "date": key.date,
287 "team_id": key.team_id,
288 "api_key": key.api_key,
289 "usage_unit": key.usage_unit,
290 "units": increment.units,
291 "cost": increment.cost,
292 "untracked_units": increment.untracked_units,
293 }
294 where: Final[_UsageUnitWhereUnique] = {
295 "guardrail_id_date_team_id_api_key_usage_unit": {
296 "guardrail_id": key.guardrail_id,
297 "date": key.date,
298 "team_id": key.team_id,
299 "api_key": key.api_key,
300 "usage_unit": key.usage_unit,
301 }
302 }
303 # A row written before the cost column has NULL cost, and NULL + x stays NULL, so it keeps reading as unknown
304 data: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsUpsertInput] = {
305 "create": row,
306 "update": {
307 "units": {"increment": increment.units},
308 "cost": {"increment": increment.cost},
309 "untracked_units": {"increment": increment.untracked_units},
310 },
311 }
312 await DailyGuardrailUsageUnitsRepository(prisma_client).table.upsert(where=where, data=data)
315async def _upsert_metrics_row(prisma_client: PrismaClient, key: _MetricsKey, agg: Mapping[str, int]) -> None:
316 n: Final = int(agg["requests_evaluated"])
317 await DailyGuardrailMetricsRepository(prisma_client).table.upsert(
318 where={"guardrail_id_date": {"guardrail_id": key.guardrail_id, "date": key.date}},
319 data={
320 "create": {
321 "guardrail_id": key.guardrail_id,
322 "date": key.date,
323 "requests_evaluated": n,
324 "passed_count": int(agg["passed_count"]),
325 "blocked_count": int(agg["blocked_count"]),
326 "flagged_count": int(agg["flagged_count"]),
327 },
328 "update": {
329 "requests_evaluated": {"increment": n},
330 "passed_count": {"increment": int(agg["passed_count"])},
331 "blocked_count": {"increment": int(agg["blocked_count"])},
332 "flagged_count": {"increment": int(agg["flagged_count"])},
333 },
334 },
335 )
338async def process_spend_logs_guardrail_usage(
339 prisma_client: PrismaClient,
340 logs_to_process: Sequence[Mapping[str, object]],
341 sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
342 pending: PendingRollups = _PENDING_ROLLUPS,
343) -> None:
344 """
345 After spend logs are written: update DailyGuardrailMetrics and insert
346 SpendLogGuardrailIndex rows from guardrail_information in each payload.
347 """
348 if not logs_to_process:
349 return
350 # Aggregate daily metrics by (guardrail_id, date). Latency/score metrics dropped.
351 daily_guardrail: Final[dict[_MetricsKey, dict[str, int]]] = defaultdict(
352 lambda: {
353 "requests_evaluated": 0,
354 "passed_count": 0,
355 "blocked_count": 0,
356 "flagged_count": 0,
357 }
358 )
359 index_rows_by_key: Final[dict[tuple[str, str], dict[str, object]]] = {}
361 for payload in logs_to_process:
362 request_id = payload.get("request_id")
363 start_time = _parse_payload_start_time(payload)
364 if not isinstance(request_id, str) or not request_id or start_time is None: 364 ↛ 365line 364 didn't jump to line 365 because the condition on line 364 was never true
365 continue
366 date_key = _date_str(start_time)
368 entries = _parse_guardrail_info_from_payload(payload)
369 ids_by_name = MappingProxyType(
370 {
371 e["guardrail_name"]: e["guardrail_id"]
372 for e in entries
373 if e.get("guardrail_id") and isinstance(e.get("guardrail_name"), str) and e["guardrail_name"]
374 }
375 )
376 for entry in entries: 376 ↛ 377line 376 didn't jump to line 377 because the loop on line 376 never started
377 raw_name = entry.get("guardrail_name")
378 guardrail_name = raw_name if isinstance(raw_name, str) else ""
379 guardrail_id = entry.get("guardrail_id") or ids_by_name.get(guardrail_name) or guardrail_name
380 if not isinstance(guardrail_id, str) or not guardrail_id:
381 continue
382 action = guardrail_status_to_action(entry.get("guardrail_status"))
383 if action != "not_run":
384 key = _MetricsKey(guardrail_id, date_key)
385 daily_guardrail[key]["requests_evaluated"] += 1
386 if action == "passed":
387 daily_guardrail[key]["passed_count"] += 1
388 elif action == "blocked":
389 daily_guardrail[key]["blocked_count"] += 1
390 else:
391 daily_guardrail[key]["flagged_count"] += 1
392 policy_id = entry.get("policy_id")
393 prior = index_rows_by_key.get((request_id, guardrail_id))
394 if prior is None or (prior["policy_id"] is None and policy_id is not None):
395 index_rows_by_key[(request_id, guardrail_id)] = {
396 "request_id": request_id,
397 "guardrail_id": guardrail_id,
398 "policy_id": policy_id,
399 "start_time": start_time,
400 }
401 index_rows: Final = tuple(index_rows_by_key.values())
403 async with pending.lock:
404 pending_metrics: Final = pending.metrics
405 pending_units: Final = pending.units
406 pending.metrics = MappingProxyType({})
407 pending.units = MappingProxyType({})
409 # Upsert daily guardrail metrics (counts only; latency/score dropped)
410 evaluated_metrics: Final = MappingProxyType(
411 {key: agg for key, agg in daily_guardrail.items() if int(agg["requests_evaluated"]) > 0}
412 )
413 metrics_rows: Final = _merged_metric_rows(pending_metrics, evaluated_metrics)
414 unit_rows: Final = _merged_unit_rows(pending_units, _sum_usage_unit_increments(logs_to_process))
416 if not metrics_rows and not index_rows and not unit_rows: 416 ↛ 419line 416 didn't jump to line 419 because the condition on line 416 was always true
417 return
419 try:
420 index_table: Final = SpendLogGuardrailIndexRepository(prisma_client).table
421 for statement_rows in spend_log_write_batches(
422 index_rows, SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS
423 ):
424 try:
425 await index_table.create_many(data=statement_rows, skip_duplicates=True)
426 except Exception as e:
427 verbose_proxy_logger.debug("Guardrail usage tracking: index create_many skipped: %s", e)
429 failed_metrics: Final = await _upsert_rows_with_retry(
430 metrics_rows, partial(_upsert_metrics_row, prisma_client), "daily metrics", sleep
431 )
432 failed_units: Final = await _upsert_rows_with_retry(
433 unit_rows, partial(_upsert_usage_unit_row, prisma_client), "usage unit", sleep
434 )
435 if failed_metrics or failed_units:
436 async with pending.lock:
437 pending.metrics = _capped(_merged_metric_rows(pending.metrics, failed_metrics), "daily metrics")
438 pending.units = _capped(_merged_unit_rows(pending.units, failed_units), "usage unit")
439 except Exception as e:
440 verbose_proxy_logger.warning("Guardrail usage tracking failed (non-fatal): %s", e)