| """Tool base class and registry. |
| |
| Following the same pattern used by the codebase this agent is modelled on, |
| every tool is a small subclass of :class:`ToolBase` that: |
| |
| * declares a unique ``name`` and ``description`` (used by the orchestrator / |
| LLM to decide which tool to call) |
| * exposes a JSON-schema-ish ``input_schema`` describing its arguments |
| * implements :meth:`run` which performs the actual work and returns a |
| dictionary result |
| |
| Tools register themselves automatically on import via the |
| ``@register_tool`` decorator (or by being added to ``ALL_TOOLS`` in |
| ``unity_tools.py``). The :class:`ToolRegistry` keeps a single global map of |
| ``name -> tool_class``. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import inspect |
| import json |
| import logging |
| from dataclasses import dataclass, field |
| from typing import Any, Callable, Dict, List, Optional, Type |
|
|
| from unity_agent.config import Settings |
| from unity_agent.transport.unity_transport import UnityTransport |
|
|
| log = logging.getLogger("unity_agent.tools") |
|
|
|
|
| |
| |
| |
| @dataclass |
| class ToolResult: |
| """Structured return value from :meth:`ToolBase.run`.""" |
|
|
| ok: bool |
| tool: str |
| summary: str |
| files: List[str] = field(default_factory=list) |
| data: Dict[str, Any] = field(default_factory=dict) |
| error: Optional[str] = None |
|
|
| def to_dict(self) -> dict: |
| return { |
| "ok": self.ok, |
| "tool": self.tool, |
| "summary": self.summary, |
| "files": self.files, |
| "data": self.data, |
| "error": self.error, |
| } |
|
|
|
|
| class ToolBase: |
| """Base class for every Unity agent tool. |
| |
| Subclasses MUST override: |
| |
| * ``name`` (str) -- unique tool name (snake_case) |
| * ``description`` (str) -- natural language description for the LLM |
| * ``input_schema`` (dict)-- JSON-schema describing parameters |
| * :meth:`run` -- the actual implementation |
| """ |
|
|
| name: str = "base_tool" |
| description: str = "Override me." |
| input_schema: Dict[str, Any] = {"type": "object", "properties": {}} |
|
|
| def __init__(self, settings: Settings, transport: UnityTransport) -> None: |
| self.settings = settings |
| self.transport = transport |
|
|
| |
| |
| |
| def _ensure_project(self, project_name: Optional[str] = None) -> None: |
| """Open a project if one is not already open.""" |
| if project_name and ( |
| self.transport.project_name is None |
| or self.transport.project_name != project_name |
| ): |
| self.transport.open_project(project_name) |
| |
| if self.transport._project_root is None: |
| self.transport.open_project(self.settings.product_name) |
|
|
| def result( |
| self, |
| ok: bool, |
| summary: str, |
| files: Optional[List[str]] = None, |
| data: Optional[Dict[str, Any]] = None, |
| error: Optional[str] = None, |
| ) -> ToolResult: |
| return ToolResult( |
| ok=ok, |
| tool=self.name, |
| summary=summary, |
| files=files or [], |
| data=data or {}, |
| error=error, |
| ) |
|
|
| |
| |
| |
| def run(self, **kwargs: Any) -> ToolResult: |
| raise NotImplementedError(f"{self.__class__.__name__}.run() not implemented") |
|
|
| def schema(self) -> Dict[str, Any]: |
| return { |
| "name": self.name, |
| "description": self.description, |
| "input_schema": self.input_schema, |
| } |
|
|
|
|
| |
| |
| |
| class ToolRegistry: |
| """Global registry mapping tool names to tool classes.""" |
|
|
| def __init__(self) -> None: |
| self._tools: Dict[str, Type[ToolBase]] = {} |
|
|
| def register(self, tool_cls: Type[ToolBase]) -> Type[ToolBase]: |
| if not issubclass(tool_cls, ToolBase): |
| raise TypeError(f"{tool_cls} must subclass ToolBase") |
| name = tool_cls.name |
| if name in self._tools: |
| log.debug("Overwriting already-registered tool %s", name) |
| self._tools[name] = tool_cls |
| return tool_cls |
|
|
| def get(self, name: str) -> Optional[Type[ToolBase]]: |
| return self._tools.get(name) |
|
|
| def all_names(self) -> List[str]: |
| return sorted(self._tools.keys()) |
|
|
| def all_schemas(self) -> List[Dict[str, Any]]: |
| return [cls(None, None).schema() for cls in self._tools.values()] |
|
|
| def instantiate( |
| self, name: str, settings: Settings, transport: UnityTransport |
| ) -> ToolBase: |
| cls = self.get(name) |
| if cls is None: |
| raise KeyError(f"Unknown tool: {name}") |
| return cls(settings, transport) |
|
|
|
|
| |
| _GLOBAL_REGISTRY = ToolRegistry() |
|
|
|
|
| def get_registry() -> ToolRegistry: |
| return _GLOBAL_REGISTRY |
|
|
|
|
| def register_tool(tool_cls: Type[ToolBase]) -> Type[ToolBase]: |
| """Class decorator that registers a tool with the global registry.""" |
| return _GLOBAL_REGISTRY.register(tool_cls) |
|
|
|
|
| def instantiate_all(settings: Settings, transport: UnityTransport) -> Dict[str, ToolBase]: |
| """Instantiate every registered tool with the given settings + transport.""" |
| return { |
| name: cls(settings, transport) |
| for name, cls in _GLOBAL_REGISTRY._tools.items() |
| } |
|
|