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