diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index c191b742268..165ef1d12f7 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -889,6 +889,19 @@ curl http://0.0.0.0:4000/v1/chat/completions \ Example of using [Bedrock Guardrails with LiteLLM](https://docs.aws.amazon.com/bedrock/latest/userguide/guardrails-use-converse-api.html) +### Selective Content Moderation with `guarded_text` + +LiteLLM supports selective content moderation using the `guarded_text` content type. This allows you to wrap only specific content that should be moderated by Bedrock Guardrails, rather than evaluating the entire conversation. + +**How it works:** +- Content with `type: "guarded_text"` gets automatically wrapped in `guardrailConverseContent` blocks +- Only the wrapped content is evaluated by Bedrock Guardrails +- Regular content with `type: "text"` bypasses guardrail evaluation + +:::note +If `guarded_text` is not used, the entire conversation history will be sent to the guardrail for evaluation, which can increase latency and costs. +::: + @@ -915,6 +928,24 @@ response = completion( "trace": "disabled", # The trace behavior for the guardrail. Can either be "disabled" or "enabled" }, ) + +# Selective guardrail usage with guarded_text - only specific content is evaluated +response_guard = completion( + model="anthropic.claude-v2", + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is the main topic of this legal document?"}, + {"type": "guarded_text", "text": "This document contains sensitive legal information that should be moderated by guardrails."} + ] + } + ], + guardrailConfig={ + "guardrailIdentifier": "gr-abc123", + "guardrailVersion": "DRAFT" + } +) ``` @@ -993,7 +1024,20 @@ response = client.chat.completions.create(model="bedrock-claude-v1", messages = temperature=0.7 ) -print(response) +# For adding selective guardrail usage with guarded_text +response_guard = client.chat.completions.create(model="bedrock-claude-v1", messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is the main topic of this legal document?"}, + {"type": "guarded_text", "text": "This document contains sensitive legal information that should be moderated by guardrails."} + ] + } +], +temperature=0.7 +) + +print(response_guard) ``` diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 65f49cf08b8..356d48dcb89 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -16,8 +16,8 @@ from litellm import verbose_logger from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client from litellm.types.files import get_file_extension_from_mime_type from litellm.types.llms.anthropic import * -from litellm.types.llms.bedrock import MessageBlock as BedrockMessageBlock from litellm.types.llms.bedrock import CachePointBlock +from litellm.types.llms.bedrock import MessageBlock as BedrockMessageBlock from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.llms.ollama import OllamaVisionModelObject from litellm.types.llms.openai import ( @@ -1067,10 +1067,10 @@ def convert_to_gemini_tool_call_invoke( if tool_calls is not None: for tool in tool_calls: if "function" in tool: - gemini_function_call: Optional[VertexFunctionCall] = ( - _gemini_tool_call_invoke_helper( - function_call_params=tool["function"] - ) + gemini_function_call: Optional[ + VertexFunctionCall + ] = _gemini_tool_call_invoke_helper( + function_call_params=tool["function"] ) if gemini_function_call is not None: _parts_list.append( @@ -1589,9 +1589,9 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_element["cache_control"] = ( - _content_element["cache_control"] - ) + _anthropic_content_element[ + "cache_control" + ] = _content_element["cache_control"] user_content.append(_anthropic_content_element) elif m.get("type", "") == "text": m = cast(ChatCompletionTextObject, m) @@ -1629,9 +1629,9 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_text_element["cache_control"] = ( - _content_element["cache_control"] - ) + _anthropic_content_text_element[ + "cache_control" + ] = _content_element["cache_control"] user_content.append(_anthropic_content_text_element) @@ -2482,8 +2482,7 @@ class BedrockImageProcessor: if is_document: return BedrockImageProcessor._get_document_format( - mime_type=mime_type, - supported_doc_formats=supported_doc_formats + mime_type=mime_type, supported_doc_formats=supported_doc_formats ) else: @@ -2495,12 +2494,9 @@ class BedrockImageProcessor: f"Unsupported image format: {image_format}. Supported formats: {supported_image_and_video_formats}" ) return image_format - + @staticmethod - def _get_document_format( - mime_type: str, - supported_doc_formats: List[str] - ) -> str: + def _get_document_format(mime_type: str, supported_doc_formats: List[str]) -> str: """ Get the document format from the mime type @@ -2519,13 +2515,9 @@ class BedrockImageProcessor: The document format """ valid_extensions: Optional[List[str]] = None - potential_extensions = mimetypes.guess_all_extensions( - mime_type, strict=False - ) + potential_extensions = mimetypes.guess_all_extensions(mime_type, strict=False) valid_extensions = [ - ext[1:] - for ext in potential_extensions - if ext[1:] in supported_doc_formats + ext[1:] for ext in potential_extensions if ext[1:] in supported_doc_formats ] # Fallback to types/files.py if mimetypes doesn't return valid extensions @@ -2689,10 +2681,12 @@ def _convert_to_bedrock_tool_call_invoke( ) bedrock_content_block = BedrockContentBlock(toolUse=bedrock_tool) _parts_list.append(bedrock_content_block) - + # Check for cache_control and add a separate cachePoint block if tool.get("cache_control", None) is not None: - cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default")) + cache_point_block = BedrockContentBlock( + cachePoint=CachePointBlock(type="default") + ) _parts_list.append(cache_point_block) return _parts_list except Exception as e: @@ -2754,7 +2748,7 @@ def _convert_to_bedrock_tool_call_result( for content in content_list: if content["type"] == "text": content_str += content["text"] - + message.get("name", "") id = str(message.get("tool_call_id", str(uuid.uuid4()))) @@ -2763,7 +2757,7 @@ def _convert_to_bedrock_tool_call_result( content=[tool_result_content_block], toolUseId=id, ) - + content_block = BedrockContentBlock(toolResult=tool_result) return content_block @@ -3085,6 +3079,7 @@ class BedrockConverseMessagesProcessor: messages.append(DEFAULT_USER_CONTINUE_MESSAGE) return messages + @staticmethod async def _bedrock_converse_messages_pt_async( # noqa: PLR0915 messages: List, @@ -3128,6 +3123,12 @@ class BedrockConverseMessagesProcessor: if element["type"] == "text": _part = BedrockContentBlock(text=element["text"]) _parts.append(_part) + elif element["type"] == "guarded_text": + # Wrap guarded_text in guardrailConverseContent block + _part = BedrockContentBlock( + guardrailConverseContent={"text": element["text"]} + ) + _parts.append(_part) elif element["type"] == "image_url": format: Optional[str] = None if isinstance(element["image_url"], dict): @@ -3170,6 +3171,7 @@ class BedrockConverseMessagesProcessor: msg_i += 1 if user_content: + if len(contents) > 0 and contents[-1]["role"] == "user": if ( assistant_continue_message is not None @@ -3199,26 +3201,29 @@ class BedrockConverseMessagesProcessor: current_message = messages[msg_i] tool_call_result = _convert_to_bedrock_tool_call_result(current_message) tool_content.append(tool_call_result) - + # Check if we need to add a separate cachePoint block has_cache_control = False - + # Check for message-level cache_control if current_message.get("cache_control", None) is not None: has_cache_control = True # Check for content-level cache_control in list content elif isinstance(current_message.get("content"), list): for content_element in current_message["content"]: - if (isinstance(content_element, dict) and - content_element.get("cache_control", None) is not None): + if ( + isinstance(content_element, dict) + and content_element.get("cache_control", None) is not None + ): has_cache_control = True break - + # Add a separate cachePoint block if cache_control is present if has_cache_control: - cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default")) + cache_point_block = BedrockContentBlock( + cachePoint=CachePointBlock(type="default") + ) tool_content.append(cache_point_block) - msg_i += 1 if tool_content: @@ -3299,7 +3304,7 @@ class BedrockConverseMessagesProcessor: image_url=image_url ) assistants_parts.append(assistants_part) - # Add cache point block for assistant content elements + # Add cache point block for assistant content elements _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast( @@ -3311,8 +3316,12 @@ class BedrockConverseMessagesProcessor: if _cache_point_block is not None: assistants_parts.append(_cache_point_block) assistant_content.extend(assistants_parts) - elif _assistant_content is not None and isinstance(_assistant_content, str): - assistant_content.append(BedrockContentBlock(text=_assistant_content)) + elif _assistant_content is not None and isinstance( + _assistant_content, str + ): + assistant_content.append( + BedrockContentBlock(text=_assistant_content) + ) # Add cache point block for assistant string content _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( @@ -3496,6 +3505,12 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 if element["type"] == "text": _part = BedrockContentBlock(text=element["text"]) _parts.append(_part) + elif element["type"] == "guarded_text": + # Wrap guarded_text in guardrailConverseContent block + _part = BedrockContentBlock( + guardrailConverseContent={"text": element["text"]} + ) + _parts.append(_part) elif element["type"] == "image_url": format: Optional[str] = None if isinstance(element["image_url"], dict): @@ -3539,6 +3554,7 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 msg_i += 1 if user_content: + if len(contents) > 0 and contents[-1]["role"] == "user": if ( assistant_continue_message is not None @@ -3565,29 +3581,33 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 while msg_i < len(messages) and messages[msg_i]["role"] == "tool": tool_call_result = _convert_to_bedrock_tool_call_result(messages[msg_i]) current_message = messages[msg_i] - + # Add the tool result first tool_content.append(tool_call_result) - + # Check if we need to add a separate cachePoint block has_cache_control = False - + # Check for message-level cache_control if current_message.get("cache_control", None) is not None: has_cache_control = True # Check for content-level cache_control in list content elif isinstance(current_message.get("content"), list): for content_element in current_message["content"]: - if (isinstance(content_element, dict) and - content_element.get("cache_control", None) is not None): + if ( + isinstance(content_element, dict) + and content_element.get("cache_control", None) is not None + ): has_cache_control = True break - + # Add a separate cachePoint block if cache_control is present if has_cache_control: - cache_point_block = BedrockContentBlock(cachePoint=CachePointBlock(type="default")) + cache_point_block = BedrockContentBlock( + cachePoint=CachePointBlock(type="default") + ) tool_content.append(cache_point_block) - + msg_i += 1 if tool_content: # if last message was a 'user' message, then add a blank assistant message (bedrock requires alternating roles) @@ -3852,10 +3872,9 @@ def function_call_prompt(messages: list, functions: list): if isinstance(message["content"], str): message["content"] += f""" {function_prompt}""" else: - message["content"].append({ - "type": "text", - "text": f""" {function_prompt}""" - }) + message["content"].append( + {"type": "text", "text": f""" {function_prompt}"""} + ) function_added_to_prompt = True if function_added_to_prompt is False: diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index e3d65be8bbf..88103a86cf3 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -501,7 +501,6 @@ class AmazonConverseConfig(BaseConfig): ) and not is_thinking_enabled ): - optional_params["tool_choice"] = ToolChoiceValuesBlock( tool=SpecificToolChoiceBlock(name=RESPONSE_FORMAT_TOOL_NAME) ) @@ -995,7 +994,9 @@ class AmazonConverseConfig(BaseConfig): return message, returned_finish_reason - def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[ + def _translate_message_content( + self, content_blocks: List[ContentBlock] + ) -> Tuple[ str, List[ChatCompletionToolCallChunk], Optional[List[BedrockConverseReasoningContentBlock]], @@ -1010,9 +1011,9 @@ class AmazonConverseConfig(BaseConfig): """ content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( - None - ) + reasoningContentBlocks: Optional[ + List[BedrockConverseReasoningContentBlock] + ] = None for idx, content in enumerate(content_blocks): """ - Content is either a tool response or text @@ -1133,9 +1134,9 @@ class AmazonConverseConfig(BaseConfig): chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( - None - ) + reasoningContentBlocks: Optional[ + List[BedrockConverseReasoningContentBlock] + ] = None if message is not None: ( @@ -1148,12 +1149,12 @@ class AmazonConverseConfig(BaseConfig): chat_completion_message["provider_specific_fields"] = { "reasoningContentBlocks": reasoningContentBlocks, } - chat_completion_message["reasoning_content"] = ( - self._transform_reasoning_content(reasoningContentBlocks) - ) - chat_completion_message["thinking_blocks"] = ( - self._transform_thinking_blocks(reasoningContentBlocks) - ) + chat_completion_message[ + "reasoning_content" + ] = self._transform_reasoning_content(reasoningContentBlocks) + chat_completion_message[ + "thinking_blocks" + ] = self._transform_thinking_blocks(reasoningContentBlocks) chat_completion_message["content"] = content_str if ( json_mode is True @@ -1171,7 +1172,6 @@ class AmazonConverseConfig(BaseConfig): # Bedrock returns the response wrapped in a "properties" object # We need to extract the actual content from this wrapper try: - response_data = json.loads(json_mode_content_str) # If Bedrock wrapped the response in "properties", extract the content diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index baa7c205204..a829a6b94b9 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -3,14 +3,9 @@ from typing import Any, List, Literal, Optional, Union from typing_extensions import ( TYPE_CHECKING, - Protocol, Required, - Self, TypedDict, - TypeGuard, - get_origin, override, - runtime_checkable, ) from .openai import ChatCompletionToolCallChunk @@ -93,6 +88,12 @@ class BedrockConverseReasoningContentBlockDelta(TypedDict, total=False): text: str +class GuardrailConverseContentBlock(TypedDict, total=False): + """Content block for selective guardrail evaluation in Bedrock Converse API""" + + text: str + + class ContentBlock(TypedDict, total=False): text: str image: ImageBlock @@ -102,6 +103,7 @@ class ContentBlock(TypedDict, total=False): toolUse: ToolUseBlock cachePoint: CachePointBlock reasoningContent: BedrockConverseReasoningContentBlock + guardrailConverseContent: GuardrailConverseContentBlock class MessageBlock(TypedDict): @@ -581,30 +583,35 @@ class AmazonDeepSeekR1StreamingResponse(TypedDict): class BedrockS3InputDataConfig(TypedDict): """S3 input data configuration for Bedrock batch jobs.""" + s3Uri: str class BedrockInputDataConfig(TypedDict): """Input data configuration for Bedrock batch jobs.""" + s3InputDataConfig: BedrockS3InputDataConfig class BedrockS3OutputDataConfig(TypedDict): """S3 output data configuration for Bedrock batch jobs.""" + s3Uri: str class BedrockOutputDataConfig(TypedDict): """Output data configuration for Bedrock batch jobs.""" + s3OutputDataConfig: BedrockS3OutputDataConfig class BedrockCreateBatchRequest(TypedDict, total=False): """ Request structure for creating a Bedrock batch inference job. - + Reference: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_CreateModelInvocationJob.html """ + jobName: str roleArn: str modelId: str @@ -616,21 +623,17 @@ class BedrockCreateBatchRequest(TypedDict, total=False): BedrockBatchJobStatus = Literal[ - "Submitted", - "InProgress", - "Completed", - "Failed", - "Stopping", - "Stopped" + "Submitted", "InProgress", "Completed", "Failed", "Stopping", "Stopped" ] class BedrockCreateBatchResponse(TypedDict): """ Response structure from creating a Bedrock batch inference job. - + Reference: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_CreateModelInvocationJob.html """ + jobArn: str jobName: str status: BedrockBatchJobStatus @@ -639,9 +642,10 @@ class BedrockCreateBatchResponse(TypedDict): class BedrockGetBatchResponse(TypedDict, total=False): """ Response structure from getting a Bedrock batch inference job. - + Reference: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_GetModelInvocationJob.html """ + jobArn: str jobName: str modelId: str diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 79e3c73dbc1..4adb751d905 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -723,6 +723,7 @@ ValidUserMessageContentTypes = [ "input_audio", "audio_url", "document", + "guarded_text", "video_url", "file", ] # used for validating user messages. Prevent users from accidentally sending anthropic messages. diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index 2fc710664e6..df003850f70 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -1589,4 +1589,282 @@ async def test_no_cache_control_no_cache_point(): # Tool message should only have tool result, no cachePoint tool_content = result[2]["content"] assert len(tool_content) == 1 - assert "toolResult" in tool_content[0] \ No newline at end of file + assert "toolResult" in tool_content[0] + + +# ============================================================================ +# Guarded Text Feature Tests +# ============================================================================ + +def test_guarded_text_wraps_in_guardrail_converse_content(): + """Test that guarded_text content type gets wrapped in guardrailConverseContent blocks.""" + from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Regular text content"}, + {"type": "guarded_text", "text": "This should be guarded"}, + {"type": "text", "text": "More regular text"} + ] + } + ] + + result = _bedrock_converse_messages_pt( + messages=messages, + model="us.amazon.nova-pro-v1:0", + llm_provider="bedrock_converse" + ) + + # Should have 1 message + assert len(result) == 1 + assert result[0]["role"] == "user" + + # Should have 3 content blocks + content = result[0]["content"] + assert len(content) == 3 + + # First and third should be regular text + assert "text" in content[0] + assert content[0]["text"] == "Regular text content" + assert "text" in content[2] + assert content[2]["text"] == "More regular text" + + # Second should be guardrailConverseContent + assert "guardrailConverseContent" in content[1] + assert content[1]["guardrailConverseContent"]["text"] == "This should be guarded" + + +def test_guarded_text_with_system_messages(): + """Test guarded_text with system messages using the full transformation.""" + config = AmazonConverseConfig() + + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + { + "role": "user", + "content": [ + {"type": "text", "text": "What is the main topic of this legal document?"}, + {"type": "guarded_text", "text": "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question."} + ] + } + ] + + optional_params = { + "guardrailConfig": { + "guardrailIdentifier": "gr-abc123", + "guardrailVersion": "DRAFT" + } + } + + result = config._transform_request( + model="us.amazon.nova-pro-v1:0", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={} + ) + + # Should have system content blocks + assert "system" in result + assert len(result["system"]) == 1 + assert result["system"][0]["text"] == "You are a helpful assistant." + + # Should have 1 message (system messages are removed) + assert "messages" in result + assert len(result["messages"]) == 1 + + # User message should have both regular text and guarded text + user_message = result["messages"][0] + assert user_message["role"] == "user" + content = user_message["content"] + assert len(content) == 2 + + # First should be regular text + assert "text" in content[0] + assert content[0]["text"] == "What is the main topic of this legal document?" + + # Second should be guardrailConverseContent + assert "guardrailConverseContent" in content[1] + assert content[1]["guardrailConverseContent"]["text"] == "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question." + + +def test_guarded_text_with_mixed_content_types(): + """Test guarded_text with mixed content types including images.""" + from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Look at this image"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,test"}}, + {"type": "guarded_text", "text": "This sensitive content should be guarded"} + ] + } + ] + + result = _bedrock_converse_messages_pt( + messages=messages, + model="us.amazon.nova-pro-v1:0", + llm_provider="bedrock_converse" + ) + + # Should have 1 message + assert len(result) == 1 + assert result[0]["role"] == "user" + + # Should have 3 content blocks + content = result[0]["content"] + assert len(content) == 3 + + # First should be regular text + assert "text" in content[0] + assert content[0]["text"] == "Look at this image" + + # Second should be image + assert "image" in content[1] + + # Third should be guardrailConverseContent + assert "guardrailConverseContent" in content[2] + assert content[2]["guardrailConverseContent"]["text"] == "This sensitive content should be guarded" + + +@pytest.mark.asyncio +async def test_async_guarded_text(): + """Test async version of guarded_text processing.""" + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello"}, + {"type": "guarded_text", "text": "This should be guarded"} + ] + } + ] + + result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="us.amazon.nova-pro-v1:0", + llm_provider="bedrock_converse" + ) + + # Should have 1 message + assert len(result) == 1 + assert result[0]["role"] == "user" + + # Should have 2 content blocks + content = result[0]["content"] + assert len(content) == 2 + + # First should be regular text + assert "text" in content[0] + assert content[0]["text"] == "Hello" + + # Second should be guardrailConverseContent + assert "guardrailConverseContent" in content[1] + assert content[1]["guardrailConverseContent"]["text"] == "This should be guarded" + + +def test_guarded_text_with_tool_calls(): + """Test guarded_text with tool calls in the conversation.""" + from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What's the weather?"}, + {"type": "guarded_text", "text": "Please be careful with sensitive information"} + ] + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": {"name": "get_weather", "arguments": "{}"} + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_123", + "content": "It's sunny and 25°C" + } + ] + + result = _bedrock_converse_messages_pt( + messages=messages, + model="us.amazon.nova-pro-v1:0", + llm_provider="bedrock_converse" + ) + + # Should have 3 messages + assert len(result) == 3 + + # First message (user) should have both text and guarded_text + user_message = result[0] + assert user_message["role"] == "user" + content = user_message["content"] + assert len(content) == 2 + + # First should be regular text + assert "text" in content[0] + assert content[0]["text"] == "What's the weather?" + + # Second should be guardrailConverseContent + assert "guardrailConverseContent" in content[1] + assert content[1]["guardrailConverseContent"]["text"] == "Please be careful with sensitive information" + + # Other messages should not have guardrailConverseContent + for i in range(1, 3): + content = result[i]["content"] + for block in content: + assert "guardrailConverseContent" not in block + + +def test_guarded_text_guardrail_config_preserved(): + """Test that guardrailConfig is preserved when using guarded_text.""" + config = AmazonConverseConfig() + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello"}, + {"type": "guarded_text", "text": "This should be guarded"} + ] + } + ] + + optional_params = { + "guardrailConfig": { + "guardrailIdentifier": "gr-abc123", + "guardrailVersion": "DRAFT" + } + } + + result = config._transform_request( + model="us.amazon.nova-pro-v1:0", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={} + ) + + # GuardrailConfig should be present at top level + assert "guardrailConfig" in result + assert result["guardrailConfig"]["guardrailIdentifier"] == "gr-abc123" + + # GuardrailConfig should also be in inferenceConfig + assert "inferenceConfig" in result + assert "guardrailConfig" in result["inferenceConfig"] + assert result["inferenceConfig"]["guardrailConfig"]["guardrailIdentifier"] == "gr-abc123" + +