mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Add Support for Bedrock Guardrails to supportive selective Guarding (#14575)
* Add Support for Bedrock Guardrails to supportive selective Guarding * Add method for better handling * Add guarded_text content type * Add guarded_text content type * Update Dockerfile * Update Dockerfile
This commit is contained in:
parent
7de8811c4c
commit
ab1fb2b2e7
6 changed files with 427 additions and 81 deletions
|
|
@ -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.
|
||||
:::
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="LiteLLM SDK">
|
||||
|
||||
|
|
@ -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"
|
||||
}
|
||||
)
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="Proxy on request">
|
||||
|
|
@ -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)
|
||||
```
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
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"
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue