mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(guardrails): show config guardrails in usage details
This commit is contained in:
parent
84c1414aef
commit
862e4e90f7
3 changed files with 198 additions and 21 deletions
1
.github/workflows/test-unit-proxy-db.yml
vendored
1
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
124
tests/proxy_unit_tests/test_guardrail_usage_config.py
Normal file
124
tests/proxy_unit_tests/test_guardrail_usage_config.py
Normal 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}
|
||||
]
|
||||
Loading…
Add table
Reference in a new issue