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,212 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, TypeVar, overload
|
||||
|
||||
import httpx
|
||||
from openai import (
|
||||
APIConnectionError,
|
||||
APIError,
|
||||
APIStatusError,
|
||||
APITimeoutError,
|
||||
AsyncOpenAI,
|
||||
RateLimitError,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
|
||||
from browser_use.llm.base import BaseChatModel
|
||||
from browser_use.llm.deepseek.serializer import DeepSeekMessageSerializer
|
||||
from browser_use.llm.exceptions import ModelProviderError, ModelRateLimitError
|
||||
from browser_use.llm.messages import BaseMessage
|
||||
from browser_use.llm.schema import SchemaOptimizer
|
||||
from browser_use.llm.views import ChatInvokeCompletion
|
||||
|
||||
T = TypeVar('T', bound=BaseModel)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatDeepSeek(BaseChatModel):
|
||||
"""DeepSeek /chat/completions wrapper (OpenAI-compatible)."""
|
||||
|
||||
model: str = 'deepseek-chat'
|
||||
|
||||
# Generation parameters
|
||||
max_tokens: int | None = None
|
||||
temperature: float | None = None
|
||||
top_p: float | None = None
|
||||
seed: int | None = None
|
||||
|
||||
# Connection parameters
|
||||
api_key: str | None = None
|
||||
base_url: str | httpx.URL | None = 'https://api.deepseek.com/v1'
|
||||
timeout: float | httpx.Timeout | None = None
|
||||
client_params: dict[str, Any] | None = None
|
||||
|
||||
@property
|
||||
def provider(self) -> str:
|
||||
return 'deepseek'
|
||||
|
||||
def _client(self) -> AsyncOpenAI:
|
||||
return AsyncOpenAI(
|
||||
api_key=self.api_key,
|
||||
base_url=self.base_url,
|
||||
timeout=self.timeout,
|
||||
**(self.client_params or {}),
|
||||
)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return self.model
|
||||
|
||||
@overload
|
||||
async def ainvoke(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
output_format: None = None,
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
stop: list[str] | None = None,
|
||||
) -> ChatInvokeCompletion[str]: ...
|
||||
|
||||
@overload
|
||||
async def ainvoke(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
output_format: type[T],
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
stop: list[str] | None = None,
|
||||
) -> ChatInvokeCompletion[T]: ...
|
||||
|
||||
async def ainvoke(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
output_format: type[T] | None = None,
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
stop: list[str] | None = None,
|
||||
) -> ChatInvokeCompletion[T] | ChatInvokeCompletion[str]:
|
||||
"""
|
||||
DeepSeek ainvoke supports:
|
||||
1. Regular text/multi-turn conversation
|
||||
2. Function Calling
|
||||
3. JSON Output (response_format)
|
||||
4. Conversation prefix continuation (beta, prefix, stop)
|
||||
"""
|
||||
client = self._client()
|
||||
ds_messages = DeepSeekMessageSerializer.serialize_messages(messages)
|
||||
common: dict[str, Any] = {}
|
||||
|
||||
if self.temperature is not None:
|
||||
common['temperature'] = self.temperature
|
||||
if self.max_tokens is not None:
|
||||
common['max_tokens'] = self.max_tokens
|
||||
if self.top_p is not None:
|
||||
common['top_p'] = self.top_p
|
||||
if self.seed is not None:
|
||||
common['seed'] = self.seed
|
||||
|
||||
# Beta conversation prefix continuation (see official documentation)
|
||||
if self.base_url and str(self.base_url).endswith('/beta'):
|
||||
# The last assistant message must have prefix
|
||||
if ds_messages and isinstance(ds_messages[-1], dict) and ds_messages[-1].get('role') == 'assistant':
|
||||
ds_messages[-1]['prefix'] = True
|
||||
if stop:
|
||||
common['stop'] = stop
|
||||
|
||||
# ① Regular multi-turn conversation/text output
|
||||
if output_format is None and not tools:
|
||||
try:
|
||||
resp = await client.chat.completions.create( # type: ignore
|
||||
model=self.model,
|
||||
messages=ds_messages, # type: ignore
|
||||
**common,
|
||||
)
|
||||
return ChatInvokeCompletion(
|
||||
completion=resp.choices[0].message.content or '',
|
||||
usage=None,
|
||||
)
|
||||
except RateLimitError as e:
|
||||
raise ModelRateLimitError(str(e), model=self.name) from e
|
||||
except (APIError, APIConnectionError, APITimeoutError, APIStatusError) as e:
|
||||
raise ModelProviderError(str(e), model=self.name) from e
|
||||
except Exception as e:
|
||||
raise ModelProviderError(str(e), model=self.name) from e
|
||||
|
||||
# ② Function Calling path (with tools or output_format)
|
||||
if tools or (output_format is not None and hasattr(output_format, 'model_json_schema')):
|
||||
try:
|
||||
call_tools = tools
|
||||
tool_choice = None
|
||||
if output_format is not None and hasattr(output_format, 'model_json_schema'):
|
||||
tool_name = output_format.__name__
|
||||
schema = SchemaOptimizer.create_optimized_json_schema(output_format)
|
||||
schema.pop('title', None)
|
||||
call_tools = [
|
||||
{
|
||||
'type': 'function',
|
||||
'function': {
|
||||
'name': tool_name,
|
||||
'description': f'Return a JSON object of type {tool_name}',
|
||||
'parameters': schema,
|
||||
},
|
||||
}
|
||||
]
|
||||
tool_choice = {'type': 'function', 'function': {'name': tool_name}}
|
||||
resp = await client.chat.completions.create( # type: ignore
|
||||
model=self.model,
|
||||
messages=ds_messages, # type: ignore
|
||||
tools=call_tools, # type: ignore
|
||||
tool_choice=tool_choice, # type: ignore
|
||||
**common,
|
||||
)
|
||||
msg = resp.choices[0].message
|
||||
if not msg.tool_calls:
|
||||
raise ValueError('Expected tool_calls in response but got none')
|
||||
raw_args = msg.tool_calls[0].function.arguments
|
||||
if isinstance(raw_args, str):
|
||||
parsed = json.loads(raw_args)
|
||||
else:
|
||||
parsed = raw_args
|
||||
# --------- Fix: only use model_validate when output_format is not None ----------
|
||||
if output_format is not None:
|
||||
return ChatInvokeCompletion(
|
||||
completion=output_format.model_validate(parsed),
|
||||
usage=None,
|
||||
)
|
||||
else:
|
||||
# If no output_format, return dict directly
|
||||
return ChatInvokeCompletion(
|
||||
completion=parsed,
|
||||
usage=None,
|
||||
)
|
||||
except RateLimitError as e:
|
||||
raise ModelRateLimitError(str(e), model=self.name) from e
|
||||
except (APIError, APIConnectionError, APITimeoutError, APIStatusError) as e:
|
||||
raise ModelProviderError(str(e), model=self.name) from e
|
||||
except Exception as e:
|
||||
raise ModelProviderError(str(e), model=self.name) from e
|
||||
|
||||
# ③ JSON Output path (official response_format)
|
||||
if output_format is not None and hasattr(output_format, 'model_json_schema'):
|
||||
try:
|
||||
resp = await client.chat.completions.create( # type: ignore
|
||||
model=self.model,
|
||||
messages=ds_messages, # type: ignore
|
||||
response_format={'type': 'json_object'},
|
||||
**common,
|
||||
)
|
||||
content = resp.choices[0].message.content
|
||||
if not content:
|
||||
raise ModelProviderError('Empty JSON content in DeepSeek response', model=self.name)
|
||||
parsed = output_format.model_validate_json(content)
|
||||
return ChatInvokeCompletion(
|
||||
completion=parsed,
|
||||
usage=None,
|
||||
)
|
||||
except RateLimitError as e:
|
||||
raise ModelRateLimitError(str(e), model=self.name) from e
|
||||
except (APIError, APIConnectionError, APITimeoutError, APIStatusError) as e:
|
||||
raise ModelProviderError(str(e), model=self.name) from e
|
||||
except Exception as e:
|
||||
raise ModelProviderError(str(e), model=self.name) from e
|
||||
|
||||
raise ModelProviderError('No valid ainvoke execution path for DeepSeek LLM', model=self.name)
|
||||
@@ -0,0 +1,109 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, overload
|
||||
|
||||
from browser_use.llm.messages import (
|
||||
AssistantMessage,
|
||||
BaseMessage,
|
||||
ContentPartImageParam,
|
||||
ContentPartTextParam,
|
||||
SystemMessage,
|
||||
ToolCall,
|
||||
UserMessage,
|
||||
)
|
||||
|
||||
MessageDict = dict[str, Any]
|
||||
|
||||
|
||||
class DeepSeekMessageSerializer:
|
||||
"""Serializer for converting browser-use messages to DeepSeek messages."""
|
||||
|
||||
# -------- content 处理 --------------------------------------------------
|
||||
@staticmethod
|
||||
def _serialize_text_part(part: ContentPartTextParam) -> str:
|
||||
return part.text
|
||||
|
||||
@staticmethod
|
||||
def _serialize_image_part(part: ContentPartImageParam) -> dict[str, Any]:
|
||||
url = part.image_url.url
|
||||
if url.startswith('data:'):
|
||||
return {'type': 'image_url', 'image_url': {'url': url}}
|
||||
return {'type': 'image_url', 'image_url': {'url': url}}
|
||||
|
||||
@staticmethod
|
||||
def _serialize_content(content: Any) -> str | list[dict[str, Any]]:
|
||||
if content is None:
|
||||
return ''
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
serialized: list[dict[str, Any]] = []
|
||||
for part in content:
|
||||
if part.type == 'text':
|
||||
serialized.append({'type': 'text', 'text': DeepSeekMessageSerializer._serialize_text_part(part)})
|
||||
elif part.type == 'image_url':
|
||||
serialized.append(DeepSeekMessageSerializer._serialize_image_part(part))
|
||||
elif part.type == 'refusal':
|
||||
serialized.append({'type': 'text', 'text': f'[Refusal] {part.refusal}'})
|
||||
return serialized
|
||||
|
||||
# -------- Tool-call 处理 -------------------------------------------------
|
||||
@staticmethod
|
||||
def _serialize_tool_calls(tool_calls: list[ToolCall]) -> list[dict[str, Any]]:
|
||||
deepseek_tool_calls: list[dict[str, Any]] = []
|
||||
for tc in tool_calls:
|
||||
try:
|
||||
arguments = json.loads(tc.function.arguments)
|
||||
except json.JSONDecodeError:
|
||||
arguments = {'arguments': tc.function.arguments}
|
||||
deepseek_tool_calls.append(
|
||||
{
|
||||
'id': tc.id,
|
||||
'type': 'function',
|
||||
'function': {
|
||||
'name': tc.function.name,
|
||||
'arguments': arguments,
|
||||
},
|
||||
}
|
||||
)
|
||||
return deepseek_tool_calls
|
||||
|
||||
# -------- 单条消息序列化 -------------------------------------------------
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: UserMessage) -> MessageDict: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: SystemMessage) -> MessageDict: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def serialize(message: AssistantMessage) -> MessageDict: ...
|
||||
|
||||
@staticmethod
|
||||
def serialize(message: BaseMessage) -> MessageDict:
|
||||
if isinstance(message, UserMessage):
|
||||
return {
|
||||
'role': 'user',
|
||||
'content': DeepSeekMessageSerializer._serialize_content(message.content),
|
||||
}
|
||||
if isinstance(message, SystemMessage):
|
||||
return {
|
||||
'role': 'system',
|
||||
'content': DeepSeekMessageSerializer._serialize_content(message.content),
|
||||
}
|
||||
if isinstance(message, AssistantMessage):
|
||||
msg: MessageDict = {
|
||||
'role': 'assistant',
|
||||
'content': DeepSeekMessageSerializer._serialize_content(message.content),
|
||||
}
|
||||
if message.tool_calls:
|
||||
msg['tool_calls'] = DeepSeekMessageSerializer._serialize_tool_calls(message.tool_calls)
|
||||
return msg
|
||||
raise ValueError(f'Unknown message type: {type(message)}')
|
||||
|
||||
# -------- 列表序列化 -----------------------------------------------------
|
||||
@staticmethod
|
||||
def serialize_messages(messages: list[BaseMessage]) -> list[MessageDict]:
|
||||
return [DeepSeekMessageSerializer.serialize(m) for m in messages]
|
||||
Reference in New Issue
Block a user