mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(bedrock): handle thinking with tool calls for Claude 4 models
This commit is contained in:
parent
e285e2b91d
commit
81fefc69c9
4 changed files with 398 additions and 220 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue