Update aws bedrock (#3334)
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
from typing import Optional, Dict, Any, List
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
import os
|
||||
|
||||
|
||||
class AWSBedrockConfig(BaseLlmConfig):
|
||||
"""
|
||||
Configuration class for AWS Bedrock LLM integration.
|
||||
|
||||
Supports all available Bedrock models with automatic provider detection.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: Optional[str] = None,
|
||||
temperature: float = 0.1,
|
||||
max_tokens: int = 2000,
|
||||
top_p: float = 0.9,
|
||||
top_k: int = 1,
|
||||
aws_access_key_id: Optional[str] = None,
|
||||
aws_secret_access_key: Optional[str] = None,
|
||||
aws_region: str = "us-west-2",
|
||||
aws_session_token: Optional[str] = None,
|
||||
aws_profile: Optional[str] = None,
|
||||
model_kwargs: Optional[Dict[str, Any]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize AWS Bedrock configuration.
|
||||
|
||||
Args:
|
||||
model: Bedrock model identifier (e.g., "amazon.nova-3-mini-20241119-v1:0")
|
||||
temperature: Controls randomness (0.0 to 2.0)
|
||||
max_tokens: Maximum tokens to generate
|
||||
top_p: Nucleus sampling parameter (0.0 to 1.0)
|
||||
top_k: Top-k sampling parameter (1 to 40)
|
||||
aws_access_key_id: AWS access key (optional, uses env vars if not provided)
|
||||
aws_secret_access_key: AWS secret key (optional, uses env vars if not provided)
|
||||
aws_region: AWS region for Bedrock service
|
||||
aws_session_token: AWS session token for temporary credentials
|
||||
aws_profile: AWS profile name for credentials
|
||||
model_kwargs: Additional model-specific parameters
|
||||
**kwargs: Additional arguments passed to base class
|
||||
"""
|
||||
super().__init__(
|
||||
model=model or "anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
self.aws_access_key_id = aws_access_key_id
|
||||
self.aws_secret_access_key = aws_secret_access_key
|
||||
self.aws_region = aws_region
|
||||
self.aws_session_token = aws_session_token
|
||||
self.aws_profile = aws_profile
|
||||
self.model_kwargs = model_kwargs or {}
|
||||
|
||||
@property
|
||||
def provider(self) -> str:
|
||||
"""Get the provider from the model identifier."""
|
||||
if not self.model or "." not in self.model:
|
||||
return "unknown"
|
||||
return self.model.split(".")[0]
|
||||
|
||||
@property
|
||||
def model_name(self) -> str:
|
||||
"""Get the model name without provider prefix."""
|
||||
if not self.model or "." not in self.model:
|
||||
return self.model
|
||||
return ".".join(self.model.split(".")[1:])
|
||||
|
||||
def get_model_config(self) -> Dict[str, Any]:
|
||||
"""Get model-specific configuration parameters."""
|
||||
base_config = {
|
||||
"temperature": self.temperature,
|
||||
"max_tokens": self.max_tokens,
|
||||
"top_p": self.top_p,
|
||||
"top_k": self.top_k,
|
||||
}
|
||||
|
||||
# Add custom model kwargs
|
||||
base_config.update(self.model_kwargs)
|
||||
|
||||
return base_config
|
||||
|
||||
def get_aws_config(self) -> Dict[str, Any]:
|
||||
"""Get AWS configuration parameters."""
|
||||
config = {
|
||||
"region_name": self.aws_region,
|
||||
}
|
||||
|
||||
if self.aws_access_key_id:
|
||||
config["aws_access_key_id"] = self.aws_access_key_id or os.getenv("AWS_ACCESS_KEY_ID")
|
||||
|
||||
if self.aws_secret_access_key:
|
||||
config["aws_secret_access_key"] = self.aws_secret_access_key or os.getenv("AWS_SECRET_ACCESS_KEY")
|
||||
|
||||
if self.aws_session_token:
|
||||
config["aws_session_token"] = self.aws_session_token or os.getenv("AWS_SESSION_TOKEN")
|
||||
|
||||
if self.aws_profile:
|
||||
config["profile_name"] = self.aws_profile or os.getenv("AWS_PROFILE")
|
||||
|
||||
return config
|
||||
|
||||
def validate_model_format(self) -> bool:
|
||||
"""
|
||||
Validate that the model identifier follows Bedrock naming convention.
|
||||
|
||||
Returns:
|
||||
True if valid, False otherwise
|
||||
"""
|
||||
if not self.model:
|
||||
return False
|
||||
|
||||
# Check if model follows provider.model-name format
|
||||
if "." not in self.model:
|
||||
return False
|
||||
|
||||
provider, model_name = self.model.split(".", 1)
|
||||
|
||||
# Validate provider
|
||||
valid_providers = [
|
||||
"ai21", "amazon", "anthropic", "cohere", "meta", "mistral",
|
||||
"stability", "writer", "deepseek", "gpt-oss", "perplexity",
|
||||
"snowflake", "titan", "command", "j2", "llama"
|
||||
]
|
||||
|
||||
if provider not in valid_providers:
|
||||
return False
|
||||
|
||||
# Validate model name is not empty
|
||||
if not model_name:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def get_supported_regions(self) -> List[str]:
|
||||
"""Get list of AWS regions that support Bedrock."""
|
||||
return [
|
||||
"us-east-1",
|
||||
"us-west-2",
|
||||
"us-east-2",
|
||||
"eu-west-1",
|
||||
"ap-southeast-1",
|
||||
"ap-northeast-1",
|
||||
]
|
||||
|
||||
def get_model_capabilities(self) -> Dict[str, Any]:
|
||||
"""Get model capabilities based on provider."""
|
||||
capabilities = {
|
||||
"supports_tools": False,
|
||||
"supports_vision": False,
|
||||
"supports_streaming": False,
|
||||
"supports_multimodal": False,
|
||||
}
|
||||
|
||||
if self.provider == "anthropic":
|
||||
capabilities.update({
|
||||
"supports_tools": True,
|
||||
"supports_vision": True,
|
||||
"supports_streaming": True,
|
||||
"supports_multimodal": True,
|
||||
})
|
||||
elif self.provider == "amazon":
|
||||
capabilities.update({
|
||||
"supports_tools": True,
|
||||
"supports_vision": True,
|
||||
"supports_streaming": True,
|
||||
"supports_multimodal": True,
|
||||
})
|
||||
elif self.provider == "cohere":
|
||||
capabilities.update({
|
||||
"supports_tools": True,
|
||||
"supports_streaming": True,
|
||||
})
|
||||
elif self.provider == "meta":
|
||||
capabilities.update({
|
||||
"supports_vision": True,
|
||||
"supports_streaming": True,
|
||||
})
|
||||
elif self.provider == "mistral":
|
||||
capabilities.update({
|
||||
"supports_vision": True,
|
||||
"supports_streaming": True,
|
||||
})
|
||||
|
||||
return capabilities
|
||||
+526
-196
@@ -1,20 +1,28 @@
|
||||
import json
|
||||
import os
|
||||
import logging
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
try:
|
||||
import boto3
|
||||
from botocore.exceptions import ClientError, NoCredentialsError
|
||||
except ImportError:
|
||||
raise ImportError("The 'boto3' library is required. Please install it using 'pip install boto3'.")
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.configs.llms.aws_bedrock import AWSBedrockConfig
|
||||
from mem0.llms.base import LLMBase
|
||||
|
||||
PROVIDERS = ["ai21", "amazon", "anthropic", "cohere", "meta", "mistral", "stability", "writer"]
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
PROVIDERS = [
|
||||
"ai21", "amazon", "anthropic", "cohere", "meta", "mistral", "stability", "writer",
|
||||
"deepseek", "gpt-oss", "perplexity", "snowflake", "titan", "command", "j2", "llama"
|
||||
]
|
||||
|
||||
|
||||
def extract_provider(model: str) -> str:
|
||||
"""Extract provider from model identifier."""
|
||||
for provider in PROVIDERS:
|
||||
if re.search(rf"\b{re.escape(provider)}\b", model):
|
||||
return provider
|
||||
@@ -22,51 +30,192 @@ def extract_provider(model: str) -> str:
|
||||
|
||||
|
||||
class AWSBedrockLLM(LLMBase):
|
||||
def __init__(self, config: Optional[BaseLlmConfig] = None):
|
||||
super().__init__(config)
|
||||
"""
|
||||
AWS Bedrock LLM integration for Mem0.
|
||||
|
||||
if not self.config.model:
|
||||
self.config.model = "anthropic.claude-3-5-sonnet-20240620-v1:0"
|
||||
Supports all available Bedrock models with automatic provider detection.
|
||||
"""
|
||||
|
||||
# Get AWS config from environment variables or use defaults
|
||||
aws_access_key = os.environ.get("AWS_ACCESS_KEY_ID", "")
|
||||
aws_secret_key = os.environ.get("AWS_SECRET_ACCESS_KEY", "")
|
||||
aws_region = os.environ.get("AWS_REGION", "us-west-2")
|
||||
|
||||
# Check if AWS config is provided in the config
|
||||
if hasattr(self.config, "aws_access_key_id"):
|
||||
aws_access_key = self.config.aws_access_key_id
|
||||
if hasattr(self.config, "aws_secret_access_key"):
|
||||
aws_secret_key = self.config.aws_secret_access_key
|
||||
if hasattr(self.config, "aws_region"):
|
||||
aws_region = self.config.aws_region
|
||||
|
||||
self.client = boto3.client(
|
||||
"bedrock-runtime",
|
||||
region_name=aws_region,
|
||||
aws_access_key_id=aws_access_key if aws_access_key else None,
|
||||
aws_secret_access_key=aws_secret_key if aws_secret_key else None,
|
||||
)
|
||||
|
||||
self.model_kwargs = {
|
||||
"temperature": self.config.temperature,
|
||||
"max_tokens_to_sample": self.config.max_tokens,
|
||||
"top_p": self.config.top_p,
|
||||
}
|
||||
|
||||
def _format_messages(self, messages: List[Dict[str, str]]) -> str:
|
||||
def __init__(self, config: Optional[Union[AWSBedrockConfig, BaseLlmConfig, Dict]] = None):
|
||||
"""
|
||||
Formats a list of messages into the required prompt structure for the model.
|
||||
Initialize AWS Bedrock LLM.
|
||||
|
||||
Args:
|
||||
messages (List[Dict[str, str]]): A list of dictionaries where each dictionary represents a message.
|
||||
Each dictionary contains 'role' and 'content' keys.
|
||||
|
||||
Returns:
|
||||
str: A formatted string combining all messages, structured with roles capitalized and separated by newlines.
|
||||
config: AWS Bedrock configuration object
|
||||
"""
|
||||
# Convert to AWSBedrockConfig if needed
|
||||
if config is None:
|
||||
config = AWSBedrockConfig()
|
||||
elif isinstance(config, dict):
|
||||
config = AWSBedrockConfig(**config)
|
||||
elif isinstance(config, BaseLlmConfig) and not isinstance(config, AWSBedrockConfig):
|
||||
# Convert BaseLlmConfig to AWSBedrockConfig
|
||||
config = AWSBedrockConfig(
|
||||
model=config.model,
|
||||
temperature=config.temperature,
|
||||
max_tokens=config.max_tokens,
|
||||
top_p=config.top_p,
|
||||
top_k=config.top_k,
|
||||
enable_vision=getattr(config, "enable_vision", False),
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
# Initialize AWS client
|
||||
self._initialize_aws_client()
|
||||
|
||||
# Get model configuration
|
||||
self.model_config = self.config.get_model_config()
|
||||
self.provider = extract_provider(self.config.model)
|
||||
|
||||
# Initialize provider-specific settings
|
||||
self._initialize_provider_settings()
|
||||
|
||||
def _initialize_aws_client(self):
|
||||
"""Initialize AWS Bedrock client with proper credentials."""
|
||||
try:
|
||||
aws_config = self.config.get_aws_config()
|
||||
|
||||
# Create Bedrock runtime client
|
||||
self.client = boto3.client("bedrock-runtime", **aws_config)
|
||||
|
||||
# Test connection
|
||||
self._test_connection()
|
||||
|
||||
except NoCredentialsError:
|
||||
raise ValueError(
|
||||
"AWS credentials not found. Please set AWS_ACCESS_KEY_ID, "
|
||||
"AWS_SECRET_ACCESS_KEY, and AWS_REGION environment variables, "
|
||||
"or provide them in the config."
|
||||
)
|
||||
except ClientError as e:
|
||||
if e.response["Error"]["Code"] == "UnauthorizedOperation":
|
||||
raise ValueError(
|
||||
f"Unauthorized access to Bedrock. Please ensure your AWS credentials "
|
||||
f"have permission to access Bedrock in region {self.config.aws_region}."
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"AWS Bedrock error: {e}")
|
||||
|
||||
def _test_connection(self):
|
||||
"""Test connection to AWS Bedrock service."""
|
||||
try:
|
||||
# List available models to test connection
|
||||
bedrock_client = boto3.client("bedrock", **self.config.get_aws_config())
|
||||
response = bedrock_client.list_foundation_models()
|
||||
self.available_models = [model["modelId"] for model in response["modelSummaries"]]
|
||||
|
||||
# Check if our model is available
|
||||
if self.config.model not in self.available_models:
|
||||
logger.warning(f"Model {self.config.model} may not be available in region {self.config.aws_region}")
|
||||
logger.info(f"Available models: {', '.join(self.available_models[:5])}...")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Could not verify model availability: {e}")
|
||||
self.available_models = []
|
||||
|
||||
def _initialize_provider_settings(self):
|
||||
"""Initialize provider-specific settings and capabilities."""
|
||||
# Determine capabilities based on provider and model
|
||||
self.supports_tools = self.provider in ["anthropic", "cohere", "amazon"]
|
||||
self.supports_vision = self.provider in ["anthropic", "amazon", "meta", "mistral"]
|
||||
self.supports_streaming = self.provider in ["anthropic", "cohere", "mistral", "amazon", "meta"]
|
||||
|
||||
# Set message formatting method
|
||||
if self.provider == "anthropic":
|
||||
self._format_messages = self._format_messages_anthropic
|
||||
elif self.provider == "cohere":
|
||||
self._format_messages = self._format_messages_cohere
|
||||
elif self.provider == "amazon":
|
||||
self._format_messages = self._format_messages_amazon
|
||||
elif self.provider == "meta":
|
||||
self._format_messages = self._format_messages_meta
|
||||
elif self.provider == "mistral":
|
||||
self._format_messages = self._format_messages_mistral
|
||||
else:
|
||||
self._format_messages = self._format_messages_generic
|
||||
|
||||
def _format_messages_anthropic(self, messages: List[Dict[str, str]]) -> List[Dict[str, Any]]:
|
||||
"""Format messages for Anthropic models."""
|
||||
formatted_messages = []
|
||||
|
||||
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
|
||||
elif role == "user":
|
||||
formatted_messages.append({"role": "user", "content": [{"type": "text", "text": content}]})
|
||||
elif role == "assistant":
|
||||
formatted_messages.append({"role": "assistant", "content": [{"type": "text", "text": content}]})
|
||||
|
||||
return formatted_messages
|
||||
|
||||
def _format_messages_cohere(self, messages: List[Dict[str, str]]) -> str:
|
||||
"""Format messages for Cohere models."""
|
||||
formatted_messages = []
|
||||
|
||||
for message in messages:
|
||||
role = message["role"].capitalize()
|
||||
content = message["content"]
|
||||
formatted_messages.append(f"{role}: {content}")
|
||||
|
||||
return "\n".join(formatted_messages)
|
||||
|
||||
def _format_messages_amazon(self, messages: List[Dict[str, str]]) -> List[Dict[str, Any]]:
|
||||
"""Format messages for Amazon models (including Nova)."""
|
||||
formatted_messages = []
|
||||
|
||||
for message in messages:
|
||||
role = message["role"]
|
||||
content = message["content"]
|
||||
|
||||
if role == "system":
|
||||
# Amazon models support system messages
|
||||
formatted_messages.append({"role": "system", "content": content})
|
||||
elif role == "user":
|
||||
formatted_messages.append({"role": "user", "content": content})
|
||||
elif role == "assistant":
|
||||
formatted_messages.append({"role": "assistant", "content": content})
|
||||
|
||||
return formatted_messages
|
||||
|
||||
def _format_messages_meta(self, messages: List[Dict[str, str]]) -> str:
|
||||
"""Format messages for Meta models."""
|
||||
formatted_messages = []
|
||||
|
||||
for message in messages:
|
||||
role = message["role"].capitalize()
|
||||
content = message["content"]
|
||||
formatted_messages.append(f"{role}: {content}")
|
||||
|
||||
return "\n".join(formatted_messages)
|
||||
|
||||
def _format_messages_mistral(self, messages: List[Dict[str, str]]) -> List[Dict[str, Any]]:
|
||||
"""Format messages for Mistral models."""
|
||||
formatted_messages = []
|
||||
|
||||
for message in messages:
|
||||
role = message["role"]
|
||||
content = message["content"]
|
||||
|
||||
if role == "system":
|
||||
# Mistral supports system messages
|
||||
formatted_messages.append({"role": "system", "content": content})
|
||||
elif role == "user":
|
||||
formatted_messages.append({"role": "user", "content": content})
|
||||
elif role == "assistant":
|
||||
formatted_messages.append({"role": "assistant", "content": content})
|
||||
|
||||
return formatted_messages
|
||||
|
||||
def _format_messages_generic(self, messages: List[Dict[str, str]]) -> str:
|
||||
"""Generic message formatting for other providers."""
|
||||
formatted_messages = []
|
||||
|
||||
for message in messages:
|
||||
role = message["role"].capitalize()
|
||||
content = message["content"]
|
||||
@@ -74,21 +223,145 @@ class AWSBedrockLLM(LLMBase):
|
||||
|
||||
return "\n\nHuman: " + "".join(formatted_messages) + "\n\nAssistant:"
|
||||
|
||||
def _parse_response(self, response, tools) -> str:
|
||||
def _prepare_input(self, prompt: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Process the response based on whether tools are used or not.
|
||||
Prepare input for the current provider's model.
|
||||
|
||||
Args:
|
||||
response: The raw response from API.
|
||||
tools: The list of tools provided in the request.
|
||||
prompt: Text prompt to process
|
||||
|
||||
Returns:
|
||||
str or dict: The processed response.
|
||||
Prepared input dictionary
|
||||
"""
|
||||
# Base configuration
|
||||
input_body = {"prompt": prompt}
|
||||
|
||||
# Provider-specific parameter mappings
|
||||
provider_mappings = {
|
||||
"meta": {"max_tokens": "max_gen_len"},
|
||||
"ai21": {"max_tokens": "maxTokens", "top_p": "topP"},
|
||||
"mistral": {"max_tokens": "max_tokens"},
|
||||
"cohere": {"max_tokens": "max_tokens", "top_p": "p"},
|
||||
"amazon": {"max_tokens": "maxTokenCount", "top_p": "topP"},
|
||||
"anthropic": {"max_tokens": "max_tokens", "top_p": "top_p"},
|
||||
}
|
||||
|
||||
# Apply provider mappings
|
||||
if self.provider in provider_mappings:
|
||||
for old_key, new_key in provider_mappings[self.provider].items():
|
||||
if old_key in self.model_config:
|
||||
input_body[new_key] = self.model_config[old_key]
|
||||
|
||||
# Special handling for specific providers
|
||||
if self.provider == "cohere" and "cohere.command" in self.config.model:
|
||||
input_body["message"] = input_body.pop("prompt")
|
||||
elif self.provider == "amazon":
|
||||
# Amazon Nova and other Amazon models
|
||||
if "nova" in self.config.model.lower():
|
||||
# Nova models use the converse API format
|
||||
input_body = {
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"max_tokens": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
else:
|
||||
# Legacy Amazon models
|
||||
input_body = {
|
||||
"inputText": prompt,
|
||||
"textGenerationConfig": {
|
||||
"maxTokenCount": self.model_config.get("max_tokens", 5000),
|
||||
"topP": self.model_config.get("top_p", 0.9),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
},
|
||||
}
|
||||
# Remove None values
|
||||
input_body["textGenerationConfig"] = {
|
||||
k: v for k, v in input_body["textGenerationConfig"].items() if v is not None
|
||||
}
|
||||
elif self.provider == "anthropic":
|
||||
input_body = {
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": prompt}]}],
|
||||
"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",
|
||||
}
|
||||
elif self.provider == "meta":
|
||||
input_body = {
|
||||
"prompt": prompt,
|
||||
"max_gen_len": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
elif self.provider == "mistral":
|
||||
input_body = {
|
||||
"prompt": prompt,
|
||||
"max_tokens": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
else:
|
||||
# Generic case - add all model config parameters
|
||||
input_body.update(self.model_config)
|
||||
|
||||
return input_body
|
||||
|
||||
def _convert_tool_format(self, original_tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Convert tools to Bedrock-compatible format.
|
||||
|
||||
Args:
|
||||
original_tools: List of tool definitions
|
||||
|
||||
Returns:
|
||||
Converted tools in Bedrock format
|
||||
"""
|
||||
new_tools = []
|
||||
|
||||
for tool in original_tools:
|
||||
if tool["type"] == "function":
|
||||
function = tool["function"]
|
||||
new_tool = {
|
||||
"toolSpec": {
|
||||
"name": function["name"],
|
||||
"description": function.get("description", ""),
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": function["parameters"].get("required", []),
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
# Add properties
|
||||
for prop, details in function["parameters"].get("properties", {}).items():
|
||||
new_tool["toolSpec"]["inputSchema"]["json"]["properties"][prop] = details
|
||||
|
||||
new_tools.append(new_tool)
|
||||
|
||||
return new_tools
|
||||
|
||||
def _parse_response(
|
||||
self, response: Dict[str, Any], tools: Optional[List[Dict]] = None
|
||||
) -> Union[str, Dict[str, Any]]:
|
||||
"""
|
||||
Parse response from Bedrock API.
|
||||
|
||||
Args:
|
||||
response: Raw API response
|
||||
tools: List of tools if used
|
||||
|
||||
Returns:
|
||||
Parsed response
|
||||
"""
|
||||
if tools:
|
||||
# Handle tool-enabled responses
|
||||
processed_response = {"tool_calls": []}
|
||||
|
||||
if response["output"]["message"]["content"]:
|
||||
if response.get("output", {}).get("message", {}).get("content"):
|
||||
for item in response["output"]["message"]["content"]:
|
||||
if "toolUse" in item:
|
||||
processed_response["tool_calls"].append(
|
||||
@@ -100,171 +373,228 @@ class AWSBedrockLLM(LLMBase):
|
||||
|
||||
return processed_response
|
||||
|
||||
response_body = response.get("body").read().decode()
|
||||
response_json = json.loads(response_body)
|
||||
return response_json.get("content", [{"text": ""}])[0].get("text", "")
|
||||
# Handle regular text responses
|
||||
try:
|
||||
response_body = response.get("body").read().decode()
|
||||
response_json = json.loads(response_body)
|
||||
|
||||
def _prepare_input(
|
||||
self,
|
||||
provider: str,
|
||||
model: str,
|
||||
prompt: str,
|
||||
model_kwargs: Optional[Dict[str, Any]] = {},
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Prepares the input dictionary for the specified provider's model by mapping and renaming
|
||||
keys in the input based on the provider's requirements.
|
||||
# Provider-specific response parsing
|
||||
if self.provider == "anthropic":
|
||||
return response_json.get("content", [{"text": ""}])[0].get("text", "")
|
||||
elif self.provider == "amazon":
|
||||
# Handle both Nova and legacy Amazon models
|
||||
if "nova" in self.config.model.lower():
|
||||
# Nova models return content in a different format
|
||||
if "content" in response_json:
|
||||
return response_json["content"][0]["text"]
|
||||
elif "completion" in response_json:
|
||||
return response_json["completion"]
|
||||
else:
|
||||
# Legacy Amazon models
|
||||
return response_json.get("completion", "")
|
||||
elif self.provider == "meta":
|
||||
return response_json.get("generation", "")
|
||||
elif self.provider == "mistral":
|
||||
return response_json.get("outputs", [{"text": ""}])[0].get("text", "")
|
||||
elif self.provider == "cohere":
|
||||
return response_json.get("generations", [{"text": ""}])[0].get("text", "")
|
||||
elif self.provider == "ai21":
|
||||
return response_json.get("completions", [{"data", {"text": ""}}])[0].get("data", {}).get("text", "")
|
||||
else:
|
||||
# Generic parsing - try common response fields
|
||||
for field in ["content", "text", "completion", "generation"]:
|
||||
if field in response_json:
|
||||
if isinstance(response_json[field], list) and response_json[field]:
|
||||
return response_json[field][0].get("text", "")
|
||||
elif isinstance(response_json[field], str):
|
||||
return response_json[field]
|
||||
|
||||
Args:
|
||||
provider (str): The name of the service provider (e.g., "meta", "ai21", "mistral", "cohere", "amazon").
|
||||
model (str): The name or identifier of the model being used.
|
||||
prompt (str): The text prompt to be processed by the model.
|
||||
model_kwargs (Dict[str, Any]): Additional keyword arguments specific to the model's requirements.
|
||||
# Fallback
|
||||
return str(response_json)
|
||||
|
||||
Returns:
|
||||
Dict[str, Any]: The prepared input dictionary with the correct keys and values for the specified provider.
|
||||
"""
|
||||
|
||||
input_body = {"prompt": prompt, **model_kwargs}
|
||||
|
||||
provider_mappings = {
|
||||
"meta": {"max_tokens_to_sample": "max_gen_len"},
|
||||
"ai21": {"max_tokens_to_sample": "maxTokens", "top_p": "topP"},
|
||||
"mistral": {"max_tokens_to_sample": "max_tokens"},
|
||||
"cohere": {"max_tokens_to_sample": "max_tokens", "top_p": "p"},
|
||||
}
|
||||
|
||||
if provider in provider_mappings:
|
||||
for old_key, new_key in provider_mappings[provider].items():
|
||||
if old_key in input_body:
|
||||
input_body[new_key] = input_body.pop(old_key)
|
||||
|
||||
if provider == "cohere" and "cohere.command-r" in model:
|
||||
input_body["message"] = input_body.pop("prompt")
|
||||
|
||||
if provider == "amazon":
|
||||
input_body = {
|
||||
"inputText": prompt,
|
||||
"textGenerationConfig": {
|
||||
"maxTokenCount": self.model_kwargs["max_tokens_to_sample"]
|
||||
or self.model_kwargs["max_tokens"]
|
||||
or 5000,
|
||||
"topP": self.model_kwargs["top_p"] or 0.9,
|
||||
"temperature": self.model_kwargs["temperature"] or 0.1,
|
||||
},
|
||||
}
|
||||
input_body["textGenerationConfig"] = {
|
||||
k: v for k, v in input_body["textGenerationConfig"].items() if v is not None
|
||||
}
|
||||
|
||||
return input_body
|
||||
|
||||
def _convert_tool_format(self, original_tools):
|
||||
"""
|
||||
Converts a list of tools from their original format to a new standardized format.
|
||||
|
||||
Args:
|
||||
original_tools (list): A list of dictionaries representing the original tools, each containing a 'type' key and corresponding details.
|
||||
|
||||
Returns:
|
||||
list: A list of dictionaries representing the tools in the new standardized format.
|
||||
"""
|
||||
new_tools = []
|
||||
|
||||
for tool in original_tools:
|
||||
if tool["type"] == "function":
|
||||
function = tool["function"]
|
||||
new_tool = {
|
||||
"toolSpec": {
|
||||
"name": function["name"],
|
||||
"description": function["description"],
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": function["parameters"].get("required", []),
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
for prop, details in function["parameters"].get("properties", {}).items():
|
||||
new_tool["toolSpec"]["inputSchema"]["json"]["properties"][prop] = details
|
||||
|
||||
new_tools.append(new_tool)
|
||||
|
||||
return new_tools
|
||||
except Exception as e:
|
||||
logger.warning(f"Could not parse response: {e}")
|
||||
return "Error parsing response"
|
||||
|
||||
def generate_response(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
response_format=None,
|
||||
response_format: Optional[str] = None,
|
||||
tools: Optional[List[Dict]] = None,
|
||||
tool_choice: str = "auto",
|
||||
):
|
||||
stream: bool = False,
|
||||
**kwargs,
|
||||
) -> Union[str, Dict[str, Any]]:
|
||||
"""
|
||||
Generate a response based on the given messages using AWS Bedrock.
|
||||
Generate response using AWS Bedrock.
|
||||
|
||||
Args:
|
||||
messages (list): List of message dicts containing 'role' and 'content'.
|
||||
tools (list, optional): List of tools that the model can call. Defaults to None.
|
||||
tool_choice (str, optional): Tool choice method. Defaults to "auto".
|
||||
messages: List of message dictionaries
|
||||
response_format: Response format specification
|
||||
tools: List of tools for function calling
|
||||
tool_choice: Tool choice method
|
||||
stream: Whether to stream the response
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
str: The generated response.
|
||||
Generated response
|
||||
"""
|
||||
|
||||
if tools:
|
||||
# Use converse method when tools are provided
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"text": message["content"]} for message in messages],
|
||||
}
|
||||
]
|
||||
inference_config = {
|
||||
"temperature": self.model_kwargs["temperature"],
|
||||
"maxTokens": self.model_kwargs["max_tokens_to_sample"],
|
||||
"topP": self.model_kwargs["top_p"],
|
||||
}
|
||||
tools_config = {"tools": self._convert_tool_format(tools)}
|
||||
|
||||
response = self.client.converse(
|
||||
modelId=self.config.model,
|
||||
messages=messages,
|
||||
inferenceConfig=inference_config,
|
||||
toolConfig=tools_config,
|
||||
)
|
||||
else:
|
||||
# Use invoke_model method when no tools are provided
|
||||
prompt = self._format_messages(messages)
|
||||
provider = extract_provider(self.config.model)
|
||||
input_body = self._prepare_input(provider, self.config.model, prompt, model_kwargs=self.model_kwargs)
|
||||
body = json.dumps(input_body)
|
||||
|
||||
if provider == "anthropic" or provider == "deepseek":
|
||||
input_body = {
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": prompt}]}],
|
||||
"max_tokens": self.model_kwargs["max_tokens_to_sample"] or self.model_kwargs["max_tokens"] or 5000,
|
||||
"temperature": self.model_kwargs["temperature"] or 0.1,
|
||||
"top_p": self.model_kwargs["top_p"] or 0.9,
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
}
|
||||
|
||||
body = json.dumps(input_body)
|
||||
|
||||
response = self.client.invoke_model(
|
||||
body=body,
|
||||
modelId=self.config.model,
|
||||
accept="application/json",
|
||||
contentType="application/json",
|
||||
)
|
||||
try:
|
||||
if tools and self.supports_tools:
|
||||
# Use converse method for tool-enabled models
|
||||
return self._generate_with_tools(messages, tools, stream)
|
||||
else:
|
||||
response = self.client.invoke_model(
|
||||
body=body,
|
||||
modelId=self.config.model,
|
||||
accept="application/json",
|
||||
contentType="application/json",
|
||||
)
|
||||
# Use standard invoke_model method
|
||||
return self._generate_standard(messages, stream)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to generate response: {e}")
|
||||
raise RuntimeError(f"Failed to generate response: {e}")
|
||||
|
||||
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."""
|
||||
# Format messages for tool-enabled models
|
||||
if self.provider == "anthropic":
|
||||
formatted_messages = 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"]}]
|
||||
|
||||
# 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 tools configuration
|
||||
tools_config = {"tools": self._convert_tool_format(tools)}
|
||||
|
||||
# Make API call
|
||||
response = self.client.converse(
|
||||
modelId=self.config.model,
|
||||
messages=formatted_messages,
|
||||
inferenceConfig=inference_config,
|
||||
toolConfig=tools_config,
|
||||
)
|
||||
|
||||
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
|
||||
if self.provider == "anthropic":
|
||||
formatted_messages = self._format_messages_anthropic(messages)
|
||||
input_body = {
|
||||
"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",
|
||||
}
|
||||
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)
|
||||
input_body = {
|
||||
"messages": formatted_messages,
|
||||
"max_tokens": self.model_config.get("max_tokens", 5000),
|
||||
"temperature": self.model_config.get("temperature", 0.1),
|
||||
"top_p": self.model_config.get("top_p", 0.9),
|
||||
}
|
||||
|
||||
# Use converse API for Nova models
|
||||
response = self.client.converse(
|
||||
modelId=self.config.model,
|
||||
messages=input_body["messages"],
|
||||
inferenceConfig={
|
||||
"maxTokens": input_body["max_tokens"],
|
||||
"temperature": input_body["temperature"],
|
||||
"topP": input_body["top_p"],
|
||||
}
|
||||
)
|
||||
|
||||
return self._parse_response(response)
|
||||
else:
|
||||
prompt = self._format_messages(messages)
|
||||
input_body = self._prepare_input(prompt)
|
||||
|
||||
# Convert to JSON
|
||||
body = json.dumps(input_body)
|
||||
|
||||
# Make API call
|
||||
response = self.client.invoke_model(
|
||||
body=body,
|
||||
modelId=self.config.model,
|
||||
accept="application/json",
|
||||
contentType="application/json",
|
||||
)
|
||||
|
||||
return self._parse_response(response)
|
||||
|
||||
def list_available_models(self) -> List[Dict[str, Any]]:
|
||||
"""List all available models in the current region."""
|
||||
try:
|
||||
bedrock_client = boto3.client("bedrock", **self.config.get_aws_config())
|
||||
response = bedrock_client.list_foundation_models()
|
||||
|
||||
models = []
|
||||
for model in response["modelSummaries"]:
|
||||
provider = extract_provider(model["modelId"])
|
||||
models.append(
|
||||
{
|
||||
"model_id": model["modelId"],
|
||||
"provider": provider,
|
||||
"model_name": model["modelId"].split(".", 1)[1]
|
||||
if "." in model["modelId"]
|
||||
else model["modelId"],
|
||||
"modelArn": model.get("modelArn", ""),
|
||||
"providerName": model.get("providerName", ""),
|
||||
"inputModalities": model.get("inputModalities", []),
|
||||
"outputModalities": model.get("outputModalities", []),
|
||||
"responseStreamingSupported": model.get("responseStreamingSupported", False),
|
||||
}
|
||||
)
|
||||
|
||||
return models
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Could not list models: {e}")
|
||||
return []
|
||||
|
||||
def get_model_capabilities(self) -> Dict[str, Any]:
|
||||
"""Get capabilities of the current model."""
|
||||
return {
|
||||
"model_id": self.config.model,
|
||||
"provider": self.provider,
|
||||
"model_name": self.config.model_name,
|
||||
"supports_tools": self.supports_tools,
|
||||
"supports_vision": self.supports_vision,
|
||||
"supports_streaming": self.supports_streaming,
|
||||
"max_tokens": self.model_config.get("max_tokens", 2000),
|
||||
}
|
||||
|
||||
def validate_model_access(self) -> bool:
|
||||
"""Validate if the model is accessible."""
|
||||
try:
|
||||
# Try to invoke the model with a minimal request
|
||||
if self.provider == "amazon" and "nova" in self.config.model.lower():
|
||||
# Test Nova model with converse API
|
||||
test_messages = [{"role": "user", "content": "test"}]
|
||||
self.client.converse(
|
||||
modelId=self.config.model,
|
||||
messages=test_messages,
|
||||
inferenceConfig={"maxTokens": 10}
|
||||
)
|
||||
else:
|
||||
# Test other models with invoke_model
|
||||
test_body = json.dumps({"prompt": "test"})
|
||||
self.client.invoke_model(
|
||||
body=test_body,
|
||||
modelId=self.config.model,
|
||||
accept="application/json",
|
||||
contentType="application/json",
|
||||
)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
Reference in New Issue
Block a user