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,273 @@
|
||||
from collections.abc import Iterable, Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal, TypeVar, overload
|
||||
|
||||
import httpx
|
||||
from openai import APIConnectionError, APIStatusError, AsyncOpenAI, RateLimitError
|
||||
from openai.types.chat import ChatCompletionContentPartTextParam
|
||||
from openai.types.chat.chat_completion import ChatCompletion
|
||||
from openai.types.shared.chat_model import ChatModel
|
||||
from openai.types.shared_params.reasoning_effort import ReasoningEffort
|
||||
from openai.types.shared_params.response_format_json_schema import JSONSchema, ResponseFormatJSONSchema
|
||||
from pydantic import BaseModel
|
||||
|
||||
from browser_use.llm.base import BaseChatModel
|
||||
from browser_use.llm.exceptions import ModelProviderError
|
||||
from browser_use.llm.messages import BaseMessage
|
||||
from browser_use.llm.openai.serializer import OpenAIMessageSerializer
|
||||
from browser_use.llm.schema import SchemaOptimizer
|
||||
from browser_use.llm.views import ChatInvokeCompletion, ChatInvokeUsage
|
||||
|
||||
T = TypeVar('T', bound=BaseModel)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatOpenAI(BaseChatModel):
|
||||
"""
|
||||
A wrapper around AsyncOpenAI that implements the BaseLLM protocol.
|
||||
|
||||
This class accepts all AsyncOpenAI parameters while adding model
|
||||
and temperature parameters for the LLM interface (if temperature it not `None`).
|
||||
"""
|
||||
|
||||
# Model configuration
|
||||
model: ChatModel | str
|
||||
|
||||
# Model params
|
||||
temperature: float | None = 0.2
|
||||
frequency_penalty: float | None = 0.3 # this avoids infinite generation of \t for models like 4.1-mini
|
||||
reasoning_effort: ReasoningEffort = 'low'
|
||||
seed: int | None = None
|
||||
service_tier: Literal['auto', 'default', 'flex', 'priority', 'scale'] | None = None
|
||||
top_p: float | None = None
|
||||
add_schema_to_system_prompt: bool = False # Add JSON schema to system prompt instead of using response_format
|
||||
|
||||
# Client initialization parameters
|
||||
api_key: str | None = None
|
||||
organization: str | None = None
|
||||
project: str | None = None
|
||||
base_url: str | httpx.URL | None = None
|
||||
websocket_base_url: str | httpx.URL | None = None
|
||||
timeout: float | httpx.Timeout | None = None
|
||||
max_retries: int = 5 # Increase default retries for automation reliability
|
||||
default_headers: Mapping[str, str] | None = None
|
||||
default_query: Mapping[str, object] | None = None
|
||||
http_client: httpx.AsyncClient | None = None
|
||||
_strict_response_validation: bool = False
|
||||
max_completion_tokens: int | None = 4096
|
||||
reasoning_models: list[ChatModel | str] | None = field(
|
||||
default_factory=lambda: [
|
||||
'o4-mini',
|
||||
'o3',
|
||||
'o3-mini',
|
||||
'o1',
|
||||
'o1-pro',
|
||||
'o3-pro',
|
||||
'gpt-5',
|
||||
'gpt-5-mini',
|
||||
'gpt-5-nano',
|
||||
]
|
||||
)
|
||||
|
||||
# Static
|
||||
@property
|
||||
def provider(self) -> str:
|
||||
return 'openai'
|
||||
|
||||
def _get_client_params(self) -> dict[str, Any]:
|
||||
"""Prepare client parameters dictionary."""
|
||||
# Define base client params
|
||||
base_params = {
|
||||
'api_key': self.api_key,
|
||||
'organization': self.organization,
|
||||
'project': self.project,
|
||||
'base_url': self.base_url,
|
||||
'websocket_base_url': self.websocket_base_url,
|
||||
'timeout': self.timeout,
|
||||
'max_retries': self.max_retries,
|
||||
'default_headers': self.default_headers,
|
||||
'default_query': self.default_query,
|
||||
'_strict_response_validation': self._strict_response_validation,
|
||||
}
|
||||
|
||||
# Create client_params dict with non-None values
|
||||
client_params = {k: v for k, v in base_params.items() if v is not None}
|
||||
|
||||
# Add http_client if provided
|
||||
if self.http_client is not None:
|
||||
client_params['http_client'] = self.http_client
|
||||
|
||||
return client_params
|
||||
|
||||
def get_client(self) -> AsyncOpenAI:
|
||||
"""
|
||||
Returns an AsyncOpenAI client.
|
||||
|
||||
Returns:
|
||||
AsyncOpenAI: An instance of the AsyncOpenAI client.
|
||||
"""
|
||||
client_params = self._get_client_params()
|
||||
return AsyncOpenAI(**client_params)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return str(self.model)
|
||||
|
||||
def _get_usage(self, response: ChatCompletion) -> ChatInvokeUsage | None:
|
||||
if response.usage is not None:
|
||||
completion_tokens = response.usage.completion_tokens
|
||||
completion_token_details = response.usage.completion_tokens_details
|
||||
if completion_token_details is not None:
|
||||
reasoning_tokens = completion_token_details.reasoning_tokens
|
||||
if reasoning_tokens is not None:
|
||||
completion_tokens += reasoning_tokens
|
||||
|
||||
usage = ChatInvokeUsage(
|
||||
prompt_tokens=response.usage.prompt_tokens,
|
||||
prompt_cached_tokens=response.usage.prompt_tokens_details.cached_tokens
|
||||
if response.usage.prompt_tokens_details is not None
|
||||
else None,
|
||||
prompt_cache_creation_tokens=None,
|
||||
prompt_image_tokens=None,
|
||||
# Completion
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=response.usage.total_tokens,
|
||||
)
|
||||
else:
|
||||
usage = 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]:
|
||||
"""
|
||||
Invoke the 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
|
||||
"""
|
||||
|
||||
openai_messages = OpenAIMessageSerializer.serialize_messages(messages)
|
||||
|
||||
try:
|
||||
model_params: dict[str, Any] = {}
|
||||
|
||||
if self.temperature is not None:
|
||||
model_params['temperature'] = self.temperature
|
||||
|
||||
if self.frequency_penalty is not None:
|
||||
model_params['frequency_penalty'] = self.frequency_penalty
|
||||
|
||||
if self.max_completion_tokens is not None:
|
||||
model_params['max_completion_tokens'] = self.max_completion_tokens
|
||||
|
||||
if self.top_p is not None:
|
||||
model_params['top_p'] = self.top_p
|
||||
|
||||
if self.seed is not None:
|
||||
model_params['seed'] = self.seed
|
||||
|
||||
if self.service_tier is not None:
|
||||
model_params['service_tier'] = self.service_tier
|
||||
|
||||
if self.reasoning_models and any(str(m).lower() in str(self.model).lower() for m in self.reasoning_models):
|
||||
model_params['reasoning_effort'] = self.reasoning_effort
|
||||
del model_params['temperature']
|
||||
del model_params['frequency_penalty']
|
||||
|
||||
if output_format is None:
|
||||
# Return string response
|
||||
response = await self.get_client().chat.completions.create(
|
||||
model=self.model,
|
||||
messages=openai_messages,
|
||||
**model_params,
|
||||
)
|
||||
|
||||
usage = self._get_usage(response)
|
||||
return ChatInvokeCompletion(
|
||||
completion=response.choices[0].message.content or '',
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
else:
|
||||
response_format: JSONSchema = {
|
||||
'name': 'agent_output',
|
||||
'strict': True,
|
||||
'schema': SchemaOptimizer.create_optimized_json_schema(output_format),
|
||||
}
|
||||
|
||||
# Add JSON schema to system prompt if requested
|
||||
if self.add_schema_to_system_prompt and openai_messages and openai_messages[0]['role'] == 'system':
|
||||
schema_text = f'\n<json_schema>\n{response_format}\n</json_schema>'
|
||||
if isinstance(openai_messages[0]['content'], str):
|
||||
openai_messages[0]['content'] += schema_text
|
||||
elif isinstance(openai_messages[0]['content'], Iterable):
|
||||
openai_messages[0]['content'] = list(openai_messages[0]['content']) + [
|
||||
ChatCompletionContentPartTextParam(text=schema_text, type='text')
|
||||
]
|
||||
|
||||
# Return structured response
|
||||
response = await self.get_client().chat.completions.create(
|
||||
model=self.model,
|
||||
messages=openai_messages,
|
||||
response_format=ResponseFormatJSONSchema(json_schema=response_format, type='json_schema'),
|
||||
**model_params,
|
||||
)
|
||||
|
||||
if response.choices[0].message.content is None:
|
||||
raise ModelProviderError(
|
||||
message='Failed to parse structured output from model response',
|
||||
status_code=500,
|
||||
model=self.name,
|
||||
)
|
||||
|
||||
usage = self._get_usage(response)
|
||||
|
||||
parsed = output_format.model_validate_json(response.choices[0].message.content)
|
||||
|
||||
return ChatInvokeCompletion(
|
||||
completion=parsed,
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
except RateLimitError as e:
|
||||
error_message = e.response.json().get('error', {})
|
||||
error_message = (
|
||||
error_message.get('message', 'Unknown model error') if isinstance(error_message, dict) else error_message
|
||||
)
|
||||
raise ModelProviderError(
|
||||
message=error_message,
|
||||
status_code=e.response.status_code,
|
||||
model=self.name,
|
||||
) from e
|
||||
|
||||
except APIConnectionError as e:
|
||||
raise ModelProviderError(message=str(e), model=self.name) from e
|
||||
|
||||
except APIStatusError as e:
|
||||
try:
|
||||
error_message = e.response.json().get('error', {})
|
||||
except Exception:
|
||||
error_message = e.response.text
|
||||
error_message = (
|
||||
error_message.get('message', 'Unknown model error') if isinstance(error_message, dict) else error_message
|
||||
)
|
||||
raise ModelProviderError(
|
||||
message=error_message,
|
||||
status_code=e.response.status_code,
|
||||
model=self.name,
|
||||
) from e
|
||||
|
||||
except Exception as e:
|
||||
raise ModelProviderError(message=str(e), model=self.name) from e
|
||||
@@ -0,0 +1,15 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from browser_use.llm.openai.chat import ChatOpenAI
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatOpenAILike(ChatOpenAI):
|
||||
"""
|
||||
A class for to interact with any provider using the OpenAI API schema.
|
||||
|
||||
Args:
|
||||
model (str): The name of the OpenAI model to use.
|
||||
"""
|
||||
|
||||
model: str
|
||||
@@ -0,0 +1,165 @@
|
||||
from typing import overload
|
||||
|
||||
from openai.types.chat import (
|
||||
ChatCompletionAssistantMessageParam,
|
||||
ChatCompletionContentPartImageParam,
|
||||
ChatCompletionContentPartRefusalParam,
|
||||
ChatCompletionContentPartTextParam,
|
||||
ChatCompletionMessageFunctionToolCallParam,
|
||||
ChatCompletionMessageParam,
|
||||
ChatCompletionSystemMessageParam,
|
||||
ChatCompletionUserMessageParam,
|
||||
)
|
||||
from openai.types.chat.chat_completion_content_part_image_param import ImageURL
|
||||
from openai.types.chat.chat_completion_message_function_tool_call_param import Function
|
||||
|
||||
from browser_use.llm.messages import (
|
||||
AssistantMessage,
|
||||
BaseMessage,
|
||||
ContentPartImageParam,
|
||||
ContentPartRefusalParam,
|
||||
ContentPartTextParam,
|
||||
SystemMessage,
|
||||
ToolCall,
|
||||
UserMessage,
|
||||
)
|
||||
|
||||
|
||||
class OpenAIMessageSerializer:
|
||||
"""Serializer for converting between custom message types and OpenAI message param types."""
|
||||
|
||||
@staticmethod
|
||||
def _serialize_content_part_text(part: ContentPartTextParam) -> ChatCompletionContentPartTextParam:
|
||||
return ChatCompletionContentPartTextParam(text=part.text, type='text')
|
||||
|
||||
@staticmethod
|
||||
def _serialize_content_part_image(part: ContentPartImageParam) -> ChatCompletionContentPartImageParam:
|
||||
return ChatCompletionContentPartImageParam(
|
||||
image_url=ImageURL(url=part.image_url.url, detail=part.image_url.detail),
|
||||
type='image_url',
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _serialize_content_part_refusal(part: ContentPartRefusalParam) -> ChatCompletionContentPartRefusalParam:
|
||||
return ChatCompletionContentPartRefusalParam(refusal=part.refusal, type='refusal')
|
||||
|
||||
@staticmethod
|
||||
def _serialize_user_content(
|
||||
content: str | list[ContentPartTextParam | ContentPartImageParam],
|
||||
) -> str | list[ChatCompletionContentPartTextParam | ChatCompletionContentPartImageParam]:
|
||||
"""Serialize content for user messages (text and images allowed)."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
|
||||
serialized_parts: list[ChatCompletionContentPartTextParam | ChatCompletionContentPartImageParam] = []
|
||||
for part in content:
|
||||
if part.type == 'text':
|
||||
serialized_parts.append(OpenAIMessageSerializer._serialize_content_part_text(part))
|
||||
elif part.type == 'image_url':
|
||||
serialized_parts.append(OpenAIMessageSerializer._serialize_content_part_image(part))
|
||||
return serialized_parts
|
||||
|
||||
@staticmethod
|
||||
def _serialize_system_content(
|
||||
content: str | list[ContentPartTextParam],
|
||||
) -> str | list[ChatCompletionContentPartTextParam]:
|
||||
"""Serialize content for system messages (text only)."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
|
||||
serialized_parts: list[ChatCompletionContentPartTextParam] = []
|
||||
for part in content:
|
||||
if part.type == 'text':
|
||||
serialized_parts.append(OpenAIMessageSerializer._serialize_content_part_text(part))
|
||||
return serialized_parts
|
||||
|
||||
@staticmethod
|
||||
def _serialize_assistant_content(
|
||||
content: str | list[ContentPartTextParam | ContentPartRefusalParam] | None,
|
||||
) -> str | list[ChatCompletionContentPartTextParam | ChatCompletionContentPartRefusalParam] | None:
|
||||
"""Serialize content for assistant messages (text and refusal allowed)."""
|
||||
if content is None:
|
||||
return None
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
|
||||
serialized_parts: list[ChatCompletionContentPartTextParam | ChatCompletionContentPartRefusalParam] = []
|
||||
for part in content:
|
||||
if part.type == 'text':
|
||||
serialized_parts.append(OpenAIMessageSerializer._serialize_content_part_text(part))
|
||||
elif part.type == 'refusal':
|
||||
serialized_parts.append(OpenAIMessageSerializer._serialize_content_part_refusal(part))
|
||||
return serialized_parts
|
||||
|
||||
@staticmethod
|
||||
def _serialize_tool_call(tool_call: ToolCall) -> ChatCompletionMessageFunctionToolCallParam:
|
||||
return ChatCompletionMessageFunctionToolCallParam(
|
||||
id=tool_call.id,
|
||||
function=Function(name=tool_call.function.name, arguments=tool_call.function.arguments),
|
||||
type='function',
|
||||
)
|
||||
|
||||
# endregion
|
||||
|
||||
# region - Serialize overloads
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: UserMessage) -> ChatCompletionUserMessageParam: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: SystemMessage) -> ChatCompletionSystemMessageParam: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: AssistantMessage) -> ChatCompletionAssistantMessageParam: ...
|
||||
|
||||
@staticmethod
|
||||
def serialize(message: BaseMessage) -> ChatCompletionMessageParam:
|
||||
"""Serialize a custom message to an OpenAI message param."""
|
||||
|
||||
if isinstance(message, UserMessage):
|
||||
user_result: ChatCompletionUserMessageParam = {
|
||||
'role': 'user',
|
||||
'content': OpenAIMessageSerializer._serialize_user_content(message.content),
|
||||
}
|
||||
if message.name is not None:
|
||||
user_result['name'] = message.name
|
||||
return user_result
|
||||
|
||||
elif isinstance(message, SystemMessage):
|
||||
system_result: ChatCompletionSystemMessageParam = {
|
||||
'role': 'system',
|
||||
'content': OpenAIMessageSerializer._serialize_system_content(message.content),
|
||||
}
|
||||
if message.name is not None:
|
||||
system_result['name'] = message.name
|
||||
return system_result
|
||||
|
||||
elif isinstance(message, AssistantMessage):
|
||||
# Handle content serialization
|
||||
content = None
|
||||
if message.content is not None:
|
||||
content = OpenAIMessageSerializer._serialize_assistant_content(message.content)
|
||||
|
||||
assistant_result: ChatCompletionAssistantMessageParam = {'role': 'assistant'}
|
||||
|
||||
# Only add content if it's not None
|
||||
if content is not None:
|
||||
assistant_result['content'] = content
|
||||
|
||||
if message.name is not None:
|
||||
assistant_result['name'] = message.name
|
||||
if message.refusal is not None:
|
||||
assistant_result['refusal'] = message.refusal
|
||||
if message.tool_calls:
|
||||
assistant_result['tool_calls'] = [OpenAIMessageSerializer._serialize_tool_call(tc) for tc in message.tool_calls]
|
||||
|
||||
return assistant_result
|
||||
|
||||
else:
|
||||
raise ValueError(f'Unknown message type: {type(message)}')
|
||||
|
||||
@staticmethod
|
||||
def serialize_messages(messages: list[BaseMessage]) -> list[ChatCompletionMessageParam]:
|
||||
return [OpenAIMessageSerializer.serialize(m) for m in messages]
|
||||
Reference in New Issue
Block a user