From 6a3ea83b23adf0645b1e48d33aa5f3a17bb95229 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 7 May 2025 18:30:57 -0700 Subject: [PATCH] [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 --- .../docs/proxy/guardrails/bedrock.md | 39 ++ litellm/integrations/custom_guardrail.py | 6 + .../guardrail_hooks/bedrock_guardrails.py | 346 ++++++++++++++++-- .../guardrails/guardrail_initializers.py | 2 + litellm/proxy/proxy_config.yaml | 16 +- litellm/types/guardrails.py | 21 +- .../guardrail_hooks/bedrock_guardrails.py | 121 ++++++ .../test_bedrock_guardrails.py | 79 ++++ 8 files changed, 587 insertions(+), 43 deletions(-) create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py create mode 100644 tests/guardrails_tests/test_bedrock_guardrails.py diff --git a/docs/my-website/docs/proxy/guardrails/bedrock.md b/docs/my-website/docs/proxy/guardrails/bedrock.md index 81c561fcfc0..017a1e76c31 100644 --- a/docs/my-website/docs/proxy/guardrails/bedrock.md +++ b/docs/my-website/docs/proxy/guardrails/bedrock.md @@ -135,3 +135,42 @@ curl -i http://localhost:4000/v1/chat/completions \ +## 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. + diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 41a3800116e..3737422ce19 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 5c6b53be251..408a451464c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index c32d75f9868..5b5ab23d7f2 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -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) diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 39e7cfd07a7..cc7f15b8f65 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -6,4 +6,18 @@ model_list: litellm_settings: - callbacks: ["resend_email"] \ No newline at end of file + 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 \ No newline at end of file diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index b7018fe29f9..b4ae6606034 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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] diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py new file mode 100644 index 00000000000..2bef70c381d --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -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]] diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py new file mode 100644 index 00000000000..9252a0dffc2 --- /dev/null +++ b/tests/guardrails_tests/test_bedrock_guardrails.py @@ -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?" + +