From 036b358636023dbae83b17a0c614976ab958ff44 Mon Sep 17 00:00:00 2001 From: "hgyun.lee" Date: Wed, 20 Aug 2025 18:27:36 +0900 Subject: [PATCH] Synchronize cache behavior between acompletion and completion --- .../prompt_templates/factory.py | 51 ++++++++-- .../chat/test_converse_transformation.py | 92 +++++++++++++++++-- 2 files changed, 129 insertions(+), 14 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 0d8c3bacbf5..77cbe4c9a8e 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -3193,9 +3193,30 @@ class BedrockConverseMessagesProcessor: ## MERGE CONSECUTIVE TOOL CALL MESSAGES ## 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] + 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): + 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) @@ -3275,13 +3296,29 @@ class BedrockConverseMessagesProcessor: 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) + 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/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index b60a30bee80..1c91cc0fe8b 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -938,9 +938,11 @@ def test_transform_request_with_function_tool(): assert request_data["toolConfig"]["tools"][0]["toolSpec"]["name"] == "get_weather" -def test_assistant_message_cache_control(): +@pytest.mark.asyncio +async 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 + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor # Test assistant message with string content and cache_control messages = [ @@ -957,6 +959,22 @@ def test_assistant_message_cache_control(): model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse" ) + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result # Should have user message and assistant message assert len(result) == 2 @@ -971,9 +989,11 @@ def test_assistant_message_cache_control(): assert assistant_content[1]["cachePoint"]["type"] == "default" -def test_assistant_message_list_content_cache_control(): +@pytest.mark.asyncio +async 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 + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor messages = [ {"role": "user", "content": "Hello"}, @@ -994,6 +1014,14 @@ def test_assistant_message_list_content_cache_control(): model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse" ) + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result # Assistant message should have text content and cachePoint assistant_content = result[1]["content"] @@ -1003,9 +1031,11 @@ def test_assistant_message_list_content_cache_control(): assert assistant_content[1]["cachePoint"]["type"] == "default" -def test_tool_message_cache_control(): +@pytest.mark.asyncio +async 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 + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor messages = [ {"role": "user", "content": "What's the weather?"}, @@ -1038,6 +1068,14 @@ def test_tool_message_cache_control(): model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse" ) + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result # Should have user, assistant, and user (tool results) messages assert len(result) == 3 @@ -1055,9 +1093,11 @@ def test_tool_message_cache_control(): assert tool_message_content[1]["cachePoint"]["type"] == "default" -def test_tool_message_string_content_cache_control(): +@pytest.mark.asyncio +async 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 + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor messages = [ {"role": "user", "content": "What's the weather?"}, @@ -1085,6 +1125,14 @@ def test_tool_message_string_content_cache_control(): model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse" ) + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result # Last message should contain tool result and cachePoint tool_message_content = result[2]["content"] @@ -1099,9 +1147,11 @@ def test_tool_message_string_content_cache_control(): assert tool_message_content[1]["cachePoint"]["type"] == "default" -def test_assistant_tool_calls_cache_control(): +@pytest.mark.asyncio +async 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 + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor messages = [ {"role": "user", "content": "Calculate 2+2"}, @@ -1124,6 +1174,14 @@ def test_assistant_tool_calls_cache_control(): model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse" ) + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result # Assistant message should have tool use and cachePoint assistant_content = result[1]["content"] @@ -1139,9 +1197,11 @@ def test_assistant_tool_calls_cache_control(): assert assistant_content[1]["cachePoint"]["type"] == "default" -def test_multiple_tool_calls_with_mixed_cache_control(): +@pytest.mark.asyncio +async 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 + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor messages = [ {"role": "user", "content": "Do multiple calculations"}, @@ -1170,6 +1230,14 @@ def test_multiple_tool_calls_with_mixed_cache_control(): model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse" ) + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result # Assistant message should have: toolUse1, cachePoint, toolUse2 assistant_content = result[1]["content"] @@ -1188,9 +1256,11 @@ def test_multiple_tool_calls_with_mixed_cache_control(): assert assistant_content[2]["toolUse"]["toolUseId"] == "call_2" -def test_no_cache_control_no_cache_point(): +@pytest.mark.asyncio +async 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 + from litellm.litellm_core_utils.prompt_templates.factory import BedrockConverseMessagesProcessor messages = [ {"role": "user", "content": "Hello"}, @@ -1207,6 +1277,14 @@ def test_no_cache_control_no_cache_point(): model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_provider="bedrock_converse" ) + + async_result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + llm_provider="bedrock_converse" + ) + + assert result == async_result # Assistant message should only have text content, no cachePoint assistant_content = result[1]["content"]