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()
    }