122 lines
3.5 KiB
Python
122 lines
3.5 KiB
Python
|
|
"""Global tool registry for AI Assistant plugin tools.
|
||
|
|
|
||
|
|
Plugins can register tools that AI agents can call during chat sessions.
|
||
|
|
Each tool declares a name, description, JSON schema for parameters,
|
||
|
|
and an async handler. Tools can optionally require specific RBAC permissions.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import logging
|
||
|
|
from dataclasses import dataclass, field
|
||
|
|
from typing import Any, Awaitable, Callable, Protocol
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
|
||
|
|
class ToolHandler(Protocol):
|
||
|
|
async def __call__(
|
||
|
|
self,
|
||
|
|
arguments: dict[str, Any],
|
||
|
|
context: dict[str, Any],
|
||
|
|
) -> str: ...
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass
|
||
|
|
class AITool:
|
||
|
|
"""Represents a tool that an AI agent can call."""
|
||
|
|
|
||
|
|
name: str
|
||
|
|
description: str
|
||
|
|
parameters: dict[str, Any] # JSON Schema for parameters
|
||
|
|
handler: ToolHandler
|
||
|
|
plugin_name: str = ""
|
||
|
|
required_permission: str | None = None # e.g. "mail:send"
|
||
|
|
category: str = "general"
|
||
|
|
|
||
|
|
def to_openai_schema(self) -> dict[str, Any]:
|
||
|
|
"""Convert to OpenAI function-calling tool schema."""
|
||
|
|
return {
|
||
|
|
"type": "function",
|
||
|
|
"function": {
|
||
|
|
"name": self.name,
|
||
|
|
"description": self.description,
|
||
|
|
"parameters": self.parameters,
|
||
|
|
},
|
||
|
|
}
|
||
|
|
|
||
|
|
|
||
|
|
class ToolRegistry:
|
||
|
|
"""Singleton registry for AI tools."""
|
||
|
|
|
||
|
|
_instance: ToolRegistry | None = None
|
||
|
|
|
||
|
|
def __new__(cls) -> ToolRegistry:
|
||
|
|
if cls._instance is None:
|
||
|
|
cls._instance = super().__new__(cls)
|
||
|
|
cls._instance._tools: dict[str, AITool] = {}
|
||
|
|
return cls._instance
|
||
|
|
|
||
|
|
def register(
|
||
|
|
self,
|
||
|
|
name: str,
|
||
|
|
description: str,
|
||
|
|
parameters: dict[str, Any],
|
||
|
|
handler: ToolHandler,
|
||
|
|
plugin_name: str = "",
|
||
|
|
required_permission: str | None = None,
|
||
|
|
category: str = "general",
|
||
|
|
) -> None:
|
||
|
|
"""Register a tool."""
|
||
|
|
tool = AITool(
|
||
|
|
name=name,
|
||
|
|
description=description,
|
||
|
|
parameters=parameters,
|
||
|
|
handler=handler,
|
||
|
|
plugin_name=plugin_name,
|
||
|
|
required_permission=required_permission,
|
||
|
|
category=category,
|
||
|
|
)
|
||
|
|
self._tools[name] = tool
|
||
|
|
logger.info("AI tool registered: %s (plugin=%s)", name, plugin_name)
|
||
|
|
|
||
|
|
def unregister(self, name: str) -> None:
|
||
|
|
"""Unregister a tool by name."""
|
||
|
|
self._tools.pop(name, None)
|
||
|
|
|
||
|
|
def unregister_plugin(self, plugin_name: str) -> None:
|
||
|
|
"""Unregister all tools from a plugin."""
|
||
|
|
to_remove = [
|
||
|
|
name for name, tool in self._tools.items() if tool.plugin_name == plugin_name
|
||
|
|
]
|
||
|
|
for name in to_remove:
|
||
|
|
self._tools.pop(name, None)
|
||
|
|
|
||
|
|
def get(self, name: str) -> AITool | None:
|
||
|
|
return self._tools.get(name)
|
||
|
|
|
||
|
|
def get_all(self) -> list[AITool]:
|
||
|
|
return list(self._tools.values())
|
||
|
|
|
||
|
|
def get_by_names(self, names: list[str]) -> list[AITool]:
|
||
|
|
return [self._tools[name] for name in names if name in self._tools]
|
||
|
|
|
||
|
|
def list_for_api(self) -> list[dict[str, Any]]:
|
||
|
|
"""Return tool list for API response."""
|
||
|
|
return [
|
||
|
|
{
|
||
|
|
"name": tool.name,
|
||
|
|
"description": tool.description,
|
||
|
|
"parameters": tool.parameters,
|
||
|
|
"plugin_name": tool.plugin_name,
|
||
|
|
"required_permission": tool.required_permission,
|
||
|
|
"category": tool.category,
|
||
|
|
}
|
||
|
|
for tool in self._tools.values()
|
||
|
|
]
|
||
|
|
|
||
|
|
|
||
|
|
def get_tool_registry() -> ToolRegistry:
|
||
|
|
"""Get the global tool registry singleton."""
|
||
|
|
return ToolRegistry()
|