fix(bedrock): handle thinking with tool calls for Claude 4 models

This commit is contained in:
Benedikt Óskarsson 2026-01-08 00:44:59 +00:00
parent e285e2b91d
commit 81fefc69c9
No known key found for this signature in database
4 changed files with 398 additions and 220 deletions

View file

@ -53,7 +53,12 @@ from litellm.types.utils import (
PromptTokensDetailsWrapper,
Usage,
)
from litellm.utils import add_dummy_tool, has_tool_call_blocks, supports_reasoning
from litellm.utils import (
add_dummy_tool,
has_tool_call_blocks,
last_assistant_with_tool_calls_has_no_thinking_blocks,
supports_reasoning,
)
from ..common_utils import (
BedrockError,
@ -729,7 +734,7 @@ class AmazonConverseConfig(BaseConfig):
return optional_params
"""
Follow similar approach to anthropic - translate to a single tool call.
Follow similar approach to anthropic - translate to a single tool call.
When using tools in this way: - https://docs.anthropic.com/en/docs/build-with-claude/tool-use#json-mode
- You usually want to provide a single tool
@ -912,16 +917,16 @@ class AmazonConverseConfig(BaseConfig):
inference_params = {
k: v for k, v in inference_params.items() if k in total_supported_params
}
# Only set the topK value in for models that support it
additional_request_params.update(
self._handle_top_k_value(model, inference_params)
)
# Filter out internal/MCP-related parameters that shouldn't be sent to the API
# These are LiteLLM internal parameters, not API parameters
additional_request_params = filter_internal_params(additional_request_params)
# Filter out non-serializable objects (exceptions, callables, logging objects, etc.)
# from additional_request_params to prevent JSON serialization errors
# This filters: Exception objects, callable objects (functions), Logging objects, etc.
@ -1021,9 +1026,24 @@ class AmazonConverseConfig(BaseConfig):
llm_provider="bedrock",
)
# Drop thinking param if thinking is enabled but thinking_blocks are missing
# This prevents the error: "Expected thinking or redacted_thinking, but found tool_use"
# Related issues: https://github.com/BerriAI/litellm/issues/14194
if (
optional_params.get("thinking") is not None
and messages is not None
and last_assistant_with_tool_calls_has_no_thinking_blocks(messages)
):
if litellm.modify_params:
optional_params.pop("thinking", None)
litellm.verbose_logger.warning(
"Dropping 'thinking' param because the last assistant message with tool_calls "
"has no thinking_blocks. The model won't use extended thinking for this turn."
)
# Prepare and separate parameters
inference_params, additional_request_params, request_metadata = (
self._prepare_request_params(optional_params, model)
inference_params, additional_request_params, request_metadata = self._prepare_request_params(
optional_params, model
)
original_tools = inference_params.pop("tools", [])
@ -1410,11 +1430,11 @@ class AmazonConverseConfig(BaseConfig):
)
"""
Bedrock Response Object has optional message block
Bedrock Response Object has optional message block
completion_response["output"].get("message", None)
A message block looks like this (Example 1):
A message block looks like this (Example 1):
"output": {
"message": {
"role": "assistant",

View file

@ -374,6 +374,29 @@ class BedrockLLM(BaseAWSLLM):
def __init__(self) -> None:
super().__init__()
@staticmethod
def is_claude_messages_api_model(model: str) -> bool:
"""
Check if the model uses the Claude Messages API (Claude 3+).
Handles:
- Regional prefixes: eu.anthropic.claude-*, us.anthropic.claude-*
- Claude 3 models: claude-3-haiku, claude-3-sonnet, claude-3-opus, claude-3-5-*, claude-3-7-*
- Claude 4 models: claude-opus-4, claude-sonnet-4, claude-haiku-4
"""
# Normalize model string to lowercase for matching
model_lower = model.lower()
# Claude 3+ indicators (all use Messages API)
messages_api_indicators = [
"claude-3", # Claude 3.x models
"claude-opus-4", # Claude Opus 4
"claude-sonnet-4", # Claude Sonnet 4
"claude-haiku-4", # Claude Haiku 4
]
return any(indicator in model_lower for indicator in messages_api_indicators)
def convert_messages_to_prompt(
self, model, messages, provider, custom_prompt_dict
) -> Tuple[str, Optional[list]]:
@ -465,7 +488,7 @@ class BedrockLLM(BaseAWSLLM):
completion_response["generations"][0]["finish_reason"]
)
elif provider == "anthropic":
if model.startswith("anthropic.claude-3"):
if self.is_claude_messages_api_model(model):
json_schemas: dict = {}
_is_function_call = False
## Handle Tool Calling
@ -595,13 +618,12 @@ class BedrockLLM(BaseAWSLLM):
outputText = choice["message"].get("content")
elif "text" in choice: # fallback for completion format
outputText = choice["text"]
# Set finish reason
if "finish_reason" in choice:
model_response.choices[0].finish_reason = map_finish_reason(
choice["finish_reason"]
)
# Set usage if available
if "usage" in completion_response:
usage = completion_response["usage"]
@ -842,7 +864,7 @@ class BedrockLLM(BaseAWSLLM):
] = True # cohere requires stream = True in inference params
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "anthropic":
if model.startswith("anthropic.claude-3"):
if self.is_claude_messages_api_model(model):
# Separate system prompt from rest of message
system_prompt_idx: list[int] = []
system_messages: list[str] = []
@ -940,13 +962,13 @@ class BedrockLLM(BaseAWSLLM):
# Use AmazonBedrockOpenAIConfig for proper OpenAI transformation
openai_config = AmazonBedrockOpenAIConfig()
supported_params = openai_config.get_supported_openai_params(model=model)
# Filter to only supported OpenAI params
filtered_params = {
k: v for k, v in inference_params.items()
k: v for k, v in inference_params.items()
if k in supported_params
}
# OpenAI uses messages format, not prompt
data = json.dumps({"messages": messages, **filtered_params})
else:

View file

@ -621,7 +621,7 @@ def load_credentials_from_list(kwargs: dict):
"""
# Access CredentialAccessor via module to trigger lazy loading if needed
CredentialAccessor = getattr(sys.modules[__name__], 'CredentialAccessor')
credential_name = kwargs.get("litellm_credential_name")
if credential_name and litellm.credential_list:
credential_accessor = CredentialAccessor.get_credential_values(credential_name)
@ -648,7 +648,7 @@ def _is_gemini_model(model: Optional[str], custom_llm_provider: Optional[str]) -
if custom_llm_provider in ["vertex_ai", "vertex_ai_beta"]:
return model is not None and "gemini" in model.lower()
return True
# Check if model name contains gemini
return model is not None and "gemini" in model.lower()
@ -670,7 +670,7 @@ def _process_assistant_message_tool_calls(
"""
role = msg_copy.get("role")
tool_calls = msg_copy.get("tool_calls")
if role == "assistant" and isinstance(tool_calls, list):
new_tool_calls = []
for tc in tool_calls:
@ -683,17 +683,17 @@ def _process_assistant_message_tool_calls(
else:
new_tool_calls.append(tc)
continue
# Remove thought signature from ID if present
if isinstance(tc_dict.get("id"), str):
if thought_signature_separator in tc_dict["id"]:
tc_dict["id"] = _remove_thought_signature_from_id(
tc_dict["id"], thought_signature_separator
)
new_tool_calls.append(tc_dict)
msg_copy["tool_calls"] = new_tool_calls
return msg_copy
@ -708,7 +708,7 @@ def _process_tool_message_id(msg_copy: dict, thought_signature_separator: str) -
msg_copy["tool_call_id"] = _remove_thought_signature_from_id(
msg_copy["tool_call_id"], thought_signature_separator
)
return msg_copy
@ -719,7 +719,7 @@ def _remove_thought_signatures_from_messages(
Remove thought signatures from tool call IDs in all messages.
"""
processed_messages = []
for msg in messages:
# Handle Pydantic models (convert to dict)
if hasattr(msg, "model_dump"):
@ -730,17 +730,17 @@ def _remove_thought_signatures_from_messages(
# Unknown type, keep as is
processed_messages.append(msg)
continue
# Process assistant messages with tool_calls
msg_dict = _process_assistant_message_tool_calls(
msg_dict, thought_signature_separator
)
# Process tool messages with tool_call_id
msg_dict = _process_tool_message_id(msg_dict, thought_signature_separator)
processed_messages.append(msg_dict)
return processed_messages
@ -960,7 +960,7 @@ def function_setup( # noqa: PLR0915
input=buffer.getvalue(),
model=model,
)
### REMOVE THOUGHT SIGNATURES FROM TOOL CALL IDS FOR NON-GEMINI MODELS ###
# Gemini models embed thought signatures in tool call IDs. When sending
# messages with tool calls to non-Gemini providers, we need to remove these
@ -976,7 +976,7 @@ def function_setup( # noqa: PLR0915
# Get custom_llm_provider to determine target provider
custom_llm_provider = kwargs.get("custom_llm_provider")
# If custom_llm_provider not in kwargs, try to determine it from the model
if not custom_llm_provider and model:
try:
@ -987,18 +987,18 @@ def function_setup( # noqa: PLR0915
except Exception:
# If we can't determine the provider, skip this processing
pass
# Only process if target is NOT a Gemini model
if not _is_gemini_model(model, custom_llm_provider):
verbose_logger.debug(
"Removing thought signatures from tool call IDs for non-Gemini model"
)
# Process messages to remove thought signatures
processed_messages = _remove_thought_signatures_from_messages(
messages, THOUGHT_SIGNATURE_SEPARATOR
)
# Update messages in kwargs or args
if "messages" in kwargs:
kwargs["messages"] = processed_messages
@ -2977,7 +2977,7 @@ def get_optional_params_embeddings( # noqa: PLR0915
):
# Lazy load get_supported_openai_params
get_supported_openai_params = getattr(sys.modules[__name__], 'get_supported_openai_params')
# retrieve all parameters passed to the function
passed_params = locals()
custom_llm_provider = passed_params.pop("custom_llm_provider", None)
@ -4063,7 +4063,14 @@ def get_optional_params( # noqa: PLR0915
),
)
elif "anthropic" in bedrock_base_model and bedrock_route == "invoke":
if bedrock_base_model.startswith("anthropic.claude-3"):
# Check for Claude 3+ models (Messages API) including regional prefixes and Claude 4
# Models like eu.anthropic.claude-opus-4-5, us.anthropic.claude-3-5-sonnet, etc.
bedrock_base_model_lower = bedrock_base_model.lower()
is_messages_api_model = any(
indicator in bedrock_base_model_lower
for indicator in ["claude-3", "claude-opus-4", "claude-sonnet-4", "claude-haiku-4"]
)
if is_messages_api_model:
optional_params = (
litellm.AmazonAnthropicClaudeConfig().map_openai_params(
non_default_params=non_default_params,
@ -6911,7 +6918,7 @@ def get_valid_models(
# init litellm_params
#################################
from litellm.types.router import LiteLLM_Params
if litellm_params is None:
litellm_params = LiteLLM_Params(model="")
if api_key is not None:
@ -7513,7 +7520,7 @@ class ProviderConfigManager:
return litellm.IBMWatsonXAIConfig()
elif litellm.LlmProviders.EMPOWER == provider:
return litellm.EmpowerChatConfig()
elif litellm.LlmProviders.MINIMAX == provider:
elif litellm.LlmProviders.MINIMAX == provider:
return litellm.MinimaxChatConfig()
elif litellm.LlmProviders.GITHUB == provider:
return litellm.GithubChatConfig()
@ -8314,8 +8321,7 @@ class ProviderConfigManager:
from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config
return get_vertex_ai_ocr_config(model=model)
MistralOCRConfig = getattr(sys.modules[__name__], 'MistralOCRConfig')
MistralOCRConfig = getattr(sys.modules[__name__], 'MistralOCRConfig')
PROVIDER_TO_CONFIG_MAP = {
litellm.LlmProviders.MISTRAL: MistralOCRConfig,
}
@ -8752,12 +8758,12 @@ def __getattr__(name: str) -> Any:
"""Lazy import handler for utils module with cached registry for improved performance."""
# Use cached registry from _lazy_imports instead of importing tuples every time
from litellm._lazy_imports import _get_lazy_import_registry
registry = _get_lazy_import_registry()
# Check if name is in registry and call the cached handler function
if name in registry:
handler_func = registry[name]
return handler_func(name)
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")