120 lines
3.2 KiB
Python
120 lines
3.2 KiB
Python
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,
|
|
)
|