mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[Fix] List Guardrails - Show config.yaml guardrails on litellm ui (#10959)
* fix: listing guardrails defined on litellm config * fix: list guardrails on litellm config * fix: list guardrails on litellm config * test: list guardrails on litellm config * fix: linting * Update litellm/proxy/guardrails/guardrail_endpoints.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * fix: GuardrailInfoLiteLLMParamsResponse --------- Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
parent
d099092bc1
commit
7d8ed6f362
6 changed files with 242 additions and 13 deletions
|
|
@ -144,6 +144,7 @@ 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:
|
||||
|
|
@ -155,6 +156,7 @@ async def list_guardrails_v2():
|
|||
)
|
||||
|
||||
guardrail_configs: List[GuardrailInfoResponse] = []
|
||||
seen_guardrail_ids = set()
|
||||
for guardrail in guardrails:
|
||||
guardrail_configs.append(
|
||||
GuardrailInfoResponse(
|
||||
|
|
@ -164,8 +166,26 @@ async def list_guardrails_v2():
|
|||
guardrail_info=guardrail.get("guardrail_info"),
|
||||
created_at=guardrail.get("created_at"),
|
||||
updated_at=guardrail.get("updated_at"),
|
||||
guardrail_definition_location="db",
|
||||
)
|
||||
)
|
||||
seen_guardrail_ids.add(guardrail.get("guardrail_id"))
|
||||
|
||||
# get guardrails initialized on litellm config.yaml
|
||||
in_memory_guardrails = IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails()
|
||||
for guardrail in in_memory_guardrails:
|
||||
# only add guardrails that are not in DB guardrail list already
|
||||
if guardrail.get("guardrail_id") not in seen_guardrail_ids:
|
||||
guardrail_configs.append(
|
||||
GuardrailInfoResponse(
|
||||
guardrail_id=guardrail.get("guardrail_id"),
|
||||
guardrail_name=guardrail.get("guardrail_name"),
|
||||
litellm_params=dict(guardrail.get("litellm_params") or {}),
|
||||
guardrail_info=dict(guardrail.get("guardrail_info") or {}),
|
||||
guardrail_definition_location="config",
|
||||
)
|
||||
)
|
||||
seen_guardrail_ids.add(guardrail.get("guardrail_id"))
|
||||
|
||||
return ListGuardrailsResponse(guardrails=guardrail_configs)
|
||||
except Exception as e:
|
||||
|
|
@ -291,6 +311,7 @@ async def get_guardrail(guardrail_id: str):
|
|||
result = await GUARDRAIL_REGISTRY.get_guardrail_by_id_from_db(
|
||||
guardrail_id=guardrail_id, prisma_client=prisma_client
|
||||
)
|
||||
|
||||
if result is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Guardrail with ID {guardrail_id} not found"
|
||||
|
|
@ -605,6 +626,7 @@ async def get_guardrail_info(guardrail_id: str):
|
|||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
|
|
@ -614,6 +636,11 @@ async def get_guardrail_info(guardrail_id: str):
|
|||
result = await GUARDRAIL_REGISTRY.get_guardrail_by_id_from_db(
|
||||
guardrail_id=guardrail_id, prisma_client=prisma_client
|
||||
)
|
||||
if result is None:
|
||||
result = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id(
|
||||
guardrail_id=guardrail_id
|
||||
)
|
||||
|
||||
if result is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Guardrail with ID {guardrail_id} not found"
|
||||
|
|
@ -622,8 +649,8 @@ async def get_guardrail_info(guardrail_id: str):
|
|||
return GuardrailInfoResponse(
|
||||
guardrail_id=result.get("guardrail_id"),
|
||||
guardrail_name=result.get("guardrail_name"),
|
||||
litellm_params=result.get("litellm_params"),
|
||||
guardrail_info=result.get("guardrail_info"),
|
||||
litellm_params=dict(result.get("litellm_params") or {}),
|
||||
guardrail_info=dict(result.get("guardrail_info") or {}),
|
||||
created_at=result.get("created_at"),
|
||||
updated_at=result.get("updated_at"),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -235,6 +235,7 @@ class InMemoryGuardrailHandler:
|
|||
Returns a Guardrail object if the guardrail is initialized successfully
|
||||
"""
|
||||
guardrail_id = guardrail.get("guardrail_id") or str(uuid.uuid4())
|
||||
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"
|
||||
|
|
@ -273,7 +274,7 @@ class InMemoryGuardrailHandler:
|
|||
if initializer:
|
||||
custom_guardrail_callback = initializer(litellm_params, guardrail)
|
||||
elif isinstance(guardrail_type, str) and "." in guardrail_type:
|
||||
self.initialize_custom_guardrail(
|
||||
custom_guardrail_callback = self.initialize_custom_guardrail(
|
||||
guardrail=guardrail,
|
||||
guardrail_type=guardrail_type,
|
||||
litellm_params=litellm_params,
|
||||
|
|
@ -288,8 +289,6 @@ class InMemoryGuardrailHandler:
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
guardrail_id = parsed_guardrail.get("guardrail_id") or str(uuid.uuid4())
|
||||
|
||||
# store references to the guardrail in memory
|
||||
self.IN_MEMORY_GUARDRAILS[guardrail_id] = parsed_guardrail
|
||||
self.guardrail_id_to_custom_guardrail[guardrail_id] = custom_guardrail_callback
|
||||
|
|
@ -302,7 +301,7 @@ class InMemoryGuardrailHandler:
|
|||
guardrail_type: str,
|
||||
litellm_params: LitellmParams,
|
||||
config_file_path: Optional[str] = None,
|
||||
) -> None:
|
||||
) -> Optional[CustomGuardrail]:
|
||||
"""
|
||||
Initialize a Custom Guardrail from a python file
|
||||
|
||||
|
|
@ -348,6 +347,8 @@ class InMemoryGuardrailHandler:
|
|||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_guardrail_callback) # type: ignore
|
||||
|
||||
return _guardrail_callback
|
||||
|
||||
def update_in_memory_guardrail(
|
||||
self, guardrail_id: str, guardrail: Guardrail
|
||||
) -> None:
|
||||
|
|
@ -376,6 +377,18 @@ class InMemoryGuardrailHandler:
|
|||
"""
|
||||
self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None)
|
||||
|
||||
def list_in_memory_guardrails(self) -> List[Guardrail]:
|
||||
"""
|
||||
List all guardrails in memory
|
||||
"""
|
||||
return list(self.IN_MEMORY_GUARDRAILS.values())
|
||||
|
||||
def get_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]:
|
||||
"""
|
||||
Get a guardrail by its ID from memory
|
||||
"""
|
||||
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
|
||||
|
||||
|
||||
########################################################
|
||||
# In Memory Guardrail Handler for LiteLLM Proxy
|
||||
|
|
|
|||
|
|
@ -7,4 +7,20 @@ model_list:
|
|||
|
||||
|
||||
general_settings:
|
||||
store_prompts_in_spend_logs: true
|
||||
store_prompts_in_spend_logs: true
|
||||
|
||||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "custom-pre-guard"
|
||||
litellm_params:
|
||||
guardrail: custom_guardrail.myCustomGuardrail # 👈 Key change
|
||||
mode: "pre_call" # runs async_pre_call_hook
|
||||
- guardrail_name: "custom-during-guard"
|
||||
litellm_params:
|
||||
guardrail: custom_guardrail.myCustomGuardrail
|
||||
mode: "during_call" # runs async_moderation_hook
|
||||
- guardrail_name: "custom-post-guard"
|
||||
litellm_params:
|
||||
guardrail: custom_guardrail.myCustomGuardrail
|
||||
mode: "post_call" # runs async_post_call_success_hook
|
||||
|
|
@ -390,12 +390,12 @@ class DynamicGuardrailParams(TypedDict):
|
|||
extra_body: Dict[str, Any]
|
||||
|
||||
|
||||
class GuardrailLiteLLMParamsResponse(BaseModel):
|
||||
class GuardrailInfoLiteLLMParamsResponse(BaseModel):
|
||||
"""The returned LiteLLM Params object for /guardrails/list"""
|
||||
|
||||
guardrail: str
|
||||
mode: Union[str, List[str]]
|
||||
default_on: bool = Field(default=False)
|
||||
default_on: Optional[bool] = False
|
||||
pii_entities_config: Optional[Dict[PiiEntityType, PiiAction]] = None
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
|
|
@ -409,10 +409,11 @@ class GuardrailLiteLLMParamsResponse(BaseModel):
|
|||
class GuardrailInfoResponse(BaseModel):
|
||||
guardrail_id: Optional[str] = None
|
||||
guardrail_name: str
|
||||
litellm_params: GuardrailLiteLLMParamsResponse
|
||||
guardrail_info: Optional[Dict]
|
||||
litellm_params: Optional[GuardrailInfoLiteLLMParamsResponse] = None
|
||||
guardrail_info: Optional[Dict] = None
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
guardrail_definition_location: Literal["config", "db"] = "config"
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
|
|
|||
|
|
@ -93,11 +93,11 @@ def test_guardrail_list_of_event_hooks():
|
|||
|
||||
|
||||
def test_guardrail_info_response():
|
||||
from litellm.types.guardrails import GuardrailInfoResponse, LitellmParams, GuardrailLiteLLMParamsResponse
|
||||
from litellm.types.guardrails import GuardrailInfoResponse, LitellmParams, GuardrailInfoLiteLLMParamsResponse
|
||||
|
||||
guardrail_info = GuardrailInfoResponse(
|
||||
guardrail_name="aporia-pre-guard",
|
||||
litellm_params=GuardrailLiteLLMParamsResponse(
|
||||
litellm_params=GuardrailInfoLiteLLMParamsResponse(
|
||||
guardrail="aporia",
|
||||
mode="pre_call",
|
||||
),
|
||||
|
|
|
|||
172
tests/litellm/proxy/guardrails/test_guardrail_endpoints.py
Normal file
172
tests/litellm/proxy/guardrails/test_guardrail_endpoints.py
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from typing import Dict, List, Optional
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import (
|
||||
get_guardrail_info,
|
||||
list_guardrails_v2,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_registry import (
|
||||
IN_MEMORY_GUARDRAIL_HANDLER,
|
||||
InMemoryGuardrailHandler,
|
||||
)
|
||||
from litellm.types.guardrails import (
|
||||
GuardrailInfoLiteLLMParamsResponse,
|
||||
GuardrailInfoResponse,
|
||||
)
|
||||
|
||||
# Mock data for testing
|
||||
MOCK_DB_GUARDRAIL = {
|
||||
"guardrail_id": "test-db-guardrail",
|
||||
"guardrail_name": "Test DB Guardrail",
|
||||
"litellm_params": {
|
||||
"guardrail": "test.guardrail",
|
||||
"mode": "pre_call",
|
||||
},
|
||||
"guardrail_info": {"description": "Test guardrail from DB"},
|
||||
"created_at": datetime.now(),
|
||||
"updated_at": datetime.now(),
|
||||
}
|
||||
|
||||
MOCK_CONFIG_GUARDRAIL = {
|
||||
"guardrail_id": "test-config-guardrail",
|
||||
"guardrail_name": "Test Config Guardrail",
|
||||
"litellm_params": {
|
||||
"guardrail": "custom_guardrail.myCustomGuardrail",
|
||||
"mode": "during_call",
|
||||
},
|
||||
"guardrail_info": {"description": "Test guardrail from config"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_prisma_client(mocker):
|
||||
"""Mock Prisma client for testing"""
|
||||
mock_client = mocker.Mock()
|
||||
# Create async mocks for the database methods
|
||||
mock_client.db = mocker.Mock()
|
||||
mock_client.db.litellm_guardrailstable = mocker.Mock()
|
||||
mock_client.db.litellm_guardrailstable.find_many = AsyncMock(
|
||||
return_value=[MOCK_DB_GUARDRAIL]
|
||||
)
|
||||
mock_client.db.litellm_guardrailstable.find_unique = AsyncMock(
|
||||
return_value=MOCK_DB_GUARDRAIL
|
||||
)
|
||||
return mock_client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_in_memory_handler(mocker):
|
||||
"""Mock InMemoryGuardrailHandler for testing"""
|
||||
mock_handler = mocker.Mock(spec=InMemoryGuardrailHandler)
|
||||
mock_handler.list_in_memory_guardrails.return_value = [MOCK_CONFIG_GUARDRAIL]
|
||||
mock_handler.get_guardrail_by_id.return_value = MOCK_CONFIG_GUARDRAIL
|
||||
return mock_handler
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_guardrails_v2_with_db_and_config(
|
||||
mocker, mock_prisma_client, mock_in_memory_handler
|
||||
):
|
||||
"""Test listing guardrails from both DB and config"""
|
||||
# Mock the prisma client
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
# Mock the in-memory handler
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
response = await list_guardrails_v2()
|
||||
|
||||
assert len(response.guardrails) == 2
|
||||
|
||||
# Check DB guardrail
|
||||
db_guardrail = next(
|
||||
g for g in response.guardrails if g.guardrail_id == "test-db-guardrail"
|
||||
)
|
||||
assert db_guardrail.guardrail_name == "Test DB Guardrail"
|
||||
assert db_guardrail.guardrail_definition_location == "db"
|
||||
assert isinstance(db_guardrail.litellm_params, GuardrailInfoLiteLLMParamsResponse)
|
||||
|
||||
# Check config guardrail
|
||||
config_guardrail = next(
|
||||
g for g in response.guardrails if g.guardrail_id == "test-config-guardrail"
|
||||
)
|
||||
assert config_guardrail.guardrail_name == "Test Config Guardrail"
|
||||
assert config_guardrail.guardrail_definition_location == "config"
|
||||
assert isinstance(
|
||||
config_guardrail.litellm_params, GuardrailInfoLiteLLMParamsResponse
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_info_from_db(mocker, mock_prisma_client):
|
||||
"""Test getting guardrail info from DB"""
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
response = await get_guardrail_info("test-db-guardrail")
|
||||
|
||||
assert response.guardrail_id == "test-db-guardrail"
|
||||
assert response.guardrail_name == "Test DB Guardrail"
|
||||
assert isinstance(response.litellm_params, GuardrailInfoLiteLLMParamsResponse)
|
||||
assert response.guardrail_info == {"description": "Test guardrail from DB"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_info_from_config(
|
||||
mocker, mock_prisma_client, mock_in_memory_handler
|
||||
):
|
||||
"""Test getting guardrail info from config when not found in DB"""
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
# Mock DB to return None
|
||||
mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
response = await get_guardrail_info("test-config-guardrail")
|
||||
|
||||
assert response.guardrail_id == "test-config-guardrail"
|
||||
assert response.guardrail_name == "Test Config Guardrail"
|
||||
assert isinstance(response.litellm_params, GuardrailInfoLiteLLMParamsResponse)
|
||||
assert response.guardrail_info == {"description": "Test guardrail from config"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_info_not_found(
|
||||
mocker, mock_prisma_client, mock_in_memory_handler
|
||||
):
|
||||
"""Test getting guardrail info when not found in either DB or config"""
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
# Mock both DB and in-memory handler to return None
|
||||
mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
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
|
||||
assert "not found" in str(exc_info.value.detail)
|
||||
Loading…
Add table
Reference in a new issue