From 4ff9c751b56c9e9f7327f767354bd7d9b15ef70e Mon Sep 17 00:00:00 2001 From: Saicharan Ramineni <84414237+GodlyDonuts@users.noreply.github.com> Date: Tue, 23 Jun 2026 09:17:54 -0400 Subject: [PATCH] fix(guardrails): include config guardrails in usage details (#30911) * fix(guardrails): show config guardrails in usage details * test(guardrails): cover config guardrail helper branches * fix(guardrails): persist guardrail_info and resolve config guardrails in logs Address review findings on the config-guardrail usage work: - initialize_guardrail dropped guardrail_info when building the in-memory Guardrail, so description and type were always empty for config guardrails in production; persist the field. - guardrails_usage_logs only resolved a logical name for DB-backed guardrails, so logs for a config guardrail queried by UUID were always empty; fall back to the in-memory list like the detail endpoint does. - _get_config_loaded_guardrails now expresses an explicit allow (source == "config") instead of a double-negative skip, and _get_guardrail_dict_field dispatches on type so a falsy-but-valid value (e.g. {}) is not dropped. Tests exercise the real handler path (initialize_guardrail through _get_config_loaded_guardrails) rather than mocking it, so they fail if guardrail_info is dropped or the logs fallback is removed. --- .github/workflows/test-unit-proxy-db.yml | 1 + .../proxy/guardrails/guardrail_registry.py | 1 + litellm/proxy/guardrails/usage_endpoints.py | 90 +++-- .../test_guardrail_usage_config.py | 309 ++++++++++++++++++ .../guardrails/test_guardrail_registry.py | 26 ++ 5 files changed, 408 insertions(+), 19 deletions(-) create mode 100644 tests/proxy_unit_tests/test_guardrail_usage_config.py 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/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 8962073fe7a..b00e8565123 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -478,6 +478,7 @@ class InMemoryGuardrailHandler: guardrail_id=guardrail.get("guardrail_id"), guardrail_name=guardrail["guardrail_name"], litellm_params=litellm_params, + guardrail_info=guardrail.get("guardrail_info"), ) # store references to the guardrail in memory diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index e03bdbb95d2..be6c141f920 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -144,6 +144,62 @@ 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: + if isinstance(g, dict): + return g.get(field_name) + return getattr(g, field_name, 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 + + return [ + guardrail + for guardrail in IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails() + if IN_MEMORY_GUARDRAIL_HANDLER.get_source(guardrail.get("guardrail_id") or "") + == "config" + ] + + +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]], @@ -163,9 +219,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: @@ -175,7 +231,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, @@ -264,7 +320,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( @@ -323,9 +381,11 @@ async def guardrails_usage_detail( guardrail = await GuardrailsRepository(prisma_client).table.find_unique(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 ( @@ -367,17 +427,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, @@ -550,9 +602,9 @@ async def guardrails_usage_logs( if guardrail_id: guardrail = await GuardrailsRepository(prisma_client).table.find_unique( where={"guardrail_id": guardrail_id} - ) + ) or _find_config_loaded_guardrail(guardrail_id) if guardrail: - logical_name = getattr(guardrail, "guardrail_name", None) + logical_name = _get_guardrail_dict_field(guardrail, "guardrail_name") if logical_name and logical_name not in effective_guardrail_ids: effective_guardrail_ids.append(logical_name) 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..559b3721773 --- /dev/null +++ b/tests/proxy_unit_tests/test_guardrail_usage_config.py @@ -0,0 +1,309 @@ +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()], + ) + + +def _register_real_config_guardrail(monkeypatch, *, guardrail_id, name, guardrail_info): + from litellm.proxy.guardrails import guardrail_registry + from litellm.types.guardrails import Guardrail, LitellmParams + + handler = guardrail_registry.InMemoryGuardrailHandler() + handler.initialize_guardrail( + guardrail=Guardrail( + guardrail_id=guardrail_id, + guardrail_name=name, + litellm_params=LitellmParams( + guardrail="bedrock", mode="pre_call", default_on=False + ), + guardrail_info=guardrail_info, + ), + source="config", + ) + monkeypatch.setattr(guardrail_registry, "IN_MEMORY_GUARDRAIL_HANDLER", handler) + + +@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} + ] + + +@pytest.mark.asyncio +async def test_usage_detail_raises_404_when_guardrail_absent(monkeypatch): + from fastapi import HTTPException + + _patch_common(monkeypatch, metrics=[], detail_unique=None) + monkeypatch.setattr(usage_endpoints, "_get_config_loaded_guardrails", lambda: []) + + with pytest.raises(HTTPException) as exc_info: + await usage_endpoints.guardrails_usage_detail( + guardrail_id="missing-guardrail", + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert exc_info.value.status_code == 404 + + +def test_get_config_loaded_guardrails_filters_by_source(monkeypatch): + from litellm.proxy.guardrails import guardrail_registry + + handler = guardrail_registry.InMemoryGuardrailHandler() + handler.IN_MEMORY_GUARDRAILS = { + "cfg-1": {"guardrail_id": "cfg-1", "guardrail_name": "config-one"}, + "db-1": {"guardrail_id": "db-1", "guardrail_name": "db-one"}, + "no-source": {"guardrail_id": "no-source", "guardrail_name": "no-source"}, + } + handler._sources = {"cfg-1": "config", "db-1": "db"} + monkeypatch.setattr(guardrail_registry, "IN_MEMORY_GUARDRAIL_HANDLER", handler) + + result = usage_endpoints._get_config_loaded_guardrails() + + returned_ids = {g["guardrail_id"] for g in result} + assert "cfg-1" in returned_ids + assert "db-1" not in returned_ids + assert "no-source" not in returned_ids + + +def test_merge_config_loaded_guardrails_dedupes_and_appends(monkeypatch): + db_guardrails = [ + SimpleNamespace(guardrail_id="shared-id", guardrail_name="shared-name") + ] + monkeypatch.setattr( + usage_endpoints, + "_get_config_loaded_guardrails", + lambda: [ + {"guardrail_id": "shared-id", "guardrail_name": "shared-name"}, + {"guardrail_id": "cfg-only", "guardrail_name": "cfg-only-name"}, + ], + ) + + merged = usage_endpoints._merge_config_loaded_guardrails(db_guardrails) + + ids = [usage_endpoints._get_guardrail_attrs(g)[0] for g in merged] + assert ids == ["shared-id", "cfg-only"] + + +def test_find_config_loaded_guardrail_returns_none_when_absent(monkeypatch): + monkeypatch.setattr( + usage_endpoints, + "_get_config_loaded_guardrails", + lambda: [{"guardrail_id": "cfg-1", "guardrail_name": "config-one"}], + ) + + assert usage_endpoints._find_config_loaded_guardrail("does-not-exist") is None + + +def test_get_guardrail_litellm_params_handles_model_dump_and_missing(): + pydantic_like = SimpleNamespace( + litellm_params=SimpleNamespace( + model_dump=lambda exclude_none: {"guardrail": "presidio"} + ) + ) + assert usage_endpoints._get_guardrail_litellm_params(pydantic_like) == { + "guardrail": "presidio" + } + + assert usage_endpoints._get_guardrail_litellm_params(SimpleNamespace()) == {} + + +@pytest.mark.asyncio +async def test_usage_detail_surfaces_guardrail_info_through_real_handler(monkeypatch): + """End-to-end through the real in-memory handler (no mocking of + _get_config_loaded_guardrails): description, type and provider must come + from the config guardrail's persisted guardrail_info / litellm_params.""" + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + SimpleNamespace(prisma_client=object()), + ) + monkeypatch.setattr( + usage_endpoints, "GuardrailsRepository", _repo(_Table(rows=[], unique=None)) + ) + monkeypatch.setattr( + usage_endpoints, + "DailyGuardrailMetricsRepository", + _repo(_Table(rows=[_metric(guardrail_id="bedrock-guard")])), + ) + _register_real_config_guardrail( + monkeypatch, + guardrail_id="bedrock-guard-uuid", + name="bedrock-guard", + guardrail_info={"description": "blocks PII", "type": "bedrock"}, + ) + + response = await usage_endpoints.guardrails_usage_detail( + guardrail_id="bedrock-guard", + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert response.guardrail_name == "bedrock-guard" + assert response.description == "blocks PII" + assert response.type == "bedrock" + assert response.provider == "bedrock" + + +@pytest.mark.asyncio +async def test_usage_logs_resolves_config_guardrail_by_uuid(monkeypatch): + """A config guardrail queried by UUID has no DB row, so the logs endpoint + must fall back to the in-memory list and add its logical name to the + SpendLogs index filter; otherwise the logs tab is always empty.""" + captured = {} + + class _IndexTable: + async def find_many(self, *args, **kwargs): + captured["where"] = kwargs.get("where") + return [] + + async def count(self, *args, **kwargs): + return 0 + + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + SimpleNamespace(prisma_client=object()), + ) + monkeypatch.setattr( + usage_endpoints, "GuardrailsRepository", _repo(_Table(rows=[], unique=None)) + ) + monkeypatch.setattr( + usage_endpoints, + "SpendLogGuardrailIndexRepository", + _repo(_IndexTable()), + ) + _register_real_config_guardrail( + monkeypatch, + guardrail_id="bedrock-guard-uuid", + name="bedrock-guard", + guardrail_info={"description": "blocks PII"}, + ) + + response = await usage_endpoints.guardrails_usage_logs( + guardrail_id="bedrock-guard-uuid", + policy_id=None, + page=1, + page_size=50, + action=None, + start_date=None, + end_date=None, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert response.total == 0 + assert captured["where"]["guardrail_id"] == { + "in": ["bedrock-guard-uuid", "bedrock-guard"] + } diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 0ef9ad857f9..c0ea3aa291c 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -166,6 +166,32 @@ def test_initialize_guardrail_early_return_updates_source_marker(): assert handler.get_source("collide") == "db" +def test_initialize_guardrail_preserves_guardrail_info(): + """ + guardrail_info from the config must survive into IN_MEMORY_GUARDRAILS so + downstream consumers (e.g. the usage dashboard) can read the description and + type. Dropping it makes those fields silently empty for config guardrails. + """ + handler = InMemoryGuardrailHandler() + g = Guardrail( + guardrail_id="with-info", + guardrail_name="bedrock", + litellm_params=LitellmParams( + guardrail="bedrock", mode="pre_call", default_on=False + ), + guardrail_info={"description": "blocks PII", "type": "bedrock"}, + ) + + handler.initialize_guardrail(guardrail=g, source="config") + + stored = handler.get_guardrail_by_id("with-info") + assert stored is not None + assert stored.get("guardrail_info") == { + "description": "blocks PII", + "type": "bedrock", + } + + def test_sync_guardrail_from_db_marks_source_db_when_unchanged(): """ sync_guardrail_from_db must enforce source='db' even when params are