Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 10 additions & 9 deletions litellm/proxy/guardrails/guardrail_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@ def _get_guardrails_list_response(
)
guardrail_configs.append(
GuardrailInfoResponse(
guardrail_id=guardrail.get("guardrail_id"),
guardrail_name=guardrail.get("guardrail_name"),
litellm_params=masked_params,
guardrail_info=guardrail.get("guardrail_info"),
Expand Down Expand Up @@ -178,13 +179,14 @@ async def list_guardrails_v2(
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
from litellm.proxy.proxy_server import prisma_client

if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")

is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN

try:
guardrails = await GUARDRAIL_REGISTRY.get_all_guardrails_from_db(prisma_client=prisma_client)
guardrails = (
await GUARDRAIL_REGISTRY.get_all_guardrails_from_db(prisma_client=prisma_client)
if prisma_client is not None
else []
)

excluded_guardrail_ids: set = set()
if not is_admin:
Expand Down Expand Up @@ -1228,13 +1230,12 @@ async def get_guardrail_info(guardrail_id: str):
from litellm.proxy.proxy_server import prisma_client
from litellm.types.guardrails import GUARDRAIL_DEFINITION_LOCATION

if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")

try:
guardrail_definition_location: GUARDRAIL_DEFINITION_LOCATION = GUARDRAIL_DEFINITION_LOCATION.DB
result = await GUARDRAIL_REGISTRY.get_guardrail_by_id_from_db(
guardrail_id=guardrail_id, prisma_client=prisma_client
result = (
await GUARDRAIL_REGISTRY.get_guardrail_by_id_from_db(guardrail_id=guardrail_id, prisma_client=prisma_client)
if prisma_client is not None
else None
)
if result is None:
in_memory = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id(guardrail_id=guardrail_id)
Expand Down
10 changes: 9 additions & 1 deletion litellm/proxy/guardrails/guardrail_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import importlib
import os
from datetime import datetime, timezone
from itertools import chain, count
from typing import Any, Dict, List, Literal, Optional, Set, Type, cast

from pydantic import ValidationError
Expand Down Expand Up @@ -65,6 +66,8 @@
SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: initialize_llm_as_a_judge,
}

CONFIG_GUARDRAIL_ID_NAMESPACE = uuid.UUID("625f63f4-935a-50e5-98b5-fbe77babc74a")

guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = {
SupportedGuardrailIntegrations.BEDROCK.value: BedrockGuardrail,
SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail,
Expand Down Expand Up @@ -407,6 +410,11 @@ def __init__(self):
and never deleted by reconciliation.
"""

def _stable_guardrail_id(self, guardrail_name: str) -> str:
seeds = chain((guardrail_name,), (f"{guardrail_name}:{occurrence}" for occurrence in count(1)))
candidate_ids = (str(uuid.uuid5(CONFIG_GUARDRAIL_ID_NAMESPACE, seed.encode("utf-8"))) for seed in seeds)
return next(candidate_id for candidate_id in candidate_ids if candidate_id not in self.IN_MEMORY_GUARDRAILS)

def initialize_guardrail(
self,
guardrail: Guardrail,
Expand All @@ -419,7 +427,7 @@ def initialize_guardrail(

Returns a Guardrail object if the guardrail is initialized successfully
"""
guardrail_id = guardrail.get("guardrail_id") or str(uuid.uuid4())
guardrail_id = guardrail.get("guardrail_id") or self._stable_guardrail_id(guardrail["guardrail_name"])
guardrail["guardrail_id"] = guardrail_id
if guardrail_id in self.IN_MEMORY_GUARDRAILS:
verbose_proxy_logger.debug("guardrail_id already exists in IN_MEMORY_GUARDRAILS")
Expand Down
108 changes: 108 additions & 0 deletions tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -400,6 +400,114 @@ async def test_get_guardrail_info_not_found(
assert "not found" in str(exc_info.value.detail)


@pytest.mark.asyncio
async def test_list_guardrails_v2_without_prisma_returns_config_guardrails(
mocker, mock_in_memory_handler
):
"""
A proxy without a DB must still list config-defined guardrails instead of
raising 500 'Prisma client not initialized'.
"""
mocker.patch("litellm.proxy.proxy_server.prisma_client", None)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)

response = await list_guardrails_v2(user_api_key_dict=MOCK_ADMIN_USER)

assert len(response.guardrails) == 1
config_guardrail = response.guardrails[0]
assert config_guardrail.guardrail_id == "test-config-guardrail"
assert config_guardrail.guardrail_name == "Test Config Guardrail"
assert config_guardrail.guardrail_definition_location == "config"


@pytest.mark.asyncio
async def test_list_guardrails_v2_without_prisma_non_admin_sees_unrestricted_config_guardrails(
mocker, mock_in_memory_handler
):
"""
A non-admin caller on a no-DB proxy must see config guardrails that carry
no team_id restriction; the team lookup must not blow up without a DB.
"""
mocker.patch("litellm.proxy.proxy_server.prisma_client", None)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)

non_admin_auth = UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="internal-user-1"
)
response = await list_guardrails_v2(user_api_key_dict=non_admin_auth)

assert [g.guardrail_id for g in response.guardrails] == ["test-config-guardrail"]


@pytest.mark.asyncio
async def test_get_guardrail_info_without_prisma_returns_config_guardrail(
mocker, mock_in_memory_handler
):
"""
The info endpoint must serve config-defined guardrails from the in-memory
registry when no DB is attached instead of raising 500.
"""
mocker.patch("litellm.proxy.proxy_server.prisma_client", None)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)

response = await get_guardrail_info("test-config-guardrail")

assert response.guardrail_id == "test-config-guardrail"
assert response.guardrail_name == "Test Config Guardrail"
assert response.guardrail_definition_location == "config"


@pytest.mark.asyncio
async def test_get_guardrail_info_without_prisma_404s_unknown_id(
mocker, mock_in_memory_handler
):
mocker.patch("litellm.proxy.proxy_server.prisma_client", None)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)
mock_in_memory_handler.get_guardrail_by_id.return_value = None

