from __future__ import annotations from typing import Any, cast from uuid import UUID from agentscope.tool import Toolkit from agentscope.types import JSONSerializableObject from core.agentscope.tools.custom.calendar import ( calendar_read, calendar_share, calendar_write, ) from core.agentscope.tools.custom.user_lookup import user_lookup from core.agentscope.tools.tool_config import ( TOOL_CONFIGS, ToolGroup, resolve_tool_names_by_groups, ) from core.agentscope.tools.tool_middleware import register_tool_middlewares from sqlalchemy.ext.asyncio import AsyncSession from schemas.agent.system_agent import AgentType TOOL_FUNCTIONS: dict[str, Any] = { "calendar_read": calendar_read, "calendar_write": calendar_write, "calendar_share": calendar_share, "user_lookup": user_lookup, } AGENT_TYPE_TO_GROUPS: dict[AgentType, set[ToolGroup]] = { AgentType.ROUTER: {ToolGroup.READ}, AgentType.WORKER: {ToolGroup.READ, ToolGroup.WRITE}, } def _resolve_enabled_tools( *, groups: set[ToolGroup] | None, enabled_tool_names: set[str] | None, ) -> set[str]: if enabled_tool_names is not None: unknown = enabled_tool_names - set(TOOL_FUNCTIONS) if unknown: raise ValueError(f"unknown tools in enabled_tool_names: {sorted(unknown)}") return set(enabled_tool_names) if groups is None: return set(TOOL_FUNCTIONS) resolved = resolve_tool_names_by_groups(groups) unknown = resolved - set(TOOL_FUNCTIONS) if unknown: raise ValueError(f"tool config contains unknown tools: {sorted(unknown)}") return resolved def build_toolkit( *, session: AsyncSession, owner_id: UUID, groups: set[ToolGroup] | None = None, enabled_tool_names: set[str] | None = None, enable_hitl: bool | None = None, ): toolkit = Toolkit() enabled_names = _resolve_enabled_tools( groups=groups, enabled_tool_names=enabled_tool_names, ) preset_kwargs = cast( dict[str, JSONSerializableObject], { "session": session, "owner_id": owner_id, }, ) for tool_name in sorted(enabled_names): tool_func = TOOL_FUNCTIONS[tool_name] toolkit.register_tool_function( tool_func, func_name=tool_name, preset_kwargs=preset_kwargs, ) approval_enabled = enable_hitl if enable_hitl is not None else True if approval_enabled: register_tool_middlewares(toolkit=toolkit, config_by_name=TOOL_CONFIGS) return toolkit def build_stage_toolkit( *, agent_type: AgentType, session: AsyncSession, owner_id: UUID, enabled_tool_names: set[str] | None = None, enable_hitl: bool | None = None, ): groups = AGENT_TYPE_TO_GROUPS.get(agent_type) if groups is None: raise ValueError(f"unknown agent_type: {agent_type}") stage_enabled_names = resolve_tool_names_by_groups(set(groups)) selected_names = ( stage_enabled_names if enabled_tool_names is None else stage_enabled_names | set(enabled_tool_names) ) return build_toolkit( session=session, owner_id=owner_id, enabled_tool_names=selected_names, enable_hitl=enable_hitl, )