From bdb04849d98d0d1af4e9b2ac1e553c82f1fdc6e0 Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Thu, 24 Sep 2026 17:23:07 +0000 Subject: [PATCH] fix(bedrock): strip Anthropic native extensions Bedrock Invoke rejects from Messages bodies --- litellm/llms/bedrock/common_utils.py | 95 ++++ .../bedrock/count_tokens/transformation.py | 24 +- .../anthropic_claude3_transformation.py | 50 ++ .../bedrock/messages/mantle_transformation.py | 7 + .../coverage_registry/llm_conversational.yaml | 1 + tests/e2e/coverage_registry/schema.py | 1 + ..._messages_bedrock_native_extensions_e2e.py | 125 +++++ ..._invoke_messages_native_extensions_wire.py | 453 ++++++++++++++++++ .../test_anthropic_claude3_transformation.py | 158 ++++++ .../test_litellm/llms/bedrock/test_mantle.py | 29 ++ ...est_bedrock_count_tokens_transformation.py | 85 ++++ 11 files changed, 1026 insertions(+), 2 deletions(-) create mode 100644 tests/e2e/llm_translation/test_messages_bedrock_native_extensions_e2e.py create mode 100644 tests/integration/providers/test_bedrock_invoke_messages_native_extensions_wire.py diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 9b52f531cbb..5179aa738e1 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -10,6 +10,7 @@ import json import os import re from collections.abc import Mapping, Sequence +from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict if TYPE_CHECKING: @@ -323,6 +324,100 @@ def strip_unsupported_bedrock_invoke_output_config_keys( request_body["output_config"] = {"format": preserved_format} # rebind-ok: out-param # mutable-ok: json +BEDROCK_INVOKE_UNSUPPORTED_MESSAGE_KEYS: Final = frozenset({"output_config"}) +BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES: Final = frozenset({"tool_addition"}) +BEDROCK_INVOKE_SUPPORTED_THINKING_DISPLAY_VALUES: Final = frozenset({"summarized", "omitted"}) + + +@dataclass(frozen=True, slots=True) +class SanitizedBedrockInvokeMessages: + messages: tuple[object, ...] + offenders: tuple[str, ...] + emptied: tuple[str, ...] + + +def _is_unsupported_bedrock_invoke_block(block: object, unsupported_block_types: frozenset[str]) -> bool: + if not isinstance(block, dict): + return False + block_type: Final = block.get("type") + return isinstance(block_type, str) and block_type in unsupported_block_types + + +def _bedrock_invoke_message_offenders( + index: int, message: object, unsupported_keys: frozenset[str], unsupported_block_types: frozenset[str] +) -> tuple[str, ...]: + if not isinstance(message, dict): + return () + key_paths: Final = tuple(f"messages[{index}].{key}" for key in sorted(unsupported_keys & message.keys())) + content: Final = message.get("content") + block_paths: Final = ( + tuple( + f"messages[{index}].content[{block_index}] (type '{block['type']}')" + for block_index, block in enumerate(content) + if _is_unsupported_bedrock_invoke_block(block, unsupported_block_types) + ) + if isinstance(content, list) + else () + ) + return key_paths + block_paths + + +def _sanitized_bedrock_invoke_message( + message: object, unsupported_keys: frozenset[str], unsupported_block_types: frozenset[str] +) -> object: + if not isinstance(message, dict): + return message + return { # mutable-ok: outbound JSON message, same plain dict shape as the caller's input + key: ( + [ # mutable-ok: outbound JSON content list + b for b in value if not _is_unsupported_bedrock_invoke_block(b, unsupported_block_types) + ] + if key == "content" and isinstance(value, list) + else value + ) + for key, value in message.items() + if key not in unsupported_keys + } + + +def sanitize_bedrock_invoke_messages( + messages: Sequence[object], + unsupported_keys: frozenset[str], + unsupported_block_types: frozenset[str], +) -> SanitizedBedrockInvokeMessages: + offenders: Final = tuple( + path + for i, m in enumerate(messages) + for path in _bedrock_invoke_message_offenders(i, m, unsupported_keys, unsupported_block_types) + ) + if not offenders: + return SanitizedBedrockInvokeMessages(messages=tuple(messages), offenders=(), emptied=()) + emptied: Final = tuple( + f"messages[{i}]" + for i, m in enumerate(messages) + if isinstance(m, dict) + and isinstance(m.get("content"), list) + and len(m["content"]) > 0 + and all(_is_unsupported_bedrock_invoke_block(b, unsupported_block_types) for b in m["content"]) + ) + return SanitizedBedrockInvokeMessages( + messages=tuple( + _sanitized_bedrock_invoke_message(m, unsupported_keys, unsupported_block_types) for m in messages + ), + offenders=offenders, + emptied=emptied, + ) + + +def normalize_bedrock_invoke_thinking_display(thinking: object, supported_display_values: frozenset[str]) -> object: + if not isinstance(thinking, dict): + return thinking + display: Final = thinking.get("display") + if not isinstance(display, str) or display in supported_display_values: + return thinking + return {**thinking, "display": "summarized"} # mutable-ok: outbound JSON thinking object + + def normalize_custom_field_on_tools(request_body: dict) -> None: """ Drop the ``custom`` field from each tool, first hoisting a boolean diff --git a/litellm/llms/bedrock/count_tokens/transformation.py b/litellm/llms/bedrock/count_tokens/transformation.py index 48fc41ed12b..f1cc1a33f94 100644 --- a/litellm/llms/bedrock/count_tokens/transformation.py +++ b/litellm/llms/bedrock/count_tokens/transformation.py @@ -12,7 +12,12 @@ from typing import Final, Literal from pydantic import JsonValue from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM -from litellm.llms.bedrock.common_utils import get_bedrock_base_model +from litellm.llms.bedrock.common_utils import ( + BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES, + BEDROCK_INVOKE_UNSUPPORTED_MESSAGE_KEYS, + get_bedrock_base_model, + sanitize_bedrock_invoke_messages, +) # Placeholder satisfying the Anthropic InvokeModel schema's required # max_tokens field; CountTokens only counts input, so it has no effect @@ -190,7 +195,22 @@ class BedrockCountTokensConfig(BaseAWSLLM): # For InvokeModel, we need to provide the raw body that would be sent to the model # Remove the 'model' field from the body as it's not part of the model input - body_data: Final = {k: v for k, v in request_data.items() if k != "model"} + messages: Final = request_data.get("messages") + sanitized_messages: Final = ( + list( # mutable-ok: outbound JSON messages array + sanitize_bedrock_invoke_messages( + messages, + unsupported_keys=BEDROCK_INVOKE_UNSUPPORTED_MESSAGE_KEYS, + unsupported_block_types=BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES, + ).messages + ) + if isinstance(messages, list) + else messages + ) + body_data: Final = { # mutable-ok: outbound JSON body, defaults are set below like before + **{k: v for k, v in request_data.items() if k not in ("model", "messages")}, # mutable-ok: spread source + **({"messages": sanitized_messages} if "messages" in request_data else {}), # mutable-ok: spread source + } if "messages" in body_data: # Bedrock validates the body against the model's InvokeModel schema; diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 14bd2bee6cf..d687fa79234 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -29,15 +29,20 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation AmazonInvokeConfig, ) from litellm.llms.bedrock.common_utils import ( + BEDROCK_INVOKE_SUPPORTED_THINKING_DISPLAY_VALUES, + BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES, + BEDROCK_INVOKE_UNSUPPORTED_MESSAGE_KEYS, BedrockError, apply_bedrock_invoke_structured_output, bedrock_supports_tool_search, ensure_bedrock_anthropic_messages_tool_names, get_anthropic_beta_from_headers, is_claude_4_5_on_bedrock, + normalize_bedrock_invoke_thinking_display, normalize_bedrock_opus_output_config_effort, normalize_custom_field_on_tools, normalize_tool_input_schema_types_for_bedrock_invoke, + sanitize_bedrock_invoke_messages, strip_unsupported_bedrock_invoke_output_config_keys, tools_without_eager_input_streaming, ) @@ -455,6 +460,9 @@ class AmazonAnthropicClaudeMessagesConfig( "clear_tool_uses_20250919": ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value, } ) + BEDROCK_INVOKE_UNSUPPORTED_MESSAGE_KEYS: frozenset[str] = BEDROCK_INVOKE_UNSUPPORTED_MESSAGE_KEYS + BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES: frozenset[str] = BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES + BEDROCK_INVOKE_SUPPORTED_THINKING_DISPLAY_VALUES: frozenset[str] = BEDROCK_INVOKE_SUPPORTED_THINKING_DISPLAY_VALUES @classmethod def _filter_context_management_for_bedrock_invoke( @@ -616,6 +624,47 @@ class AmazonAnthropicClaudeMessagesConfig( llm_provider="bedrock", ) + def _apply_bedrock_invoke_native_extension_policy( + self, + anthropic_messages_request: dict, # mutable-ok: outbound body edited in place like the sibling sanitizers + model: str, + ) -> None: + messages: Final = anthropic_messages_request.get("messages") + if isinstance(messages, list): + sanitized: Final = sanitize_bedrock_invoke_messages( + messages, + unsupported_keys=self.BEDROCK_INVOKE_UNSUPPORTED_MESSAGE_KEYS, + unsupported_block_types=self.BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES, + ) + if sanitized.emptied: + raise litellm.BadRequestError( + message=( + f"{', '.join(sanitized.emptied)} would be left with empty content after removing " + f"{', '.join(sanitized.offenders)}, which Bedrock Invoke does not accept. " + "Remove or rewrite the message." + ), + model=model, + llm_provider="bedrock", + ) + if sanitized.offenders: + verbose_logger.warning( + "Bedrock Invoke: stripped unsupported native extensions at %s for model=%s", + sanitized.offenders, + model, + ) + anthropic_messages_request["messages"] = list(sanitized.messages) # mutable-ok: outbound JSON array + thinking: Final = anthropic_messages_request.get("thinking") + normalized_thinking: Final = normalize_bedrock_invoke_thinking_display( + thinking, supported_display_values=self.BEDROCK_INVOKE_SUPPORTED_THINKING_DISPLAY_VALUES + ) + if normalized_thinking is not thinking: + verbose_logger.warning( + "Bedrock Invoke: mapping unsupported thinking.display %r to 'summarized' for model=%s", + thinking.get("display") if isinstance(thinking, dict) else thinking, + model, + ) + anthropic_messages_request["thinking"] = normalized_thinking + def _strip_unsupported_bedrock_invoke_fields( self, anthropic_messages_request: dict, @@ -695,6 +744,7 @@ class AmazonAnthropicClaudeMessagesConfig( # 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it for older models) self._remove_ttl_from_cache_control(anthropic_messages_request=anthropic_messages_request, model=model) + self._apply_bedrock_invoke_native_extension_policy(anthropic_messages_request, model=model) # 5. Route structured-output params (`output_format` / # `output_config.format`) to native enforcement or the inline-schema diff --git a/litellm/llms/bedrock/messages/mantle_transformation.py b/litellm/llms/bedrock/messages/mantle_transformation.py index 66744275778..3c7160aef58 100644 --- a/litellm/llms/bedrock/messages/mantle_transformation.py +++ b/litellm/llms/bedrock/messages/mantle_transformation.py @@ -63,6 +63,13 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig): def should_filter_anthropic_beta_headers(self) -> bool: return False + def _apply_bedrock_invoke_native_extension_policy( + self, + anthropic_messages_request: dict, # mutable-ok: signature shared with the Invoke parent + model: str, + ) -> None: + return + def get_complete_url( self, api_base: str | None, diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index e4e1ac2c7b6..47ee178d494 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -60,6 +60,7 @@ - {id: llm.messages.bedrock_invoke.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Flagged Claude 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (#32578/#32831/#32882)", fail_before_fix: proven} - {id: llm.messages.bedrock_invoke.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Claude <= 4.7 rejects role system inside messages; unflagged models must convert reminders to user turns in place (hoisting collapses the prompt cache) or every Claude Code session 400s (#32831)", fail_before_fix: proven} - {id: llm.messages.bedrock_invoke.web_search_server_tool.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: web_search_server_tool, streaming: nonstream, assertions: [works], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Bedrock hosts no web_search server tool, so this only works because interception rewrites it before the upstream call and the agentic loop feeds the results back in native shape; a regression that short-circuits or forwards it instead yields raw text or AWS's 400", fail_before_fix: unproven} +- {id: llm.messages.bedrock_invoke.native_extensions.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: native_extensions, streaming: nonstream, assertions: [works, rejects_with_actionable_error], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Claude Code emits per-message output_config, thinking.display values Bedrock rejects, and tool_addition content blocks; Bedrock Invoke must always strip the fields and blocks and map the display value before the upstream call, and a message left empty by that strip gets a LiteLLM 400 naming it instead of Bedrock's 400 (LIT-8439)", fail_before_fix: proven} - {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (customer RCA gap)", fail_before_fix: proven} - {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry Claude <= 4.7 rejects role system inside messages; unflagged models must convert reminders to user turns in place (hoisting collapses the prompt cache) or every Claude Code session 400s (customer RCA gap)", fail_before_fix: proven} - {id: llm.messages.vertex.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (customer RCA gap)", fail_before_fix: proven} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index e009b02b69c..c9a720a90e0 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -73,6 +73,7 @@ LlmCapability = Literal[ "long_context_1m", "mid_conversation_system", "multi_turn", + "native_extensions", "pdf_input", "prompt_cache_1h", "prompt_cache_5m", diff --git a/tests/e2e/llm_translation/test_messages_bedrock_native_extensions_e2e.py b/tests/e2e/llm_translation/test_messages_bedrock_native_extensions_e2e.py new file mode 100644 index 00000000000..82cc854a57b --- /dev/null +++ b/tests/e2e/llm_translation/test_messages_bedrock_native_extensions_e2e.py @@ -0,0 +1,125 @@ +from __future__ import annotations + +from typing import cast + +import anthropic +import pytest +from anthropic.types import Message, MessageParam, TextBlock, ThinkingConfigParam +from e2e_config import unique_marker +from lifecycle import ResourceManager +from models import LiteLLMParamsBody +from proxy_client import ProxyClient +from sdk_clients import NO_PROXY_CACHE, SdkClients + +pytestmark = pytest.mark.e2e + +BEDROCK_INVOKE_BACKEND = "bedrock/invoke/us.anthropic.claude-sonnet-5" +TOOL_ADDITION_BLOCK = {"type": "tool_addition", "tool_reference": {"type": "tool_reference", "tool_name": "Read"}} + + +def _register(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, str]: + model = f"e2e-bedrock-msgs-ext-{unique_marker()}" + model_id = proxy.create_model( + model, + LiteLLMParamsBody( + model=BEDROCK_INVOKE_BACKEND, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="us-east-1", + ), + ) + resources.defer(lambda: proxy.delete_model(model_id)) + return model, resources.key() + + +def _output_config_messages() -> list[MessageParam]: + return [ + {"role": "user", "content": "read the file /tmp/a.txt"}, + cast( + MessageParam, + { + "role": "assistant", + "output_config": {"effort": "high"}, + "content": [ + {"type": "text", "text": "Reading it now."}, + {"type": "tool_use", "id": "toolu_1", "name": "Read", "input": {"path": "/tmp/a.txt"}}, + ], + }, + ), + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "hello world"}], + }, + ] + + +def _tool_addition_messages() -> list[MessageParam]: + return [ + {"role": "user", "content": "read /tmp/a.txt"}, + cast(MessageParam, {"role": "assistant", "content": [TOOL_ADDITION_BLOCK, {"type": "text", "text": "ok"}]}), + {"role": "user", "content": "continue"}, + ] + + +def _only_tool_addition_messages() -> list[MessageParam]: + return [ + {"role": "user", "content": "read /tmp/a.txt"}, + cast(MessageParam, {"role": "assistant", "content": [TOOL_ADDITION_BLOCK, TOOL_ADDITION_BLOCK]}), + {"role": "user", "content": "continue"}, + ] + + +def _text(message: Message) -> str: + return "".join(block.text for block in message.content if isinstance(block, TextBlock)) + + +def _assert_answered(message: Message) -> None: + assert message.role == "assistant", f"unexpected role: {message.role!r}" + assert _text(message).strip(), f"/v1/messages returned no text: {message.content!r}" + + +class TestBedrockMessagesNativeExtensions: + @pytest.mark.covers("llm.messages.bedrock_invoke.native_extensions.nonstream.works") + def test_nested_output_config_is_stripped_before_bedrock_invoke( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model, key = _register(proxy, resources) + message = sdk.anthropic(key).messages.create( + model=model, max_tokens=300, messages=_output_config_messages(), extra_body=NO_PROXY_CACHE + ) + _assert_answered(message) + + @pytest.mark.covers("llm.messages.bedrock_invoke.native_extensions.nonstream.works") + def test_tool_addition_block_is_stripped_before_bedrock_invoke( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model, key = _register(proxy, resources) + message = sdk.anthropic(key).messages.create( + model=model, max_tokens=300, messages=_tool_addition_messages(), extra_body=NO_PROXY_CACHE + ) + _assert_answered(message) + + @pytest.mark.covers("llm.messages.bedrock_invoke.native_extensions.nonstream.works") + def test_thinking_display_updates_is_mapped_before_bedrock_invoke( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model, key = _register(proxy, resources) + message = sdk.anthropic(key).messages.create( + model=model, + max_tokens=300, + thinking=cast(ThinkingConfigParam, {"type": "adaptive", "display": "updates"}), + messages=[{"role": "user", "content": "what is 2+2? think briefly"}], + extra_body=NO_PROXY_CACHE, + ) + _assert_answered(message) + + @pytest.mark.covers("llm.messages.bedrock_invoke.native_extensions.nonstream.works") + def test_message_emptied_by_stripping_is_rejected_naming_the_message( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model, key = _register(proxy, resources) + with pytest.raises(anthropic.BadRequestError) as exc_info: + sdk.anthropic(key).messages.create( + model=model, max_tokens=300, messages=_only_tool_addition_messages(), extra_body=NO_PROXY_CACHE + ) + assert "messages[1]" in str(exc_info.value), str(exc_info.value) diff --git a/tests/integration/providers/test_bedrock_invoke_messages_native_extensions_wire.py b/tests/integration/providers/test_bedrock_invoke_messages_native_extensions_wire.py new file mode 100644 index 00000000000..044aec898e6 --- /dev/null +++ b/tests/integration/providers/test_bedrock_invoke_messages_native_extensions_wire.py @@ -0,0 +1,453 @@ +import base64 +import json +import os +import signal +import threading +import uuid +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import psutil +import pytest +from integration._support.client import Gateway, Scenario, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.upstream import _aws_event_frame +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +MODEL_ID: Final = "us.anthropic.claude-sonnet-4-6" +TOKEN: Final = "synthetic-bedrock-bearer" +INVOKE: Final = f"/model/{MODEL_ID}/invoke" +INVOKE_STREAM: Final = f"/model/{MODEL_ID}/invoke-with-response-stream" +BODY: Final = TypeAdapter(dict[str, JsonValue]) +MESSAGES: Final = TypeAdapter(list[dict[str, JsonValue]]) +TOOL_ADDITION: Final = {"type": "tool_addition", "tool_reference": {"type": "tool_reference", "tool_name": "Read"}} +TOOL_USE: Final = {"type": "tool_use", "id": "toolu_1", "name": "Read", "input": {"path": "/tmp/a.txt"}} +TOOL_RESULT: Final = {"type": "tool_result", "tool_use_id": "toolu_1", "content": "hello world"} +READING: Final = {"type": "text", "text": "Reading it now."} +REJECTED: Final = "Bedrock Invoke rejects the extension: " + + +def _output_config_turns(tag: str) -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": f"read the file /tmp/a.txt {tag}"}, + {"role": "assistant", "output_config": {"effort": "high"}, "content": [READING, TOOL_USE]}, + {"role": "user", "content": [TOOL_RESULT]}, + ] + + +def _tool_addition_turns(tag: str) -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": f"read /tmp/a.txt {tag}"}, + {"role": "assistant", "content": [TOOL_ADDITION, {"type": "text", "text": "ok"}]}, + {"role": "user", "content": "continue"}, + ] + + +def _without_extensions(messages: list[dict[str, JsonValue]]) -> list[dict[str, JsonValue]]: + return [ + { + key: ( + [block for block in value if not (isinstance(block, dict) and block.get("type") == "tool_addition")] + if key == "content" and isinstance(value, list) + else value + ) + for key, value in message.items() + if key != "output_config" + } + for message in messages + ] + + +def _first_tag(body: dict[str, JsonValue]) -> str: + first: Final = MESSAGES.validate_python(body["messages"])[0]["content"] + text: Final = first if isinstance(first, str) else MESSAGES.validate_python(first)[0]["text"] + assert isinstance(text, str), first + return text.rsplit(" ", 1)[-1] + + +def _rejection(body: dict[str, JsonValue]) -> str | None: + thinking: Final = body.get("thinking") + if isinstance(thinking, dict) and thinking.get("display") not in (None, "summarized", "omitted"): + return f"thinking.display={thinking['display']!r}" + messages: Final = MESSAGES.validate_python(body["messages"]) + if any("output_config" in message for message in messages): + return "message.output_config" + if any( + isinstance(block, dict) and block.get("type") == "tool_addition" + for message in messages + if isinstance(message["content"], list) + for block in message["content"] + ): + return "tool_addition block" + return None + + +def _message(tag: str) -> bytes: + return json.dumps( + { + "id": f"msg_{tag}", + "type": "message", + "role": "assistant", + "model": MODEL_ID, + "content": [{"type": "text", "text": f"answer {tag}"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 4}, + } + ).encode() + + +def _chunk(payload: dict[str, JsonValue]) -> bytes: + encoded: Final = base64.b64encode(json.dumps(payload, separators=(",", ":")).encode()).decode() + return _aws_event_frame("chunk", {"bytes": encoded}, "", "") + + +def _stream(tag: str) -> bytes: + return ( + _chunk( + { + "type": "message_start", + "message": { + "id": f"msg_{tag}", + "type": "message", + "role": "assistant", + "model": MODEL_ID, + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 0}, + }, + } + ) + + _chunk({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}) + + _chunk({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": f"answer {tag}"}}) + + _chunk({"type": "content_block_stop", "index": 0}) + + _chunk({"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 4}}) + + _chunk({"type": "message_stop"}) + ) + + +def lenient_bedrock_peer(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {TOKEN}" + tag: Final = _first_tag(BODY.validate_python(json.loads(request.body))) + if request.target == INVOKE_STREAM: + return Reply(content_type="application/vnd.amazon.eventstream", chunks=(_stream(tag),)) + assert request.target == INVOKE, request.target + return Reply(body=_message(tag)) + + +def bedrock_peer(request: Request) -> Reply: + rejected: Final = _rejection(BODY.validate_python(json.loads(request.body))) + if rejected is not None: + return Reply(status=400, body=json.dumps({"message": REJECTED + rejected}).encode()) + return lenient_bedrock_peer(request) + + +def _register(scenario: Scenario, api_base: str) -> str: + return scenario.model( + model=f"bedrock/invoke/{MODEL_ID}", + api_key=TOKEN, + aws_region_name="us-east-1", + api_base=api_base, + ) + + +def _sent(wire_requests: tuple[Request, ...], target: str) -> dict[str, JsonValue]: + assert len(wire_requests) == 1, [request.target for request in wire_requests] + assert wire_requests[0].target == target, wire_requests[0].target + return BODY.validate_python(json.loads(wire_requests[0].body)) + + +def test_per_message_output_config_is_dropped_before_bedrock_invoke(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + response: Final = gateway.request( + "POST", "/v1/messages", {"model": model, "max_tokens": 300, "messages": _output_config_turns(tag)} + ) + assert response.status_code == 200, response.text + assert response.json()["content"] == [{"type": "text", "text": f"answer {tag}"}], response.text + sent: Final = _sent(wire.drain(), INVOKE) + assert sent["messages"] == _without_extensions(_output_config_turns(tag)), sent + assert "output_config" not in json.dumps(sent), sent + assert sent["max_tokens"] == 300 and "model" not in sent, sent + + +def test_thinking_display_updates_is_mapped_to_summarized_for_bedrock_invoke(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + client: Final = anthropic.Anthropic( + api_key=gateway.key, base_url=str(gateway.client.base_url).rstrip("/"), max_retries=0 + ) + message: Final = client.messages.create( + model=model, + max_tokens=300, + thinking={"type": "adaptive", "display": "updates"}, # pyright: ignore[reportArgumentType] # Claude Code sends this shape; the SDK types lag + messages=[{"role": "user", "content": f"what is 2+2? think briefly {tag}"}], + ) + assert message.role == "assistant" and message.id == f"msg_{tag}", message + sent: Final = _sent(wire.drain(), INVOKE) + assert sent["thinking"] == {"type": "adaptive", "display": "summarized"}, sent + assert sent["messages"] == [{"role": "user", "content": f"what is 2+2? think briefly {tag}"}], sent + + +@pytest.mark.asyncio +async def test_tool_addition_block_is_dropped_before_bedrock_invoke(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + async with anthropic.AsyncAnthropic( + api_key=gateway.key, base_url=str(gateway.client.base_url).rstrip("/"), max_retries=0 + ) as client: + message: Final = await client.messages.create( + model=model, + max_tokens=300, + messages=_tool_addition_turns(tag), # pyright: ignore[reportArgumentType] # tool_addition is not in the SDK's block union + ) + assert message.id == f"msg_{tag}", message + sent: Final = _sent(wire.drain(), INVOKE) + assert sent["messages"] == [ + {"role": "user", "content": f"read /tmp/a.txt {tag}"}, + {"role": "assistant", "content": [{"type": "text", "text": "ok"}]}, + {"role": "user", "content": "continue"}, + ], sent + + +def test_all_three_extensions_are_sanitized_on_the_streaming_invoke_path(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + turns: Final = _output_config_turns(tag) + _tool_addition_turns(tag)[1:] + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 300, + "stream": True, + "thinking": {"type": "adaptive", "display": "updates"}, + "messages": turns, + }, + ) + assert response.status_code == 200, response.text + events: Final = tuple( + json.loads(line.removeprefix("data:")) for line in response.text.splitlines() if line.startswith("data:") + ) + assert events[-1]["type"] == "message_stop", response.text + assert ( + "".join(event["delta"]["text"] for event in events if event["type"] == "content_block_delta") + == f"answer {tag}" + ), response.text + sent: Final = _sent(wire.drain(), INVOKE_STREAM) + assert sent["messages"] == _without_extensions(turns), sent + assert sent["thinking"] == {"type": "adaptive", "display": "summarized"}, sent + assert "tool_addition" not in json.dumps(sent) and "output_config" not in json.dumps(sent), sent + + +def test_message_holding_only_tool_addition_blocks_is_rejected_before_bedrock(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + turns: Final = [ + {"role": "user", "content": f"read /tmp/a.txt {tag}"}, + {"role": "assistant", "content": [TOOL_ADDITION, TOOL_ADDITION]}, + {"role": "user", "content": "continue"}, + ] + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + response: Final = gateway.request( + "POST", "/v1/messages", {"model": model, "max_tokens": 300, "messages": turns} + ) + assert response.status_code == 400, response.text + assert "messages[1]" in response.text, response.text + assert wire.drain() == (), "rejected request must not reach Bedrock" + + +@pytest.mark.parametrize("display", ["summarized", "omitted", None]) +def test_supported_or_absent_thinking_display_reaches_bedrock_unchanged(gateway: Gateway, display: str | None) -> None: + tag: Final = uuid.uuid4().hex + thinking: Final = {"type": "adaptive", **({"display": display} if display is not None else {})} + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 300, "thinking": thinking, "messages": [{"role": "user", "content": tag}]}, + ) + assert response.status_code == 200, response.text + sent: Final = _sent(wire.drain(), INVOKE) + assert sent["thinking"] == thinking, sent + + +@pytest.mark.parametrize( + "output_config", [7, ["high"], "", "x" * 5000], ids=["int", "list", "empty_string", "five_kb_string"] +) +def test_malformed_per_message_output_config_is_dropped_on_every_message( + gateway: Gateway, output_config: JsonValue +) -> None: + tag: Final = uuid.uuid4().hex + turns: Final = [ + {"role": "user", "content": f"read {tag}", "output_config": output_config}, + {"role": "assistant", "content": [READING], "output_config": output_config}, + {"role": "user", "content": "continue", "output_config": output_config}, + ] + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + response: Final = gateway.request( + "POST", "/v1/messages", {"model": model, "max_tokens": 300, "messages": turns} + ) + assert response.status_code == 200, response.text + sent: Final = _sent(wire.drain(), INVOKE) + assert sent["messages"] == _without_extensions(turns), sent + + +def test_non_string_thinking_display_and_bare_string_blocks_are_forwarded_as_sent(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + turns: Final = [ + {"role": "user", "content": f"read {tag}"}, + {"role": "assistant", "content": ["ok"]}, + {"role": "user", "content": "continue"}, + ] + with wire_server(lenient_bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 300, "thinking": {"type": "adaptive", "display": 7}, "messages": turns}, + ) + assert response.status_code == 200, response.text + sent: Final = _sent(wire.drain(), INVOKE) + assert sent["thinking"] == {"type": "adaptive", "display": 7}, sent + assert sent["messages"] == turns, sent + + +def test_anthropic_direct_deployment_forwards_the_extensions_verbatim(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + turns: Final = _output_config_turns(tag) + _tool_addition_turns(tag)[1:] + + def anthropic_peer(request: Request) -> Reply: + assert request.target == "/v1/messages", request.target + return Reply(body=_message(tag)) + + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-6", api_base=wire.url, api_key="synthetic-anthropic-key" + ) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 300, + "thinking": {"type": "adaptive", "display": "updates"}, + "messages": turns, + }, + ) + assert response.status_code == 200, response.text + sent: Final = _sent(wire.drain(), "/v1/messages") + assert sent["messages"] == turns, sent + assert sent["thinking"] == {"type": "adaptive", "display": "updates"}, sent + + +def test_bedrock_invoke_chat_completions_are_untouched_by_the_messages_sanitizer(gateway: Gateway) -> None: + tag: Final = uuid.uuid4().hex + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = _register(scenario, wire.url) + body: Final = gateway.chat(model, text=f"chat control {tag}") + assert body["choices"][0]["message"]["content"] == f"answer {tag}", body + sent: Final = _sent(wire.drain(), INVOKE) + assert sent["messages"] == [{"role": "user", "content": [{"type": "text", "text": f"chat control {tag}"}]}], ( + sent + ) + + +def _burst_request(gateway: Gateway, model: str, tag: str, index: int) -> tuple[str, int, str]: + shape: Final = index % 3 + stream: Final = index % 2 == 1 + body: Final = { + "model": model, + "max_tokens": 300, + "stream": stream, + **({"thinking": {"type": "adaptive", "display": "updates"}} if shape == 1 else {}), + "messages": ( + _output_config_turns(tag) + if shape == 0 + else [{"role": "user", "content": f"think {tag}"}] + if shape == 1 + else _tool_addition_turns(tag) + ), + } + try: + response: Final = gateway.request("POST", "/v1/messages", body) + except httpx.TransportError as error: + return tag, 0, repr(error) + return tag, response.status_code, response.text + + +def _workers(root: psutil.Process) -> tuple[psutil.Process, ...]: + return tuple(child for child in root.children(recursive=True) if child.status() != psutil.STATUS_ZOMBIE) + + +@pytest.mark.timeout(300) +def test_burst_survives_a_stalled_upstream_and_a_killed_worker_with_exactly_one_upstream_call_per_request( + gateway: Gateway, tmp_path: Path +) -> None: + burst: Final = 30 + gate: Final = threading.Event() + stalled: Final = threading.Semaphore(10) + + def stalling_peer(request: Request) -> Reply: + if stalled.acquire(blocking=False): + assert gate.wait(timeout=60), "burst gate never released" + return bedrock_peer(request) + + with ( + wire_server(stalling_peer) as wire, + owned_proxy_process(gateway, tmp_path, {}, workers=2) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = _register(scenario, wire.url) + tags: Final = tuple(uuid.uuid4().hex for _ in range(burst)) + with ThreadPoolExecutor(max_workers=burst) as pool: + futures: Final = tuple( + pool.submit(_burst_request, owned.gateway, model, tag, index) for index, tag in enumerate(tags) + ) + victim: Final = eventually( + lambda: _workers(psutil.Process(owned.process.pid)), lambda found: len(found) >= 2, seconds=30 + )[-1] + os.kill(victim.pid, signal.SIGTERM) + gate.set() + results: Final = tuple(future.result(timeout=120) for future in futures) + assert not eventually( + lambda: victim.is_running() and victim.status() != psutil.STATUS_ZOMBIE, lambda alive: not alive, seconds=30 + ) + assert _workers(psutil.Process(owned.process.pid)), "no worker survived the kill" + after: Final = owned.gateway.request( + "POST", "/v1/messages", {"model": model, "max_tokens": 300, "messages": _output_config_turns("after")} + ) + assert after.status_code == 200, after.text + received: Final = wire.drain() + seen: Final = tuple(_first_tag(BODY.validate_python(json.loads(request.body))) for request in received) + assert sorted(seen) == sorted(set(seen)), seen + assert set(seen) <= set(tags) | {"after"}, seen + assert all(_rejection(BODY.validate_python(json.loads(request.body))) is None for request in received), seen + answered: Final = tuple(result for result in results if result[1] == 200) + assert all(f"answer {tag}" in text for tag, _, text in answered), answered + assert {tag for tag, _, _ in answered} <= set(seen), (answered, seen) + assert len(answered) >= burst - 10, [(tag, status, text[:80]) for tag, status, text in results] + assert all(status in (0, 200) or status >= 500 for _, status, _ in results), [ + (tag, status, text[:80]) for tag, status, text in results + ] + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)', + ([f"msg_{tag}" for tag, _, _ in answered],), + ), + lambda values: len(values) == len(answered), + seconds=90, + ) + assert sorted(row["request_id"] for row in rows) == sorted(f"msg_{tag}" for tag, _, _ in answered), rows diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index f4d51d975bb..0a03fa0e157 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -3494,3 +3494,161 @@ async def test_get_async_streaming_response_iterator_yields_small_frame_before_u remaining: Final = tuple([chunk async for chunk in iterator]) assert any(chunk.startswith(b"event: message_stop\n") for chunk in remaining), remaining await iterator.aclose() + + +_NATIVE_EXTENSIONS_MODEL: Final = "us.anthropic.claude-sonnet-4-6" +_TOOL_ADDITION_BLOCK: Final = { + "type": "tool_addition", + "tool_reference": {"type": "tool_reference", "tool_name": "Read"}, +} +_TOOL_USE_BLOCK: Final = {"type": "tool_use", "id": "toolu_01", "name": "Read", "input": {"path": "a.txt"}} + + +def _transform_for_bedrock_invoke( + messages: list[dict], + optional_params: dict | None = None, + config: AmazonAnthropicClaudeMessagesConfig | None = None, +) -> dict: + from litellm.types.router import GenericLiteLLMParams + + return (config or AmazonAnthropicClaudeMessagesConfig()).transform_anthropic_messages_request( + model=_NATIVE_EXTENSIONS_MODEL, + messages=messages, + anthropic_messages_optional_request_params={"max_tokens": 64, **(optional_params or {})}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + +def test_bedrock_invoke_strips_nested_message_output_config_by_default(): + m0: Final = {"role": "user", "content": "hi"} + m1: Final = {"role": "assistant", "content": "hello"} + messages: Final = [m0, m1, {"role": "user", "content": "go", "output_config": {"effort": "low"}}] + + result: Final = _transform_for_bedrock_invoke(messages) + + assert result["messages"] == [m0, m1, {"role": "user", "content": "go"}] + + +def test_bedrock_invoke_strips_tool_addition_blocks_and_keeps_siblings_in_order(): + text_a: Final = {"type": "text", "text": "a"} + text_b: Final = {"type": "text", "text": "b"} + messages: Final = [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": [_TOOL_ADDITION_BLOCK, text_a, _TOOL_USE_BLOCK, _TOOL_ADDITION_BLOCK, text_b]}, + {"role": "user", "content": "go"}, + ] + + result: Final = _transform_for_bedrock_invoke(messages) + + assert result["messages"][1]["content"] == [text_a, _TOOL_USE_BLOCK, text_b] + + +def test_bedrock_invoke_rejects_message_emptied_by_stripping(): + import litellm + + with pytest.raises(litellm.BadRequestError) as exc: + _transform_for_bedrock_invoke( + [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": [_TOOL_ADDITION_BLOCK]}, + {"role": "user", "content": "go"}, + ] + ) + assert "messages[1]" in str(exc.value) + + already_empty: Final = _transform_for_bedrock_invoke( + [{"role": "user", "content": "hi"}, {"role": "assistant", "content": []}, {"role": "user", "content": "go"}] + ) + assert already_empty["messages"][1]["content"] == [] + + +def test_bedrock_invoke_maps_thinking_display_updates_to_summarized(): + optional_params: Final = {"thinking": {"type": "adaptive", "display": "updates"}} + + result: Final = _transform_for_bedrock_invoke([{"role": "user", "content": "hi"}], optional_params) + + assert result["thinking"] == {"type": "adaptive", "display": "summarized"} + assert optional_params["thinking"]["display"] == "updates" + + +@pytest.mark.parametrize( + "thinking", + [{"type": "adaptive", "display": "omitted"}, {"type": "adaptive", "display": "summarized"}, {"type": "adaptive"}], +) +def test_bedrock_invoke_passes_supported_thinking_display_through(thinking: dict): + result: Final = _transform_for_bedrock_invoke([{"role": "user", "content": "hi"}], {"thinking": dict(thinking)}) + + assert result["thinking"] == thinking + + +def test_bedrock_invoke_leaves_clean_body_unchanged(): + messages: Final = [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": [{"type": "text", "text": "hello"}, _TOOL_USE_BLOCK]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "x"}]}, + ] + thinking: Final = {"type": "adaptive", "display": "summarized"} + + result: Final = _transform_for_bedrock_invoke(copy.deepcopy(messages), {"thinking": dict(thinking)}) + + assert result["messages"] == messages + assert result["thinking"] == thinking + + +def test_bedrock_invoke_keeps_unknown_content_block_types(): + unknown_block: Final = {"type": "connector_text", "text": "x"} + messages: Final = [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": [{"type": "text", "text": "a"}, unknown_block, _TOOL_ADDITION_BLOCK]}, + {"role": "user", "content": "go"}, + ] + + result: Final = _transform_for_bedrock_invoke(messages) + + assert result["messages"][1]["content"][1] == unknown_block + + +def test_bedrock_invoke_tolerates_unhashable_discriminator_values(): + messages: Final = [ + {"role": "user", "content": [{"type": "text", "text": "hi"}]}, + {"role": "assistant", "content": [{"type": ["bogus"]}, {"type": "text", "text": "ok"}]}, + {"role": "user", "content": "go"}, + ] + thinking: Final = {"type": "adaptive", "display": ["updates"]} + + result: Final = _transform_for_bedrock_invoke(copy.deepcopy(messages), {"thinking": copy.deepcopy(thinking)}) + + assert result["messages"] == messages + assert result["thinking"] == thinking + + +def test_bedrock_invoke_does_not_mutate_caller_messages(): + messages: Final = [ + {"role": "user", "content": "hi", "output_config": {"effort": "low"}}, + {"role": "assistant", "content": [_TOOL_ADDITION_BLOCK, {"type": "text", "text": "ok"}]}, + {"role": "user", "content": "go"}, + ] + optional_params: Final = {"thinking": {"type": "adaptive", "display": "updates"}} + messages_snapshot: Final = copy.deepcopy(messages) + optional_params_snapshot: Final = copy.deepcopy(optional_params) + + _transform_for_bedrock_invoke(messages, optional_params) + + assert messages == messages_snapshot + assert optional_params["thinking"] == optional_params_snapshot["thinking"] + + +def test_bedrock_invoke_native_extension_policy_is_class_resolved(): + class _KeepsToolAdditions(AmazonAnthropicClaudeMessagesConfig): + BEDROCK_INVOKE_UNSUPPORTED_CONTENT_BLOCK_TYPES: frozenset[str] = frozenset() + + messages: Final = [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": [_TOOL_ADDITION_BLOCK, {"type": "text", "text": "ok"}]}, + {"role": "user", "content": "go"}, + ] + + result: Final = _transform_for_bedrock_invoke(messages, config=_KeepsToolAdditions()) + + assert result["messages"][1]["content"] == [_TOOL_ADDITION_BLOCK, {"type": "text", "text": "ok"}] diff --git a/tests/test_litellm/llms/bedrock/test_mantle.py b/tests/test_litellm/llms/bedrock/test_mantle.py index 37cf49a85ec..22fa3e930a3 100644 --- a/tests/test_litellm/llms/bedrock/test_mantle.py +++ b/tests/test_litellm/llms/bedrock/test_mantle.py @@ -792,3 +792,32 @@ async def test_mantle_anthropic_messages_streaming_sends_stream_and_passes_throu assert "event: message_start" in text assert '"text": "pong"' in text assert "event: message_stop" in text + + +def test_mantle_messages_keep_per_message_output_config_tool_addition_and_display_updates(): + from litellm.llms.bedrock_mantle.messages.transformation import BedrockMantleAnthropicMessagesConfig + from litellm.types.router import GenericLiteLLMParams + + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": [ + {"type": "tool_addition", "tool_reference": {"type": "tool_reference", "tool_name": "Read"}}, + {"type": "text", "text": "ok"}, + ], + }, + {"role": "user", "content": "go", "output_config": {"effort": "low"}}, + ] + thinking = {"type": "adaptive", "display": "updates"} + + for config in (AmazonMantleMessagesConfig(), BedrockMantleAnthropicMessagesConfig()): + result = config.transform_anthropic_messages_request( + model="mantle/claude-mythos-preview", + messages=json.loads(json.dumps(messages)), + anthropic_messages_optional_request_params={"max_tokens": 64, "thinking": dict(thinking)}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert result["messages"] == messages, type(config).__name__ + assert result["thinking"] == thinking, type(config).__name__ diff --git a/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py b/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py index b357c5ac126..fde18b51a50 100644 --- a/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py +++ b/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py @@ -253,3 +253,88 @@ def test_count_tokens_endpoint_encodes_model_id(monkeypatch): endpoint == "https://bedrock-runtime.us-east-1.amazonaws.com/model/..%2F..%2Fmodel%2Fother%3Fx%3D1%23frag/count-tokens" ) + + +def _decoded_invoke_body(result: dict) -> dict: + return json.loads(base64.b64decode(result["input"]["invokeModel"]["body"])) + + +def test_transform_to_invoke_model_format_strips_per_message_output_config(): + config = BedrockCountTokensConfig() + request = { + "model": "us.anthropic.claude-sonnet-4-6", + "messages": [ + {"role": "user", "content": [{"type": "text", "text": "hi"}]}, + {"role": "assistant", "content": [{"type": "text", "text": "hello"}]}, + {"role": "user", "content": [{"type": "text", "text": "go"}], "output_config": {"effort": "low"}}, + ], + } + + body = _decoded_invoke_body(config.transform_anthropic_to_bedrock_count_tokens(request)) + + assert body["messages"] == [ + request["messages"][0], + request["messages"][1], + {"role": "user", "content": [{"type": "text", "text": "go"}]}, + ] + + +def test_transform_to_invoke_model_format_strips_tool_addition_blocks_and_keeps_siblings(): + config = BedrockCountTokensConfig() + text_a = {"type": "text", "text": "a"} + tool_use = {"type": "tool_use", "id": "toolu_01", "name": "Read", "input": {"path": "a.txt"}} + text_b = {"type": "text", "text": "b"} + tool_addition = {"type": "tool_addition", "tool_reference": {"type": "tool_reference", "tool_name": "Read"}} + request = { + "model": "us.anthropic.claude-sonnet-4-6", + "messages": [ + {"role": "user", "content": [{"type": "text", "text": "hi"}]}, + {"role": "assistant", "content": [tool_addition, text_a, tool_use, tool_addition, text_b]}, + ], + } + + body = _decoded_invoke_body(config.transform_anthropic_to_bedrock_count_tokens(request)) + + assert body["messages"][1]["content"] == [text_a, tool_use, text_b] + + +def test_transform_to_invoke_model_format_leaves_clean_anthropic_body_unchanged(): + config = BedrockCountTokensConfig() + request = { + "model": "us.anthropic.claude-sonnet-4-6", + "system": "be brief", + "messages": [ + {"role": "user", "content": [{"type": "text", "text": "hi"}]}, + {"role": "assistant", "content": [{"type": "text", "text": "hello"}]}, + ], + } + + body = _decoded_invoke_body(config.transform_anthropic_to_bedrock_count_tokens(request)) + + assert body == { + "system": "be brief", + "messages": request["messages"], + "anthropic_version": "bedrock-2023-05-31", + "max_tokens": DEFAULT_ANTHROPIC_INVOKE_MODEL_MAX_TOKENS, + } + + +def test_transform_anthropic_to_bedrock_request_string_content_still_uses_converse(): + config = BedrockCountTokensConfig() + request = { + "model": "us.anthropic.claude-sonnet-4-6", + "messages": [{"role": "user", "content": "Hello"}, {"role": "assistant", "content": "Hi"}], + } + + result = config.transform_anthropic_to_bedrock_count_tokens(request) + + assert result == { + "input": { + "converse": { + "messages": [ + {"role": "user", "content": [{"text": "Hello"}]}, + {"role": "assistant", "content": [{"text": "Hi"}]}, + ] + } + } + }