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

120 lines
3.2 KiB
Python
Raw Normal View History

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