Merge pull request #13640 from yytdfc/fix_bedrock_epc

[Bug Fix] Add cachePoint support for assistant and tool messages in Bedrock
This commit is contained in:
Krish Dholakia 2025-08-16 01:20:28 -07:00 • committed by GitHub
commit 1b2ec16eee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 336 additions and 2 deletions

View file

@ -17,6 +17,7 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_
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.custom_http import httpxSpecialProvider
from litellm.types.llms.ollama import OllamaVisionModelObject
from litellm.types.llms.openai import (
@ -2685,6 +2686,11 @@ 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"))
_parts_list.append(cache_point_block)
return _parts_list
except Exception as e:
raise Exception(
@ -2745,6 +2751,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())))
@ -2753,6 +2760,7 @@ def _convert_to_bedrock_tool_call_result(
content=[tool_result_content_block],
toolUseId=id,
)
content_block = BedrockContentBlock(toolResult=tool_result)
return content_block
@ -3516,8 +3524,30 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
tool_content: List[BedrockContentBlock] = []
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):
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"))
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)
@ -3589,9 +3619,28 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
image_url=image_url
)
assistants_parts.append(assistants_part)
# Add cache point block for assistant content elements
_cache_point_block = (
litellm.AmazonConverseConfig()._get_cache_point_block(
message_block=cast(
OpenAIMessageContentListBlock, element
),
block_type="content_block",
)
)
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))
# Add cache point block for assistant string content
_cache_point_block = (
litellm.AmazonConverseConfig()._get_cache_point_block(
assistant_message_block, block_type="content_block"
)
)
if _cache_point_block is not None:
assistant_content.append(_cache_point_block)
_tool_calls = assistant_message_block.get("tool_calls", [])
if _tool_calls:
assistant_content.extend(

View file

@ -25,6 +25,7 @@ from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMExcepti
from litellm.types.llms.bedrock import *
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionAssistantMessage,
ChatCompletionRedactedThinkingBlock,
ChatCompletionResponseMessage,
ChatCompletionSystemMessage,
@ -505,6 +506,7 @@ class AmazonConverseConfig(BaseConfig):
OpenAIMessageContentListBlock,
ChatCompletionUserMessage,
ChatCompletionSystemMessage,
ChatCompletionAssistantMessage,
],
block_type: Literal["system"],
) -> Optional[SystemContentBlock]:
@ -517,6 +519,7 @@ class AmazonConverseConfig(BaseConfig):
OpenAIMessageContentListBlock,
ChatCompletionUserMessage,
ChatCompletionSystemMessage,
ChatCompletionAssistantMessage,
],
block_type: Literal["content_block"],
) -> Optional[ContentBlock]:
@ -528,6 +531,7 @@ class AmazonConverseConfig(BaseConfig):
OpenAIMessageContentListBlock,
ChatCompletionUserMessage,
ChatCompletionSystemMessage,
ChatCompletionAssistantMessage,
],
block_type: Literal["system", "content_block"],
) -> Optional[Union[SystemContentBlock, ContentBlock]]:

View file

