diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e1b8d8b5dbc..54b8ec9fa9c 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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 diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 626c8b7f297..258601ff5a0 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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 diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index f76f86301d6..318f8c407ce 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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"] \ No newline at end of file + callbacks: ["langfuse_otel"] \ No newline at end of file diff --git a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py index 48027ec1c69..0e3b4b670f6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py @@ -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, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b8c17934b04..008b151afce 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 \ No newline at end of file + return prisma_client diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/test_guardrails_ai.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/test_guardrails_ai.py new file mode 100644 index 00000000000..a3e1b0cd4c8 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/test_guardrails_ai.py @@ -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"