01c36eb32e
- 删除 mock_api_client、mock_calendar_service、mock_history_service - 新增 fixed_length_code_input、link_button、message_composer 共享组件 - 优化登录/注册/密码重置页面使用新组件 - 简化 injection.dart 移除 mock 分支 - 更新 env.dart 配置(BACKEND_URL 替换 API_URL) - 后端 agentscope 工具和测试更新 - 重构 AGENTS.md 文档结构 - 新增 deploy/ 目录和 protocol 文档
140 lines
3.7 KiB
Python
140 lines
3.7 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from typing import Any, cast
|
|
from uuid import UUID
|
|
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from core.agentscope.tools.custom.calendar import (
|
|
calendar_share,
|
|
calendar_read,
|
|
calendar_write,
|
|
)
|
|
from core.agentscope.tools.hitl_middleware import register_tool_middlewares
|
|
from core.agentscope.tools.tool_meta import TOOL_META
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class CustomToolBinding:
|
|
name: str
|
|
func: Any
|
|
preset_kwargs: dict[str, object]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ToolGroup:
|
|
stage: str
|
|
tool_names: frozenset[str]
|
|
|
|
|
|
TOOL_GROUPS: dict[str, ToolGroup] = {
|
|
"intent": ToolGroup(stage="intent", tool_names=frozenset({"calendar_read"})),
|
|
"execution": ToolGroup(
|
|
stage="execution",
|
|
tool_names=frozenset({"calendar_read", "calendar_write", "calendar_share"}),
|
|
),
|
|
"report": ToolGroup(stage="report", tool_names=frozenset()),
|
|
}
|
|
|
|
|
|
def get_tool_group(stage: str) -> ToolGroup:
|
|
group = TOOL_GROUPS.get(stage)
|
|
if group is None:
|
|
raise ValueError(f"unknown tool group stage: {stage}")
|
|
return group
|
|
|
|
|
|
def _load_custom_tool_bindings(
|
|
*,
|
|
session: AsyncSession,
|
|
owner_id: UUID,
|
|
user_token: str | None,
|
|
) -> list[CustomToolBinding]:
|
|
return [
|
|
CustomToolBinding(
|
|
name="calendar_read",
|
|
func=calendar_read,
|
|
preset_kwargs={
|
|
"session": session,
|
|
"owner_id": owner_id,
|
|
"user_token": user_token or "",
|
|
},
|
|
),
|
|
CustomToolBinding(
|
|
name="calendar_write",
|
|
func=calendar_write,
|
|
preset_kwargs={
|
|
"session": session,
|
|
"owner_id": owner_id,
|
|
"user_token": user_token or "",
|
|
},
|
|
),
|
|
CustomToolBinding(
|
|
name="calendar_share",
|
|
func=calendar_share,
|
|
preset_kwargs={
|
|
"session": session,
|
|
"owner_id": owner_id,
|
|
"user_token": user_token or "",
|
|
},
|
|
),
|
|
]
|
|
|
|
|
|
def build_toolkit(
|
|
*,
|
|
session: AsyncSession,
|
|
owner_id: UUID,
|
|
user_token: str | None = None,
|
|
enable_hitl: bool = True,
|
|
enabled_tool_names: set[str] | None = None,
|
|
):
|
|
from agentscope.tool import Toolkit
|
|
from agentscope.types import JSONSerializableObject
|
|
|
|
toolkit = Toolkit()
|
|
bindings = _load_custom_tool_bindings(
|
|
session=session,
|
|
owner_id=owner_id,
|
|
user_token=user_token,
|
|
)
|
|
registered_tool_names: set[str] = set()
|
|
for binding in bindings:
|
|
if enabled_tool_names is not None and binding.name not in enabled_tool_names:
|
|
continue
|
|
registered_tool_names.add(binding.name)
|
|
toolkit.register_tool_function(
|
|
binding.func,
|
|
func_name=binding.name,
|
|
preset_kwargs=cast(
|
|
dict[str, JSONSerializableObject],
|
|
binding.preset_kwargs,
|
|
),
|
|
)
|
|
if enabled_tool_names is not None:
|
|
missing = enabled_tool_names - registered_tool_names
|
|
if missing:
|
|
raise ValueError(f"unknown tools in enabled_tool_names: {sorted(missing)}")
|
|
if enable_hitl:
|
|
register_tool_middlewares(toolkit=toolkit, meta_by_name=TOOL_META)
|
|
return toolkit
|
|
|
|
|
|
def build_stage_toolkit(
|
|
*,
|
|
stage: str,
|
|
session: AsyncSession,
|
|
owner_id: UUID,
|
|
user_token: str | None = None,
|
|
enable_hitl: bool = True,
|
|
):
|
|
group = get_tool_group(stage)
|
|
return build_toolkit(
|
|
session=session,
|
|
owner_id=owner_id,
|
|
user_token=user_token,
|
|
enable_hitl=enable_hitl,
|
|
enabled_tool_names=set(group.tool_names),
|
|
)
|