Fix bedrock anthropic models to use system field (#3438)
Signed-off-by: Andrew Carbonetto <andrew.carbonetto@improving.com>
This commit is contained in:
committed by
GitHub
parent
e3f0277cb9
commit
9e5810dfb7
+89
-30
@@ -136,23 +136,27 @@ class AWSBedrockLLM(LLMBase):
|
||||
else:
|
||||
self._format_messages = self._format_messages_generic
|
||||
|
||||
def _format_messages_anthropic(self, messages: List[Dict[str, str]]) -> List[Dict[str, Any]]:
|
||||
def _format_messages_anthropic(self, messages: List[Dict[str, str]]) -> tuple[List[Dict[str, Any]], Optional[str]]:
|
||||
"""Format messages for Anthropic models."""
|
||||
formatted_messages = []
|
||||
system_message = None
|
||||
|
||||
for message in messages:
|
||||
role = message["role"]
|
||||
content = message["content"]
|
||||
|
||||
if role == "system":
|
||||
# Anthropic doesn't support system messages, prepend to first user message
|
||||
continue
|
||||
# Anthropic supports system messages as a separate parameter
|
||||
# see: https://docs.anthropic.com/en/docs/build-with-claude/prompt-engineering/system-prompts
|
||||
system_message = content
|
||||
elif role == "user":
|
||||
formatted_messages.append({"role": "user", "content": [{"type": "text", "text": content}]})
|
||||
# Use Converse API format
|
||||
formatted_messages.append({"role": "user", "content": [{"text": content}]})
|
||||
elif role == "assistant":
|
||||
formatted_messages.append({"role": "assistant", "content": [{"type": "text", "text": content}]})
|
||||
# Use Converse API format
|
||||
formatted_messages.append({"role": "assistant", "content": [{"text": content}]})
|
||||
|
||||
return formatted_messages
|
||||
return formatted_messages, system_message
|
||||
|
||||
def _format_messages_cohere(self, messages: List[Dict[str, str]]) -> str:
|
||||
"""Format messages for Cohere models."""
|
||||
@@ -451,48 +455,103 @@ class AWSBedrockLLM(LLMBase):
|
||||
logger.error(f"Failed to generate response: {e}")
|
||||
raise RuntimeError(f"Failed to generate response: {e}")
|
||||
|
||||
@staticmethod
|
||||
def _convert_tools_to_converse_format(tools: List[Dict]) -> List[Dict]:
|
||||
"""Convert OpenAI-style tools to Converse API format."""
|
||||
if not tools:
|
||||
return []
|
||||
|
||||
converse_tools = []
|
||||
for tool in tools:
|
||||
if tool.get("type") == "function" and "function" in tool:
|
||||
func = tool["function"]
|
||||
converse_tool = {
|
||||
"toolSpec": {
|
||||
"name": func["name"],
|
||||
"description": func.get("description", ""),
|
||||
"inputSchema": {
|
||||
"json": func.get("parameters", {})
|
||||
}
|
||||
}
|
||||
}
|
||||
converse_tools.append(converse_tool)
|
||||
|
||||
return converse_tools
|
||||
|
||||
def _generate_with_tools(self, messages: List[Dict[str, str]], tools: List[Dict], stream: bool = False) -> Dict[str, Any]:
|
||||
"""Generate response with tool calling support."""
|
||||
"""Generate response with tool calling support using correct message format."""
|
||||
# Format messages for tool-enabled models
|
||||
system_message = None
|
||||
if self.provider == "anthropic":
|
||||
formatted_messages = self._format_messages_anthropic(messages)
|
||||
formatted_messages, system_message = self._format_messages_anthropic(messages)
|
||||
elif self.provider == "amazon":
|
||||
formatted_messages = self._format_messages_amazon(messages)
|
||||
else:
|
||||
formatted_messages = [{"role": "user", "content": messages[-1]["content"]}]
|
||||
formatted_messages = [{"role": "user", "content": [{"text": messages[-1]["content"]}]}]
|
||||
|
||||
# Prepare inference configuration
|
||||
inference_config = {
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"maxTokens": self.model_config.get("max_tokens", 2000),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
# Prepare tool configuration in Converse API format
|
||||
tool_config = None
|
||||
if tools:
|
||||
converse_tools = self._convert_tools_to_converse_format(tools)
|
||||
if converse_tools:
|
||||
tool_config = {"tools": converse_tools}
|
||||
|
||||
# Prepare converse parameters
|
||||
converse_params = {
|
||||
"modelId": self.config.model,
|
||||
"messages": formatted_messages,
|
||||
"inferenceConfig": {
|
||||
"maxTokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
}
|
||||
|
||||
# Prepare tools configuration
|
||||
tools_config = {"tools": self._convert_tool_format(tools)}
|
||||
# Add system message if present (for Anthropic)
|
||||
if system_message:
|
||||
converse_params["system"] = [{"text": system_message}]
|
||||
|
||||
# Add tool config if present
|
||||
if tool_config:
|
||||
converse_params["toolConfig"] = tool_config
|
||||
|
||||
# Make API call
|
||||
response = self.client.converse(
|
||||
modelId=self.config.model,
|
||||
messages=formatted_messages,
|
||||
inferenceConfig=inference_config,
|
||||
toolConfig=tools_config,
|
||||
)
|
||||
response = self.client.converse(**converse_params)
|
||||
|
||||
return self._parse_response(response, tools)
|
||||
|
||||
def _generate_standard(self, messages: List[Dict[str, str]], stream: bool = False) -> str:
|
||||
"""Generate standard text response."""
|
||||
# Format messages according to provider
|
||||
"""Generate standard text response using Converse API for Anthropic models."""
|
||||
# For Anthropic models, always use Converse API
|
||||
if self.provider == "anthropic":
|
||||
formatted_messages = self._format_messages_anthropic(messages)
|
||||
input_body = {
|
||||
formatted_messages, system_message = self._format_messages_anthropic(messages)
|
||||
|
||||
# Prepare converse parameters
|
||||
converse_params = {
|
||||
"modelId": self.config.model,
|
||||
"messages": formatted_messages,
|
||||
"max_tokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
"inferenceConfig": {
|
||||
"maxTokens": self.model_config.get("max_tokens", 2000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
}
|
||||
|
||||
# Add system message if present
|
||||
if system_message:
|
||||
converse_params["system"] = [{"text": system_message}]
|
||||
|
||||
# Use converse API for Anthropic models
|
||||
response = self.client.converse(**converse_params)
|
||||
|
||||
# Parse Converse API response
|
||||
if hasattr(response, 'output') and hasattr(response.output, 'message'):
|
||||
return response.output.message.content[0].text
|
||||
elif 'output' in response and 'message' in response['output']:
|
||||
return response['output']['message']['content'][0]['text']
|
||||
else:
|
||||
return str(response)
|
||||
|
||||
elif self.provider == "amazon" and "nova" in self.config.model.lower():
|
||||
# Nova models use converse API even without tools
|
||||
formatted_messages = self._format_messages_amazon(messages)
|
||||
|
||||
Reference in New Issue
Block a user