Files
social-app/backend/src/core/agentscope/tools/toolkit.py
T

140 lines
3.7 KiB
Python
Raw Normal View History

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),
)