File size: 5,879 Bytes
40c0886 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | """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")
# --------------------------------------------------------------------------- #
# Tool base
# --------------------------------------------------------------------------- #
@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
# ------------------------------------------------------------------ #
# Helpers
# ------------------------------------------------------------------ #
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 no project is open at all, default to settings.product_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,
)
# ------------------------------------------------------------------ #
# Public API
# ------------------------------------------------------------------ #
def run(self, **kwargs: Any) -> ToolResult: # noqa: D401 - simple
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,
}
# --------------------------------------------------------------------------- #
# Registry
# --------------------------------------------------------------------------- #
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()] # type: ignore[arg-type]
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)
# Single shared registry.
_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()
}
|