diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index b99ea8f14a0..d6e25c33220 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -509,6 +509,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 e06901e7ace..28f48fcc895 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -155,9 +155,9 @@ def _get_guardrail_attrs(g: Any) -> tuple[Any, str]: 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 - ) + 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]: @@ -177,16 +177,12 @@ def _get_guardrail_info(g: Any) -> dict[str, Any]: 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 + 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]: @@ -636,9 +632,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 index a753f8758be..559b3721773 100644 --- a/tests/proxy_unit_tests/test_guardrail_usage_config.py +++ b/tests/proxy_unit_tests/test_guardrail_usage_config.py @@ -78,6 +78,25 @@ def _patch_common(monkeypatch, *, metrics, detail_unique=None): ) +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( @@ -200,3 +219,91 @@ def test_get_guardrail_litellm_params_handles_model_dump_and_missing(): } 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