diff --git a/.circleci/config.yml b/.circleci/config.yml index 27ab837c9d6..3cfa15c99d6 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_guardrail.py:/app/custom_guardrail.py \ my-app:latest \ --config /app/config.yaml \ --port 4000 \ 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..09819b5dcb2 --- /dev/null +++ b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md @@ -0,0 +1,390 @@ +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 + +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 + +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` + +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 + +`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 # 👈 Key change + 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. [This runs the `async_pre_call_hook`](#1-write-a-customguardrail-class) + +```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. [This runs the `async_moderation_hook`](#1-write-a-customguardrail-class) + + +```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. [This runs the `async_post_call_success_hook`](#1-write-a-customguardrail-class) + + +```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", 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 new file mode 100644 index 00000000000..2ed989cfd38 --- /dev/null +++ b/litellm/proxy/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/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 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..2ed989cfd38 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/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/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index f0e2a9e2eca..643e135961d 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 @@ -83,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 = [] @@ -161,6 +165,43 @@ 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"] + ): + import os + + from litellm.proxy.utils import get_instance_fn + + # 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, + ) + + 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"], + ) + litellm.callbacks.append(_guardrail_callback) # type: ignore parsed_guardrail = Guardrail( guardrail_name=guardrail["guardrail_name"], @@ -169,6 +210,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 diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 320216a79b9..6be2454a2cd 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -1,17 +1,20 @@ model_list: - model_name: fake-openai-endpoint 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/fake + api_key: fake-key + api_base: https://exampleopenaiendpoint-production.up.railway.app/ 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/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: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a7701771791..09fc014d58b 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, @@ -375,14 +393,36 @@ 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 ) else: _callback = callback # type: ignore - if ( + + if _callback is not None and isinstance(_callback, CustomGuardrail): + 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 +434,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 +455,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 +780,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 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)