feat(ai): phase 0 couche d'outils partagee (backend/tools) + tests
CI / lint (push) Successful in 57s
CI / security (push) Successful in 39s
CI / test (push) Failing after 41s
CI / build (push) Skipped
CI / e2e (push) Skipped
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s

- backend/tools/: ToolContext, registry @tool + schemas JSON, audit ai_tool_call
- Services lecture/recherche: list_vaults, list_directory, read_file, search_fulltext, list_tags
- Permissions check_vault_access + resolve_safe_path, confirmation gating (two-step)
- Redaction des secrets, limites de taille, audit JSONL (args sensibles resumes)
- tests/test_tools.py: 30 tests (registry, contexte, execution, confirmation, audit)
- ROADMAP: item #79 phase A livree (A2 partiel)
This commit is contained in:
2026-09-11 12:03:21 -04:00
parent a3642caa3d
commit 5c1823d6d2
7 changed files with 988 additions and 2 deletions
+54
View File
@@ -0,0 +1,54 @@
"""Audit logging for AI tool calls.
Reuses the application audit log (``data/audit.log``, JSON lines) and adds an
``ai_tool_call`` action. Argument values that may contain sensitive payloads
(file content, prompts) are summarized rather than stored verbatim.
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
from backend.audit import _write_entry
# Argument keys whose values may contain secrets or large payloads.
_SENSITIVE_ARG_KEYS = {"content", "text", "body"}
_MAX_ARG_CHARS = 200
def _sanitize_arguments(arguments: dict[str, Any] | None) -> dict[str, Any]:
"""Return a log-safe view of tool arguments."""
safe: dict[str, Any] = {}
for key, value in (arguments or {}).items():
if key in _SENSITIVE_ARG_KEYS:
safe[key] = f"<{len(str(value))} chars>"
else:
safe[key] = str(value)[:_MAX_ARG_CHARS]
return safe
def log_tool_call(
*,
username: str,
mode: str,
tool: str,
arguments: dict[str, Any] | None = None,
ok: bool = True,
vault: str | None = None,
ip: str | None = None,
error: str | None = None,
) -> None:
"""Append an ``ai_tool_call`` entry to the audit log."""
_write_entry({
"timestamp": datetime.now(timezone.utc).isoformat(),
"action": "ai_tool_call",
"username": username,
"ip": ip or "unknown",
"mode": mode,
"tool": tool,
"vault": vault,
"ok": ok,
"error": error,
"arguments": _sanitize_arguments(arguments),
})
+169
View File
@@ -0,0 +1,169 @@
"""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
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
@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")
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 resolve_safe_path(vault_root: Path, relative_path: str) -> Path:
"""Resolve a vault-relative path, rejecting traversal outside the vault.
Mirrors ``backend.main._resolve_safe_path`` but raises a domain error
instead of ``HTTPException`` so the tool layer stays transport-agnostic.
Raises:
ToolPermissionError: When the resolved path escapes the vault root.
ToolError: When the path cannot be resolved.
"""
full_path = vault_root / (relative_path or "")
try:
resolved = full_path.resolve(strict=False)
root = vault_root.resolve(strict=False)
except Exception as e:
raise ToolError(f"Path resolution error: {e}", code="path_error") from e
try:
resolved.relative_to(root)
except ValueError:
# Case-insensitive fallback for Windows / Docker path casing.
if not str(resolved).lower().startswith(str(root).lower()):
logger.warning(f"Path outside vault: vault={root}, requested={relative_path}")
raise ToolPermissionError(
"Access denied: path outside vault",
code="path_outside_vault",
details={"path": relative_path},
) from None
return resolved
+203
View File
@@ -0,0 +1,203 @@
"""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.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 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 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)
+52
View File
@@ -0,0 +1,52 @@
"""Pydantic input/output models for the AI tool layer.
Input models double as JSON Schemas advertised to LLMs (via
``model_json_schema``) and as validation for arguments received over MCP.
"""
from __future__ import annotations
from typing import Any
from pydantic import BaseModel, Field
class ListVaultsInput(BaseModel):
"""No parameters — lists the vaults the caller can access."""
class ListDirectoryInput(BaseModel):
"""Browse a directory inside a vault."""
vault: str = Field(..., description="Vault name")
path: str = Field("", description="Vault-relative directory path (empty = vault root)")
class ReadFileInput(BaseModel):
"""Read a text file's content from a vault."""
vault: str = Field(..., description="Vault name")
path: str = Field(..., description="Vault-relative file path")
class SearchFulltextInput(BaseModel):
"""Full-text search across one or all accessible vaults."""
q: str = Field(..., min_length=1, description="Search query")
vault: str = Field("all", description="Vault name or 'all'")
tag: str | None = Field(None, description="Optional comma-separated tag filter")
limit: int = Field(50, ge=1, le=200, description="Maximum number of results")
class ListTagsInput(BaseModel):
"""List tags and their occurrence counts."""
vault: str | None = Field(None, description="Vault name or 'all' (default: all accessible)")
class ToolResult(BaseModel):
"""Uniform result returned by :func:`backend.tools.registry.call_tool`."""
ok: bool = True
data: Any = None
error: str | None = None
+166
View File
@@ -0,0 +1,166 @@
"""Built-in tool services (Phase 0 — read/search).
These functions are the single source of truth consumed by both the in-app
assistant and the MCP server. They delegate to the existing core modules
(``backend.indexer``, ``backend.search``) rather than duplicating route logic.
Mutating tools (create/edit/rename/move/delete) are added in later phases.
"""
from __future__ import annotations
import logging
from pathlib import Path
from typing import Any
from backend.tools.context import (
ToolContext,
ToolError,
ToolNotFoundError,
ToolRisk,
resolve_safe_path,
)
from backend.tools.registry import tool
from backend.tools.schemas import (
ListDirectoryInput,
ListTagsInput,
ListVaultsInput,
ReadFileInput,
SearchFulltextInput,
)
logger = logging.getLogger("obsigate.tools.service")
# Maximum file size returned by ``read_file`` (bytes).
TOOL_MAX_READ_BYTES = 200_000
def _vault_data_or_raise(vault: str) -> dict[str, Any]:
from backend.indexer import get_vault_data
data = get_vault_data(vault)
if not data:
raise ToolNotFoundError(f"Vault '{vault}' not found", details={"vault": vault})
return data
@tool(
name="list_vaults",
description="List the vaults the current user is allowed to access.",
input_model=ListVaultsInput,
risk=ToolRisk.READ,
)
def list_vaults(ctx: ToolContext, _params: ListVaultsInput) -> list[dict[str, Any]]:
"""Return accessible vaults with a file count."""
from backend.indexer import get_vault_data, get_vault_names
vaults: list[dict[str, Any]] = []
for name in get_vault_names():
if not ctx.has_vault_access(name):
continue
data = get_vault_data(name) or {}
vaults.append({"name": name, "file_count": len(data.get("files", []))})
return vaults
@tool(
name="list_directory",
description="List files and subdirectories of a directory inside a vault.",
input_model=ListDirectoryInput,
risk=ToolRisk.READ,
requires_vault=True,
)
def list_directory(ctx: ToolContext, params: ListDirectoryInput) -> list[dict[str, Any]]:
"""Return the entries of a vault directory (direct children only)."""
data = _vault_data_or_raise(params.vault)
root = Path(data["path"])
target = resolve_safe_path(root, params.path) if params.path else root.resolve()
if not target.exists() or not target.is_dir():
raise ToolNotFoundError(f"Directory not found: {params.path}", details={"vault": params.vault, "path": params.path})
from backend.vault_settings import get_vault_setting
hide_hidden = (get_vault_setting(params.vault) or {}).get("hideHiddenFiles", False)
items: list[dict[str, Any]] = []
for entry in sorted(target.iterdir(), key=lambda e: (not e.is_dir(), e.name.lower())):
if hide_hidden and entry.name.startswith("."):
continue
rel = str(entry.relative_to(root)).replace("\\", "/")
items.append({
"name": entry.name,
"path": rel,
"type": "directory" if entry.is_dir() else "file",
})
return items
@tool(
name="read_file",
description="Read the text content of a file inside a vault (secrets redacted).",
input_model=ReadFileInput,
risk=ToolRisk.READ,
requires_vault=True,
)
def read_file(ctx: ToolContext, params: ReadFileInput) -> dict[str, Any]:
"""Return the (redacted) content of a vault file."""
data = _vault_data_or_raise(params.vault)
root = Path(data["path"])
target = resolve_safe_path(root, params.path)
if not target.exists() or not target.is_file():
raise ToolNotFoundError(f"File not found: {params.path}", details={"vault": params.vault, "path": params.path})
size = target.stat().st_size
if size > TOOL_MAX_READ_BYTES:
raise ToolError(
f"File too large ({size} bytes > {TOOL_MAX_READ_BYTES})",
code="file_too_large",
details={"vault": params.vault, "path": params.path, "size": size},
)
content = target.read_text(encoding="utf-8", errors="replace")
from backend.secret_redactor import redact_file_content
content = redact_file_content(content, params.path)
return {"vault": params.vault, "path": params.path, "size": size, "content": content}
@tool(
name="search_fulltext",
description="Full-text search across one vault or all accessible vaults.",
input_model=SearchFulltextInput,
risk=ToolRisk.READ,
)
def search_fulltext(ctx: ToolContext, params: SearchFulltextInput) -> list[dict[str, Any]]:
"""Return ranked search results, filtered to accessible vaults."""
from backend.search import search
results = search(params.q, vault_filter=params.vault, tag_filter=params.tag, limit=params.limit)
return [r for r in results if ctx.has_vault_access(r["vault"])]
@tool(
name="list_tags",
description="List tags with their occurrence counts for a vault or all accessible vaults.",
input_model=ListTagsInput,
risk=ToolRisk.READ,
)
def list_tags(ctx: ToolContext, params: ListTagsInput) -> list[dict[str, Any]]:
"""Return tags sorted by descending count."""
from backend.search import get_all_tags
if params.vault and params.vault != "all":
ctx.require_vault_access(params.vault)
return [{"tag": tag, "count": count} for tag, count in get_all_tags(params.vault).items()]
from backend.indexer import get_vault_names
merged: dict[str, int] = {}
for name in get_vault_names():
if not ctx.has_vault_access(name):
continue
for tag, count in get_all_tags(name).items():
merged[tag] = merged.get(tag, 0) + count
return [{"tag": tag, "count": count} for tag, count in sorted(merged.items(), key=lambda x: -x[1])]