From 6d3ce76e1c941ec81c8ea2445d27517fb2789e2d Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Thu, 24 Jul 2025 22:38:03 -0700 Subject: [PATCH] Guardrails - support model-level guardrails (#12968) * fix(custom_guardrail.py): initial logic for model level guardrails * feat(custom_guardrail.py): working pre call guardrails * fix(custom_guardrails.py): check if custom guardrails set before running event hook * test(test_custom_guardrail.py): add unit tests for async pre call deployment hook on custom guardrail * feat(custom_guardrail.py): add post call processing support for guardrails allows model based guardrails to run on the post call event for that model only * fix(utils.py): only run if call type is in enum * test: update unit tests to work --- litellm/integrations/custom_guardrail.py | 106 +++++++++++++++++- litellm/integrations/custom_logger.py | 18 ++- litellm/proxy/_new_secret_config.yaml | 38 ++++--- .../guardrails/guardrail_hooks/presidio.py | 2 + litellm/types/utils.py | 1 + litellm/utils.py | 30 ++++- .../integrations/test_custom_guardrail.py | 84 ++++++++++++++ 7 files changed, 254 insertions(+), 25 deletions(-) create mode 100644 tests/test_litellm/integrations/test_custom_guardrail.py diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index e3376e20b3..5bf84bd8ce 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1,7 +1,8 @@ from datetime import datetime -from typing import Dict, List, Literal, Optional, Type, Union +from typing import Any, Dict, List, Literal, Optional, Type, Union, get_args from litellm._logging import verbose_logger +from litellm.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.types.guardrails import ( DynamicGuardrailParams, @@ -11,7 +12,13 @@ from litellm.types.guardrails import ( PiiEntityType, ) from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel -from litellm.types.utils import StandardLoggingGuardrailInformation +from litellm.types.utils import ( + CallTypes, + LLMResponseTypes, + StandardLoggingGuardrailInformation, +) + +dc = DualCache() class CustomGuardrail(CustomLogger): @@ -108,10 +115,14 @@ class CustomGuardrail(CustomLogger): self, data: dict ) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]: """ - Returns the guardrail(s) to be run from the metadata + Returns the guardrail(s) to be run from the metadata or root """ + if "guardrails" in data: + return data["guardrails"] metadata = data.get("metadata") or {} requested_guardrails = metadata.get("guardrails") or [] + if requested_guardrails: + return requested_guardrails return requested_guardrails def _guardrail_is_in_requested_guardrails( @@ -130,7 +141,94 @@ class CustomGuardrail(CustomLogger): return False - def should_run_guardrail(self, data, event_type: GuardrailEventHooks) -> bool: + async def async_pre_call_deployment_hook( + self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] + ) -> Optional[dict]: + + from litellm.proxy._types import UserAPIKeyAuth + + # should run guardrail + litellm_guardrails = kwargs.get("guardrails") + if litellm_guardrails is None or not isinstance(litellm_guardrails, list): + return kwargs + + if ( + self.should_run_guardrail( + data=kwargs, event_type=GuardrailEventHooks.pre_call + ) + is not True + ): + return kwargs + + # CHECK IF GUARDRAIL REJECTS THE REQUEST + if call_type == CallTypes.completion or call_type == CallTypes.acompletion: + result = await self.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth( + user_id=kwargs.get("user_api_key_user_id"), + team_id=kwargs.get("user_api_key_team_id"), + end_user_id=kwargs.get("user_api_key_end_user_id"), + api_key=kwargs.get("user_api_key_hash"), + request_route=kwargs.get("user_api_key_request_route"), + ), + cache=dc, + data=kwargs, + call_type=call_type.value or "acompletion", # type: ignore + ) + + if result is not None and isinstance(result, dict): + result_messages = result.get("messages") + if result_messages is not None: # update for any pii / masking logic + kwargs["messages"] = result_messages + + return kwargs + + async def async_post_call_success_deployment_hook( + self, + request_data: dict, + response: LLMResponseTypes, + call_type: Optional[CallTypes], + ) -> Optional[LLMResponseTypes]: + """ + Allow modifying / reviewing the response just after it's received from the deployment. + """ + from litellm.proxy._types import UserAPIKeyAuth + + # should run guardrail + litellm_guardrails = request_data.get("guardrails") + if litellm_guardrails is None or not isinstance(litellm_guardrails, list): + return response + + if ( + self.should_run_guardrail( + data=request_data, event_type=GuardrailEventHooks.post_call + ) + is not True + ): + return response + + # CHECK IF GUARDRAIL REJECTS THE REQUEST + result = await self.async_post_call_success_hook( + user_api_key_dict=UserAPIKeyAuth( + user_id=request_data.get("user_api_key_user_id"), + team_id=request_data.get("user_api_key_team_id"), + end_user_id=request_data.get("user_api_key_end_user_id"), + api_key=request_data.get("user_api_key_hash"), + request_route=request_data.get("user_api_key_request_route"), + ), + data=request_data, + response=response, + ) + + if result is None or not isinstance(result, get_args(LLMResponseTypes)): + return response + + return result + + def should_run_guardrail( + self, + data, + event_type: GuardrailEventHooks, + ) -> bool: """ Returns True if the guardrail should be run on the event_type """ diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 0cb293b6b6..08ff197c20 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -172,6 +172,17 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac def pre_call_check(self, deployment: dict) -> Optional[dict]: pass + async def async_post_call_success_deployment_hook( + self, + request_data: dict, + response: LLMResponseTypes, + call_type: Optional[CallTypes], + ) -> Optional[LLMResponseTypes]: + """ + Allow modifying / reviewing the response just after it's received from the deployment. + """ + pass + #### Fallback Events - router/proxy only #### async def log_model_group_rate_limit_error( self, exception: Exception, original_model_group: Optional[str], kwargs: dict @@ -372,18 +383,19 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac except Exception: print_verbose(f"Custom Logger Error - {traceback.format_exc()}") pass - + ######################################################### # MCP TOOL CALL HOOKS ######################################################### - async def async_post_mcp_tool_call_hook(self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time) -> Optional[MCPPostCallResponseObject]: + async def async_post_mcp_tool_call_hook( + self, kwargs, response_obj: MCPPostCallResponseObject, start_time, end_time + ) -> Optional[MCPPostCallResponseObject]: """ This log gets called after the MCP tool call is made. Useful if you want to modiy the standard logging payload after the MCP tool call is made. """ return None - # Useful helpers for custom logger classes diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 017c1d7884..c249da0534 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,21 +1,25 @@ model_list: - - model_name: openai/gpt-4o + - model_name: claude-sonnet-4 + litellm_params: + model: anthropic/claude-sonnet-4-20250514 + api_key: os.environ/ANTHROPIC_API_KEY + api_base: https://api.anthropic.com/v1 + guardrails: ["azure-text-moderation"] + - model_name: openai-gpt-4o litellm_params: model: openai/gpt-4o - api_key: os.environ/OPENAI_API_KEY -litellm_settings: - callbacks: ["prometheus"] - custom_prometheus_tags: ["User-Agent: Roo-Code"] - -router_settings: - routing_strategy: simple-shuffle - timeout: 300 - retry_policy: { - "AuthenticationErrorRetries": 0, - "BadRequestErrorRetries": 0, - "ContentPolicyViolationErrorRetries": 0, - "InternalServerErrorRetries": 1, - "RateLimitErrorRetries": 2, - "TimeoutErrorRetries": 0 - } \ No newline at end of file +guardrails: + - guardrail_name: "presidio-pii" + litellm_params: + guardrail: presidio # supported values: "aporia", "bedrock", "lakera", "presidio" + mode: "pre_call" + presidio_language: "en" # optional: set default language for PII analysis + pii_entities_config: + PERSON: "BLOCK" # Will mask credit card numbers + - guardrail_name: azure-text-moderation + litellm_params: + guardrail: azure/text_moderations + mode: "post_call" + api_key: os.environ/AZURE_GUARDRAIL_API_KEY + api_base: os.environ/AZURE_GUARDRAIL_API_BASE diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index c7ac00ee81..d6b958513b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -336,6 +336,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): presidio_config=presidio_config, request_data=request_data, ) + verbose_proxy_logger.debug("analyze_results: %s", analyze_results) #################################################### @@ -403,6 +404,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): LitellmCallTypes.completion.value, LitellmCallTypes.acompletion.value, ]: + messages = data["messages"] tasks = [] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 76e7dcc444..c5677074f9 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2093,6 +2093,7 @@ all_litellm_params = [ "metadata", "litellm_metadata", "litellm_trace_id", + "litellm_guardrails", "tags", "acompletion", "aimg_generation", diff --git a/litellm/utils.py b/litellm/utils.py index a8ed6c512f..9a618a49a1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -171,6 +171,7 @@ from litellm.types.utils import ( ImageResponse, LlmProviders, LlmProvidersSet, + LLMResponseTypes, Message, ModelInfo, ModelInfoBase, @@ -904,7 +905,6 @@ def client(original_function): # noqa: PLR0915 modified_kwargs = kwargs.copy() for callback in litellm.callbacks: - if isinstance(callback, CustomLogger): result = await callback.async_pre_call_deployment_hook( modified_kwargs, typed_call_type @@ -914,6 +914,27 @@ def client(original_function): # noqa: PLR0915 return modified_kwargs + async def async_post_call_success_deployment_hook( + request_data: dict, response: Any, call_type: Optional[CallTypes] + ) -> Optional[Any]: + """ + Allow modifying / reviewing the response just after it's received from the deployment. + """ + try: + typed_call_type = CallTypes(call_type) + except ValueError: + typed_call_type = None # unknown call type + + for callback in litellm.callbacks: + if isinstance(callback, CustomLogger): + result = await callback.async_post_call_success_deployment_hook( + request_data, cast(LLMResponseTypes, response), typed_call_type + ) + if result is not None: + return result + + return response + def post_call_processing(original_response, model, optional_params: Optional[dict]): try: if original_response is None: @@ -1443,6 +1464,13 @@ def client(original_function): # noqa: PLR0915 post_call_processing( original_response=result, model=model, optional_params=kwargs ) + # Only run if call_type is a valid value in CallTypes + if call_type in [ct.value for ct in CallTypes]: + await async_post_call_success_deployment_hook( + request_data=kwargs, + response=result, + call_type=CallTypes(call_type), + ) ## Add response to cache await _llm_caching_handler.async_set_cache( diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py new file mode 100644 index 0000000000..e71b68ab93 --- /dev/null +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -0,0 +1,84 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._types import CallTypes, UserAPIKeyAuth + + +class TestCustomGuardrailDeploymentHook: + + @pytest.mark.asyncio + async def test_async_pre_call_deployment_hook_no_guardrails(self): + """Test that method returns kwargs unchanged when no guardrails are present""" + custom_guardrail = CustomGuardrail() + + # Test with guardrails as None + kwargs = { + "messages": [{"role": "user", "content": "test message"}], + "model": "gpt-3.5-turbo", + "guardrails": None, + } + + result = await custom_guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert result == kwargs + + # Test with guardrails as non-list + kwargs["guardrails"] = "not_a_list" + + result = await custom_guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + assert result == kwargs + + @pytest.mark.asyncio + async def test_async_pre_call_deployment_hook_with_guardrails_and_message_update( + self, + ): + """Test that method processes guardrails and updates messages when result contains messages""" + custom_guardrail = CustomGuardrail() + + # Mock the async_pre_call_hook method + mock_result = {"messages": [{"role": "user", "content": "filtered message"}]} + custom_guardrail.async_pre_call_hook = AsyncMock(return_value=mock_result) + + original_messages = [{"role": "user", "content": "original message"}] + kwargs = { + "messages": original_messages, + "model": "gpt-3.5-turbo", + "guardrails": ["some_guardrail"], + "user_api_key_user_id": "test_user", + "user_api_key_team_id": "test_team", + "user_api_key_end_user_id": "test_end_user", + "user_api_key_hash": "test_hash", + "user_api_key_request_route": "test_route", + } + + result = await custom_guardrail.async_pre_call_deployment_hook( + kwargs=kwargs, call_type=CallTypes.completion + ) + + # Verify async_pre_call_hook was called with correct parameters + custom_guardrail.async_pre_call_hook.assert_called_once() + call_args = custom_guardrail.async_pre_call_hook.call_args + + # Check that UserAPIKeyAuth was created properly + user_api_key_dict = call_args[1]["user_api_key_dict"] + assert isinstance(user_api_key_dict, UserAPIKeyAuth) + assert user_api_key_dict.user_id == "test_user" + assert user_api_key_dict.team_id == "test_team" + assert user_api_key_dict.end_user_id == "test_end_user" + assert user_api_key_dict.api_key == "test_hash" + assert user_api_key_dict.request_route == "test_route" + + # Check other parameters + assert call_args[1]["data"] == kwargs + assert call_args[1]["call_type"] == "completion" + + # Verify messages were updated in result + assert result["messages"] == mock_result["messages"] + assert result["messages"] != original_messages