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:
Krish Dholakia 2025-07-10 21:41:02 -07:00 • committed by GitHub
parent 7c373140a8
commit 9d9562d0e8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 362 additions and 30 deletions

View file

@ -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

View file

@ -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

View file

@ -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"]

View file

@ -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,

View file

@ -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

View file

@ -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"