From e62d0c792254d7952185530678f9cf5afbc7e663 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 08:21:07 -0700 Subject: [PATCH 1/9] add the ability to init a custom guardrail --- litellm/proxy/guardrails/init_guardrails.py | 21 ++++++++++++++++++++- 1 file changed, 20 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index f0e2a9e2eca..4180bce2d38 100644 --- a/litellm/proxy/guardrails/init_guardrails.py +++ b/litellm/proxy/guardrails/init_guardrails.py @@ -1,3 +1,4 @@ +import importlib import traceback from typing import Dict, List, Literal @@ -161,6 +162,25 @@ def init_guardrails_v2(all_guardrails: dict): category_thresholds=litellm_params.get("category_thresholds"), ) litellm.callbacks.append(_lakera_callback) # type: ignore + elif ( + isinstance(litellm_params["guardrail"], str) + and "." in litellm_params["guardrail"] + ): + # Custom guardrail + _guardrail = litellm_params["guardrail"] + _file_name, _class_name = _guardrail.split(".") + verbose_proxy_logger.debug( + "Initializing custom guardrail: %s, file_name: %s, class_name: %s", + _guardrail, + _file_name, + _class_name, + ) + _guardrail_class = getattr(importlib.import_module(_file_name), _class_name) + _guardrail_callback = _guardrail_class( + guardrail_name=guardrail["guardrail_name"], + event_hook=litellm_params["mode"], + ) + litellm.callbacks.append(_guardrail_callback) # type: ignore parsed_guardrail = Guardrail( guardrail_name=guardrail["guardrail_name"], @@ -169,6 +189,5 @@ def init_guardrails_v2(all_guardrails: dict): guardrail_list.append(parsed_guardrail) guardrail_name = guardrail["guardrail_name"] - # pretty print guardrail_list in green print(f"\nGuardrail List:{guardrail_list}\n") # noqa From af92cff44dff493c21d83f03d6073529b25a38dc Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 08:32:07 -0700 Subject: [PATCH 2/9] add custom guardrail reference --- litellm/proxy/custom_guardrail.py | 115 ++++++++++++++++ .../guardrail_hooks/custom_guardrail.py | 115 ++++++++++++++++ litellm/proxy/proxy_config.yaml | 26 ++-- litellm/proxy/utils.py | 125 ++++++++++++++---- 4 files changed, 342 insertions(+), 39 deletions(-) create mode 100644 litellm/proxy/custom_guardrail.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py diff --git a/litellm/proxy/custom_guardrail.py b/litellm/proxy/custom_guardrail.py new file mode 100644 index 00000000000..bdcdcee1cbe --- /dev/null +++ b/litellm/proxy/custom_guardrail.py @@ -0,0 +1,115 @@ +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path +import asyncio +import json +import sys +import traceback +import uuid +from datetime import datetime +from typing import Any, Dict, List, Literal, Optional, Union + +from fastapi import HTTPException + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_helpers import should_proceed_based_on_metadata +from litellm.types.guardrails import GuardrailEventHooks + + +class myCustomGuardrail(CustomGuardrail): + def __init__( + self, + **kwargs, + ): + # store kwargs as optional_params + self.optional_params = kwargs + + super().__init__(**kwargs) + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: Literal[ + "completion", + "text_completion", + "embeddings", + "image_generation", + "moderation", + "audio_transcription", + "pass_through_endpoint", + ], + ) -> Optional[Union[Exception, str, dict]]: + # In this guardrail, if a user inputs `litellm` we will mask it. + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + _content = _content.replace("litellm", "********") + message["content"] = _content + + verbose_proxy_logger.debug( + "async_pre_call_hook: Message after masking %s", _messages + ) + + return data + + async def async_moderation_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + call_type: Literal["completion", "embeddings", "image_generation"], + ): + """ + Runs in parallel to LLM API call + Runs on only Input + """ + + # this works the same as async_pre_call_hook, but just runs in parallel as the LLM API Call + # In this guardrail, if a user inputs `litellm` we will mask it. + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + _content = _content.replace("litellm", "********") + message["content"] = _content + + verbose_proxy_logger.debug( + "async_pre_call_hook: Message after masking %s", _messages + ) + pass + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response, + ): + """ + Runs on response from LLM API call + + If a response contains the word "coffee" -> we will raise an exception + """ + verbose_proxy_logger.debug("async_pre_call_hook response: %s", response) + if isinstance(response, litellm.ModelResponse): + for choice in response.choices: + if isinstance(choice, litellm.Choices): + verbose_proxy_logger.debug("async_pre_call_hook choice: %s", choice) + if ( + choice.message.content + and isinstance(choice.message.content, str) + and "coffee" in choice.message.content + ): + raise ValueError("Guardrail failed Coffee Detected") diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py new file mode 100644 index 00000000000..bdcdcee1cbe --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py @@ -0,0 +1,115 @@ +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path +import asyncio +import json +import sys +import traceback +import uuid +from datetime import datetime +from typing import Any, Dict, List, Literal, Optional, Union + +from fastapi import HTTPException + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_helpers import should_proceed_based_on_metadata +from litellm.types.guardrails import GuardrailEventHooks + + +class myCustomGuardrail(CustomGuardrail): + def __init__( + self, + **kwargs, + ): + # store kwargs as optional_params + self.optional_params = kwargs + + super().__init__(**kwargs) + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: Literal[ + "completion", + "text_completion", + "embeddings", + "image_generation", + "moderation", + "audio_transcription", + "pass_through_endpoint", + ], + ) -> Optional[Union[Exception, str, dict]]: + # In this guardrail, if a user inputs `litellm` we will mask it. + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + _content = _content.replace("litellm", "********") + message["content"] = _content + + verbose_proxy_logger.debug( + "async_pre_call_hook: Message after masking %s", _messages + ) + + return data + + async def async_moderation_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + call_type: Literal["completion", "embeddings", "image_generation"], + ): + """ + Runs in parallel to LLM API call + Runs on only Input + """ + + # this works the same as async_pre_call_hook, but just runs in parallel as the LLM API Call + # In this guardrail, if a user inputs `litellm` we will mask it. + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + _content = _content.replace("litellm", "********") + message["content"] = _content + + verbose_proxy_logger.debug( + "async_pre_call_hook: Message after masking %s", _messages + ) + pass + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response, + ): + """ + Runs on response from LLM API call + + If a response contains the word "coffee" -> we will raise an exception + """ + verbose_proxy_logger.debug("async_pre_call_hook response: %s", response) + if isinstance(response, litellm.ModelResponse): + for choice in response.choices: + if isinstance(choice, litellm.Choices): + verbose_proxy_logger.debug("async_pre_call_hook choice: %s", choice) + if ( + choice.message.content + and isinstance(choice.message.content, str) + and "coffee" in choice.message.content + ): + raise ValueError("Guardrail failed Coffee Detected") diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 320216a79b9..acb792aec99 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -1,17 +1,19 @@ model_list: - - model_name: fake-openai-endpoint + - model_name: gpt-4 litellm_params: - model: azure/chatgpt-v-2 - api_base: https://openai-gpt-4-test-v-1.openai.azure.com/ - api_version: "2023-05-15" - tenant_id: os.environ/AZURE_TENANT_ID - client_id: os.environ/AZURE_CLIENT_ID - client_secret: os.environ/AZURE_CLIENT_SECRET + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY guardrails: - - guardrail_name: "bedrock-pre-guard" + - guardrail_name: "custom-pre-guard" litellm_params: - guardrail: bedrock # supported values: "aporia", "bedrock", "lakera" - mode: "post_call" - guardrailIdentifier: ff6ujrregl1q - guardrailVersion: "DRAFT" \ No newline at end of file + guardrail: custom_guardrail.myCustomGuardrail + mode: "pre_call" + - guardrail_name: "custom-during-guard" + litellm_params: + guardrail: custom_guardrail.myCustomGuardrail + mode: "during_call" + - guardrail_name: "custom-post-guard" + litellm_params: + guardrail: custom_guardrail.myCustomGuardrail + mode: "post_call" \ No newline at end of file diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a7701771791..7bdad862bc7 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -30,6 +30,7 @@ from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes from litellm.caching import DualCache, RedisCache from litellm.exceptions import RejectedRequestError +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.slack_alerting import SlackAlerting from litellm.litellm_core_utils.core_helpers import ( @@ -344,6 +345,23 @@ class ProxyLogging: ttl=alerting_threshold, ) + async def process_pre_call_hook_response(self, response, data, call_type): + if isinstance(response, Exception): + raise response + if isinstance(response, dict): + return response + if isinstance(response, str): + if call_type in ["completion", "text_completion"]: + raise RejectedRequestError( + message=response, + model=data.get("model", ""), + llm_provider="", + request_data=data, + ) + else: + raise HTTPException(status_code=400, detail={"error": response}) + return data + # The actual implementation of the function async def pre_call_hook( self, @@ -382,7 +400,33 @@ class ProxyLogging: ) else: _callback = callback # type: ignore + if ( + _callback is not None + and isinstance(_callback, CustomGuardrail) + and "pre_call_hook" in vars(_callback.__class__) + ): + from litellm.types.guardrails import GuardrailEventHooks + + if ( + _callback.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.pre_call + ) + is not True + ): + continue + response = await _callback.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=self.call_details["user_api_key_cache"], + data=data, + call_type=call_type, + ) + if response is not None: + data = await self.process_pre_call_hook_response( + response=response, data=data, call_type=call_type + ) + + elif ( _callback is not None and isinstance(_callback, CustomLogger) and "async_pre_call_hook" in vars(_callback.__class__) @@ -394,25 +438,9 @@ class ProxyLogging: call_type=call_type, ) if response is not None: - if isinstance(response, Exception): - raise response - elif isinstance(response, dict): - data = response - elif isinstance(response, str): - if ( - call_type == "completion" - or call_type == "text_completion" - ): - raise RejectedRequestError( - message=response, - model=data.get("model", ""), - llm_provider="", - request_data=data, - ) - else: - raise HTTPException( - status_code=400, detail={"error": response} - ) + data = await self.process_pre_call_hook_response( + response=response, data=data, call_type=call_type + ) return data except Exception as e: @@ -431,11 +459,30 @@ class ProxyLogging: ], ): """ - Runs the CustomLogger's async_moderation_hook() + Runs the CustomGuardrail's async_moderation_hook() """ for callback in litellm.callbacks: try: - if isinstance(callback, CustomLogger): + if isinstance(callback, CustomGuardrail): + ################################################################ + # Check if guardrail should be run for GuardrailEventHooks.during_call hook + ################################################################ + + # V1 implementation - backwards compatibility + if callback.event_hook is None: + if callback.moderation_check == "pre_call": + return + else: + # Main - V2 Guardrails implementation + from litellm.types.guardrails import GuardrailEventHooks + + if ( + callback.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.during_call + ) + is not True + ): + continue await callback.async_moderation_hook( data=data, user_api_key_dict=user_api_key_dict, @@ -737,12 +784,36 @@ class ProxyLogging: ) else: _callback = callback # type: ignore - if _callback is not None and isinstance(_callback, CustomLogger): - await _callback.async_post_call_success_hook( - user_api_key_dict=user_api_key_dict, - data=data, - response=response, - ) + + if _callback is not None: + ############## Handle Guardrails ######################################## + ############################################################################# + if isinstance(callback, CustomGuardrail): + # Main - V2 Guardrails implementation + from litellm.types.guardrails import GuardrailEventHooks + + if ( + callback.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.post_call + ) + is not True + ): + continue + + await callback.async_post_call_success_hook( + user_api_key_dict=user_api_key_dict, + data=data, + response=response, + ) + + ############ Handle CustomLogger ############################### + ################################################################# + elif isinstance(_callback, CustomLogger): + await _callback.async_post_call_success_hook( + user_api_key_dict=user_api_key_dict, + data=data, + response=response, + ) except Exception as e: raise e return response From a99258440c3bed2ff4ce7cf6015b4b9314634b19 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 09:34:08 -0700 Subject: [PATCH 3/9] fix use guardrail for pre call hook --- litellm/integrations/custom_guardrail.py | 8 ++--- litellm/proxy/custom_guardrail.py | 34 +++++++------------ .../guardrail_hooks/custom_guardrail.py | 34 +++++++------------ litellm/proxy/utils.py | 8 ++--- 4 files changed, 30 insertions(+), 54 deletions(-) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 047d1b6d374..25512716cd9 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -18,16 +18,16 @@ class CustomGuardrail(CustomLogger): super().__init__(**kwargs) def should_run_guardrail(self, data, event_type: GuardrailEventHooks) -> bool: + metadata = data.get("metadata") or {} + requested_guardrails = metadata.get("guardrails") or [] verbose_logger.debug( - "inside should_run_guardrail for guardrail=%s event_type= %s guardrail_supported_event_hooks= %s", + "inside should_run_guardrail for guardrail=%s event_type= %s guardrail_supported_event_hooks= %s requested_guardrails= %s", self.guardrail_name, event_type, self.event_hook, + requested_guardrails, ) - metadata = data.get("metadata") or {} - requested_guardrails = metadata.get("guardrails") or [] - if self.guardrail_name not in requested_guardrails: return False diff --git a/litellm/proxy/custom_guardrail.py b/litellm/proxy/custom_guardrail.py index bdcdcee1cbe..2ed989cfd38 100644 --- a/litellm/proxy/custom_guardrail.py +++ b/litellm/proxy/custom_guardrail.py @@ -1,19 +1,5 @@ -import os -import sys - -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path -import asyncio -import json -import sys -import traceback -import uuid -from datetime import datetime from typing import Any, Dict, List, Literal, Optional, Union -from fastapi import HTTPException - import litellm from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache @@ -48,7 +34,13 @@ class myCustomGuardrail(CustomGuardrail): "pass_through_endpoint", ], ) -> Optional[Union[Exception, str, dict]]: - # In this guardrail, if a user inputs `litellm` we will mask it. + """ + Runs before the LLM API call + Runs on only Input + Use this if you want to MODIFY the input + """ + + # In this guardrail, if a user inputs `litellm` we will mask it and then send it to the LLM _messages = data.get("messages") if _messages: for message in _messages: @@ -73,6 +65,8 @@ class myCustomGuardrail(CustomGuardrail): """ Runs in parallel to LLM API call Runs on only Input + + This can NOT modify the input, only used to reject or accept a call before going to LLM API """ # this works the same as async_pre_call_hook, but just runs in parallel as the LLM API Call @@ -83,13 +77,7 @@ class myCustomGuardrail(CustomGuardrail): _content = message.get("content") if isinstance(_content, str): if "litellm" in _content.lower(): - _content = _content.replace("litellm", "********") - message["content"] = _content - - verbose_proxy_logger.debug( - "async_pre_call_hook: Message after masking %s", _messages - ) - pass + raise ValueError("Guardrail failed words - `litellm` detected") async def async_post_call_success_hook( self, @@ -100,6 +88,8 @@ class myCustomGuardrail(CustomGuardrail): """ Runs on response from LLM API call + It can be used to reject a response + If a response contains the word "coffee" -> we will raise an exception """ verbose_proxy_logger.debug("async_pre_call_hook response: %s", response) diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py index bdcdcee1cbe..2ed989cfd38 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py @@ -1,19 +1,5 @@ -import os -import sys - -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path -import asyncio -import json -import sys -import traceback -import uuid -from datetime import datetime from typing import Any, Dict, List, Literal, Optional, Union -from fastapi import HTTPException - import litellm from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache @@ -48,7 +34,13 @@ class myCustomGuardrail(CustomGuardrail): "pass_through_endpoint", ], ) -> Optional[Union[Exception, str, dict]]: - # In this guardrail, if a user inputs `litellm` we will mask it. + """ + Runs before the LLM API call + Runs on only Input + Use this if you want to MODIFY the input + """ + + # In this guardrail, if a user inputs `litellm` we will mask it and then send it to the LLM _messages = data.get("messages") if _messages: for message in _messages: @@ -73,6 +65,8 @@ class myCustomGuardrail(CustomGuardrail): """ Runs in parallel to LLM API call Runs on only Input + + This can NOT modify the input, only used to reject or accept a call before going to LLM API """ # this works the same as async_pre_call_hook, but just runs in parallel as the LLM API Call @@ -83,13 +77,7 @@ class myCustomGuardrail(CustomGuardrail): _content = message.get("content") if isinstance(_content, str): if "litellm" in _content.lower(): - _content = _content.replace("litellm", "********") - message["content"] = _content - - verbose_proxy_logger.debug( - "async_pre_call_hook: Message after masking %s", _messages - ) - pass + raise ValueError("Guardrail failed words - `litellm` detected") async def async_post_call_success_hook( self, @@ -100,6 +88,8 @@ class myCustomGuardrail(CustomGuardrail): """ Runs on response from LLM API call + It can be used to reject a response + If a response contains the word "coffee" -> we will raise an exception """ verbose_proxy_logger.debug("async_pre_call_hook response: %s", response) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 7bdad862bc7..09fc014d58b 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -393,7 +393,7 @@ class ProxyLogging: try: for callback in litellm.callbacks: - _callback: Optional[CustomLogger] = None + _callback = None if isinstance(callback, str): _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( callback @@ -401,11 +401,7 @@ class ProxyLogging: else: _callback = callback # type: ignore - if ( - _callback is not None - and isinstance(_callback, CustomGuardrail) - and "pre_call_hook" in vars(_callback.__class__) - ): + if _callback is not None and isinstance(_callback, CustomGuardrail): from litellm.types.guardrails import GuardrailEventHooks if ( From d10430c8812aa3a2c36bf99b74ec96c672dcc5c2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 09:41:54 -0700 Subject: [PATCH 4/9] doc custom guardrail --- .../docs/proxy/guardrails/custom_guardrail.md | 381 ++++++++++++++++++ docs/my-website/sidebars.js | 2 +- 2 files changed, 382 insertions(+), 1 deletion(-) create mode 100644 docs/my-website/docs/proxy/guardrails/custom_guardrail.md diff --git a/docs/my-website/docs/proxy/guardrails/custom_guardrail.md b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md new file mode 100644 index 00000000000..51a0a0a78dc --- /dev/null +++ b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md @@ -0,0 +1,381 @@ +import Image from '@theme/IdealImage'; +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Custom Guardrail + +Use this is you want to write code to run a custom guardrail + +## Quick Start + +### 1. Write a `CustomGuardrail` Class + +```python +from typing import Any, Dict, List, Literal, Optional, Union + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_helpers import should_proceed_based_on_metadata +from litellm.types.guardrails import GuardrailEventHooks + + +class myCustomGuardrail(CustomGuardrail): + def __init__( + self, + **kwargs, + ): + # store kwargs as optional_params + self.optional_params = kwargs + + super().__init__(**kwargs) + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: Literal[ + "completion", + "text_completion", + "embeddings", + "image_generation", + "moderation", + "audio_transcription", + "pass_through_endpoint", + ], + ) -> Optional[Union[Exception, str, dict]]: + """ + Runs before the LLM API call + Runs on only Input + Use this if you want to MODIFY the input + """ + + # In this guardrail, if a user inputs `litellm` we will mask it and then send it to the LLM + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + _content = _content.replace("litellm", "********") + message["content"] = _content + + verbose_proxy_logger.debug( + "async_pre_call_hook: Message after masking %s", _messages + ) + + return data + + async def async_moderation_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + call_type: Literal["completion", "embeddings", "image_generation"], + ): + """ + Runs in parallel to LLM API call + Runs on only Input + + This can NOT modify the input, only used to reject or accept a call before going to LLM API + """ + + # this works the same as async_pre_call_hook, but just runs in parallel as the LLM API Call + # In this guardrail, if a user inputs `litellm` we will mask it. + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + raise ValueError("Guardrail failed words - `litellm` detected") + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response, + ): + """ + Runs on response from LLM API call + + It can be used to reject a response + + If a response contains the word "coffee" -> we will raise an exception + """ + verbose_proxy_logger.debug("async_pre_call_hook response: %s", response) + if isinstance(response, litellm.ModelResponse): + for choice in response.choices: + if isinstance(choice, litellm.Choices): + verbose_proxy_logger.debug("async_pre_call_hook choice: %s", choice) + if ( + choice.message.content + and isinstance(choice.message.content, str) + and "coffee" in choice.message.content + ): + raise ValueError("Guardrail failed Coffee Detected") + + +``` + +### 2. Pass your custom guardrail class in LiteLLM `config.yaml` + +We pass the custom callback class defined in **Step1** to the config.yaml. +Set `callbacks` to `python_filename.logger_instance_name` + +In the config below, we pass + +- Python Filename: `custom_guardrail.py` +- Guardrail class name : `myCustomGuardrail`. This is defined in Step 1 + +`guardrail: custom_guardrail.myCustomGuardrail` + +```yaml +model_list: + - model_name: gpt-4 + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + +guardrails: + - guardrail_name: "custom-pre-guard" + litellm_params: + guardrail: custom_guardrail.myCustomGuardrail + 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 +``` + +### 3. Start LiteLLM Gateway + + +```shell +litellm --config config.yaml --detailed_debug +``` + + +### 4. Test it + +#### Test `"custom-pre-guard"` + + +**[Langchain, OpenAI SDK Usage Examples](../proxy/user_keys##request-format)** + + + + +Expect this to mask the word `litellm` before sending the request to the LLM API + +```shell +curl -i -X POST http://localhost:4000/v1/chat/completions \ +-H "Content-Type: application/json" \ +-H "Authorization: Bearer sk-1234" \ +-d '{ + "model": "gpt-4", + "messages": [ + { + "role": "user", + "content": "say the word - `litellm`" + } + ], + "guardrails": ["custom-pre-guard"] +}' +``` + +Expected response after pre-guard + +```json +{ + "id": "chatcmpl-9zREDkBIG20RJB4pMlyutmi1hXQWc", + "choices": [ + { + "finish_reason": "stop", + "index": 0, + "message": { + "content": "It looks like you've chosen a string of asterisks. This could be a way to censor or hide certain text. However, without more context, I can't provide a specific word or phrase. If there's something specific you'd like me to say or if you need help with a topic, feel free to let me know!", + "role": "assistant", + "tool_calls": null, + "function_call": null + } + } + ], + "created": 1724429701, + "model": "gpt-4o-2024-05-13", + "object": "chat.completion", + "system_fingerprint": "fp_3aa7262c27", + "usage": { + "completion_tokens": 65, + "prompt_tokens": 14, + "total_tokens": 79 + }, + "service_tier": null +} + +``` + + + + + +```shell +curl -i http://localhost:4000/v1/chat/completions \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-npnwjPQciVRok5yNZgKmFQ" \ + -d '{ + "model": "gpt-3.5-turbo", + "messages": [ + {"role": "user", "content": "hi what is the weather"} + ], + "guardrails": ["custom-pre-guard"] + }' +``` + + + + + + + +#### Test `"custom-during-guard"` + + +**[Langchain, OpenAI SDK Usage Examples](../proxy/user_keys##request-format)** + + + + +Expect this to fail since since `litellm` is in the message content + +```shell +curl -i -X POST http://localhost:4000/v1/chat/completions \ +-H "Content-Type: application/json" \ +-H "Authorization: Bearer sk-1234" \ +-d '{ + "model": "gpt-4", + "messages": [ + { + "role": "user", + "content": "say the word - `litellm`" + } + ], + "guardrails": ["custom-during-guard"] +}' +``` + +Expected response after running during-guard + +```json +{ + "error": { + "message": "Guardrail failed words - `litellm` detected", + "type": "None", + "param": "None", + "code": "500" + } +} +``` + + + + + +```shell +curl -i http://localhost:4000/v1/chat/completions \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-npnwjPQciVRok5yNZgKmFQ" \ + -d '{ + "model": "gpt-3.5-turbo", + "messages": [ + {"role": "user", "content": "hi what is the weather"} + ], + "guardrails": ["custom-during-guard"] + }' +``` + + + + + + + +#### Test `"custom-post-guard"` + + + +**[Langchain, OpenAI SDK Usage Examples](../proxy/user_keys##request-format)** + + + + +Expect this to fail since since `coffee` will be in the response content + +```shell +curl -i -X POST http://localhost:4000/v1/chat/completions \ +-H "Content-Type: application/json" \ +-H "Authorization: Bearer sk-1234" \ +-d '{ + "model": "gpt-4", + "messages": [ + { + "role": "user", + "content": "what is coffee" + } + ], + "guardrails": ["custom-post-guard"] +}' +``` + +Expected response after running during-guard + +```json +{ + "error": { + "message": "Guardrail failed Coffee Detected", + "type": "None", + "param": "None", + "code": "500" + } +} +``` + + + + + +```shell + curl -i -X POST http://localhost:4000/v1/chat/completions \ +-H "Content-Type: application/json" \ +-H "Authorization: Bearer sk-1234" \ +-d '{ + "model": "gpt-4", + "messages": [ + { + "role": "user", + "content": "what is tea" + } + ], + "guardrails": ["custom-post-guard"] +}' +``` + + + + + + + +## **CustomGuardrail methods** + +| Component | Description | Optional | Checked Data | Can Modify Input | Can Modify Output | Can Fail Call | +|-----------|-------------|----------|--------------|------------------|-------------------|----------------| +| `async_pre_call_hook` | A hook that runs before the LLM API call | ✅ | INPUT | ✅ | ❌ | ✅ | +| `async_moderation_hook` | A hook that runs during the LLM API call| ✅ | INPUT | ❌ | ❌ | ✅ | +| `async_post_call_success_hook` | A hook that runs after a successful LLM API call| ✅ | INPUT, OUTPUT | ❌ | ✅ | ✅ | diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 339647dfa1f..8c8c87fb8fa 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -66,7 +66,7 @@ const sidebars = { { type: "category", label: "🛡️ [Beta] Guardrails", - items: ["proxy/guardrails/quick_start", "proxy/guardrails/aporia_api", "proxy/guardrails/lakera_ai", "proxy/guardrails/bedrock", "prompt_injection"], + items: ["proxy/guardrails/quick_start", "proxy/guardrails/aporia_api", "proxy/guardrails/lakera_ai", "proxy/guardrails/bedrock", "proxy/guardrails/custom_guardrail", "prompt_injection"], }, { type: "category", From d40695b979b264aa77cc68fd30742c11cf274a3f Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 09:50:19 -0700 Subject: [PATCH 5/9] docs custom guardrails --- .../docs/proxy/guardrails/custom_guardrail.md | 25 +++++++++++++------ 1 file changed, 17 insertions(+), 8 deletions(-) diff --git a/docs/my-website/docs/proxy/guardrails/custom_guardrail.md b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md index 51a0a0a78dc..09819b5dcb2 100644 --- a/docs/my-website/docs/proxy/guardrails/custom_guardrail.md +++ b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md @@ -10,6 +10,16 @@ Use this is you want to write code to run a custom guardrail ### 1. Write a `CustomGuardrail` Class +A CustomGuardrail has 3 methods to enforce guardrails +- `async_pre_call_hook` - (Optional) modify input or reject request before making LLM API call +- `async_moderation_hook` - (Optional) reject request, runs while making LLM API call (help to lower latency) +- `async_post_call_success_hook`- (Optional) apply guardrail on input/output, runs after making LLM API call + +**[See detailed spec of methods here](#customguardrail-methods)** + +**Example `CustomGuardrail` Class** + +Create a new file called `custom_guardrail.py` and add this code to it ```python from typing import Any, Dict, List, Literal, Optional, Union @@ -122,10 +132,7 @@ class myCustomGuardrail(CustomGuardrail): ### 2. Pass your custom guardrail class in LiteLLM `config.yaml` -We pass the custom callback class defined in **Step1** to the config.yaml. -Set `callbacks` to `python_filename.logger_instance_name` - -In the config below, we pass +In the config below, we point the guardrail to our custom guardrail by setting `guardrail: custom_guardrail.myCustomGuardrail` - Python Filename: `custom_guardrail.py` - Guardrail class name : `myCustomGuardrail`. This is defined in Step 1 @@ -142,7 +149,7 @@ model_list: guardrails: - guardrail_name: "custom-pre-guard" litellm_params: - guardrail: custom_guardrail.myCustomGuardrail + guardrail: custom_guardrail.myCustomGuardrail # 👈 Key change mode: "pre_call" # runs async_pre_call_hook - guardrail_name: "custom-during-guard" litellm_params: @@ -172,7 +179,7 @@ litellm --config config.yaml --detailed_debug -Expect this to mask the word `litellm` before sending the request to the LLM API +Expect this to mask the word `litellm` before sending the request to the LLM API. [This runs the `async_pre_call_hook`](#1-write-a-customguardrail-class) ```shell curl -i -X POST http://localhost:4000/v1/chat/completions \ @@ -252,7 +259,8 @@ curl -i http://localhost:4000/v1/chat/completions \ -Expect this to fail since since `litellm` is in the message content +Expect this to fail since since `litellm` is in the message content. [This runs the `async_moderation_hook`](#1-write-a-customguardrail-class) + ```shell curl -i -X POST http://localhost:4000/v1/chat/completions \ @@ -315,7 +323,8 @@ curl -i http://localhost:4000/v1/chat/completions \ -Expect this to fail since since `coffee` will be in the response content +Expect this to fail since since `coffee` will be in the response content. [This runs the `async_post_call_success_hook`](#1-write-a-customguardrail-class) + ```shell curl -i -X POST http://localhost:4000/v1/chat/completions \ From 7d30188f84ab482677d1ab2bc251393f7fa251e8 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 09:52:52 -0700 Subject: [PATCH 6/9] custom_callbacks --- .circleci/config.yml | 1 + .../example_config_yaml/custom_guardrail.py | 105 ++++++++++++++++++ .../example_config_yaml/otel_test_config.yaml | 14 ++- 3 files changed, 119 insertions(+), 1 deletion(-) create mode 100644 litellm/proxy/example_config_yaml/custom_guardrail.py diff --git a/.circleci/config.yml b/.circleci/config.yml index 27ab837c9d6..b562fbdd50c 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -328,6 +328,7 @@ jobs: -e APORIA_API_KEY_1=$APORIA_API_KEY_1 \ --name my-app \ -v $(pwd)/litellm/proxy/example_config_yaml/otel_test_config.yaml:/app/config.yaml \ + -v $(pwd)/litellm/proxy/example_config_yaml/custom_callbacks.py:/app/custom_callbacks.py \ my-app:latest \ --config /app/config.yaml \ --port 4000 \ diff --git a/litellm/proxy/example_config_yaml/custom_guardrail.py b/litellm/proxy/example_config_yaml/custom_guardrail.py new file mode 100644 index 00000000000..2ed989cfd38 --- /dev/null +++ b/litellm/proxy/example_config_yaml/custom_guardrail.py @@ -0,0 +1,105 @@ +from typing import Any, Dict, List, Literal, Optional, Union + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_helpers import should_proceed_based_on_metadata +from litellm.types.guardrails import GuardrailEventHooks + + +class myCustomGuardrail(CustomGuardrail): + def __init__( + self, + **kwargs, + ): + # store kwargs as optional_params + self.optional_params = kwargs + + super().__init__(**kwargs) + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: Literal[ + "completion", + "text_completion", + "embeddings", + "image_generation", + "moderation", + "audio_transcription", + "pass_through_endpoint", + ], + ) -> Optional[Union[Exception, str, dict]]: + """ + Runs before the LLM API call + Runs on only Input + Use this if you want to MODIFY the input + """ + + # In this guardrail, if a user inputs `litellm` we will mask it and then send it to the LLM + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + _content = _content.replace("litellm", "********") + message["content"] = _content + + verbose_proxy_logger.debug( + "async_pre_call_hook: Message after masking %s", _messages + ) + + return data + + async def async_moderation_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + call_type: Literal["completion", "embeddings", "image_generation"], + ): + """ + Runs in parallel to LLM API call + Runs on only Input + + This can NOT modify the input, only used to reject or accept a call before going to LLM API + """ + + # this works the same as async_pre_call_hook, but just runs in parallel as the LLM API Call + # In this guardrail, if a user inputs `litellm` we will mask it. + _messages = data.get("messages") + if _messages: + for message in _messages: + _content = message.get("content") + if isinstance(_content, str): + if "litellm" in _content.lower(): + raise ValueError("Guardrail failed words - `litellm` detected") + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response, + ): + """ + Runs on response from LLM API call + + It can be used to reject a response + + If a response contains the word "coffee" -> we will raise an exception + """ + verbose_proxy_logger.debug("async_pre_call_hook response: %s", response) + if isinstance(response, litellm.ModelResponse): + for choice in response.choices: + if isinstance(choice, litellm.Choices): + verbose_proxy_logger.debug("async_pre_call_hook choice: %s", choice) + if ( + choice.message.content + and isinstance(choice.message.content, str) + and "coffee" in choice.message.content + ): + raise ValueError("Guardrail failed Coffee Detected") diff --git a/litellm/proxy/example_config_yaml/otel_test_config.yaml b/litellm/proxy/example_config_yaml/otel_test_config.yaml index 8ca4f37fd6a..a041a2bd0ce 100644 --- a/litellm/proxy/example_config_yaml/otel_test_config.yaml +++ b/litellm/proxy/example_config_yaml/otel_test_config.yaml @@ -27,4 +27,16 @@ guardrails: guardrail: bedrock # supported values: "aporia", "bedrock", "lakera" mode: "pre_call" guardrailIdentifier: ff6ujrregl1q - guardrailVersion: "DRAFT" \ No newline at end of file + guardrailVersion: "DRAFT" + - guardrail_name: "custom-pre-guard" + litellm_params: + guardrail: custom_guardrail.myCustomGuardrail + mode: "pre_call" + - guardrail_name: "custom-during-guard" + litellm_params: + guardrail: custom_guardrail.myCustomGuardrail + mode: "during_call" + - guardrail_name: "custom-post-guard" + litellm_params: + guardrail: custom_guardrail.myCustomGuardrail + mode: "post_call" \ No newline at end of file From 1b1e0f2d77699a0bd6aef3e6d05a77ed242d2845 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 10:54:42 -0700 Subject: [PATCH 7/9] init custom guardrail class --- litellm/proxy/guardrails/init_guardrails.py | 25 +++++++++++++++++++-- litellm/proxy/proxy_server.py | 4 +++- 2 files changed, 26 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index 4180bce2d38..643e135961d 100644 --- a/litellm/proxy/guardrails/init_guardrails.py +++ b/litellm/proxy/guardrails/init_guardrails.py @@ -84,7 +84,10 @@ Map guardrail_name: , , during_call """ -def init_guardrails_v2(all_guardrails: dict): +def init_guardrails_v2( + all_guardrails: dict, + config_file_path: str, +): # Convert the loaded data to the TypedDict structure guardrail_list = [] @@ -166,6 +169,10 @@ def init_guardrails_v2(all_guardrails: dict): isinstance(litellm_params["guardrail"], str) and "." in litellm_params["guardrail"] ): + import os + + from litellm.proxy.utils import get_instance_fn + # Custom guardrail _guardrail = litellm_params["guardrail"] _file_name, _class_name = _guardrail.split(".") @@ -175,7 +182,21 @@ def init_guardrails_v2(all_guardrails: dict): _file_name, _class_name, ) - _guardrail_class = getattr(importlib.import_module(_file_name), _class_name) + + directory = os.path.dirname(config_file_path) + module_file_path = os.path.join(directory, _file_name) + module_file_path += ".py" + + spec = importlib.util.spec_from_file_location(_class_name, module_file_path) # type: ignore + if spec is None: + raise ImportError( + f"Could not find a module specification for {module_file_path}" + ) + + module = importlib.util.module_from_spec(spec) # type: ignore + spec.loader.exec_module(module) # type: ignore + _guardrail_class = getattr(module, _class_name) + _guardrail_callback = _guardrail_class( guardrail_name=guardrail["guardrail_name"], event_hook=litellm_params["mode"], diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3ef5609db37..f4206f726ac 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1959,7 +1959,9 @@ class ProxyConfig: # Guardrail settings guardrails_v2 = config.get("guardrails", None) if guardrails_v2: - init_guardrails_v2(all_guardrails=guardrails_v2) + init_guardrails_v2( + all_guardrails=guardrails_v2, config_file_path=config_file_path + ) return router, router.get_model_list(), general_settings def get_model_info_with_id(self, model, db_model=False) -> RouterModelInfo: From 5895f0c6158936ed12109560654d0921547ec9fe Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 11:58:20 -0700 Subject: [PATCH 8/9] fix custom guardrail test --- .circleci/config.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index b562fbdd50c..3cfa15c99d6 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -328,7 +328,7 @@ jobs: -e APORIA_API_KEY_1=$APORIA_API_KEY_1 \ --name my-app \ -v $(pwd)/litellm/proxy/example_config_yaml/otel_test_config.yaml:/app/config.yaml \ - -v $(pwd)/litellm/proxy/example_config_yaml/custom_callbacks.py:/app/custom_callbacks.py \ + -v $(pwd)/litellm/proxy/example_config_yaml/custom_guardrail.py:/app/custom_guardrail.py \ my-app:latest \ --config /app/config.yaml \ --port 4000 \ From 918e4fcfe5c3795d0b14d40d483e94d3248aa365 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 23 Aug 2024 12:01:43 -0700 Subject: [PATCH 9/9] feat add test for custom guardrails --- litellm/proxy/proxy_config.yaml | 7 ++++--- tests/otel_tests/test_guardrails.py | 21 +++++++++++++++++++++ 2 files changed, 25 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index acb792aec99..6be2454a2cd 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -1,8 +1,9 @@ model_list: - - model_name: gpt-4 + - model_name: fake-openai-endpoint litellm_params: - model: openai/gpt-4o - api_key: os.environ/OPENAI_API_KEY + model: openai/fake + api_key: fake-key + api_base: https://exampleopenaiendpoint-production.up.railway.app/ guardrails: - guardrail_name: "custom-pre-guard" diff --git a/tests/otel_tests/test_guardrails.py b/tests/otel_tests/test_guardrails.py index 34f14186e13..2b5bfc644e1 100644 --- a/tests/otel_tests/test_guardrails.py +++ b/tests/otel_tests/test_guardrails.py @@ -217,3 +217,24 @@ async def test_bedrock_guardrail_triggered(): print(e) assert "GUARDRAIL_INTERVENED" in str(e) assert "Violated guardrail policy" in str(e) + + +@pytest.mark.asyncio +async def test_custom_guardrail_during_call_triggered(): + """ + - Tests a request where our bedrock guardrail should be triggered + - Assert that the guardrails applied are returned in the response headers + """ + async with aiohttp.ClientSession() as session: + try: + response, headers = await chat_completion( + session, + "sk-1234", + model="fake-openai-endpoint", + messages=[{"role": "user", "content": f"Hello do you like litellm?"}], + guardrails=["custom-during-guard"], + ) + pytest.fail("Should have thrown an exception") + except Exception as e: + print(e) + assert "Guardrail failed words - `litellm` detected" in str(e)