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)