211 lines
7.2 KiB
Python
211 lines
7.2 KiB
Python
"""Execution context and shared primitives for the AI tool layer.
|
|
|
|
The tool layer is deliberately transport-agnostic: the same tool functions are
|
|
invoked by the in-app assistant (function calling) and by the MCP server. This
|
|
module holds the context object that carries the caller identity, the
|
|
permission helpers, and the domain error types shared across the layer.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
from dataclasses import dataclass, field
|
|
from enum import Enum
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from backend.auth.middleware import check_vault_access
|
|
from backend.services.errors import ServiceError
|
|
from backend.services.paths import resolve_safe_path as _resolve_service_path
|
|
|
|
logger = logging.getLogger("obsigate.tools")
|
|
|
|
|
|
class ToolMode(str, Enum):
|
|
"""How a tool was invoked."""
|
|
|
|
IN_APP = "in_app"
|
|
MCP = "mcp"
|
|
|
|
|
|
class ToolScope(str, Enum):
|
|
"""Which front a tool is exposed to."""
|
|
|
|
IN_APP = "in_app"
|
|
MCP = "mcp"
|
|
|
|
|
|
class ToolRisk(str, Enum):
|
|
"""Risk level of a tool.
|
|
|
|
``READ`` tools run without confirmation. ``WRITE`` and ``DANGEROUS`` tools
|
|
require an explicit confirmation (two-step ``propose``/``apply`` for MCP,
|
|
an "Apply" card in the in-app UI).
|
|
"""
|
|
|
|
READ = "read"
|
|
WRITE = "write"
|
|
DANGEROUS = "dangerous"
|
|
|
|
|
|
class ToolError(Exception):
|
|
"""Base class for tool execution errors.
|
|
|
|
Carries a stable, machine-readable ``code`` so callers (agent loop, MCP
|
|
server) can react without parsing the message.
|
|
"""
|
|
|
|
code = "tool_error"
|
|
|
|
def __init__(self, message: str, *, code: str | None = None, details: dict[str, Any] | None = None):
|
|
super().__init__(message)
|
|
self.message = message
|
|
if code:
|
|
self.code = code
|
|
self.details = details or {}
|
|
|
|
def to_dict(self) -> dict[str, Any]:
|
|
return {"ok": False, "error": {"code": self.code, "message": self.message, "details": self.details}}
|
|
|
|
|
|
class ToolNotFoundError(ToolError):
|
|
code = "not_found"
|
|
|
|
|
|
class ToolPermissionError(ToolError):
|
|
code = "permission_denied"
|
|
|
|
|
|
class ToolValidationError(ToolError):
|
|
code = "invalid_arguments"
|
|
|
|
|
|
class ToolConfirmationRequired(ToolError):
|
|
"""Raised when a mutating tool is invoked without confirmation.
|
|
|
|
This is the hook for the two-step ``propose``/``apply`` mechanism: the
|
|
``propose`` phase surfaces this payload to the user, and the ``apply``
|
|
phase re-invokes the tool with ``confirm=True``.
|
|
"""
|
|
|
|
code = "confirmation_required"
|
|
|
|
def __init__(self, tool: str, arguments: dict[str, Any], message: str | None = None):
|
|
super().__init__(message or f"Confirmation required for '{tool}'", code="confirmation_required")
|
|
self.tool = tool
|
|
self.arguments = arguments
|
|
|
|
def to_dict(self) -> dict[str, Any]:
|
|
payload = super().to_dict()
|
|
payload["error"]["tool"] = self.tool
|
|
payload["error"]["arguments"] = self.arguments
|
|
return payload
|
|
|
|
|
|
class ToolRateLimitError(ToolError):
|
|
"""Raised when an identity exceeds its tool-call rate limit (Phase F)."""
|
|
|
|
code = "rate_limited"
|
|
|
|
def __init__(self, tool: str, retry_after: int = 0, message: str | None = None):
|
|
super().__init__(
|
|
message or f"Rate limit exceeded for tool '{tool}'",
|
|
code="rate_limited",
|
|
details={"tool": tool, "retry_after": retry_after},
|
|
)
|
|
self.tool = tool
|
|
self.retry_after = retry_after
|
|
|
|
|
|
@dataclass
|
|
class ToolContext:
|
|
"""Carries the caller identity and execution options for a tool call.
|
|
|
|
Attributes:
|
|
user: Authenticated user dict (as returned by ``get_current_user``).
|
|
mode: Whether the call originates from the in-app assistant or MCP.
|
|
confirmed: True once a mutating action has been approved.
|
|
ip: Optional client IP for auditing.
|
|
audit_enabled: Set to False to skip audit logging (tests, dry runs).
|
|
metadata: Free-form caller metadata (conversation id, client name…).
|
|
"""
|
|
|
|
user: dict[str, Any]
|
|
mode: ToolMode = ToolMode.IN_APP
|
|
confirmed: bool = False
|
|
ip: str | None = None
|
|
audit_enabled: bool = True
|
|
metadata: dict[str, Any] = field(default_factory=dict)
|
|
|
|
@property
|
|
def username(self) -> str:
|
|
return self.user.get("username", "unknown")
|
|
|
|
@property
|
|
def rate_limit_identity(self) -> str:
|
|
"""Identity used by the per-token/per-tool rate limiter.
|
|
|
|
Prefers the JWT id (``_token_jti``, attached by the auth middleware),
|
|
then an explicit ``token_id`` in :attr:`metadata`, then the user id and
|
|
finally the username. This yields per-token limiting when a token id is
|
|
available and per-account limiting otherwise.
|
|
"""
|
|
token_id = self.metadata.get("token_id") or self.user.get("_token_jti")
|
|
if token_id:
|
|
return f"token:{token_id}"
|
|
user_id = self.user.get("id") or self.user.get("username")
|
|
return f"user:{user_id or 'anonymous'}"
|
|
|
|
def has_vault_access(self, vault: str) -> bool:
|
|
"""Return True if the caller may access *vault*."""
|
|
return bool(vault) and check_vault_access(vault, self.user)
|
|
|
|
def require_vault_access(self, vault: str) -> None:
|
|
"""Raise :class:`ToolPermissionError` unless the caller can access *vault*."""
|
|
if not self.has_vault_access(vault):
|
|
raise ToolPermissionError(
|
|
f"Access denied to vault '{vault}'",
|
|
code="vault_access_denied",
|
|
details={"vault": vault},
|
|
)
|
|
|
|
def destructive_tools_enabled(self, vault: str) -> bool:
|
|
"""Return whether destructive tools are allowed for *vault*.
|
|
|
|
Controlled by the per-vault ``aiDestructiveTools`` setting (default:
|
|
enabled). Disabling it blocks delete/rename/move/find-replace tools
|
|
while keeping create/edit/append available.
|
|
"""
|
|
from backend.vault_settings import get_vault_setting
|
|
|
|
settings = get_vault_setting(vault) or {}
|
|
return bool(settings.get("aiDestructiveTools", True))
|
|
|
|
def require_destructive_allowed(self, vault: str) -> None:
|
|
"""Raise :class:`ToolPermissionError` if destructive tools are disabled."""
|
|
if not self.destructive_tools_enabled(vault):
|
|
raise ToolPermissionError(
|
|
f"Destructive tools are disabled for vault '{vault}'",
|
|
code="destructive_tools_disabled",
|
|
details={"vault": vault},
|
|
)
|
|
|
|
|
|
def resolve_safe_path(vault_root: Path, relative_path: str | None) -> Path:
|
|
"""Resolve a vault-relative path, rejecting traversal outside the vault.
|
|
|
|
Delegates to the shared :func:`backend.services.paths.resolve_safe_path`
|
|
and maps its :class:`ServiceError` to tool domain errors so the tool layer
|
|
stays transport-agnostic.
|
|
|
|
Raises:
|
|
ToolPermissionError: When the resolved path escapes the vault root.
|
|
ToolError: When the path cannot be resolved.
|
|
"""
|
|
try:
|
|
return _resolve_service_path(vault_root, relative_path)
|
|
except ServiceError as e:
|
|
if e.code in ("path_outside_vault", "permission_denied"):
|
|
raise ToolPermissionError(e.message, code=e.code, details=e.details) from e
|
|
raise ToolError(e.message, code=e.code, details=e.details) from e
|