mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Guardrails AI - pre-call + logging only guardrail (pii detection/competitor names) support (#12506)
* fix(guardrails_ai.py): initial commit adding pre-call hook support for guardrails ai enables running user input through guardrails ai - if set * feat(guardrails_ai.py): working pre call guardrail enables pii detection to work via guardrails ai * feat(guardrails_ai.py): support logging hook enables masking input via guardrails ai on logging integrations * test(test_guardrails_ai.py): add unit test for new input processing function
This commit is contained in:
parent
7c373140a8
commit
9d9562d0e8
6 changed files with 362 additions and 30 deletions
|
|
@ -1076,29 +1076,35 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
"""
|
||||
from litellm.types.llms.base import HiddenParams
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
|
||||
callbacks = self.get_combined_callback_list(
|
||||
dynamic_success_callbacks=self.dynamic_success_callbacks,
|
||||
global_callbacks=litellm.success_callback,
|
||||
)
|
||||
post_mcp_tool_call_response_obj: MCPPostCallResponseObject = MCPPostCallResponseObject(
|
||||
mcp_tool_call_response=response_obj,
|
||||
hidden_params=HiddenParams()
|
||||
post_mcp_tool_call_response_obj: MCPPostCallResponseObject = (
|
||||
MCPPostCallResponseObject(
|
||||
mcp_tool_call_response=response_obj, hidden_params=HiddenParams()
|
||||
)
|
||||
)
|
||||
for callback in callbacks:
|
||||
try:
|
||||
if isinstance(callback, CustomLogger):
|
||||
response: Optional[MCPPostCallResponseObject] = await callback.async_post_mcp_tool_call_hook(
|
||||
kwargs=kwargs,
|
||||
response_obj=post_mcp_tool_call_response_obj,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
response: Optional[MCPPostCallResponseObject] = (
|
||||
await callback.async_post_mcp_tool_call_hook(
|
||||
kwargs=kwargs,
|
||||
response_obj=post_mcp_tool_call_response_obj,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
######################################################################
|
||||
# if any of the callbacks modify the response, use the modified response
|
||||
# current implementation returns the first modified response
|
||||
######################################################################
|
||||
if response is not None:
|
||||
response_obj = self._parse_post_mcp_call_hook_response(response=response)
|
||||
response_obj = self._parse_post_mcp_call_hook_response(
|
||||
response=response
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format(
|
||||
|
|
@ -1107,7 +1113,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
return response_obj
|
||||
|
||||
def _parse_post_mcp_call_hook_response(self, response: Optional[MCPPostCallResponseObject]) -> Any:
|
||||
def _parse_post_mcp_call_hook_response(
|
||||
self, response: Optional[MCPPostCallResponseObject]
|
||||
) -> Any:
|
||||
"""
|
||||
Parse the response from the post_mcp_tool_call_hook
|
||||
|
||||
|
|
@ -1404,7 +1412,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
and result is not None
|
||||
and self.stream is not True
|
||||
):
|
||||
if self._is_recognized_call_type_for_logging(logging_result=logging_result):
|
||||
if self._is_recognized_call_type_for_logging(
|
||||
logging_result=logging_result
|
||||
):
|
||||
## HIDDEN PARAMS ##
|
||||
hidden_params = getattr(logging_result, "_hidden_params", {})
|
||||
if hidden_params:
|
||||
|
|
@ -1500,7 +1510,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
return start_time, end_time, result
|
||||
except Exception as e:
|
||||
raise Exception(f"[Non-Blocking] LiteLLM.Success_Call Error: {str(e)}")
|
||||
|
||||
|
||||
def _is_recognized_call_type_for_logging(
|
||||
self,
|
||||
logging_result: Any,
|
||||
|
|
@ -1523,9 +1533,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
or isinstance(logging_result, OpenAIFileObject)
|
||||
or isinstance(logging_result, LiteLLMRealtimeStreamLoggingObject)
|
||||
or isinstance(logging_result, OpenAIModerationResponse)
|
||||
or (
|
||||
self.call_type == CallTypes.call_mcp_tool.value
|
||||
)
|
||||
or (self.call_type == CallTypes.call_mcp_tool.value)
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -520,11 +520,13 @@ def unpack_defs(schema: dict, defs: dict) -> None:
|
|||
|
||||
# Use iterative approach with queue to avoid recursion
|
||||
# Each item in queue is (node, parent_container, key/index, active_defs, seen_ids)
|
||||
queue: deque[tuple[Any, Union[dict, list, None], Union[str, int, None], dict, set]] = deque([(schema, None, None, root_defs, set())])
|
||||
|
||||
queue: deque[
|
||||
tuple[Any, Union[dict, list, None], Union[str, int, None], dict, set]
|
||||
] = deque([(schema, None, None, root_defs, set())])
|
||||
|
||||
while queue:
|
||||
node, parent, key, active_defs, seen = queue.popleft()
|
||||
|
||||
|
||||
# Avoid infinite loops on self-referential schemas
|
||||
if id(node) in seen:
|
||||
continue
|
||||
|
|
@ -560,7 +562,7 @@ def unpack_defs(schema: dict, defs: dict) -> None:
|
|||
schema.clear()
|
||||
schema.update(resolved)
|
||||
resolved = schema
|
||||
|
||||
|
||||
# Add resolved node to queue for further processing
|
||||
queue.append((resolved, parent, key, child_defs, seen))
|
||||
continue
|
||||
|
|
@ -750,3 +752,73 @@ def migrate_file_to_image_url(
|
|||
if format and isinstance(image_url_object["image_url"], dict):
|
||||
image_url_object["image_url"]["format"] = format
|
||||
return image_url_object
|
||||
|
||||
|
||||
def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]:
|
||||
"""
|
||||
Get the last consecutive block of messages from the user.
|
||||
|
||||
Example:
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
{"role": "assistant", "content": "I'm good, thank you!"},
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"},
|
||||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
# Iterate from the end to find the last consecutive block of user messages
|
||||
user_messages = []
|
||||
for message in reversed(messages):
|
||||
if message.get("role") == "user":
|
||||
user_messages.append(message)
|
||||
else:
|
||||
# Stop when we hit a non-user message
|
||||
break
|
||||
|
||||
if not user_messages:
|
||||
return None
|
||||
|
||||
# Reverse to get the messages in chronological order
|
||||
user_messages.reverse()
|
||||
|
||||
user_prompt = ""
|
||||
for message in user_messages:
|
||||
text_content = convert_content_list_to_str(message)
|
||||
user_prompt += text_content + "\n"
|
||||
|
||||
result = user_prompt.strip()
|
||||
return result if result else None
|
||||
|
||||
|
||||
def set_last_user_message(
|
||||
messages: List[AllMessageValues], content: str
|
||||
) -> List[AllMessageValues]:
|
||||
"""
|
||||
Set the last user message
|
||||
|
||||
1. remove all the last consecutive user messages (FROM THE END)
|
||||
2. add the new message
|
||||
"""
|
||||
idx_to_remove = []
|
||||
for idx, message in enumerate(reversed(messages)):
|
||||
if message.get("role") == "user":
|
||||
idx_to_remove.append(idx)
|
||||
else:
|
||||
# Stop when we hit a non-user message
|
||||
break
|
||||
if idx_to_remove:
|
||||
messages = [
|
||||
message
|
||||
for idx, message in enumerate(reversed(messages))
|
||||
if idx not in idx_to_remove
|
||||
]
|
||||
messages.reverse()
|
||||
messages.append({"role": "user", "content": content})
|
||||
return messages
|
||||
|
|
|
|||
|
|
@ -1,13 +1,17 @@
|
|||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: openai/gpt-3.5-turbo
|
||||
model: gpt-3.5-turbo
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
- model_name: langfuse-model
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "guardrails_ai-guard"
|
||||
litellm_params:
|
||||
model: langfuse/langfuse-model
|
||||
prompt_id: test-chat-prompt
|
||||
prompt_version: 4
|
||||
guardrail: guardrails_ai
|
||||
guard_name: "pii_detect" # 👈 Guardrail AI guard name
|
||||
mode: "logging_only"
|
||||
api_base: os.environ/GUARDRAILS_AI_API_BASE # 👈 Guardrails AI API Base. Defaults to "http://0.0.0.0:8000"
|
||||
default_on: true
|
||||
|
||||
litellm_settings:
|
||||
public_model_groups: ["gpt-3.5-turbo"]
|
||||
callbacks: ["langfuse_otel"]
|
||||
|
|
@ -7,7 +7,17 @@
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Optional, Type, TypedDict
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
TypedDict,
|
||||
Union,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -34,6 +44,19 @@ class GuardrailsAIResponse(TypedDict):
|
|||
validationPassed: bool
|
||||
|
||||
|
||||
class InferenceData(TypedDict):
|
||||
name: str
|
||||
shape: List[int]
|
||||
data: List
|
||||
datatype: str
|
||||
|
||||
|
||||
class GuardrailsAIResponsePreCall(TypedDict):
|
||||
modelname: str
|
||||
modelversion: str
|
||||
outputs: List[InferenceData]
|
||||
|
||||
|
||||
class GuardrailsAI(CustomGuardrail):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -51,7 +74,11 @@ class GuardrailsAI(CustomGuardrail):
|
|||
)
|
||||
self.guardrails_ai_guard_name = guard_name
|
||||
self.optional_params = kwargs
|
||||
supported_event_hooks = [GuardrailEventHooks.post_call]
|
||||
supported_event_hooks = [
|
||||
GuardrailEventHooks.post_call,
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.logging_only,
|
||||
]
|
||||
super().__init__(supported_event_hooks=supported_event_hooks, **kwargs)
|
||||
|
||||
async def make_guardrails_ai_api_request(self, llm_output: str, request_data: dict):
|
||||
|
|
@ -85,6 +112,98 @@ class GuardrailsAI(CustomGuardrail):
|
|||
)
|
||||
return _json_response
|
||||
|
||||
async def make_guardrails_ai_api_request_pre_call_request(
|
||||
self, text_input: str, request_data: dict
|
||||
) -> str:
|
||||
from httpx import URL
|
||||
|
||||
data = {
|
||||
"inputs": [
|
||||
{
|
||||
"name": "text",
|
||||
"shape": [1],
|
||||
"data": [text_input],
|
||||
"datatype": "BYTES", # not sure what this should be, but Guardrail's response sets BYTES for text response - https://github.com/guardrails-ai/detect_pii/blob/e4719a95a26f6caacb78d46ebb4768317032bee5/app.py#L40C31-L40C36
|
||||
}
|
||||
]
|
||||
}
|
||||
_json_data = json.dumps(data)
|
||||
response = await litellm.module_level_aclient.post(
|
||||
url=str(
|
||||
URL(self.guardrails_ai_api_base).join(
|
||||
f"guards/{self.guardrails_ai_guard_name}/validate"
|
||||
)
|
||||
),
|
||||
data=_json_data,
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
)
|
||||
verbose_proxy_logger.debug("guardrails_ai response: %s", response)
|
||||
if response.status_code == 400:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Violated guardrail policy",
|
||||
"guardrails_ai_response": response.json(),
|
||||
},
|
||||
)
|
||||
|
||||
_json_response = GuardrailsAIResponsePreCall(**response.json()) # type: ignore
|
||||
response = _json_response.get("outputs", [])[0].get("data", [])[0]
|
||||
return response
|
||||
|
||||
async def process_input(self, data: dict, call_type: str) -> dict:
|
||||
if call_type == "acompletion" or call_type == "completion":
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_last_user_message,
|
||||
set_last_user_message,
|
||||
)
|
||||
|
||||
if "messages" not in data: # invalid request
|
||||
return data
|
||||
|
||||
text = get_last_user_message(data["messages"])
|
||||
if text is None:
|
||||
return data
|
||||
updated_text = await self.make_guardrails_ai_api_request_pre_call_request(
|
||||
text_input=text, request_data=data
|
||||
)
|
||||
data["messages"] = set_last_user_message(data["messages"], updated_text)
|
||||
|
||||
return data
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: litellm.DualCache,
|
||||
data: dict,
|
||||
call_type: Literal[
|
||||
"completion",
|
||||
"text_completion",
|
||||
"embeddings",
|
||||
"image_generation",
|
||||
"moderation",
|
||||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
],
|
||||
) -> Optional[
|
||||
Union[Exception, str, dict]
|
||||
]: # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm
|
||||
|
||||
return await self.process_input(data=data, call_type=call_type)
|
||||
|
||||
async def async_logging_hook(
|
||||
self, kwargs: dict, result: Any, call_type: str
|
||||
) -> Tuple[dict, Any]:
|
||||
|
||||
if call_type == "acompletion" or call_type == "completion":
|
||||
kwargs = await self.process_input(data=kwargs, call_type=call_type)
|
||||
|
||||
return kwargs, result
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -363,7 +363,11 @@ class ProxyLogging:
|
|||
if self.alerting is not None and "slack" in self.alerting:
|
||||
# NOTE: ENSURE we only add callbacks when alerting is on
|
||||
# We should NOT add callbacks when alerting is off
|
||||
if "daily_reports" in self.alert_types or "outage_alerts" in self.alert_types or "region_outage_alerts" in self.alert_types:
|
||||
if (
|
||||
"daily_reports" in self.alert_types
|
||||
or "outage_alerts" in self.alert_types
|
||||
or "region_outage_alerts" in self.alert_types
|
||||
):
|
||||
litellm.logging_callback_manager.add_litellm_callback(self.slack_alerting_instance) # type: ignore
|
||||
litellm.logging_callback_manager.add_litellm_success_callback(
|
||||
self.slack_alerting_instance.response_taking_too_long_callback
|
||||
|
|
@ -1770,7 +1774,6 @@ class PrismaClient:
|
|||
WHERE v.token = '{token}'
|
||||
"""
|
||||
|
||||
print_verbose("sql_query being made={}".format(sql_query))
|
||||
response = await self.db.query_first(query=sql_query)
|
||||
|
||||
if response is not None:
|
||||
|
|
@ -3180,4 +3183,4 @@ def get_prisma_client_or_throw(message: str):
|
|||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": message},
|
||||
)
|
||||
return prisma_client
|
||||
return prisma_client
|
||||
|
|
|
|||
|
|
@ -0,0 +1,126 @@
|
|||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.guardrails_ai.guardrails_ai import (
|
||||
GuardrailsAI,
|
||||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrails_ai_process_input():
|
||||
"""Test the process_input method of GuardrailsAI with various scenarios"""
|
||||
|
||||
# Initialize the GuardrailsAI instance
|
||||
guardrails_ai_guardrail = GuardrailsAI(
|
||||
guardrail_name="test_guard",
|
||||
api_base="http://test.example.com",
|
||||
guard_name="gibberish-guard",
|
||||
)
|
||||
|
||||
# Test case 1: Valid completion call with messages
|
||||
with patch.object(
|
||||
guardrails_ai_guardrail,
|
||||
"make_guardrails_ai_api_request_pre_call_request",
|
||||
return_value="processed text",
|
||||
) as mock_api_request:
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
]
|
||||
}
|
||||
|
||||
result = await guardrails_ai_guardrail.process_input(data, "completion")
|
||||
|
||||
# Verify the API was called with the user message
|
||||
mock_api_request.assert_called_once_with(
|
||||
text_input="Hello, how are you?", request_data=data
|
||||
)
|
||||
|
||||
# Verify the message was updated
|
||||
assert result["messages"][1]["content"] == "processed text"
|
||||
# System message should remain unchanged
|
||||
assert result["messages"][0]["content"] == "You are a helpful assistant"
|
||||
|
||||
# Test case 2: Valid acompletion call with messages
|
||||
with patch.object(
|
||||
guardrails_ai_guardrail,
|
||||
"make_guardrails_ai_api_request_pre_call_request",
|
||||
return_value="async processed text",
|
||||
) as mock_api_request:
|
||||
|
||||
data = {"messages": [{"role": "user", "content": "What is the weather?"}]}
|
||||
|
||||
result = await guardrails_ai_guardrail.process_input(data, "acompletion")
|
||||
|
||||
mock_api_request.assert_called_once_with(
|
||||
text_input="What is the weather?", request_data=data
|
||||
)
|
||||
|
||||
assert result["messages"][0]["content"] == "async processed text"
|
||||
|
||||
# Test case 3: Invalid request without messages
|
||||
data_no_messages = {"model": "gpt-3.5-turbo"}
|
||||
|
||||
result = await guardrails_ai_guardrail.process_input(data_no_messages, "completion")
|
||||
|
||||
# Should return data unchanged
|
||||
assert result == data_no_messages
|
||||
|
||||
# Test case 4: Messages with no user text (get_last_user_message returns None)
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.common_utils.get_last_user_message",
|
||||
return_value=None,
|
||||
):
|
||||
data = {
|
||||
"messages": [{"role": "system", "content": "You are a helpful assistant"}]
|
||||
}
|
||||
|
||||
result = await guardrails_ai_guardrail.process_input(data, "completion")
|
||||
|
||||
# Should return data unchanged when no user message found
|
||||
assert result == data
|
||||
|
||||
# Test case 5: Different call_type that should not be processed
|
||||
data = {"messages": [{"role": "user", "content": "Hello"}]}
|
||||
|
||||
result = await guardrails_ai_guardrail.process_input(data, "embeddings")
|
||||
|
||||
# Should return data unchanged for non-completion call types
|
||||
assert result == data
|
||||
|
||||
# Test case 6: Complex conversation with multiple messages
|
||||
with patch.object(
|
||||
guardrails_ai_guardrail,
|
||||
"make_guardrails_ai_api_request_pre_call_request",
|
||||
return_value="sanitized message",
|
||||
) as mock_api_request:
|
||||
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "First question"},
|
||||
{"role": "assistant", "content": "First answer"},
|
||||
{"role": "user", "content": "Second question"},
|
||||
]
|
||||
}
|
||||
|
||||
result = await guardrails_ai_guardrail.process_input(data, "completion")
|
||||
|
||||
# Should process the last user message
|
||||
mock_api_request.assert_called_once_with(
|
||||
text_input="Second question", request_data=data
|
||||
)
|
||||
|
||||
# Only the last user message should be updated
|
||||
assert result["messages"][0]["content"] == "You are a helpful assistant"
|
||||
assert result["messages"][1]["content"] == "First question"
|
||||
assert result["messages"][2]["content"] == "First answer"
|
||||
assert result["messages"][3]["content"] == "sanitized message"
|
||||
Loading…
Add table
Reference in a new issue