diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 2ac9a3b7c1c..942bc821eb1 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -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 diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index d8457cf9c86..e06901e7ace 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -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, diff --git a/tests/proxy_unit_tests/test_guardrail_usage_config.py b/tests/proxy_unit_tests/test_guardrail_usage_config.py new file mode 100644 index 00000000000..e462341f734 --- /dev/null +++ b/tests/proxy_unit_tests/test_guardrail_usage_config.py @@ -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} + ]