[Feat] Bedrock Guardrails - Add support for PII Masking with bedrock guardrails (#10642)

* allow defining mask_request_content for guardrails

* allow pii masking with bedrock

* implement bedrock pre call hook

* docs bedrock pii masking

* fix linting error

* fix code quality checks
This commit is contained in:
Ishaan Jaff 2025-05-07 18:30:57 -07:00 • committed by GitHub
parent 959ba21585
commit 6a3ea83b23
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 587 additions and 43 deletions

View file

@ -135,3 +135,42 @@ curl -i http://localhost:4000/v1/chat/completions \
</Tabs>
## PII Masking with Bedrock Guardrails
Bedrock guardrails support PII detection and masking capabilities. To enable this feature, you need to:
1. Set `mode` to `pre_call` to run the guardrail check before the LLM call
2. Enable masking by setting `mask_request_content` and/or `mask_response_content` to `true`
Here's how to configure it in your config.yaml:
```yaml showLineNumbers title="litellm bedrock guardrailconfig.yaml"
guardrails:
- guardrail_name: "bedrock-pre-guard"
litellm_params:
guardrail: bedrock
mode: "pre_call" # Important: must use pre_call mode for masking
guardrailIdentifier: wf0hkdb5x07f
guardrailVersion: "DRAFT"
mask_request_content: true # Enable masking in user requests
mask_response_content: true # Enable masking in model responses
```
With this configuration, when the bedrock guardrail intervenes, litellm will read the masked output from the guardrail and send it to the model.
### Example Usage
When enabled, PII will be automatically masked in the text. For example, if a user sends:
```
My email is john.doe@example.com and my phone number is 555-123-4567
```
The text sent to the model might be masked as:
```
My email is [EMAIL] and my phone number is [PHONE_NUMBER]
```
This helps protect sensitive information while still allowing the model to understand the context of the request.

View file

@ -15,6 +15,8 @@ class CustomGuardrail(CustomLogger):
Union[GuardrailEventHooks, List[GuardrailEventHooks]]
] = None,
default_on: bool = False,
mask_request_content: bool = False,
mask_response_content: bool = False,
**kwargs,
):
"""
@ -25,6 +27,8 @@ class CustomGuardrail(CustomLogger):
supported_event_hooks: The event hooks that the guardrail supports
event_hook: The event hook to run the guardrail on
default_on: If True, the guardrail will be run by default on all requests
mask_request_content: If True, the guardrail will mask the request content
mask_response_content: If True, the guardrail will mask the response content
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
@ -32,6 +36,8 @@ class CustomGuardrail(CustomLogger):
Union[GuardrailEventHooks, List[GuardrailEventHooks]]
] = event_hook
self.default_on: bool = default_on
self.mask_request_content: bool = mask_request_content
self.mask_response_content: bool = mask_response_content
if supported_event_hooks:
## validate event_hook is in supported_event_hooks

View file

@ -13,19 +13,17 @@ sys.path.insert(
) # Adds the parent directory to the system path
import json
import sys
from typing import Any, List, Literal, Optional, Union
from typing import Any, List, Literal, Optional, Tuple, 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,
log_guardrail_information,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -33,13 +31,15 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.secret_managers.main import get_secret
from litellm.types.guardrails import (
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockContentItem,
BedrockGuardrailOutput,
BedrockGuardrailResponse,
BedrockRequest,
BedrockTextContent,
GuardrailEventHooks,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
GUARDRAIL_NAME = "bedrock"
@ -74,12 +74,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if messages:
for message in messages:
bedrock_content_item = BedrockContentItem(
text=BedrockTextContent(
text=convert_content_list_to_str(message=message)
)
message_text_content: Optional[List[str]] = (
self.get_content_for_message(message=message)
)
bedrock_request_content.append(bedrock_content_item)
if message_text_content is None:
continue
for text_content in message_text_content:
bedrock_content_item = BedrockContentItem(
text=BedrockTextContent(text=text_content)
)
bedrock_request_content.append(bedrock_content_item)
bedrock_request["content"] = bedrock_request_content
if response:
@ -191,13 +195,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
async def make_bedrock_api_request(
self, kwargs: dict, response: Optional[Union[Any, litellm.ModelResponse]] = None
):
) -> BedrockGuardrailResponse:
credentials, aws_region_name = self._load_credentials()
bedrock_request_data: dict = dict(
self.convert_to_bedrock_format(
messages=kwargs.get("messages"), response=response
)
)
bedrock_guardrail_response: BedrockGuardrailResponse = (
BedrockGuardrailResponse()
)
bedrock_request_data.update(
self.get_guardrail_dynamic_request_body_params(request_data=kwargs)
)
@ -223,7 +230,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if response.status_code == 200:
# check if the response was flagged
_json_response = response.json()
if _json_response.get("action") == "GUARDRAIL_INTERVENED":
bedrock_guardrail_response = BedrockGuardrailResponse(**_json_response)
if self._should_raise_guardrail_blocked_exception(
bedrock_guardrail_response
):
raise HTTPException(
status_code=400,
detail={
@ -238,6 +248,87 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
response.text,
)
return bedrock_guardrail_response
def _should_raise_guardrail_blocked_exception(
self, response: BedrockGuardrailResponse
) -> bool:
"""
By default always raise an exception when a guardrail intervention is detected.
If `self.mask_request_content` or `self.mask_response_content` is set to `True`, then use the output from the guardrail to mask the request or response content.
"""
# if user opted into masking, return False. since we'll use the masked output from the guardrail
if self.mask_request_content or self.mask_response_content:
return False
# if intervention, return True
if response.get("action") == "GUARDRAIL_INTERVENED":
return True
# if no intervention, return False
return False
@log_guardrail_information
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",
"rerank",
],
) -> Union[Exception, str, dict, None]:
verbose_proxy_logger.debug("Inside AIM Pre-Call Hook")
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
event_type: GuardrailEventHooks = GuardrailEventHooks.pre_call
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return data
new_messages: Optional[List[AllMessageValues]] = data.get("messages")
if new_messages is None:
verbose_proxy_logger.warning(
"Bedrock AI: not running guardrail. No messages in data"
)
return data
#########################################################
########## 1. Make the Bedrock API request ##########
#########################################################
bedrock_guardrail_response = await self.make_bedrock_api_request(kwargs=data)
#########################################################
#########################################################
########## 2. Update the messages with the guardrail response ##########
#########################################################
data["messages"] = (
self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
)
#########################################################
########## 3. Add the guardrail to the applied guardrails header ##########
#########################################################
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
return data
@log_guardrail_information
async def async_moderation_hook(
self,
@ -260,17 +351,37 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return
new_messages: Optional[List[dict]] = data.get("messages")
if new_messages is not None:
await self.make_bedrock_api_request(kwargs=data)
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
else:
new_messages: Optional[List[AllMessageValues]] = data.get("messages")
if new_messages is None:
verbose_proxy_logger.warning(
"Bedrock AI: not running guardrail. No messages in data"
)
pass
return
#########################################################
########## 1. Make the Bedrock API request ##########
#########################################################
bedrock_guardrail_response = await self.make_bedrock_api_request(kwargs=data)
#########################################################
#########################################################
########## 2. Update the messages with the guardrail response ##########
#########################################################
data["messages"] = (
self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
)
#########################################################
########## 3. Add the guardrail to the applied guardrails header ##########
#########################################################
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
return data
@log_guardrail_information
async def async_post_call_success_hook(
@ -292,13 +403,190 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
):
return
new_messages: Optional[List[dict]] = data.get("messages")
if new_messages is not None:
await self.make_bedrock_api_request(kwargs=data, response=response)
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
else:
new_messages: Optional[List[AllMessageValues]] = data.get("messages")
if new_messages is None:
verbose_proxy_logger.warning(
"Bedrock AI: not running guardrail. No messages in data"
)
return
#########################################################
########## 1. Make the Bedrock API request ##########
#########################################################
bedrock_guardrail_response = await self.make_bedrock_api_request(
kwargs=data, response=response
)
#########################################################
#########################################################
########## 2. Update the messages with the guardrail response ##########
#########################################################
data["messages"] = (
self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
)
#########################################################
########## 3. Add the guardrail to the applied guardrails header ##########
#########################################################
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
########### HELPER FUNCTIONS for bedrock guardrails ############################
##############################################################################
##############################################################################
def _update_messages_with_updated_bedrock_guardrail_response(
self,
messages: List[AllMessageValues],
bedrock_guardrail_response: BedrockGuardrailResponse,
) -> List[AllMessageValues]:
"""
Use the output from the bedrock guardrail to mask sensitive content in messages.
Args:
messages: Original list of messages
bedrock_guardrail_response: Response from Bedrock guardrail containing masked content
Returns:
List of messages with content masked according to guardrail response
"""
# Skip processing if masking is not enabled
if not (self.mask_request_content or self.mask_response_content):
return messages
# Get masked texts from guardrail response
masked_texts = self._extract_masked_texts_from_response(
bedrock_guardrail_response
)
if not masked_texts:
return messages
# Apply masking to messages using index tracking
return self._apply_masking_to_messages(
messages=messages, masked_texts=masked_texts
)
def _extract_masked_texts_from_response(
self, bedrock_guardrail_response: BedrockGuardrailResponse
) -> List[str]:
"""
Extract all masked text outputs from the guardrail response.
Args:
bedrock_guardrail_response: Response from Bedrock guardrail
Returns:
List of masked text strings
"""
masked_output_text: List[str] = []
masked_outputs: Optional[List[BedrockGuardrailOutput]] = (
bedrock_guardrail_response.get("outputs", []) or []
)
if not masked_outputs:
verbose_proxy_logger.debug("No masked outputs found in guardrail response")
return []
for output in masked_outputs:
text_content: Optional[str] = output.get("text")
if text_content is not None:
masked_output_text.append(text_content)
return masked_output_text
def _apply_masking_to_messages(
self, messages: List[AllMessageValues], masked_texts: List[str]
) -> List[AllMessageValues]:
"""
Apply masked texts to message content using index tracking.
Args:
messages: Original messages
masked_texts: List of masked text strings from guardrail
Returns:
Updated messages with masked content
"""
updated_messages = []
masking_index = 0
for message in messages:
new_message = message.copy()
content = new_message.get("content")
# Skip messages with no content
if content is None:
updated_messages.append(new_message)
continue
# Handle string content
if isinstance(content, str):
if masking_index < len(masked_texts):
new_message["content"] = masked_texts[masking_index]
masking_index += 1
# Handle list content
elif isinstance(content, list):
new_message["content"], masking_index = self._mask_content_list(
content_list=content,
masked_texts=masked_texts,
masking_index=masking_index,
)
updated_messages.append(new_message)
return updated_messages
def _mask_content_list(
self, content_list: List[Any], masked_texts: List[str], masking_index: int
) -> Tuple[List[Any], int]:
"""
Apply masking to a list of content items.
Args:
content_list: List of content items
masked_texts: List of masked text strings
starting_index: Starting index in the masked_texts list
Returns:
Updated content list with masked items
"""
new_content: List[Union[dict, str]] = []
for item in content_list:
if isinstance(item, dict) and "text" in item:
new_item = item.copy()
if masking_index < len(masked_texts):
new_item["text"] = masked_texts[masking_index]
masking_index += 1
new_content.append(new_item)
elif isinstance(item, str):
if masking_index < len(masked_texts):
item = masked_texts[masking_index]
masking_index += 1
if item is not None:
new_content.append(item)
return new_content, masking_index
def get_content_for_message(self, message: AllMessageValues) -> Optional[List[str]]:
"""
Get the content for a message.
For bedrock guardrails we create a list of all the text content in the message.
If a message has a list of content items, we flatten the list and return a list of text content.
"""
message_text_content = []
content = message.get("content")
if content is None:
return None
if isinstance(content, str):
message_text_content.append(content)
elif isinstance(content, list):
for item in content:
if isinstance(item, dict) and "text" in item:
message_text_content.append(item["text"])
elif isinstance(item, str):
message_text_content.append(item)
return message_text_content

View file

@ -27,6 +27,8 @@ def initialize_bedrock(litellm_params, guardrail):
guardrailIdentifier=litellm_params["guardrailIdentifier"],
guardrailVersion=litellm_params["guardrailVersion"],
default_on=litellm_params["default_on"],
mask_request_content=litellm_params.get("mask_request_content", None),
mask_response_content=litellm_params.get("mask_response_content", None),
)
litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback)

View file

@ -6,4 +6,18 @@ model_list:
litellm_settings:
callbacks: ["resend_email"]
callbacks: ["resend_email"]
guardrails:
- guardrail_name: "bedrock-pre-guard"
litellm_params:
guardrail: bedrock # supported values: "aporia", "bedrock", "lakera"
mode: "pre_call"
guardrailIdentifier: wf0hkdb5x07f # your guardrail ID on bedrock
guardrailVersion: "DRAFT" # your guardrail version on bedrock
mask_request_content: true
mask_response_content: true
general_settings:
store_model_in_db: true
store_prompts_in_spend_logs: true

View file

@ -105,6 +105,14 @@ class LitellmParams(TypedDict):
guard_name: Optional[str]
default_on: Optional[bool]
# PII control params
mask_request_content: Optional[
bool
] # will mask request content if guardrail makes any changes
mask_response_content: Optional[
bool
] # will mask response content if guardrail makes any changes
class Guardrail(TypedDict, total=False):
guardrail_name: str
@ -123,19 +131,6 @@ class GuardrailEventHooks(str, Enum):
logging_only = "logging_only"
class BedrockTextContent(TypedDict, total=False):
text: str
class BedrockContentItem(TypedDict, total=False):
text: BedrockTextContent
class BedrockRequest(TypedDict, total=False):
source: Literal["INPUT", "OUTPUT"]
content: List[BedrockContentItem]
class DynamicGuardrailParams(TypedDict):
extra_body: Dict[str, Any]

View file

@ -0,0 +1,121 @@
from typing import Any, Dict, List, Literal, Optional, TypedDict, Union
class BedrockTextContent(TypedDict, total=False):
text: str
class BedrockContentItem(TypedDict, total=False):
text: BedrockTextContent
class BedrockRequest(TypedDict, total=False):
source: Literal["INPUT", "OUTPUT"]
content: List[BedrockContentItem]
class BedrockGuardrailUsage(TypedDict, total=False):
topicPolicyUnits: Optional[int]
contentPolicyUnits: Optional[int]
wordPolicyUnits: Optional[int]
sensitiveInformationPolicyUnits: Optional[int]
sensitiveInformationPolicyFreeUnits: Optional[int]
contextualGroundingPolicyUnits: Optional[int]
class BedrockGuardrailOutput(TypedDict, total=False):
text: Optional[str]
class BedrockGuardrailTopicPolicyItem(TypedDict, total=False):
name: Optional[str]
type: Optional[str]
action: Optional[str]
class BedrockGuardrailTopicPolicy(TypedDict, total=False):
topics: List[BedrockGuardrailTopicPolicyItem]
class BedrockGuardrailContentPolicyFilter(TypedDict, total=False):
type: Optional[str]
confidence: Optional[str]
filterStrength: Optional[str]
action: Optional[str]
class BedrockGuardrailContentPolicy(TypedDict, total=False):
filters: List[BedrockGuardrailContentPolicyFilter]
class BedrockGuardrailWordPolicyCustomWord(TypedDict, total=False):
match: str
action: str
class BedrockGuardrailWordPolicyManagedWord(TypedDict, total=False):
match: Optional[str]
type: Optional[str] # Note: There might be more types
action: Optional[str]
class BedrockGuardrailWordPolicy(TypedDict, total=False):
customWords: List[BedrockGuardrailWordPolicyCustomWord]
managedWordLists: List[BedrockGuardrailWordPolicyManagedWord]
class BedrockGuardrailPiiEntity(TypedDict, total=False):
type: Optional[str] # Many PII types available per AWS docs
match: Optional[str]
action: Optional[str]
class BedrockGuardrailRegex(TypedDict, total=False):
name: Optional[str]
regex: Optional[str]
match: Optional[str]
action: Optional[str]
class BedrockGuardrailSensitiveInformationPolicy(TypedDict, total=False):
piiEntities: Optional[List[BedrockGuardrailPiiEntity]]
regexes: Optional[List[BedrockGuardrailRegex]]
class BedrockGuardrailContextualGroundingFilter(TypedDict, total=False):
type: Optional[str]
threshold: Optional[float]
score: Optional[float]
action: Optional[str]
class BedrockGuardrailContextualGroundingPolicy(TypedDict, total=False):
filters: List[BedrockGuardrailContextualGroundingFilter]
class BedrockGuardrailCoverage(TypedDict, total=False):
textCharacters: Dict[str, int]
class BedrockGuardrailInvocationMetrics(TypedDict, total=False):
guardrailProcessingLatency: int
usage: BedrockGuardrailUsage
guardrailCoverage: BedrockGuardrailCoverage
class BedrockGuardrailAssessment(TypedDict, total=False):
topicPolicy: Optional[BedrockGuardrailTopicPolicy]
contentPolicy: Optional[BedrockGuardrailContentPolicy]
wordPolicy: Optional[BedrockGuardrailWordPolicy]
sensitiveInformationPolicy: Optional[BedrockGuardrailSensitiveInformationPolicy]
contextualGroundingPolicy: Optional[BedrockGuardrailContextualGroundingPolicy]
invocationMetrics: BedrockGuardrailInvocationMetrics
guardrailCoverage: BedrockGuardrailCoverage
class BedrockGuardrailResponse(TypedDict, total=False):
usage: Optional[BedrockGuardrailUsage]
action: Optional[str]
output: Optional[List[BedrockGuardrailOutput]]
outputs: Optional[List[BedrockGuardrailOutput]]
assessments: Optional[List[BedrockGuardrailAssessment]]

View file

@ -0,0 +1,79 @@
import sys
import os
import io, asyncio
import pytest
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail
@pytest.mark.asyncio
async def test_bedrock_guardrails():
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailVersion="DRAFT",
mask_request_content=True,
)
request_data = {
"model": "gpt-4o",
"messages": [
{"role": "user", "content": "Hello, my phone number is +1 412 555 1212"},
{"role": "assistant", "content": "Hello, how can I help you today?"},
{"role": "user", "content": "I need to cancel my order"},
{"role": "user", "content": "ok, my credit card number is 1234-5678-9012-3456"},
],
}
response = await guardrail.async_moderation_hook(
data=request_data,
user_api_key_dict={},
call_type="completion"
)
print(response)
assert response["messages"][0]["content"] == "Hello, my phone number is {PHONE}"
assert response["messages"][1]["content"] == "Hello, how can I help you today?"
assert response["messages"][2]["content"] == "I need to cancel my order"
assert response["messages"][3]["content"] == "ok, my credit card number is {CREDIT_DEBIT_CARD_NUMBER}"
@pytest.mark.asyncio
async def test_bedrock_guardrails_content_list():
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailVersion="DRAFT",
mask_request_content=True,
)
request_data = {
"model": "gpt-4o",
"messages": [
{"role": "user", "content": [
{"type": "text", "text": "Hello, my phone number is +1 412 555 1212"},
{"type": "text", "text": "what time is it?"},
]},
{"role": "assistant", "content": "Hello, how can I help you today?"},
{
"role": "user",
"content": "who is the president of the united states?"
}
],
}
response = await guardrail.async_moderation_hook(
data=request_data,
user_api_key_dict={},
call_type="completion"
)
print(response)
# Verify that the list content is properly masked
assert isinstance(response["messages"][0]["content"], list)
assert response["messages"][0]["content"][0]["text"] == "Hello, my phone number is {PHONE}"
assert response["messages"][0]["content"][1]["text"] == "what time is it?"
assert response["messages"][1]["content"] == "Hello, how can I help you today?"
assert response["messages"][2]["content"] == "who is the president of the united states?"