Files
ObsiGate/backend/tools/registry.py
T
bruno 4c4e415975
CI / lint (push) Successful in 58s
CI / security (push) Successful in 40s
CI / test (push) Successful in 1m15s
CI / build (push) Successful in 37s
CI / e2e (push) Successful in 10m15s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
feat(ai): services partages A2 + SSE streaming B4 + confirmations UI B5 (#79)
2026-09-11 17:06:40 -04:00

220 lines
7.1 KiB
Python

"""Tool registry — declarative registration and uniform execution.
Tools are registered with the :func:`tool` decorator. Each tool declares its
input model (used both for JSON Schema exposure and argument validation), its
risk level, and the fronts (in-app / MCP) it is exposed to.
Execution goes through :func:`call_tool`, which centralizes argument
validation, vault permission checks, confirmation gating, and audit logging.
"""
from __future__ import annotations
import logging
import time
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
from pydantic import BaseModel, ValidationError
from backend.services.errors import ServiceError
from backend.tools.audit import log_tool_call
from backend.tools.context import (
ToolConfirmationRequired,
ToolContext,
ToolError,
ToolNotFoundError,
ToolPermissionError,
ToolRisk,
ToolScope,
ToolValidationError,
)
from backend.tools.schemas import ToolResult
logger = logging.getLogger("obsigate.tools.registry")
@dataclass(frozen=True)
class ToolSpec:
"""Static description of a registered tool."""
name: str
description: str
input_model: type[BaseModel]
handler: Callable[[ToolContext, Any], Any]
risk: ToolRisk = ToolRisk.READ
scopes: tuple[ToolScope, ...] = (ToolScope.IN_APP, ToolScope.MCP)
requires_vault: bool = False
@property
def requires_confirmation(self) -> bool:
"""Mutating tools always require an explicit confirmation."""
return self.risk != ToolRisk.READ
def parameters_schema(self) -> dict[str, Any]:
"""JSON Schema of the tool arguments (for LLM function calling)."""
schema = self.input_model.model_json_schema()
schema.pop("title", None)
return schema
def openai_schema(self) -> dict[str, Any]:
"""OpenAI-compatible function-calling schema."""
return {
"type": "function",
"function": {
"name": self.name,
"description": self.description,
"parameters": self.parameters_schema(),
},
}
_REGISTRY: dict[str, ToolSpec] = {}
def tool(
*,
name: str,
description: str,
input_model: type[BaseModel],
risk: ToolRisk = ToolRisk.READ,
scopes: tuple[ToolScope, ...] = (ToolScope.IN_APP, ToolScope.MCP),
requires_vault: bool = False,
) -> Callable[[Callable[[ToolContext, Any], Any]], Callable[[ToolContext, Any], Any]]:
"""Register a tool. Returns the original handler unchanged."""
def decorator(func: Callable[[ToolContext, Any], Any]) -> Callable[[ToolContext, Any], Any]:
if name in _REGISTRY:
raise ValueError(f"Duplicate tool name: {name}")
_REGISTRY[name] = ToolSpec(
name=name,
description=description,
input_model=input_model,
handler=func,
risk=risk,
scopes=scopes,
requires_vault=requires_vault,
)
return func
return decorator
def get_tool(name: str) -> ToolSpec | None:
"""Return the spec for *name*, or ``None``."""
return _REGISTRY.get(name)
def list_tools(*, scope: ToolScope | None = None) -> list[ToolSpec]:
"""Return registered tools, optionally filtered by scope."""
specs = list(_REGISTRY.values())
if scope is not None:
specs = [s for s in specs if scope in s.scopes]
return specs
def get_tool_schemas(*, scope: ToolScope | None = None) -> list[dict[str, Any]]:
"""Return OpenAI-compatible schemas for registered tools."""
return [spec.openai_schema() for spec in list_tools(scope=scope)]
def _audit(ctx: ToolContext, spec: ToolSpec, arguments: dict[str, Any], *, ok: bool, error: str | None = None) -> None:
if not ctx.audit_enabled:
return
try:
log_tool_call(
username=ctx.username,
mode=ctx.mode.value,
tool=spec.name,
arguments=arguments,
ok=ok,
vault=arguments.get("vault"),
ip=ctx.ip,
error=error,
)
except Exception as e:
logger.debug(f"Tool audit failed for '{spec.name}': {e}")
def _map_service_error(e: ServiceError) -> ToolError:
"""Map a shared-layer :class:`ServiceError` to a tool domain error."""
if e.code == "not_found":
return ToolNotFoundError(e.message, details=e.details)
if e.code in ("permission_denied", "path_outside_vault", "vault_access_denied"):
return ToolPermissionError(e.message, code=e.code, details=e.details)
if e.code == "invalid_arguments":
return ToolValidationError(e.message, details=e.details)
return ToolError(e.message, code=e.code, details=e.details)
def call_tool(
name: str,
ctx: ToolContext,
arguments: dict[str, Any] | None = None,
*,
confirm: bool = False,
) -> ToolResult:
"""Validate, authorize, execute and audit a tool call.
Args:
name: Registered tool name.
ctx: Caller context (identity, mode, confirmation state).
arguments: Raw arguments to validate against the tool input model.
confirm: Explicit one-shot confirmation (two-step ``apply`` phase).
Raises:
ToolNotFoundError: Unknown tool.
ToolValidationError: Arguments fail validation.
ToolPermissionError: Vault access denied or path escapes the vault.
ToolConfirmationRequired: Mutating tool invoked without confirmation.
ToolError: Any other execution error.
"""
spec = _REGISTRY.get(name)
if spec is None:
raise ToolNotFoundError(f"Unknown tool: {name}", details={"tool": name})
arguments = arguments or {}
try:
params = spec.input_model.model_validate(arguments)
except ValidationError as e:
error = ToolValidationError(
f"Invalid arguments for '{name}'",
details={"errors": e.errors(include_url=False)},
)
_audit(ctx, spec, arguments, ok=False, error=error.code)
raise error from e
vault = getattr(params, "vault", None)
if spec.requires_vault and vault:
try:
ctx.require_vault_access(vault)
except ToolPermissionError as e:
_audit(ctx, spec, arguments, ok=False, error=e.code)
raise
if spec.requires_confirmation and not (confirm or ctx.confirmed):
raise ToolConfirmationRequired(spec.name, arguments)
started = time.perf_counter()
try:
data = spec.handler(ctx, params)
except ToolError as e:
_audit(ctx, spec, arguments, ok=False, error=e.code)
raise
except ServiceError as e:
mapped = _map_service_error(e)
_audit(ctx, spec, arguments, ok=False, error=mapped.code)
raise mapped from e
except Exception as e:
logger.error(f"Tool '{name}' failed: {e}")
exec_error = ToolError(f"Tool '{name}' failed: {e}", code="tool_execution_error")
_audit(ctx, spec, arguments, ok=False, error=exec_error.code)
raise exec_error from e
duration_ms = round((time.perf_counter() - started) * 1000, 2)
logger.debug(f"Tool '{name}' ok in {duration_ms}ms")
_audit(ctx, spec, arguments, ok=True)
return ToolResult(ok=True, data=data)