ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
This commit is contained in:
@@ -0,0 +1,36 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
# Type stubs for lazy imports
|
||||
if TYPE_CHECKING:
|
||||
from browser_use.llm.aws.chat_anthropic import ChatAnthropicBedrock
|
||||
from browser_use.llm.aws.chat_bedrock import ChatAWSBedrock
|
||||
|
||||
# Lazy imports mapping for AWS chat models
|
||||
_LAZY_IMPORTS = {
|
||||
'ChatAnthropicBedrock': ('browser_use.llm.aws.chat_anthropic', 'ChatAnthropicBedrock'),
|
||||
'ChatAWSBedrock': ('browser_use.llm.aws.chat_bedrock', 'ChatAWSBedrock'),
|
||||
}
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
"""Lazy import mechanism for AWS chat models."""
|
||||
if name in _LAZY_IMPORTS:
|
||||
module_path, attr_name = _LAZY_IMPORTS[name]
|
||||
try:
|
||||
from importlib import import_module
|
||||
|
||||
module = import_module(module_path)
|
||||
attr = getattr(module, attr_name)
|
||||
# Cache the imported attribute in the module's globals
|
||||
globals()[name] = attr
|
||||
return attr
|
||||
except ImportError as e:
|
||||
raise ImportError(f'Failed to import {name} from {module_path}: {e}') from e
|
||||
|
||||
raise AttributeError(f"module '{__name__}' has no attribute '{name}'")
|
||||
|
||||
|
||||
__all__ = [
|
||||
'ChatAWSBedrock',
|
||||
'ChatAnthropicBedrock',
|
||||
]
|
||||
@@ -0,0 +1,242 @@
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, TypeVar, overload
|
||||
|
||||
from anthropic import (
|
||||
NOT_GIVEN,
|
||||
APIConnectionError,
|
||||
APIStatusError,
|
||||
AsyncAnthropicBedrock,
|
||||
RateLimitError,
|
||||
)
|
||||
from anthropic.types import CacheControlEphemeralParam, Message, ToolParam
|
||||
from anthropic.types.text_block import TextBlock
|
||||
from anthropic.types.tool_choice_tool_param import ToolChoiceToolParam
|
||||
from pydantic import BaseModel
|
||||
|
||||
from browser_use.llm.anthropic.serializer import AnthropicMessageSerializer
|
||||
from browser_use.llm.aws.chat_bedrock import ChatAWSBedrock
|
||||
from browser_use.llm.exceptions import ModelProviderError, ModelRateLimitError
|
||||
from browser_use.llm.messages import BaseMessage
|
||||
from browser_use.llm.views import ChatInvokeCompletion, ChatInvokeUsage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from boto3.session import Session # pyright: ignore
|
||||
|
||||
|
||||
T = TypeVar('T', bound=BaseModel)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatAnthropicBedrock(ChatAWSBedrock):
|
||||
"""
|
||||
AWS Bedrock Anthropic Claude chat model.
|
||||
|
||||
This is a convenience class that provides Claude-specific defaults
|
||||
for the AWS Bedrock service. It inherits all functionality from
|
||||
ChatAWSBedrock but sets Anthropic Claude as the default model.
|
||||
"""
|
||||
|
||||
# Anthropic Claude specific defaults
|
||||
model: str = 'anthropic.claude-3-5-sonnet-20240620-v1:0'
|
||||
max_tokens: int = 8192
|
||||
temperature: float | None = None
|
||||
top_p: float | None = None
|
||||
top_k: int | None = None
|
||||
stop_sequences: list[str] | None = None
|
||||
seed: int | None = None
|
||||
|
||||
# AWS credentials and configuration
|
||||
aws_access_key: str | None = None
|
||||
aws_secret_key: str | None = None
|
||||
aws_session_token: str | None = None
|
||||
aws_region: str | None = None
|
||||
session: 'Session | None' = None
|
||||
|
||||
# Client initialization parameters
|
||||
max_retries: int = 10
|
||||
default_headers: Mapping[str, str] | None = None
|
||||
default_query: Mapping[str, object] | None = None
|
||||
|
||||
@property
|
||||
def provider(self) -> str:
|
||||
return 'anthropic_bedrock'
|
||||
|
||||
def _get_client_params(self) -> dict[str, Any]:
|
||||
"""Prepare client parameters dictionary for Bedrock."""
|
||||
client_params: dict[str, Any] = {}
|
||||
|
||||
if self.session:
|
||||
credentials = self.session.get_credentials()
|
||||
client_params.update(
|
||||
{
|
||||
'aws_access_key': credentials.access_key,
|
||||
'aws_secret_key': credentials.secret_key,
|
||||
'aws_session_token': credentials.token,
|
||||
'aws_region': self.session.region_name,
|
||||
}
|
||||
)
|
||||
else:
|
||||
# Use individual credentials
|
||||
if self.aws_access_key:
|
||||
client_params['aws_access_key'] = self.aws_access_key
|
||||
if self.aws_secret_key:
|
||||
client_params['aws_secret_key'] = self.aws_secret_key
|
||||
if self.aws_region:
|
||||
client_params['aws_region'] = self.aws_region
|
||||
if self.aws_session_token:
|
||||
client_params['aws_session_token'] = self.aws_session_token
|
||||
|
||||
# Add optional parameters
|
||||
if self.max_retries:
|
||||
client_params['max_retries'] = self.max_retries
|
||||
if self.default_headers:
|
||||
client_params['default_headers'] = self.default_headers
|
||||
if self.default_query:
|
||||
client_params['default_query'] = self.default_query
|
||||
|
||||
return client_params
|
||||
|
||||
def _get_client_params_for_invoke(self) -> dict[str, Any]:
|
||||
"""Prepare client parameters dictionary for invoke."""
|
||||
client_params = {}
|
||||
|
||||
if self.temperature is not None:
|
||||
client_params['temperature'] = self.temperature
|
||||
if self.max_tokens is not None:
|
||||
client_params['max_tokens'] = self.max_tokens
|
||||
if self.top_p is not None:
|
||||
client_params['top_p'] = self.top_p
|
||||
if self.top_k is not None:
|
||||
client_params['top_k'] = self.top_k
|
||||
if self.seed is not None:
|
||||
client_params['seed'] = self.seed
|
||||
if self.stop_sequences is not None:
|
||||
client_params['stop_sequences'] = self.stop_sequences
|
||||
|
||||
return client_params
|
||||
|
||||
def get_client(self) -> AsyncAnthropicBedrock:
|
||||
"""
|
||||
Returns an AsyncAnthropicBedrock client.
|
||||
|
||||
Returns:
|
||||
AsyncAnthropicBedrock: An instance of the AsyncAnthropicBedrock client.
|
||||
"""
|
||||
client_params = self._get_client_params()
|
||||
return AsyncAnthropicBedrock(**client_params)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return str(self.model)
|
||||
|
||||
def _get_usage(self, response: Message) -> ChatInvokeUsage | None:
|
||||
"""Extract usage information from the response."""
|
||||
usage = ChatInvokeUsage(
|
||||
prompt_tokens=response.usage.input_tokens
|
||||
+ (
|
||||
response.usage.cache_read_input_tokens or 0
|
||||
), # Total tokens in Anthropic are a bit fucked, you have to add cached tokens to the prompt tokens
|
||||
completion_tokens=response.usage.output_tokens,
|
||||
total_tokens=response.usage.input_tokens + response.usage.output_tokens,
|
||||
prompt_cached_tokens=response.usage.cache_read_input_tokens,
|
||||
prompt_cache_creation_tokens=response.usage.cache_creation_input_tokens,
|
||||
prompt_image_tokens=None,
|
||||
)
|
||||
return usage
|
||||
|
||||
@overload
|
||||
async def ainvoke(self, messages: list[BaseMessage], output_format: None = None) -> ChatInvokeCompletion[str]: ...
|
||||
|
||||
@overload
|
||||
async def ainvoke(self, messages: list[BaseMessage], output_format: type[T]) -> ChatInvokeCompletion[T]: ...
|
||||
|
||||
async def ainvoke(
|
||||
self, messages: list[BaseMessage], output_format: type[T] | None = None
|
||||
) -> ChatInvokeCompletion[T] | ChatInvokeCompletion[str]:
|
||||
anthropic_messages, system_prompt = AnthropicMessageSerializer.serialize_messages(messages)
|
||||
|
||||
try:
|
||||
if output_format is None:
|
||||
# Normal completion without structured output
|
||||
response = await self.get_client().messages.create(
|
||||
model=self.model,
|
||||
messages=anthropic_messages,
|
||||
system=system_prompt or NOT_GIVEN,
|
||||
**self._get_client_params_for_invoke(),
|
||||
)
|
||||
|
||||
usage = self._get_usage(response)
|
||||
|
||||
# Extract text from the first content block
|
||||
first_content = response.content[0]
|
||||
if isinstance(first_content, TextBlock):
|
||||
response_text = first_content.text
|
||||
else:
|
||||
# If it's not a text block, convert to string
|
||||
response_text = str(first_content)
|
||||
|
||||
return ChatInvokeCompletion(
|
||||
completion=response_text,
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
else:
|
||||
# Use tool calling for structured output
|
||||
# Create a tool that represents the output format
|
||||
tool_name = output_format.__name__
|
||||
schema = output_format.model_json_schema()
|
||||
|
||||
# Remove title from schema if present (Anthropic doesn't like it in parameters)
|
||||
if 'title' in schema:
|
||||
del schema['title']
|
||||
|
||||
tool = ToolParam(
|
||||
name=tool_name,
|
||||
description=f'Extract information in the format of {tool_name}',
|
||||
input_schema=schema,
|
||||
cache_control=CacheControlEphemeralParam(type='ephemeral'),
|
||||
)
|
||||
|
||||
# Force the model to use this tool
|
||||
tool_choice = ToolChoiceToolParam(type='tool', name=tool_name)
|
||||
|
||||
response = await self.get_client().messages.create(
|
||||
model=self.model,
|
||||
messages=anthropic_messages,
|
||||
tools=[tool],
|
||||
system=system_prompt or NOT_GIVEN,
|
||||
tool_choice=tool_choice,
|
||||
**self._get_client_params_for_invoke(),
|
||||
)
|
||||
|
||||
usage = self._get_usage(response)
|
||||
|
||||
# Extract the tool use block
|
||||
for content_block in response.content:
|
||||
if hasattr(content_block, 'type') and content_block.type == 'tool_use':
|
||||
# Parse the tool input as the structured output
|
||||
try:
|
||||
return ChatInvokeCompletion(completion=output_format.model_validate(content_block.input), usage=usage)
|
||||
except Exception as e:
|
||||
# If validation fails, try to parse it as JSON first
|
||||
if isinstance(content_block.input, str):
|
||||
data = json.loads(content_block.input)
|
||||
return ChatInvokeCompletion(
|
||||
completion=output_format.model_validate(data),
|
||||
usage=usage,
|
||||
)
|
||||
raise e
|
||||
|
||||
# If no tool use block found, raise an error
|
||||
raise ValueError('Expected tool use in response but none found')
|
||||
|
||||
except APIConnectionError as e:
|
||||
raise ModelProviderError(message=e.message, model=self.name) from e
|
||||
except RateLimitError as e:
|
||||
raise ModelRateLimitError(message=e.message, model=self.name) from e
|
||||
except APIStatusError as e:
|
||||
raise ModelProviderError(message=e.message, status_code=e.status_code, model=self.name) from e
|
||||
except Exception as e:
|
||||
raise ModelProviderError(message=str(e), model=self.name) from e
|
||||
@@ -0,0 +1,289 @@
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from os import getenv
|
||||
from typing import TYPE_CHECKING, Any, TypeVar, overload
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from browser_use.llm.aws.serializer import AWSBedrockMessageSerializer
|
||||
from browser_use.llm.base import BaseChatModel
|
||||
from browser_use.llm.exceptions import ModelProviderError, ModelRateLimitError
|
||||
from browser_use.llm.messages import BaseMessage
|
||||
from browser_use.llm.views import ChatInvokeCompletion, ChatInvokeUsage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from boto3 import client as AwsClient # type: ignore
|
||||
from boto3.session import Session # type: ignore
|
||||
|
||||
T = TypeVar('T', bound=BaseModel)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatAWSBedrock(BaseChatModel):
|
||||
"""
|
||||
AWS Bedrock chat model supporting multiple providers (Anthropic, Meta, etc.).
|
||||
|
||||
This class provides access to various models via AWS Bedrock,
|
||||
supporting both text generation and structured output via tool calling.
|
||||
|
||||
To use this model, you need to either:
|
||||
1. Set the following environment variables:
|
||||
- AWS_ACCESS_KEY_ID
|
||||
- AWS_SECRET_ACCESS_KEY
|
||||
- AWS_SESSION_TOKEN (only required when using temporary credentials)
|
||||
- AWS_REGION
|
||||
2. Or provide a boto3 Session object
|
||||
3. Or use AWS SSO authentication
|
||||
"""
|
||||
|
||||
# Model configuration
|
||||
model: str = 'anthropic.claude-3-5-sonnet-20240620-v1:0'
|
||||
max_tokens: int | None = 4096
|
||||
temperature: float | None = None
|
||||
top_p: float | None = None
|
||||
seed: int | None = None
|
||||
stop_sequences: list[str] | None = None
|
||||
|
||||
# AWS credentials and configuration
|
||||
aws_access_key_id: str | None = None
|
||||
aws_secret_access_key: str | None = None
|
||||
aws_session_token: str | None = None
|
||||
aws_region: str | None = None
|
||||
aws_sso_auth: bool = False
|
||||
session: 'Session | None' = None
|
||||
|
||||
# Request parameters
|
||||
request_params: dict[str, Any] | None = None
|
||||
|
||||
# Static
|
||||
@property
|
||||
def provider(self) -> str:
|
||||
return 'aws_bedrock'
|
||||
|
||||
def _get_client(self) -> 'AwsClient': # type: ignore
|
||||
"""Get the AWS Bedrock client."""
|
||||
try:
|
||||
from boto3 import client as AwsClient # type: ignore
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
'`boto3` not installed. Please install using `pip install browser-use[aws] or pip install browser-use[all]`'
|
||||
)
|
||||
|
||||
if self.session:
|
||||
return self.session.client('bedrock-runtime')
|
||||
|
||||
# Get credentials from environment or instance parameters
|
||||
access_key = self.aws_access_key_id or getenv('AWS_ACCESS_KEY_ID')
|
||||
secret_key = self.aws_secret_access_key or getenv('AWS_SECRET_ACCESS_KEY')
|
||||
session_token = self.aws_session_token or getenv('AWS_SESSION_TOKEN')
|
||||
region = self.aws_region or getenv('AWS_REGION') or getenv('AWS_DEFAULT_REGION')
|
||||
|
||||
if self.aws_sso_auth:
|
||||
return AwsClient(service_name='bedrock-runtime', region_name=region)
|
||||
else:
|
||||
if not access_key or not secret_key:
|
||||
raise ModelProviderError(
|
||||
message='AWS credentials not found. Please set AWS_ACCESS_KEY_ID and AWS_SECRET_ACCESS_KEY environment variables (and AWS_SESSION_TOKEN if using temporary credentials) or provide a boto3 session.',
|
||||
model=self.name,
|
||||
)
|
||||
|
||||
return AwsClient(
|
||||
service_name='bedrock-runtime',
|
||||
region_name=region,
|
||||
aws_access_key_id=access_key,
|
||||
aws_secret_access_key=secret_key,
|
||||
aws_session_token=session_token,
|
||||
)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return str(self.model)
|
||||
|
||||
def _get_inference_config(self) -> dict[str, Any]:
|
||||
"""Get the inference configuration for the request."""
|
||||
config = {}
|
||||
if self.max_tokens is not None:
|
||||
config['maxTokens'] = self.max_tokens
|
||||
if self.temperature is not None:
|
||||
config['temperature'] = self.temperature
|
||||
if self.top_p is not None:
|
||||
config['topP'] = self.top_p
|
||||
if self.stop_sequences is not None:
|
||||
config['stopSequences'] = self.stop_sequences
|
||||
if self.seed is not None:
|
||||
config['seed'] = self.seed
|
||||
return config
|
||||
|
||||
def _format_tools_for_request(self, output_format: type[BaseModel]) -> list[dict[str, Any]]:
|
||||
"""Format a Pydantic model as a tool for structured output."""
|
||||
schema = output_format.model_json_schema()
|
||||
|
||||
# Convert Pydantic schema to Bedrock tool format
|
||||
properties = {}
|
||||
required = []
|
||||
|
||||
for prop_name, prop_info in schema.get('properties', {}).items():
|
||||
properties[prop_name] = {
|
||||
'type': prop_info.get('type', 'string'),
|
||||
'description': prop_info.get('description', ''),
|
||||
}
|
||||
|
||||
# Add required fields
|
||||
required = schema.get('required', [])
|
||||
|
||||
return [
|
||||
{
|
||||
'toolSpec': {
|
||||
'name': f'extract_{output_format.__name__.lower()}',
|
||||
'description': f'Extract information in the format of {output_format.__name__}',
|
||||
'inputSchema': {'json': {'type': 'object', 'properties': properties, 'required': required}},
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
def _get_usage(self, response: dict[str, Any]) -> ChatInvokeUsage | None:
|
||||
"""Extract usage information from the response."""
|
||||
if 'usage' not in response:
|
||||
return None
|
||||
|
||||
usage_data = response['usage']
|
||||
return ChatInvokeUsage(
|
||||
prompt_tokens=usage_data.get('inputTokens', 0),
|
||||
completion_tokens=usage_data.get('outputTokens', 0),
|
||||
total_tokens=usage_data.get('totalTokens', 0),
|
||||
prompt_cached_tokens=None, # Bedrock doesn't provide this
|
||||
prompt_cache_creation_tokens=None,
|
||||
prompt_image_tokens=None,
|
||||
)
|
||||
|
||||
@overload
|
||||
async def ainvoke(self, messages: list[BaseMessage], output_format: None = None) -> ChatInvokeCompletion[str]: ...
|
||||
|
||||
@overload
|
||||
async def ainvoke(self, messages: list[BaseMessage], output_format: type[T]) -> ChatInvokeCompletion[T]: ...
|
||||
|
||||
async def ainvoke(
|
||||
self, messages: list[BaseMessage], output_format: type[T] | None = None
|
||||
) -> ChatInvokeCompletion[T] | ChatInvokeCompletion[str]:
|
||||
"""
|
||||
Invoke the AWS Bedrock model with the given messages.
|
||||
|
||||
Args:
|
||||
messages: List of chat messages
|
||||
output_format: Optional Pydantic model class for structured output
|
||||
|
||||
Returns:
|
||||
Either a string response or an instance of output_format
|
||||
"""
|
||||
try:
|
||||
from botocore.exceptions import ClientError # type: ignore
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
'`boto3` not installed. Please install using `pip install browser-use[aws] or pip install browser-use[all]`'
|
||||
)
|
||||
|
||||
bedrock_messages, system_message = AWSBedrockMessageSerializer.serialize_messages(messages)
|
||||
|
||||
try:
|
||||
# Prepare the request body
|
||||
body: dict[str, Any] = {}
|
||||
|
||||
if system_message:
|
||||
body['system'] = system_message
|
||||
|
||||
inference_config = self._get_inference_config()
|
||||
if inference_config:
|
||||
body['inferenceConfig'] = inference_config
|
||||
|
||||
# Handle structured output via tool calling
|
||||
if output_format is not None:
|
||||
tools = self._format_tools_for_request(output_format)
|
||||
body['toolConfig'] = {'tools': tools}
|
||||
|
||||
# Add any additional request parameters
|
||||
if self.request_params:
|
||||
body.update(self.request_params)
|
||||
|
||||
# Filter out None values
|
||||
body = {k: v for k, v in body.items() if v is not None}
|
||||
|
||||
# Make the API call
|
||||
client = self._get_client()
|
||||
response = client.converse(modelId=self.model, messages=bedrock_messages, **body)
|
||||
|
||||
usage = self._get_usage(response)
|
||||
|
||||
# Extract the response content
|
||||
if 'output' in response and 'message' in response['output']:
|
||||
message = response['output']['message']
|
||||
content = message.get('content', [])
|
||||
|
||||
if output_format is None:
|
||||
# Return text response
|
||||
text_content = []
|
||||
for item in content:
|
||||
if 'text' in item:
|
||||
text_content.append(item['text'])
|
||||
|
||||
response_text = '\n'.join(text_content) if text_content else ''
|
||||
return ChatInvokeCompletion(
|
||||
completion=response_text,
|
||||
usage=usage,
|
||||
)
|
||||
else:
|
||||
# Handle structured output from tool calls
|
||||
for item in content:
|
||||
if 'toolUse' in item:
|
||||
tool_use = item['toolUse']
|
||||
tool_input = tool_use.get('input', {})
|
||||
|
||||
try:
|
||||
# Validate and return the structured output
|
||||
return ChatInvokeCompletion(
|
||||
completion=output_format.model_validate(tool_input),
|
||||
usage=usage,
|
||||
)
|
||||
except Exception as e:
|
||||
# If validation fails, try to parse as JSON first
|
||||
if isinstance(tool_input, str):
|
||||
try:
|
||||
data = json.loads(tool_input)
|
||||
return ChatInvokeCompletion(
|
||||
completion=output_format.model_validate(data),
|
||||
usage=usage,
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
raise ModelProviderError(
|
||||
message=f'Failed to validate structured output: {str(e)}',
|
||||
model=self.name,
|
||||
) from e
|
||||
|
||||
# If no tool use found but output_format was requested
|
||||
raise ModelProviderError(
|
||||
message='Expected structured output but no tool use found in response',
|
||||
model=self.name,
|
||||
)
|
||||
|
||||
# If no valid content found
|
||||
if output_format is None:
|
||||
return ChatInvokeCompletion(
|
||||
completion='',
|
||||
usage=usage,
|
||||
)
|
||||
else:
|
||||
raise ModelProviderError(
|
||||
message='No valid content found in response',
|
||||
model=self.name,
|
||||
)
|
||||
|
||||
except ClientError as e:
|
||||
error_code = e.response.get('Error', {}).get('Code', 'Unknown')
|
||||
error_message = e.response.get('Error', {}).get('Message', str(e))
|
||||
|
||||
if error_code in ['ThrottlingException', 'TooManyRequestsException']:
|
||||
raise ModelRateLimitError(message=error_message, model=self.name) from e
|
||||
else:
|
||||
raise ModelProviderError(message=error_message, model=self.name) from e
|
||||
except Exception as e:
|
||||
raise ModelProviderError(message=str(e), model=self.name) from e
|
||||
@@ -0,0 +1,257 @@
|
||||
import base64
|
||||
import json
|
||||
import re
|
||||
from typing import Any, overload
|
||||
|
||||
from browser_use.llm.messages import (
|
||||
AssistantMessage,
|
||||
BaseMessage,
|
||||
ContentPartImageParam,
|
||||
ContentPartRefusalParam,
|
||||
ContentPartTextParam,
|
||||
SystemMessage,
|
||||
ToolCall,
|
||||
UserMessage,
|
||||
)
|
||||
|
||||
|
||||
class AWSBedrockMessageSerializer:
|
||||
"""Serializer for converting between custom message types and AWS Bedrock message format."""
|
||||
|
||||
@staticmethod
|
||||
def _is_base64_image(url: str) -> bool:
|
||||
"""Check if the URL is a base64 encoded image."""
|
||||
return url.startswith('data:image/')
|
||||
|
||||
@staticmethod
|
||||
def _is_url_image(url: str) -> bool:
|
||||
"""Check if the URL is a regular HTTP/HTTPS image URL."""
|
||||
return url.startswith(('http://', 'https://')) and any(
|
||||
url.lower().endswith(ext) for ext in ['.jpg', '.jpeg', '.png', '.gif', '.webp', '.bmp']
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_base64_url(url: str) -> tuple[str, bytes]:
|
||||
"""Parse a base64 data URL to extract format and raw bytes."""
|
||||
# Format: data:image/jpeg;base64,<data>
|
||||
if not url.startswith('data:'):
|
||||
raise ValueError(f'Invalid base64 URL: {url}')
|
||||
|
||||
header, data = url.split(',', 1)
|
||||
|
||||
# Extract format from mime type
|
||||
mime_match = re.search(r'image/(\w+)', header)
|
||||
if mime_match:
|
||||
format_name = mime_match.group(1).lower()
|
||||
# Map common formats
|
||||
format_mapping = {'jpg': 'jpeg', 'jpeg': 'jpeg', 'png': 'png', 'gif': 'gif', 'webp': 'webp'}
|
||||
image_format = format_mapping.get(format_name, 'jpeg')
|
||||
else:
|
||||
image_format = 'jpeg' # Default format
|
||||
|
||||
# Decode base64 data
|
||||
try:
|
||||
image_bytes = base64.b64decode(data)
|
||||
except Exception as e:
|
||||
raise ValueError(f'Failed to decode base64 image data: {e}')
|
||||
|
||||
return image_format, image_bytes
|
||||
|
||||
@staticmethod
|
||||
def _download_and_convert_image(url: str) -> tuple[str, bytes]:
|
||||
"""Download an image from URL and convert to base64 bytes."""
|
||||
try:
|
||||
import httpx
|
||||
except ImportError:
|
||||
raise ImportError('httpx not available. Please install it to use URL images with AWS Bedrock.')
|
||||
|
||||
try:
|
||||
response = httpx.get(url, timeout=30)
|
||||
response.raise_for_status()
|
||||
|
||||
# Detect format from content type or URL
|
||||
content_type = response.headers.get('content-type', '').lower()
|
||||
if 'jpeg' in content_type or url.lower().endswith(('.jpg', '.jpeg')):
|
||||
image_format = 'jpeg'
|
||||
elif 'png' in content_type or url.lower().endswith('.png'):
|
||||
image_format = 'png'
|
||||
elif 'gif' in content_type or url.lower().endswith('.gif'):
|
||||
image_format = 'gif'
|
||||
elif 'webp' in content_type or url.lower().endswith('.webp'):
|
||||
image_format = 'webp'
|
||||
else:
|
||||
image_format = 'jpeg' # Default format
|
||||
|
||||
return image_format, response.content
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(f'Failed to download image from {url}: {e}')
|
||||
|
||||
@staticmethod
|
||||
def _serialize_content_part_text(part: ContentPartTextParam) -> dict[str, Any]:
|
||||
"""Convert a text content part to AWS Bedrock format."""
|
||||
return {'text': part.text}
|
||||
|
||||
@staticmethod
|
||||
def _serialize_content_part_image(part: ContentPartImageParam) -> dict[str, Any]:
|
||||
"""Convert an image content part to AWS Bedrock format."""
|
||||
url = part.image_url.url
|
||||
|
||||
if AWSBedrockMessageSerializer._is_base64_image(url):
|
||||
# Handle base64 encoded images
|
||||
image_format, image_bytes = AWSBedrockMessageSerializer._parse_base64_url(url)
|
||||
elif AWSBedrockMessageSerializer._is_url_image(url):
|
||||
# Download and convert URL images
|
||||
image_format, image_bytes = AWSBedrockMessageSerializer._download_and_convert_image(url)
|
||||
else:
|
||||
raise ValueError(f'Unsupported image URL format: {url}')
|
||||
|
||||
return {
|
||||
'image': {
|
||||
'format': image_format,
|
||||
'source': {
|
||||
'bytes': image_bytes,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _serialize_user_content(
|
||||
content: str | list[ContentPartTextParam | ContentPartImageParam],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Serialize content for user messages."""
|
||||
if isinstance(content, str):
|
||||
return [{'text': content}]
|
||||
|
||||
content_blocks: list[dict[str, Any]] = []
|
||||
for part in content:
|
||||
if part.type == 'text':
|
||||
content_blocks.append(AWSBedrockMessageSerializer._serialize_content_part_text(part))
|
||||
elif part.type == 'image_url':
|
||||
content_blocks.append(AWSBedrockMessageSerializer._serialize_content_part_image(part))
|
||||
|
||||
return content_blocks
|
||||
|
||||
@staticmethod
|
||||
def _serialize_system_content(
|
||||
content: str | list[ContentPartTextParam],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Serialize content for system messages."""
|
||||
if isinstance(content, str):
|
||||
return [{'text': content}]
|
||||
|
||||
content_blocks: list[dict[str, Any]] = []
|
||||
for part in content:
|
||||
if part.type == 'text':
|
||||
content_blocks.append(AWSBedrockMessageSerializer._serialize_content_part_text(part))
|
||||
|
||||
return content_blocks
|
||||
|
||||
@staticmethod
|
||||
def _serialize_assistant_content(
|
||||
content: str | list[ContentPartTextParam | ContentPartRefusalParam] | None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Serialize content for assistant messages."""
|
||||
if content is None:
|
||||
return []
|
||||
if isinstance(content, str):
|
||||
return [{'text': content}]
|
||||
|
||||
content_blocks: list[dict[str, Any]] = []
|
||||
for part in content:
|
||||
if part.type == 'text':
|
||||
content_blocks.append(AWSBedrockMessageSerializer._serialize_content_part_text(part))
|
||||
# Skip refusal content parts - AWS Bedrock doesn't need them
|
||||
|
||||
return content_blocks
|
||||
|
||||
@staticmethod
|
||||
def _serialize_tool_call(tool_call: ToolCall) -> dict[str, Any]:
|
||||
"""Convert a tool call to AWS Bedrock format."""
|
||||
try:
|
||||
arguments = json.loads(tool_call.function.arguments)
|
||||
except json.JSONDecodeError:
|
||||
# If arguments aren't valid JSON, wrap them
|
||||
arguments = {'arguments': tool_call.function.arguments}
|
||||
|
||||
return {
|
||||
'toolUse': {
|
||||
'toolUseId': tool_call.id,
|
||||
'name': tool_call.function.name,
|
||||
'input': arguments,
|
||||
}
|
||||
}
|
||||
|
||||
# region - Serialize overloads
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: UserMessage) -> dict[str, Any]: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: SystemMessage) -> SystemMessage: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: AssistantMessage) -> dict[str, Any]: ...
|
||||
|
||||
@staticmethod
|
||||
def serialize(message: BaseMessage) -> dict[str, Any] | SystemMessage:
|
||||
"""Serialize a custom message to AWS Bedrock format."""
|
||||
|
||||
if isinstance(message, UserMessage):
|
||||
return {
|
||||
'role': 'user',
|
||||
'content': AWSBedrockMessageSerializer._serialize_user_content(message.content),
|
||||
}
|
||||
|
||||
elif isinstance(message, SystemMessage):
|
||||
# System messages are handled separately in AWS Bedrock
|
||||
return message
|
||||
|
||||
elif isinstance(message, AssistantMessage):
|
||||
content_blocks: list[dict[str, Any]] = []
|
||||
|
||||
# Add content blocks if present
|
||||
if message.content is not None:
|
||||
content_blocks.extend(AWSBedrockMessageSerializer._serialize_assistant_content(message.content))
|
||||
|
||||
# Add tool use blocks if present
|
||||
if message.tool_calls:
|
||||
for tool_call in message.tool_calls:
|
||||
content_blocks.append(AWSBedrockMessageSerializer._serialize_tool_call(tool_call))
|
||||
|
||||
# AWS Bedrock requires at least one content block
|
||||
if not content_blocks:
|
||||
content_blocks = [{'text': ''}]
|
||||
|
||||
return {
|
||||
'role': 'assistant',
|
||||
'content': content_blocks,
|
||||
}
|
||||
|
||||
else:
|
||||
raise ValueError(f'Unknown message type: {type(message)}')
|
||||
|
||||
@staticmethod
|
||||
def serialize_messages(messages: list[BaseMessage]) -> tuple[list[dict[str, Any]], list[dict[str, Any]] | None]:
|
||||
"""
|
||||
Serialize a list of messages, extracting any system message.
|
||||
|
||||
Returns:
|
||||
Tuple of (bedrock_messages, system_message) where system_message is extracted
|
||||
from any SystemMessage in the list.
|
||||
"""
|
||||
bedrock_messages: list[dict[str, Any]] = []
|
||||
system_message: list[dict[str, Any]] | None = None
|
||||
|
||||
for message in messages:
|
||||
if isinstance(message, SystemMessage):
|
||||
# Extract system message content
|
||||
system_message = AWSBedrockMessageSerializer._serialize_system_content(message.content)
|
||||
else:
|
||||
# Serialize and add to regular messages
|
||||
serialized = AWSBedrockMessageSerializer.serialize(message)
|
||||
bedrock_messages.append(serialized)
|
||||
|
||||
return bedrock_messages, system_message
|
||||
Reference in New Issue
Block a user