@ -935,4 +935,285 @@ def test_transform_request_with_function_tool():
assert "toolConfig" in request_data
assert "tools" in request_data["toolConfig"]
assert len(request_data["toolConfig"]["tools"]) == 1
assert request_data["toolConfig"]["tools"][0]["toolSpec"]["name"] == "get_weather"
assert request_data["toolConfig"]["tools"][0]["toolSpec"]["name"] == "get_weather"
def test_assistant_message_cache_control():
"""Test that assistant messages with cache_control generate cachePoint blocks."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
# Test assistant message with string content and cache_control
messages = [
{"role": "user", "content": "Hello"},
{
"role": "assistant",
"content": "Hi there!",
"cache_control": {"type": "ephemeral"}
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_provider="bedrock_converse"
)
# Should have user message and assistant message
assert len(result) == 2
assert result[0]["role"] == "user"
assert result[1]["role"] == "assistant"
# Assistant message should have text content and cachePoint
assistant_content = result[1]["content"]
assert len(assistant_content) == 2
assert assistant_content[0]["text"] == "Hi there!"
assert "cachePoint" in assistant_content[1]
assert assistant_content[1]["cachePoint"]["type"] == "default"
def test_assistant_message_list_content_cache_control():
"""Test assistant messages with list content and cache_control."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
messages = [
{"role": "user", "content": "Hello"},
{
"role": "assistant",
"content": [
{
"type": "text",
"text": "This should be cached",
"cache_control": {"type": "ephemeral"}
}
]
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_provider="bedrock_converse"
)
# Assistant message should have text content and cachePoint
assistant_content = result[1]["content"]
assert len(assistant_content) == 2
assert assistant_content[0]["text"] == "This should be cached"
assert "cachePoint" in assistant_content[1]
assert assistant_content[1]["cachePoint"]["type"] == "default"
def test_tool_message_cache_control():
"""Test that tool messages with cache_control generate cachePoint blocks."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"}
}
]
},
{
"role": "tool",
"tool_call_id": "call_123",
"content": [
{
"type": "text",
"text": "Weather data: sunny, 25°C",
"cache_control": {"type": "ephemeral"}
}
]
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_provider="bedrock_converse"
)
# Should have user, assistant, and user (tool results) messages
assert len(result) == 3
# Last message should contain tool result and cachePoint
tool_message_content = result[2]["content"]
assert len(tool_message_content) == 2
# First should be tool result
assert "toolResult" in tool_message_content[0]
assert tool_message_content[0]["toolResult"]["content"][0]["text"] == "Weather data: sunny, 25°C"
# Second should be cachePoint
assert "cachePoint" in tool_message_content[1]
assert tool_message_content[1]["cachePoint"]["type"] == "default"
def test_tool_message_string_content_cache_control():
"""Test tool messages with string content and message-level cache_control."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"}
}
]
},
{
"role": "tool",
"tool_call_id": "call_123",
"content": "Weather: sunny, 25°C",
"cache_control": {"type": "ephemeral"}
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_provider="bedrock_converse"
)
# Last message should contain tool result and cachePoint
tool_message_content = result[2]["content"]
assert len(tool_message_content) == 2
# First should be tool result
assert "toolResult" in tool_message_content[0]
assert tool_message_content[0]["toolResult"]["content"][0]["text"] == "Weather: sunny, 25°C"
# Second should be cachePoint
assert "cachePoint" in tool_message_content[1]
assert tool_message_content[1]["cachePoint"]["type"] == "default"
def test_assistant_tool_calls_cache_control():
"""Test that assistant tool_calls with cache_control generate cachePoint blocks."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
messages = [
{"role": "user", "content": "Calculate 2+2"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_proxy_123",
"type": "function",
"function": {"name": "calc", "arguments": "{}"},
"cache_control": {"type": "ephemeral"}
}
]
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_provider="bedrock_converse"
)
# Assistant message should have tool use and cachePoint
assistant_content = result[1]["content"]
assert len(assistant_content) == 2
# First should be tool use
assert "toolUse" in assistant_content[0]
assert assistant_content[0]["toolUse"]["name"] == "calc"
assert assistant_content[0]["toolUse"]["toolUseId"] == "call_proxy_123"
# Second should be cachePoint
assert "cachePoint" in assistant_content[1]
assert assistant_content[1]["cachePoint"]["type"] == "default"
def test_multiple_tool_calls_with_mixed_cache_control():
"""Test multiple tool calls where only some have cache_control."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
messages = [
{"role": "user", "content": "Do multiple calculations"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "calc", "arguments": '{"expr": "2+2"}'},
"cache_control": {"type": "ephemeral"}
},
{
"id": "call_2",
"type": "function",
"function": {"name": "calc", "arguments": '{"expr": "3+3"}'}
# No cache_control
}
]
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_provider="bedrock_converse"
)
# Assistant message should have: toolUse1, cachePoint, toolUse2
assistant_content = result[1]["content"]
assert len(assistant_content) == 3
# First tool use with cache
assert "toolUse" in assistant_content[0]
assert assistant_content[0]["toolUse"]["toolUseId"] == "call_1"
# Cache point for first tool
assert "cachePoint" in assistant_content[1]
assert assistant_content[1]["cachePoint"]["type"] == "default"
# Second tool use without cache
assert "toolUse" in assistant_content[2]
assert assistant_content[2]["toolUse"]["toolUseId"] == "call_2"
def test_no_cache_control_no_cache_point():
"""Test that messages without cache_control don't generate cachePoint blocks."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there!"}, # No cache_control
{
"role": "tool",
"tool_call_id": "call_123",
"content": "Tool result" # No cache_control
}
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_provider="bedrock_converse"
)
# Assistant message should only have text content, no cachePoint
assistant_content = result[1]["content"]
assert len(assistant_content) == 1
assert assistant_content[0]["text"] == "Hi there!"
# 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]