ryzerrr's picture
Upload folder using huggingface_hub
40c0886 verified
Raw
History Blame Contribute Delete
5.88 kB
"""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()
}