From 151fa0417811377677e44731521199b3c8141881 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 15 Sep 2026 14:40:58 -0700 Subject: [PATCH] fix(bedrock): preserve cached prefixes for appended system messages --- .../prompt_templates/factory.py | 22 +++ .../messages/transformation.py | 5 +- .../bedrock/chat/converse_transformation.py | 55 ++++-- litellm/types/llms/bedrock.py | 2 +- litellm/utils.py | 6 + .../chat/test_converse_transformation.py | 187 +++++++++++++++++- 6 files changed, 251 insertions(+), 26 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 21ae8b001dd..b379e2d2909 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -32,6 +32,7 @@ from litellm.types.llms.openai import ( ChatCompletionFileObject, ChatCompletionFunctionMessage, ChatCompletionImageObject, + ChatCompletionSystemMessage, ChatCompletionTextObject, ChatCompletionToolCallFunctionChunk, ChatCompletionToolMessage, @@ -4282,6 +4283,13 @@ def get_assistant_message_block_or_continue_message( class BedrockConverseMessagesProcessor: + @staticmethod + def system_message(message: ChatCompletionSystemMessage, model: str) -> BedrockMessageBlock | None: + blocks: Final = litellm.AmazonConverseConfig().transform_system_message_content(message, model=model) + if not blocks: + return None + return BedrockMessageBlock(role="system", content=[BedrockContentBlock(**block) for block in blocks]) + @staticmethod def _initial_message_setup( messages: list, @@ -4335,6 +4343,13 @@ class BedrockConverseMessagesProcessor: ) while msg_i < len(messages): + if messages[msg_i]["role"] == "system": + if system_message := BedrockConverseMessagesProcessor.system_message( + cast(ChatCompletionSystemMessage, messages[msg_i]), model + ): + contents.append(system_message) + msg_i += 1 + continue user_content: list[BedrockContentBlock] = [] init_msg_i = msg_i ## MERGE CONSECUTIVE USER CONTENT ## @@ -4707,6 +4722,13 @@ def _bedrock_converse_messages_pt( ) while msg_i < len(messages): + if messages[msg_i]["role"] == "system": + if system_message := BedrockConverseMessagesProcessor.system_message( + cast(ChatCompletionSystemMessage, messages[msg_i]), model + ): + contents.append(system_message) + msg_i += 1 + continue user_content: list[BedrockContentBlock] = [] init_msg_i = msg_i ## MERGE CONSECUTIVE USER CONTENT ## diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index 27cdac34116..3638f79bfa7 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -248,7 +248,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): Subclasses whose upstream rejects the role opt in by calling this from their ``transform_anthropic_messages_request``; the first-party Anthropic path forwards ``messages`` untouched and never calls it.""" - from litellm.utils import _supports_factory + from litellm.utils import supports_mid_conversation_system messages: Final = anthropic_messages_request.get("messages") if not isinstance(messages, list): @@ -260,10 +260,9 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): hoisted: Final = messages[:leading_count] remaining: Final = ( messages[leading_count:] - if _supports_factory( + if supports_mid_conversation_system( model=model, custom_llm_provider=self.custom_llm_provider, - key="supports_mid_conversation_system", ) else [ self._system_role_message_as_user(m) if self._is_system_role_message(m) else m diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index fa18361e44c..9cf4aae8b94 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1225,30 +1225,43 @@ class AmazonConverseConfig(BaseConfig): cache_point["ttl"] = ttl return cache_point + def transform_system_message_content( + self, message: ChatCompletionSystemMessage, model: str | None = None + ) -> list[SystemContentBlock]: + system_content_blocks: Final[list[SystemContentBlock]] = [] + if isinstance(message["content"], str) and message["content"]: + system_content_blocks.append(SystemContentBlock(text=message["content"])) + cache_block = self.get_cache_point_block(message, block_type="system", model=model) + if cache_block: + system_content_blocks.append(cache_block) + elif isinstance(message["content"], list): + for m in message["content"]: + if m.get("type") == "text" and m.get("text"): + system_content_blocks.append(SystemContentBlock(text=m["text"])) + cache_block = self.get_cache_point_block(m, block_type="system", model=model) + if cache_block: + system_content_blocks.append(cache_block) + return system_content_blocks + def _transform_system_message( self, messages: list[AllMessageValues], model: str | None = None ) -> tuple[list[AllMessageValues], list[SystemContentBlock]]: - system_prompt_indices: Final = [] - system_content_blocks: Final[list[SystemContentBlock]] = [] - for idx, message in enumerate(messages): - if message["role"] == "system": - system_prompt_indices.append(idx) - if isinstance(message["content"], str) and message["content"]: - system_content_blocks.append(SystemContentBlock(text=message["content"])) - cache_block = self.get_cache_point_block(message, block_type="system", model=model) - if cache_block: - system_content_blocks.append(cache_block) - elif isinstance(message["content"], list): - for m in message["content"]: - if m.get("type") == "text" and m.get("text"): - system_content_blocks.append(SystemContentBlock(text=m["text"])) - cache_block = self.get_cache_point_block(m, block_type="system", model=model) - if cache_block: - system_content_blocks.append(cache_block) - if len(system_prompt_indices) > 0: - for idx in reversed(system_prompt_indices): - messages.pop(idx) - return messages, system_content_blocks + from litellm.utils import supports_mid_conversation_system + + hoist_count: Final = ( + next((idx for idx, message in enumerate(messages) if message["role"] != "system"), len(messages)) + if model is not None and supports_mid_conversation_system(model=model, custom_llm_provider="bedrock") + else len(messages) + ) + system_content_blocks: Final = [ + block + for message in messages[:hoist_count] + if message["role"] == "system" + for block in self.transform_system_message_content(message, model=model) + ] + return [ + message for idx, message in enumerate(messages) if idx >= hoist_count or message["role"] != "system" + ], system_content_blocks def _transform_inference_params(self, inference_params: dict) -> InferenceConfig: if "top_k" in inference_params: diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 76756ac35bb..b5fb2fa0f0d 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -210,7 +210,7 @@ class ContentBlock(TypedDict, total=False): class MessageBlock(TypedDict): content: list[ContentBlock] - role: Literal["user", "assistant"] + role: ReadOnly[Literal["user", "assistant", "system"]] class ConverseMetricsBlock(TypedDict): diff --git a/litellm/utils.py b/litellm/utils.py index af22b11224b..e3a50839816 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2612,6 +2612,12 @@ def _supports_provider_info_factory(model: str, custom_llm_provider: str | None, return None +def supports_mid_conversation_system(model: str, custom_llm_provider: str | None = None) -> bool: + return _supports_factory( + model=model, custom_llm_provider=custom_llm_provider, key="supports_mid_conversation_system" + ) + + def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> bool: """ Check if the given model supports function calling and return a boolean value. 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 2e9ea90f3b8..38351083bd0 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -1,6 +1,8 @@ import asyncio import json import os +from copy import deepcopy +from typing import Final import httpx import pytest @@ -11,7 +13,190 @@ from unittest.mock import MagicMock, patch import litellm from litellm import ModelResponse, RateLimitError, completion from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig -from litellm.types.llms.bedrock import ConverseTokenUsageBlock +from litellm.types.llms.bedrock import ConverseTokenUsageBlock, RequestObject +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionSystemMessage, + ChatCompletionToolParam, +) + + +async def _system_append_request( + messages: list[AllMessageValues], model: str, use_async: bool, + tools: list[ChatCompletionToolParam] | None = None, +) -> RequestObject: + config: Final = AmazonConverseConfig() + optional_params: Final = {"tools": tools} if tools is not None else {} + if use_async: + return await config._async_transform_request( + model=model, messages=deepcopy(messages), optional_params=optional_params, + litellm_params={}, headers={}, + ) + return config._transform_request( + model=model, messages=deepcopy(messages), optional_params=optional_params, + litellm_params={}, headers={}, + ) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("local_model_cost_map") +@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) +@pytest.mark.parametrize( + "late_system", + [ + {"role": "system", "content": "New instruction", + "cache_control": {"type": "ephemeral"}}, + {"role": "system", "content": [ + {"type": "text", "text": "New instruction", + "cache_control": {"type": "ephemeral"}}, + ]}, + ], + ids=["string-cache", "text-block-cache"], +) +async def test_system_append_preserves_cached_prefix_and_native_role( + use_async: bool, late_system: ChatCompletionSystemMessage +) -> None: + prefix: Final[list[AllMessageValues]] = [ + {"role": "system", "content": "Initial instruction"}, + {"role": "system", "content": "Second initial instruction"}, + {"role": "user", "content": [{ + "type": "text", "text": "Cached conversation", + "cache_control": {"type": "ephemeral"}, + }]}, + ] + model: Final = "us.anthropic.claude-sonnet-5" + before: Final = await _system_append_request(prefix, model, use_async) + after: Final = await _system_append_request( + [*prefix, late_system, + {"role": "assistant", "content": "Acknowledged"}, + {"role": "user", "content": "Continue"}], + model, use_async, + ) + + assert before["system"] == after["system"] == [ + {"text": "Initial instruction"}, {"text": "Second initial instruction"}, + ] + assert after["messages"][:len(before["messages"])] == before["messages"] + assert after["messages"][1] == { + "role": "system", + "content": [{"text": "New instruction"}, {"cachePoint": {"type": "default"}}], + } + assert [message["role"] for message in after["messages"]] == [ + "user", "system", "assistant", "user", + ] + assert after["messages"][-1]["content"] == [ + {"text": "Continue"}, + ] + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("local_model_cost_map") +@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) +@pytest.mark.parametrize( + "model", ["us.anthropic.claude-haiku-4-5-20251001-v1:0", + "amazon.nova-pro-v1:0", "unknown-model"] +) +async def test_system_append_keeps_legacy_hoisting( + use_async: bool, model: str +) -> None: + result: Final = await _system_append_request( + [{"role": "system", "content": "Initial instruction"}, + {"role": "user", "content": "Earlier question"}, + {"role": "system", "content": "New instruction"}, + {"role": "assistant", "content": "Acknowledged"}, + {"role": "user", "content": "Continue"}], + model, use_async, + ) + + assert result["system"] == [ + {"text": "Initial instruction"}, {"text": "New instruction"}, + ] + assert [message["role"] for message in result["messages"]] == [ + "user", "assistant", "user", + ] + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("local_model_cost_map") +@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) +async def test_system_append_native_messages_bridge(use_async: bool) -> None: + from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( + LiteLLMAnthropicMessagesAdapter, + ) + + request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai( + anthropic_message_request={ + "model": "bedrock/converse/us.anthropic.claude-sonnet-5", + "max_tokens": 16, + "system": "Initial instruction", + "messages": [ + {"role": "user", "content": "Earlier question"}, + {"role": "system", "content": [{ + "type": "text", "text": "New instruction", + "cache_control": {"type": "ephemeral"}, + }]}, + {"role": "assistant", "content": "Acknowledged"}, + {"role": "user", "content": "Continue"}, + ], + } + ) + result: Final = await _system_append_request( + request["messages"], "us.anthropic.claude-sonnet-5", use_async, + ) + + assert result["system"] == [{"text": "Initial instruction"}] + assert result["messages"][1] == { + "role": "system", + "content": [{"text": "New instruction"}, {"cachePoint": {"type": "default"}}], + } + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("local_model_cost_map") +@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) +async def test_system_append_preserves_tool_result_adjacency(use_async: bool) -> None: + result: Final = await _system_append_request( + [{"role": "user", "content": "Read the file"}, + {"role": "assistant", "content": None, "tool_calls": [{ + "id": "tool_1", "type": "function", + "function": {"name": "read_file", "arguments": "{}"}, + }]}, + {"role": "tool", "tool_call_id": "tool_1", "content": "File contents"}, + {"role": "system", "content": "Use the file contents"}, + {"role": "assistant", "content": "Acknowledged"}, + {"role": "user", "content": "Continue"}], + "us.anthropic.claude-sonnet-5", use_async, + tools=[{"type": "function", "function": { + "name": "read_file", "parameters": {"type": "object", "properties": {}}, + }}], + ) + + assert "system" not in result + assert [message["role"] for message in result["messages"]] == [ + "user", "assistant", "user", "system", "assistant", "user", + ] + assert result["messages"][1]["content"][0]["toolUse"]["toolUseId"] == "tool_1" + assert result["messages"][2]["content"][0]["toolResult"]["toolUseId"] == "tool_1" + assert result["messages"][3]["content"] == [{"text": "Use the file contents"}] + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("local_model_cost_map") +@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) +async def test_system_append_filters_empty_system_content(use_async: bool) -> None: + result: Final = await _system_append_request( + [{"role": "user", "content": "Earlier question"}, + {"role": "system", "content": ""}, + {"role": "system", "content": [{"type": "text", "text": ""}]}, + {"role": "assistant", "content": "Acknowledged"}, + {"role": "user", "content": "Continue"}], + "us.anthropic.claude-sonnet-5", use_async, + ) + + assert "system" not in result + assert [message["role"] for message in result["messages"]] == [ + "user", "assistant", "user", + ] def test_transform_usage():