238 lines
7.9 KiB
Python
238 lines
7.9 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,
|
|
ToolRateLimitError,
|
|
ToolRisk,
|
|
ToolScope,
|
|
ToolValidationError,
|
|
)
|
|
from backend.tools.ratelimit import check_and_record
|
|
from backend.tools.redaction import redact_payload
|
|
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
|
|
|
|
allowed, retry_after = check_and_record(ctx.rate_limit_identity, spec.name)
|
|
if not allowed:
|
|
rate_error = ToolRateLimitError(spec.name, retry_after)
|
|
_audit(ctx, spec, arguments, ok=False, error=rate_error.code)
|
|
raise rate_error
|
|
|
|
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.risk == ToolRisk.DANGEROUS and vault and vault != "all":
|
|
try:
|
|
ctx.require_destructive_allowed(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)
|
|
# Never forward raw secrets to the model (defense in depth: some tools —
|
|
# diffs, search snippets — return content that was not pre-redacted).
|
|
return ToolResult(ok=True, data=redact_payload(data))
|