with pytest.raises(HTTPException) as exc_info:
await get_guardrail_info("non-existent-guardrail")

assert exc_info.value.status_code == 404


def test_get_guardrails_list_response_includes_guardrail_id():
"""
The v1 list response is the UI's fallback when v2 fails; without ids every
row click requests /guardrails/undefined/info.
"""
from litellm.proxy.guardrails.guardrail_endpoints import (
_get_guardrails_list_response,
)

response = _get_guardrails_list_response(
[
{
"guardrail_id": "stable-config-id",
"guardrail_name": "tooling",
"litellm_params": {
"guardrail": "litellm_content_filter",
"mode": "pre_call",
},
}
]
)

assert response.guardrails[0].guardrail_id == "stable-config-id"


def test_get_provider_specific_params():
"""Test getting provider-specific parameters"""
from litellm.proxy.guardrails.guardrail_endpoints import _get_fields_from_model
Expand Down
89 changes: 89 additions & 0 deletions tests/test_litellm/proxy/guardrails/test_guardrail_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,95 @@ def _initializer(litellm_params, guardrail):
registry_module.guardrail_initializer_registry.pop("parallel_default_test", None)


def _register_noop_initializer(guardrail_type: str):
from litellm.proxy.guardrails import guardrail_registry as registry_module

def _initializer(litellm_params, guardrail):
return CustomGuardrail(
guardrail_name=guardrail["guardrail_name"],
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
)

registry_module.guardrail_initializer_registry[guardrail_type] = _initializer
return registry_module


def _config_guardrail(name: str, guardrail_type: str, guardrail_id=None) -> dict:
guardrail = {
"guardrail_name": name,
"litellm_params": {"guardrail": guardrail_type, "mode": "pre_call"},
}
if guardrail_id is not None:
guardrail["guardrail_id"] = guardrail_id
return guardrail


def test_config_guardrail_id_is_stable_across_boots():
"""
Config guardrails used to get a fresh uuid4 per process, so ids from a
previous boot (or another replica) 404'd on /guardrails/{id}/info even
though the guardrail was alive.
"""
registry_module = _register_noop_initializer("stable_id_test")
try:
first_boot = InMemoryGuardrailHandler().initialize_guardrail(
guardrail=_config_guardrail("tooling", "stable_id_test")
)
second_boot = InMemoryGuardrailHandler().initialize_guardrail(
guardrail=_config_guardrail("tooling", "stable_id_test")
)

assert first_boot["guardrail_id"] == second_boot["guardrail_id"]
finally:
registry_module.guardrail_initializer_registry.pop("stable_id_test", None)


def test_explicit_config_guardrail_id_wins_over_derived_id():
registry_module = _register_noop_initializer("explicit_id_test")
try:
result = InMemoryGuardrailHandler().initialize_guardrail(
guardrail=_config_guardrail(
"tooling", "explicit_id_test", guardrail_id="my-explicit-id"
)
)

assert result["guardrail_id"] == "my-explicit-id"
finally:
registry_module.guardrail_initializer_registry.pop("explicit_id_test", None)


def test_duplicate_config_guardrail_names_get_distinct_stable_ids():
"""
Duplicate guardrail_name entries are legitimate (load balancing across
deployments); each occurrence must keep its own id, stable across boots.
"""
registry_module = _register_noop_initializer("dup_name_test")
try:
handler = InMemoryGuardrailHandler()
first = handler.initialize_guardrail(
guardrail=_config_guardrail("dup", "dup_name_test")
)
second = handler.initialize_guardrail(
guardrail=_config_guardrail("dup", "dup_name_test")
)

rebooted_handler = InMemoryGuardrailHandler()
rebooted_first = rebooted_handler.initialize_guardrail(
guardrail=_config_guardrail("dup", "dup_name_test")
)
rebooted_second = rebooted_handler.initialize_guardrail(
guardrail=_config_guardrail("dup", "dup_name_test")
)

assert first["guardrail_id"] != second["guardrail_id"]
assert first["guardrail_id"] == rebooted_first["guardrail_id"]
assert second["guardrail_id"] == rebooted_second["guardrail_id"]
assert len(handler.IN_MEMORY_GUARDRAILS) == 2
finally:
registry_module.guardrail_initializer_registry.pop("dup_name_test", None)


def test_update_in_memory_guardrail():
handler = InMemoryGuardrailHandler()
handler.guardrail_id_to_custom_guardrail["123"] = CustomGuardrail(
Expand Down
Loading