diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260817143646_add_daily_guardrail_usage_units/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260817143646_add_daily_guardrail_usage_units/migration.sql new file mode 100644 index 00000000000..6838eb76f3e --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260817143646_add_daily_guardrail_usage_units/migration.sql @@ -0,0 +1,20 @@ +-- CreateTable +CREATE TABLE "LiteLLM_DailyGuardrailUsageUnits" ( + "guardrail_id" TEXT NOT NULL, + "date" TEXT NOT NULL, + "team_id" TEXT NOT NULL, + "api_key" TEXT NOT NULL, + "usage_unit" TEXT NOT NULL, + "units" BIGINT NOT NULL DEFAULT 0, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_DailyGuardrailUsageUnits_pkey" PRIMARY KEY ("guardrail_id","date","team_id","api_key","usage_unit") +); + +-- 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"); + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 71345d2ccde..d3c277278ff 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1069,6 +1069,22 @@ model LiteLLM_DailyGuardrailMetrics { @@index([guardrail_id]) } +// Daily guardrail billable usage units (one row per guardrail/day/team/key/unit type) +model LiteLLM_DailyGuardrailUsageUnits { + guardrail_id String + date String // YYYY-MM-DD + team_id String // empty string when the request had no team + api_key String // hashed virtual key; empty string when unknown + usage_unit String // provider counter name, e.g. Bedrock's contentPolicyUnits + units BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@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) model LiteLLM_DailyPolicyMetrics { policy_id String diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index e8c6eba581c..93cbb989e23 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -2053,6 +2053,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_action: Final = response.get("action") if isinstance(bedrock_action, str): tracing_detail["guardrail_action"] = bedrock_action + usage: Final = response.get("usage") + if isinstance(usage, dict): + usage_units: Final = { # mutable-ok: json.dumps'd into spend log metadata downstream + key: value for key, value in usage.items() if isinstance(value, int) + } + if usage_units: + tracing_detail["guardrail_usage"] = usage_units return tracing_detail def _extract_violation_category_names(self, response: BedrockGuardrailResponse) -> list[str]: diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index f48e20257db..be324e3d81b 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -4,8 +4,9 @@ GET /guardrails/usage/overview, /guardrails/usage/detail/:id, /guardrails/usage/ """ import json -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence from datetime import datetime, timedelta, timezone +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, overload from fastapi import APIRouter, Depends, Query @@ -16,6 +17,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.repositories.table_repositories import ( DailyGuardrailMetricsRepository, + DailyGuardrailUsageUnitsRepository, DailyPolicyMetricsRepository, GuardrailsRepository, PolicyRepository, @@ -28,6 +30,7 @@ if TYPE_CHECKING: from prisma import types as prisma_types from prisma.actions import ( LiteLLM_DailyGuardrailMetricsActions, + LiteLLM_DailyGuardrailUsageUnitsActions, LiteLLM_DailyPolicyMetricsActions, LiteLLM_GuardrailsTableActions, LiteLLM_PolicyTableActions, @@ -41,6 +44,8 @@ if TYPE_CHECKING: router: Final = APIRouter() +_EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({}) + def _guardrails_table( prisma_client: "PrismaClient", @@ -92,6 +97,38 @@ async def _find_daily_policy_metrics( return await _daily_policy_metrics_table(prisma_client).find_many(where=where) +def _daily_guardrail_usage_units_table( + prisma_client: "PrismaClient", +) -> "LiteLLM_DailyGuardrailUsageUnitsActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]": + units_table: Final[LiteLLM_DailyGuardrailUsageUnitsActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]] = ( + DailyGuardrailUsageUnitsRepository(prisma_client).table + ) + return units_table + + +async def _find_daily_guardrail_usage_units( + prisma_client: "PrismaClient", + where: "prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput", +) -> "Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits]": + return await _daily_guardrail_usage_units_table(prisma_client).find_many(where=where) + + +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) + return MappingProxyType( + {name: sum(int(r.units) for r in materialized if r.usage_unit == name) for name in counter_names} + ) + + +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}) + + # --- Response models --- @@ -140,6 +177,7 @@ class UsageOverviewRow(BaseModel): avgLatency: float | None status: str # healthy | warning | critical trend: str # up | down | stable + usageUnits: Mapping[str, int] # provider counter name -> billable units in range class UsageOverviewResponse(BaseModel): @@ -148,6 +186,12 @@ class UsageOverviewResponse(BaseModel): totalRequests: int totalBlocked: int passRate: float + totalUsageUnits: Mapping[str, int] + + +class UsageUnitsDailyPoint(BaseModel): + date: str + units: Mapping[str, int] class UsageDetailResponse(BaseModel): @@ -163,6 +207,10 @@ class UsageDetailResponse(BaseModel): trend: str description: str | None time_series: list[UsageChartPoint] + usage_units: Mapping[str, int] + usage_units_daily: Sequence[UsageUnitsDailyPoint] + usage_units_by_team: Mapping[str, Mapping[str, int]] # team_id ("" = no team) -> counter -> units + usage_units_by_key: Mapping[str, Mapping[str, int]] # hashed api key ("" = unknown) -> counter -> units class UsageLogEntry(BaseModel): @@ -278,6 +326,7 @@ def _guardrail_overview_rows( guardrails: "Sequence[_DbOrConfigGuardrail]", agg: Mapping[str, _MetricTotals], prev_agg: Mapping[str, float], + units_agg: Mapping[str, Mapping[str, int]], ) -> list[UsageOverviewRow]: rows: Final[list[UsageOverviewRow]] = [] covered_keys: Final[set[str]] = set() @@ -303,6 +352,7 @@ def _guardrail_overview_rows( prev_fail = float(prev_agg.get(k, 0.0) or 0.0) break trend = _trend_from_comparison(fail_rate, prev_fail) + row_units: Mapping[str, int] = next((units_agg[k] for k in lookup_keys if k in units_agg), _EMPTY_UNITS) rows.append( UsageOverviewRow( id=gid, @@ -315,6 +365,7 @@ def _guardrail_overview_rows( avgLatency=None, status=_status_from_fail_rate(fail_rate), trend=trend, + usageUnits=row_units, ) ) # Add rows for guardrails with metrics but not in guardrails table (e.g. MCP, config) @@ -337,6 +388,7 @@ def _guardrail_overview_rows( avgLatency=None, status=_status_from_fail_rate(fail_rate), trend=trend, + usageUnits=units_agg.get(agg_key, _EMPTY_UNITS), ) ) return rows @@ -366,6 +418,7 @@ def _policy_overview_rows( avgLatency=None, status=_status_from_fail_rate(fail_rate), trend=trend, + usageUnits=_EMPTY_UNITS, ) ) return rows @@ -386,7 +439,9 @@ async def guardrails_usage_overview( from litellm.proxy.proxy_server import prisma_client if prisma_client is None: - return UsageOverviewResponse(rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0) + return UsageOverviewResponse( + rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS + ) now: Final = datetime.now(timezone.utc) end: Final = end_date or now.strftime("%Y-%m-%d") @@ -413,19 +468,28 @@ async def guardrails_usage_overview( prisma_client, where={"date": {"gte": start_prev, "lt": start}} ) + units_where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput] = { + "date": {"gte": start, "lte": end} + } + units_rows: Final[ + Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits] + ] = await _find_daily_guardrail_usage_units(prisma_client, where=units_where) + agg: Final = _aggregate_daily_metrics(metrics, "guardrail_id") prev_agg: Final = _prev_fail_rates(metrics_prev, "guardrail_id") + units_agg: Final = _units_by(units_rows, lambda r: r.guardrail_id) chart: Final = _chart_from_metrics(metrics) total_requests: Final = sum(a["requests"] for a in agg.values()) total_blocked: Final = sum(a["blocked"] for a in agg.values()) pass_rate: Final = (100.0 * (total_requests - total_blocked) / total_requests) if total_requests else 100.0 - rows: Final = _guardrail_overview_rows(guardrails, agg, prev_agg) + rows: Final = _guardrail_overview_rows(guardrails, agg, prev_agg, units_agg) return UsageOverviewResponse( rows=rows, chart=chart, totalRequests=total_requests, totalBlocked=total_blocked, passRate=round(pass_rate, 1), + totalUsageUnits=_sum_counter_units(units_rows), ) except Exception as e: from litellm.proxy.utils import handle_exception_on_proxy @@ -485,6 +549,13 @@ async def guardrails_usage_detail( "date": {"lt": start}, }, ) + units_where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput] = { + "guardrail_id": {"in": metric_ids}, + "date": {"gte": start, "lte": end}, + } + units_rows: Final[ + Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits] + ] = await _find_daily_guardrail_usage_units(prisma_client, where=units_where) requests: Final = sum(int(m.requests_evaluated or 0) for m in metrics) blocked: Final = sum(int(m.blocked_count or 0) for m in metrics) @@ -510,6 +581,8 @@ async def guardrails_usage_detail( litellm_params: Final = _to_dict(_get_guardrail_field(guardrail, "litellm_params")) guardrail_info: Final = _to_dict(_get_guardrail_field(guardrail, "guardrail_info")) _guardrail_name: Final = _get_guardrail_field(guardrail, "guardrail_name") + daily_unit_sums: Final = sorted(_units_by(units_rows, lambda r: r.date).items()) + units_daily: Final = tuple(UsageUnitsDailyPoint(date=d, units=units) for d, units in daily_unit_sums) return UsageDetailResponse( guardrail_id=guardrail_id, @@ -524,6 +597,10 @@ async def guardrails_usage_detail( trend=trend, description=guardrail_info.get("description"), time_series=time_series, + usage_units=_sum_counter_units(units_rows), + usage_units_daily=units_daily, + usage_units_by_team=_units_by(units_rows, lambda r: r.team_id), + usage_units_by_key=_units_by(units_rows, lambda r: r.api_key), ) @@ -743,7 +820,9 @@ async def policies_usage_overview( from litellm.proxy.proxy_server import prisma_client if prisma_client is None: - return UsageOverviewResponse(rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0) + return UsageOverviewResponse( + rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS + ) now: Final = datetime.now(timezone.utc) end: Final = end_date or now.strftime("%Y-%m-%d") @@ -776,6 +855,7 @@ async def policies_usage_overview( totalRequests=total_requests, totalBlocked=total_blocked, passRate=round(pass_rate, 1), + totalUsageUnits=_EMPTY_UNITS, ) except Exception as e: from litellm.proxy.utils import handle_exception_on_proxy diff --git a/litellm/proxy/guardrails/usage_tracking.py b/litellm/proxy/guardrails/usage_tracking.py index 54dfe8eece1..cebf329b601 100644 --- a/litellm/proxy/guardrails/usage_tracking.py +++ b/litellm/proxy/guardrails/usage_tracking.py @@ -5,16 +5,25 @@ insert into SpendLogGuardrailIndex when spend logs are written. import json from collections import defaultdict +from collections.abc import Iterator, Mapping, Sequence from datetime import datetime, timezone -from typing import Any, Final +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_proxy_logger from litellm.proxy.utils import PrismaClient from litellm.repositories.table_repositories import ( DailyGuardrailMetricsRepository, + DailyGuardrailUsageUnitsRepository, SpendLogGuardrailIndexRepository, ) +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)""" + def _guardrail_status_to_action(status: str | None) -> str: """Map StandardLogging guardrail_status to blocked/passed/flagged.""" @@ -28,7 +37,7 @@ def _guardrail_status_to_action(status: str | None) -> str: return "passed" -def _parse_guardrail_info_from_payload(payload: dict[str, Any]) -> list[dict[str, Any]]: +def _parse_guardrail_info_from_payload(payload: Mapping[str, Any]) -> Sequence[Mapping[str, Any]]: """Extract guardrail_information from spend log payload metadata.""" meta = payload.get("metadata") if not meta: @@ -53,6 +62,68 @@ def _date_str(dt: datetime) -> str: return dt.astimezone(timezone.utc).strftime("%Y-%m-%d") +def _parse_payload_start_time(payload: Mapping[str, Any]) -> datetime | None: + start_time: Final = payload.get("startTime") + if isinstance(start_time, datetime): + return start_time + if not isinstance(start_time, str): + return None + try: + return datetime.fromisoformat(start_time.replace("Z", "+00:00")) + except (ValueError, TypeError): + return 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: + 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 "" + 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 + + +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}) + + +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, + "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, + } + } + data: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsUpsertInput] = { + "create": row, + "update": {"units": {"increment": units}}, + } + await DailyGuardrailUsageUnitsRepository(prisma_client).table.upsert(where=where, data=data) + + async def process_spend_logs_guardrail_usage( prisma_client: PrismaClient, logs_to_process: list[dict[str, Any]], @@ -76,14 +147,9 @@ async def process_spend_logs_guardrail_usage( for payload in logs_to_process: request_id = payload.get("request_id") - start_time = payload.get("startTime") - if not request_id or not start_time: + start_time = _parse_payload_start_time(payload) + if not request_id or start_time is None: continue - if isinstance(start_time, str): - try: - start_time = datetime.fromisoformat(start_time.replace("Z", "+00:00")) - except (ValueError, TypeError): - continue date_key = _date_str(start_time) for entry in _parse_guardrail_info_from_payload(payload): @@ -109,31 +175,17 @@ async def process_spend_logs_guardrail_usage( } ) - if not daily_guardrail and not index_rows: + usage_unit_totals: Final = _sum_usage_unit_increments(logs_to_process) + + if not daily_guardrail and not index_rows and not usage_unit_totals: return try: # Insert index rows (skip duplicates by request_id + guardrail_id) if index_rows: - index_data: Final = [] - for r in index_rows: - st = r["start_time"] - if isinstance(st, str): - try: - st = datetime.fromisoformat(st.replace("Z", "+00:00")) - except (ValueError, TypeError): - continue - index_data.append( - { - "request_id": r["request_id"], - "guardrail_id": r["guardrail_id"], - "policy_id": r.get("policy_id"), - "start_time": st, - } - ) try: await SpendLogGuardrailIndexRepository(prisma_client).table.create_many( - data=index_data, + data=index_rows, skip_duplicates=True, ) except Exception as e: @@ -168,5 +220,8 @@ async def process_spend_logs_guardrail_usage( }, }, ) + + for unit_key, units in usage_unit_totals.items(): + await _upsert_usage_unit_row(prisma_client, unit_key, units) except Exception as e: verbose_proxy_logger.warning("Guardrail usage tracking failed (non-fatal): %s", e) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 4551680e1b4..da44c06040d 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -125,6 +125,19 @@ class _ProxyDBLogger(CustomLogger): existing_metadata: Final[dict] = request_data.get("metadata", None) or {} existing_metadata.update(_metadata) + # Guardrail hooks write standard_logging_guardrail_information into the + # request's litellm_metadata bucket when one exists (get_or_create_metadata_bucket + # prefers it). Failure rows are serialized from the metadata bucket lifted below, + # so carry the guardrail info over or blocked invocations lose it in spend logs. + litellm_metadata_bucket: Final = request_data.get("litellm_metadata") + if ( + isinstance(litellm_metadata_bucket, dict) + and "standard_logging_guardrail_information" not in existing_metadata + ): + guardrail_info: Final = litellm_metadata_bucket.get("standard_logging_guardrail_information") + if guardrail_info is not None: + existing_metadata["standard_logging_guardrail_information"] = guardrail_info + if "litellm_params" not in request_data: request_data["litellm_params"] = {} diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 71345d2ccde..d3c277278ff 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1069,6 +1069,22 @@ model LiteLLM_DailyGuardrailMetrics { @@index([guardrail_id]) } +// Daily guardrail billable usage units (one row per guardrail/day/team/key/unit type) +model LiteLLM_DailyGuardrailUsageUnits { + guardrail_id String + date String // YYYY-MM-DD + team_id String // empty string when the request had no team + api_key String // hashed virtual key; empty string when unknown + usage_unit String // provider counter name, e.g. Bedrock's contentPolicyUnits + units BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@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) model LiteLLM_DailyPolicyMetrics { policy_id String diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index be19f290ba6..131f4d377ef 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -158,6 +158,10 @@ class DailyGuardrailMetricsRepository(PrismaTableRepository): table_name = "litellm_dailyguardrailmetrics" +class DailyGuardrailUsageUnitsRepository(PrismaTableRepository): + table_name = "litellm_dailyguardrailusageunits" + + class PolicyAttachmentRepository(PrismaTableRepository): table_name = "litellm_policyattachmenttable" diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 272fbabf807..9e68a5de9e5 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -39,7 +39,7 @@ from pydantic import ( field_serializer, field_validator, ) -from typing_extensions import Required, TypedDict +from typing_extensions import ReadOnly, Required, TypedDict from litellm._logging import verbose_logger from litellm._uuid import uuid @@ -3007,6 +3007,11 @@ class StandardLoggingGuardrailInformation(TypedDict, total=False): surface it as a queryable span attribute without parsing the raw guardrail_response blob.""" + guardrail_usage: ReadOnly[Mapping[str, int] | None] + """Provider-reported billable usage counters for this invocation, keyed by the + provider's counter name (e.g. Bedrock's ``contentPolicyUnits``). Kept as a + sibling of guardrail_response so spend-log prompt redaction never drops it.""" + class EvalVerdict(TypedDict, total=False): criterion_name: str @@ -3050,6 +3055,7 @@ class GuardrailTracingDetail(TypedDict, total=False): risk_score: float | None violation_categories: list[str] | None guardrail_action: str | None + guardrail_usage: ReadOnly[Mapping[str, int] | None] StandardLoggingPayloadStatus = Literal["success", "failure"] diff --git a/schema.prisma b/schema.prisma index 71345d2ccde..d3c277278ff 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1069,6 +1069,22 @@ model LiteLLM_DailyGuardrailMetrics { @@index([guardrail_id]) } +// Daily guardrail billable usage units (one row per guardrail/day/team/key/unit type) +model LiteLLM_DailyGuardrailUsageUnits { + guardrail_id String + date String // YYYY-MM-DD + team_id String // empty string when the request had no team + api_key String // hashed virtual key; empty string when unknown + usage_unit String // provider counter name, e.g. Bedrock's contentPolicyUnits + units BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@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) model LiteLLM_DailyPolicyMetrics { policy_id String diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index f4f4003d5ee..53921e7e74a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -5077,3 +5077,26 @@ async def test_apply_guardrail_failure_logs_a_dict_not_a_bare_string(): logged = mock_log.call_args.kwargs["guardrail_json_response"] assert isinstance(logged, dict), f"expected a dict, got {type(logged).__name__}" assert "error" in logged + + +def test_build_tracing_detail_surfaces_usage_counters(): + """LIT-5650: the billable usage block Bedrock returns per ApplyGuardrail call must + land on the tracing detail as guardrail_usage so it reaches spend logs as a + sibling of guardrail_response (which default redaction replaces wholesale).""" + guardrail = BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT") + + detail = guardrail._build_tracing_detail( + { + "action": "GUARDRAIL_INTERVENED", + "usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 2, "wordPolicyUnits": 0, "oddball": "not-an-int"}, + } + ) + + assert detail["guardrail_usage"] == {"topicPolicyUnits": 1, "contentPolicyUnits": 2, "wordPolicyUnits": 0} + + +def test_build_tracing_detail_omits_guardrail_usage_when_bedrock_reports_none(): + guardrail = BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT") + + assert "guardrail_usage" not in guardrail._build_tracing_detail({"action": "NONE"}) + assert "guardrail_usage" not in guardrail._build_tracing_detail({"action": "NONE", "usage": {}}) diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py index bf7b1b3b238..23c7554c286 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -79,18 +79,38 @@ def _metric(guardrail_id: str, date: str = "2026-04-25", requests: int = 10, pas return m +def _units_row( + guardrail_id: str, + date: str = "2026-04-25", + team_id: str = "", + api_key: str = "", + usage_unit: str = "contentPolicyUnits", + units: int = 1, +) -> Any: + r = MagicMock() + r.guardrail_id = guardrail_id + r.date = date + r.team_id = team_id + r.api_key = api_key + r.usage_unit = usage_unit + r.units = units + return r + + def _prisma( *, find_many=None, find_unique=None, metrics=None, index_find_many=None, + units=None, ) -> MagicMock: client = MagicMock() db = client.db db.litellm_guardrailstable.find_many = AsyncMock(return_value=find_many or []) db.litellm_guardrailstable.find_unique = AsyncMock(return_value=find_unique) db.litellm_dailyguardrailmetrics.find_many = AsyncMock(return_value=metrics or []) + db.litellm_dailyguardrailusageunits.find_many = AsyncMock(return_value=units or []) db.litellm_spendlogguardrailindex.find_many = AsyncMock(return_value=index_find_many or []) db.litellm_spendlogguardrailindex.count = AsyncMock(return_value=0) db.litellm_spendlogs.find_many = AsyncMock(return_value=[]) @@ -215,6 +235,62 @@ async def test_overview_excludes_db_sourced_in_memory_entry(): assert "stale" not in ids +@pytest.mark.asyncio +async def test_overview_reports_usage_units_per_row_and_total(): + """LIT-5650: billable units must surface per guardrail row (matched by + logical name like the daily metrics) and as a response-level total.""" + prisma = _prisma( + find_many=[], + metrics=[_metric("yaml-pii", requests=4, passed=3, blocked=1)], + units=[ + _units_row("yaml-pii", usage_unit="topicPolicyUnits", units=4), + _units_row("yaml-pii", usage_unit="contentPolicyUnits", units=3), + _units_row("yaml-pii", team_id="team-a", usage_unit="contentPolicyUnits", units=2), + _units_row("other-guard", usage_unit="topicPolicyUnits", units=7), + ], + ) + handler = _config_handler(_yaml_guardrail(guardrail_id="yaml-uuid", name="yaml-pii")) + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_overview(start_date=START, end_date=END, user_api_key_dict=ADMIN) + 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} + + +@pytest.mark.asyncio +async def test_detail_breaks_units_down_by_day_team_and_key(): + prisma = _prisma( + find_unique=None, + units=[ + _units_row("yaml-pii", date="2026-04-25", team_id="team-a", api_key="hash-1", units=2), + _units_row("yaml-pii", date="2026-04-25", team_id="", api_key="hash-2", units=1), + _units_row( + "yaml-pii", date="2026-04-24", team_id="team-a", api_key="hash-1", usage_unit="topicPolicyUnits" + ), + ], + ) + handler = _config_handler(_yaml_guardrail()) + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_detail( + guardrail_id="yaml-1", start_date=START, end_date=END, user_api_key_dict=ADMIN + ) + assert resp.usage_units == {"contentPolicyUnits": 3, "topicPolicyUnits": 1} + assert [p.model_dump() for p in resp.usage_units_daily] == [ + {"date": "2026-04-24", "units": {"topicPolicyUnits": 1}}, + {"date": "2026-04-25", "units": {"contentPolicyUnits": 3}}, + ] + assert resp.usage_units_by_team == { + "team-a": {"contentPolicyUnits": 2, "topicPolicyUnits": 1}, + "": {"contentPolicyUnits": 1}, + } + assert resp.usage_units_by_key == { + "hash-1": {"contentPolicyUnits": 2, "topicPolicyUnits": 1}, + "hash-2": {"contentPolicyUnits": 1}, + } + + # ---- logs ------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py new file mode 100644 index 00000000000..7adc1dd314d --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py @@ -0,0 +1,104 @@ +import json +from datetime import datetime, timezone +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.guardrails.usage_tracking import process_spend_logs_guardrail_usage + + +def _prisma() -> MagicMock: + client = MagicMock() + db = client.db + db.litellm_dailyguardrailmetrics.upsert = AsyncMock() + db.litellm_dailyguardrailusageunits.upsert = AsyncMock() + db.litellm_spendlogguardrailindex.create_many = AsyncMock() + return client + + +def _payload( + request_id: str, + *, + team_id: str | None = "team-a", + api_key: str = "hashed-key-1", + usage: dict[str, Any] | None = None, + guardrail_status: str = "success", +) -> dict[str, Any]: + entry: dict[str, Any] = { + "guardrail_id": "bedrock-guard", + "guardrail_status": guardrail_status, + } + if usage is not None: + entry["guardrail_usage"] = usage + return { + "request_id": request_id, + "startTime": datetime(2026, 8, 17, 12, 0, tzinfo=timezone.utc), + "team_id": team_id, + "api_key": api_key, + "metadata": json.dumps({"guardrail_information": [entry]}), + } + + +def _units_upserts(prisma: MagicMock) -> dict[tuple, int]: + calls = prisma.db.litellm_dailyguardrailusageunits.upsert.call_args_list + out: dict[tuple, int] = {} + for c in calls: + where = c.kwargs["where"]["guardrail_id_date_team_id_api_key_usage_unit"] + create = c.kwargs["data"]["create"] + assert create["units"] == c.kwargs["data"]["update"]["units"]["increment"] + assert {k: create[k] for k in where} == where + out[tuple(where[k] for k in ("guardrail_id", "date", "team_id", "api_key", "usage_unit"))] = create["units"] + return out + + +@pytest.mark.asyncio +async def test_usage_units_rolled_up_by_guardrail_team_key_and_date(): + """ + LIT-5650: billable units must aggregate per (guardrail, date, team, key, + counter): same-key payloads sum into one upsert, a team-less payload gets + its own empty-string-team row, and blocked invocations (which Bedrock + still bills for) count exactly like passed ones. + """ + prisma = _prisma() + logs = [ + _payload("r1", usage={"topicPolicyUnits": 1, "contentPolicyUnits": 1}), + _payload( + "r2", + usage={"topicPolicyUnits": 1, "contentPolicyUnits": 2}, + guardrail_status="guardrail_intervened", + ), + _payload("r3", team_id=None, api_key="hashed-key-2", 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"): 2, + ("bedrock-guard", "2026-08-17", "team-a", "hashed-key-1", "contentPolicyUnits"): 3, + ("bedrock-guard", "2026-08-17", "", "hashed-key-2", "topicPolicyUnits"): 1, + } + + +@pytest.mark.asyncio +async def test_zero_and_non_int_usage_counters_are_skipped(): + prisma = _prisma() + logs = [ + _payload( + "r1", + usage={ + "topicPolicyUnits": 1, + "wordPolicyUnits": 0, + "contentPolicyImageUnits": 0, + "oddball": "not-an-int", + "boolish": True, + }, + ), + _payload("r2", usage=None), + ] + + await process_spend_logs_guardrail_usage(prisma, logs) + + assert _units_upserts(prisma) == { + ("bedrock-guard", "2026-08-17", "team-a", "hashed-key-1", "topicPolicyUnits"): 1, + } diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 2b162774aea..bca8210baa6 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -85,6 +85,73 @@ async def test_async_post_call_failure_hook(): assert metadata["original_key"] == "original_value" +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_carries_guardrail_info_from_litellm_metadata(): + """ + LIT-5650 regression: on a pre_call guardrail block the unified guardrail + layer seeds request_data["litellm_metadata"], so the guardrail hook writes + standard_logging_guardrail_information there, while the failure spend log + is serialized from request_data["metadata"]. Blocked invocations still + consume provider usage units, so the info must be carried over or the + failure row logs guardrail_information: null. + """ + logger = _ProxyDBLogger() + guardrail_info = [ + { + "guardrail_name": "bedrock-guard", + "guardrail_status": "guardrail_intervened", + "guardrail_usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 1}, + } + ] + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {"original_key": "original_value"}, + "litellm_metadata": {"standard_logging_guardrail_information": guardrail_info}, + "proxy_server_request": {"request_id": "test_request_id"}, + } + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database: + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("Violated guardrail policy"), + user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"), + ) + + metadata = mock_update_database.call_args[1]["kwargs"]["litellm_params"]["metadata"] + assert metadata["standard_logging_guardrail_information"] == guardrail_info + assert metadata["original_key"] == "original_value" + + +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_does_not_clobber_guardrail_info_in_metadata(): + logger = _ProxyDBLogger() + metadata_bucket_info = [{"guardrail_name": "from-metadata-bucket"}] + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {"standard_logging_guardrail_information": metadata_bucket_info}, + "litellm_metadata": {"standard_logging_guardrail_information": [{"guardrail_name": "from-litellm-bucket"}]}, + "proxy_server_request": {"request_id": "test_request_id"}, + } + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database: + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("Test exception"), + user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key"), + ) + + metadata = mock_update_database.call_args[1]["kwargs"]["litellm_params"]["metadata"] + assert metadata["standard_logging_guardrail_information"] == metadata_bucket_info + + @pytest.mark.asyncio async def test_async_post_call_failure_hook_non_llm_route(): # Setup diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 0f6ac3f9b4f..33d835652cd 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -1565,6 +1565,38 @@ def test_sanitize_guardrail_information_redacts_prompt_fields_when_flag_false( } +@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +def test_sanitize_guardrail_information_preserves_guardrail_usage_when_flag_false( + mock_should_store, +): + """ + LIT-5650 regression: provider-reported billable usage counters live in + guardrail_usage, a sibling of guardrail_response, precisely so the + default spend-log redaction cannot drop them. The response blob (which + also embeds a usage copy) must still be redacted wholesale. + """ + mock_should_store.return_value = False + guardrail_info = [ + { + "guardrail_name": "bedrock-guard", + "guardrail_status": "guardrail_intervened", + "guardrail_response": { + "action": "GUARDRAIL_INTERVENED", + "outputs": [{"text": "Sorry, the model cannot answer this question."}], + "usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 1}, + }, + "guardrail_usage": {"topicPolicyUnits": 1, "contentPolicyUnits": 1, "wordPolicyUnits": 0}, + } + ] + + result = _sanitize_guardrail_information_for_spend_logs(guardrail_info) + + assert result is not None + entry = result[0] + assert entry["guardrail_response"] == REDACTED_BY_LITELM_STRING + assert entry["guardrail_usage"] == {"topicPolicyUnits": 1, "contentPolicyUnits": 1, "wordPolicyUnits": 0} + + @patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_passthrough_when_flag_true( mock_should_store, diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 6d70a6aa5f4..a37bb7ef49f 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22909 + "limit": 22906 }, "LIT002": { - "limit": 26898 + "limit": 26896 }, "LIT003": { "limit": 269