diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 26388dc2362..0d8c3bacbf5 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -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( diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index c124f1c7d8b..b93ca94bed4 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -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]]: 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 ccf22c9cada..b60a30bee80 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -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" \ No newline at end of file + 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] \ No newline at end of file