fix(guardrails): show config guardrails in usage details

This commit is contained in:
Saicharan Ramineni 2026-06-21 02:26:25 -04:00
parent 84c1414aef
commit 862e4e90f7
3 changed files with 198 additions and 21 deletions

View file

@ -181,6 +181,7 @@ jobs:
- test-group: guardrails-hooks
test-path: >-
tests/proxy_unit_tests/test_proxy_setting_guardrails.py
tests/proxy_unit_tests/test_guardrail_usage_config.py
tests/proxy_unit_tests/test_banned_keyword_list.py
tests/proxy_unit_tests/test_unit_test_proxy_hooks.py
workers: 4

View file

@ -154,6 +154,66 @@ def _get_guardrail_attrs(g: Any) -> tuple[Any, str]:
return gid, (name or gid or "")
def _get_guardrail_dict_field(g: Any, field_name: str) -> Any:
return getattr(g, field_name, None) or (
g.get(field_name) if isinstance(g, dict) else None
)
def _get_guardrail_litellm_params(g: Any) -> dict[str, Any]:
litellm_params = _get_guardrail_dict_field(g, "litellm_params")
if isinstance(litellm_params, dict):
return litellm_params
if hasattr(litellm_params, "model_dump"):
return litellm_params.model_dump(exclude_none=True)
return {}
def _get_guardrail_info(g: Any) -> dict[str, Any]:
guardrail_info = _get_guardrail_dict_field(g, "guardrail_info")
return guardrail_info if isinstance(guardrail_info, dict) else {}
def _get_config_loaded_guardrails() -> list[Any]:
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
config_guardrails: list[Any] = []
for guardrail in IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails():
guardrail_id = guardrail.get("guardrail_id")
if (
guardrail_id is not None
and IN_MEMORY_GUARDRAIL_HANDLER.get_source(guardrail_id) != "config"
):
continue
config_guardrails.append(guardrail)
return config_guardrails
def _find_config_loaded_guardrail(guardrail_id_or_name: str) -> Optional[Any]:
for guardrail in _get_config_loaded_guardrails():
gid, display_name = _get_guardrail_attrs(guardrail)
if guardrail_id_or_name in (gid, display_name):
return guardrail
return None
def _merge_config_loaded_guardrails(db_guardrails: Any) -> list[Any]:
guardrails = list(db_guardrails)
seen_keys: set[str] = set()
for guardrail in guardrails:
gid, display_name = _get_guardrail_attrs(guardrail)
seen_keys.update(str(k) for k in (gid, display_name) if k)
for guardrail in _get_config_loaded_guardrails():
gid, display_name = _get_guardrail_attrs(guardrail)
lookup_keys = [str(k) for k in (gid, display_name) if k]
if any(k in seen_keys for k in lookup_keys):
continue
guardrails.append(guardrail)
seen_keys.update(lookup_keys)
return guardrails
def _guardrail_overview_rows(
guardrails: Any,
agg: Dict[str, Dict[str, Any]],
@ -173,13 +233,9 @@ def _guardrail_overview_rows(
break
req, blocked = a["requests"], a["blocked"]
fail_rate = (100.0 * blocked / req) if req else 0.0
litellm_params = (
(g.litellm_params or {}) if isinstance(g.litellm_params, dict) else {}
)
litellm_params = _get_guardrail_litellm_params(g)
provider = str(litellm_params.get("guardrail", "Unknown"))
guardrail_info = (
(g.guardrail_info or {}) if isinstance(g.guardrail_info, dict) else {}
)
guardrail_info = _get_guardrail_info(g)
gtype = str(guardrail_info.get("type", "Guardrail"))
prev_fail = 0.0
for k in lookup_keys:
@ -189,7 +245,7 @@ def _guardrail_overview_rows(
trend = _trend_from_comparison(fail_rate, prev_fail)
rows.append(
UsageOverviewRow(
id=gid,
id=str(gid or display_name),
name=display_name or str(gid),
type=gtype,
provider=provider,
@ -280,7 +336,9 @@ async def guardrails_usage_overview(
try:
# Guardrails from DB
guardrails = await GuardrailsRepository(prisma_client).table.find_many()
guardrails = _merge_config_loaded_guardrails(
await GuardrailsRepository(prisma_client).table.find_many()
)
# Daily metrics in range
metrics = await DailyGuardrailMetricsRepository(prisma_client).table.find_many(
@ -347,9 +405,11 @@ async def guardrails_usage_detail(
where={"guardrail_id": guardrail_id}
)
if not guardrail:
from fastapi import HTTPException
guardrail = _find_config_loaded_guardrail(guardrail_id)
if not guardrail:
from fastapi import HTTPException
raise HTTPException(status_code=404, detail="Guardrail not found")
raise HTTPException(status_code=404, detail="Guardrail not found")
# Metrics are keyed by logical name (from spend log metadata), not UUID
logical_id = getattr(guardrail, "guardrail_name", None) or (
@ -391,17 +451,9 @@ async def guardrails_usage_detail(
{"date": d, "passed": v["passed"], "blocked": v["blocked"], "score": None}
for d, v in sorted(ts_by_date.items())
]
_litellm_params = getattr(guardrail, "litellm_params", None) or (
guardrail.get("litellm_params") if isinstance(guardrail, dict) else None
)
litellm_params = _litellm_params if isinstance(_litellm_params, dict) else {}
_guardrail_info = getattr(guardrail, "guardrail_info", None) or (
guardrail.get("guardrail_info") if isinstance(guardrail, dict) else None
)
guardrail_info = _guardrail_info if isinstance(_guardrail_info, dict) else {}
_guardrail_name = getattr(guardrail, "guardrail_name", None) or (
guardrail.get("guardrail_name") if isinstance(guardrail, dict) else None
)
litellm_params = _get_guardrail_litellm_params(guardrail)
guardrail_info = _get_guardrail_info(guardrail)
_guardrail_name = _get_guardrail_dict_field(guardrail, "guardrail_name")
return UsageDetailResponse(
guardrail_id=guardrail_id,

View file

@ -0,0 +1,124 @@
import sys
from types import SimpleNamespace
import pytest
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.guardrails import usage_endpoints
class _Table:
def __init__(self, *, rows=None, unique=None):
self.rows = rows or []
self.unique = unique
async def find_many(self, *args, **kwargs):
return self.rows
async def find_unique(self, *args, **kwargs):
return self.unique
def _repo(table):
return lambda prisma_client: SimpleNamespace(table=table)
def _metric(
*,
guardrail_id: str,
date: str = "2026-06-20",
requests_evaluated: int = 10,
passed_count: int = 8,
blocked_count: int = 2,
flagged_count: int = 0,
):
return SimpleNamespace(
guardrail_id=guardrail_id,
date=date,
requests_evaluated=requests_evaluated,
passed_count=passed_count,
blocked_count=blocked_count,
flagged_count=flagged_count,
)
def _config_guardrail():
return {
"guardrail_id": "config-presidio-id",
"guardrail_name": "presidio-pii-mask-test",
"litellm_params": {
"guardrail": "presidio",
"mode": ["pre_call", "post_call"],
"default_on": False,
},
"guardrail_info": {"description": "Config loaded PII guardrail"},
}
def _patch_common(monkeypatch, *, metrics, detail_unique=None):
monkeypatch.setitem(
sys.modules,
"litellm.proxy.proxy_server",
SimpleNamespace(prisma_client=object()),
)
monkeypatch.setattr(
usage_endpoints,
"GuardrailsRepository",
_repo(_Table(rows=[], unique=detail_unique)),
)
monkeypatch.setattr(
usage_endpoints,
"DailyGuardrailMetricsRepository",
_repo(_Table(rows=metrics)),
)
monkeypatch.setattr(
usage_endpoints,
"_get_config_loaded_guardrails",
lambda: [_config_guardrail()],
)
@pytest.mark.asyncio
async def test_usage_overview_uses_provider_for_config_guardrails(monkeypatch):
_patch_common(
monkeypatch,
metrics=[_metric(guardrail_id="presidio-pii-mask-test")],
)
response = await usage_endpoints.guardrails_usage_overview(
start_date="2026-06-19",
end_date="2026-06-20",
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
)
assert len(response.rows) == 1
row = response.rows[0]
assert row.id == "config-presidio-id"
assert row.name == "presidio-pii-mask-test"
assert row.provider == "presidio"
assert row.requestsEvaluated == 10
assert row.failRate == 20.0
@pytest.mark.asyncio
async def test_usage_detail_resolves_config_guardrail_by_name(monkeypatch):
_patch_common(
monkeypatch,
metrics=[_metric(guardrail_id="presidio-pii-mask-test")],
detail_unique=None,
)
response = await usage_endpoints.guardrails_usage_detail(
guardrail_id="presidio-pii-mask-test",
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert response.guardrail_id == "presidio-pii-mask-test"
assert response.guardrail_name == "presidio-pii-mask-test"
assert response.provider == "presidio"
assert response.description == "Config loaded PII guardrail"
assert response.requestsEvaluated == 10
assert response.failRate == 20.0
assert response.time_series == [
{"date": "2026-06-20", "passed": 8, "blocked": 2, "score": None}
]