From e7013764f7958801bba8bb55e46d247d55cb90de Mon Sep 17 00:00:00 2001 From: Parshva Daftari <89991302+parshvadaftari@users.noreply.github.com> Date: Tue, 19 Aug 2025 02:20:34 +0530 Subject: [PATCH] Update aws bedrock (#3334) --- mem0/configs/llms/aws_bedrock.py | 191 ++++++++ mem0/llms/aws_bedrock.py | 722 ++++++++++++++++++++++--------- 2 files changed, 717 insertions(+), 196 deletions(-) create mode 100644 mem0/configs/llms/aws_bedrock.py diff --git a/mem0/configs/llms/aws_bedrock.py b/mem0/configs/llms/aws_bedrock.py new file mode 100644 index 000000000..bbdebef0a --- /dev/null +++ b/mem0/configs/llms/aws_bedrock.py @@ -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 diff --git a/mem0/llms/aws_bedrock.py b/mem0/llms/aws_bedrock.py index 8266beef8..56fa6ef42 100644 --- a/mem0/llms/aws_bedrock.py +++ b/mem0/llms/aws_bedrock.py @@ -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