Files
ObsiGate/tests/test_tools.py
T
bruno 778fa65b4c fix: assistant IA toujours en mode agent, retrait du bouton toggle #187
- Bouton « mode agent » du panneau supprimé : l'assistant est toujours
  agent (toute requête texte part sur /api/ai/bookslm/agent, les images
  restent sur /chat multimodal). Les mutations gardent la confirmation
  two-step. Deep Research et les quick actions `agent: true` basculaient
  déjà le mode en silence : le toggle ne protégeait plus rien.
- /agent résout le provider comme /chat (req.provider brut transmis à
  l'adapter pouvait désigner un fournisseur indisponible et retomber
  silencieusement sur un autre que l'étiquette SSE affichée).
- Schémas des outils mis en cache par (scope, taille du registre) avec
  copies fraîches par appelant (~27 ms de pydantic économisées par
  requête agent/MCP).
- Aide in-app réécrite, i18n FR/EN (retrait de ai.agent_mode_*),
  tests frontend adaptés (ai.test.mjs 100/100) + non-régression pytest
  (provider résolu, cache sûr).
2026-10-07 08:16:56 -04:00

455 lines
18 KiB
Python

# tests/test_tools.py — Unit tests for the AI tool layer (Phase 0)
"""Tests for backend.tools: registry, context, permissions, execution, audit.
The index-dependent tests reuse the ``client`` fixture from conftest, which
builds the in-memory index for a ``TestVault``.
"""
import json
from pathlib import Path
import pytest
from backend.tools.api import (
ToolConfirmationRequired,
ToolContext,
ToolError,
ToolNotFoundError,
ToolPermissionError,
ToolRisk,
ToolScope,
ToolValidationError,
call_tool,
get_tool,
get_tool_schemas,
list_tools,
resolve_safe_path,
)
from backend.tools.registry import ToolSpec
from backend.tools.schemas import ListVaultsInput
BUILTIN_TOOLS = {
# Phase 0
"list_vaults",
"list_directory",
"read_file",
"search_fulltext",
"list_tags",
# Phase C — read/search catalogue
"list_all_files",
"read_file_raw",
"get_backlinks",
"list_backups",
"diff_backup",
"get_graph",
"search_advanced",
"search_paths",
"suggest_tags",
"list_recent",
# Phase D — mutations
"create_file",
"create_directory",
"edit_file",
"append_to_file",
"rename_file",
"rename_directory",
"move_path",
"replace_in_files",
"delete_file",
"delete_directory",
"restore_backup",
}
def _ctx(vaults=None, **kwargs) -> ToolContext:
user = {"username": "tester", "role": "admin", "vaults": vaults or ["*"]}
kwargs.setdefault("audit_enabled", False)
return ToolContext(user=user, **kwargs)
# ═══════════════════════════════════════════════════════════════════
# Registry
# ═══════════════════════════════════════════════════════════════════
class TestRegistry:
def test_builtin_tools_registered(self):
names = {s.name for s in list_tools()}
assert BUILTIN_TOOLS <= names
def test_get_tool_returns_spec(self):
spec = get_tool("read_file")
assert spec is not None
assert spec.risk is ToolRisk.READ
assert spec.requires_vault is True
def test_get_tool_unknown_returns_none(self):
assert get_tool("does_not_exist") is None
def test_schemas_have_openai_shape(self):
schemas = {s["function"]["name"]: s for s in get_tool_schemas()}
read = schemas["read_file"]
assert read["type"] == "function"
props = read["function"]["parameters"]["properties"]
assert "vault" in props and "path" in props
def test_schemas_cached_but_caller_safe(self):
"""#187: repeated calls hit the cache and a caller mutating its copy
(e.g. stripping keys) must never corrupt the next request."""
a = get_tool_schemas()
b = get_tool_schemas()
assert a == b
assert all(x is not y for x, y in zip(a, b)), "top-level dicts must be fresh copies"
a[0]["poisoned"] = True
assert "poisoned" not in get_tool_schemas()[0]
def test_duplicate_tool_name_raises(self):
from backend.tools.registry import tool
with pytest.raises(ValueError):
@tool(name="list_vaults", description="dup", input_model=ListVaultsInput)
def _dup(ctx, params): # pragma: no cover - never registered
return None
def test_scope_filter(self, monkeypatch):
from backend.tools import registry
spec = ToolSpec(
name="_mcp_only",
description="mcp only",
input_model=ListVaultsInput,
handler=lambda ctx, params: None,
scopes=(ToolScope.MCP,),
)
monkeypatch.setitem(registry._REGISTRY, "_mcp_only", spec)
mcp_names = {s.name for s in list_tools(scope=ToolScope.MCP)}
in_app_names = {s.name for s in list_tools(scope=ToolScope.IN_APP)}
assert "_mcp_only" in mcp_names
assert "_mcp_only" not in in_app_names
# ═══════════════════════════════════════════════════════════════════
# Context & path safety
# ═══════════════════════════════════════════════════════════════════
class TestContext:
def test_wildcard_access(self):
ctx = ToolContext(user={"vaults": ["*"]})
assert ctx.has_vault_access("Anything")
def test_specific_access(self):
ctx = ToolContext(user={"vaults": ["V1"]})
assert ctx.has_vault_access("V1")
assert not ctx.has_vault_access("V2")
def test_require_vault_access_denied(self):
ctx = ToolContext(user={"vaults": ["V1"]})
with pytest.raises(ToolPermissionError):
ctx.require_vault_access("V2")
def test_token_vaults_take_precedence(self):
ctx = ToolContext(user={"vaults": ["*"], "_token_vaults": ["V1"]})
assert ctx.has_vault_access("V1")
assert not ctx.has_vault_access("V2")
def test_resolve_safe_path_inside(self, tmp_path):
resolved = resolve_safe_path(tmp_path, "sub/note.md")
assert str(resolved).startswith(str(tmp_path.resolve()))
def test_resolve_safe_path_traversal(self, tmp_path):
with pytest.raises(ToolPermissionError):
resolve_safe_path(tmp_path, "../evil.md")
# ═══════════════════════════════════════════════════════════════════
# Execution (index-backed)
# ═══════════════════════════════════════════════════════════════════
class TestCallTool:
def test_unknown_tool(self, client):
with pytest.raises(ToolNotFoundError):
call_tool("nope", _ctx(), {})
def test_invalid_arguments(self, client):
with pytest.raises(ToolValidationError):
call_tool("read_file", _ctx(), {"vault": "TestVault"})
def test_list_vaults(self, client):
result = call_tool("list_vaults", _ctx(), {})
assert result.ok
names = {v["name"] for v in result.data}
assert "TestVault" in names
def test_list_vaults_filters_by_permission(self, client):
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
result = call_tool("list_vaults", ctx, {})
assert result.data == []
def test_list_directory(self, client):
result = call_tool("list_directory", _ctx(), {"vault": "TestVault"})
names = {e["name"] for e in result.data}
assert "note1.md" in names
assert "Projets" in names
projets = next(e for e in result.data if e["name"] == "Projets")
assert projets["type"] == "directory"
def test_list_directory_not_found(self, client):
with pytest.raises(ToolNotFoundError):
call_tool("list_directory", _ctx(), {"vault": "TestVault", "path": "nope"})
def test_read_file(self, client):
result = call_tool("read_file", _ctx(), {"vault": "TestVault", "path": "note1.md"})
assert result.ok
assert "Python" in result.data["content"]
assert result.data["path"] == "note1.md"
def test_read_file_not_found(self, client):
with pytest.raises(ToolNotFoundError):
call_tool("read_file", _ctx(), {"vault": "TestVault", "path": "missing.md"})
def test_read_file_too_large(self, client):
from backend.indexer import get_vault_data
from backend.tools.service import TOOL_MAX_READ_BYTES
vault_path = Path(get_vault_data("TestVault")["path"])
big = vault_path / "big.txt"
big.write_text("x" * (TOOL_MAX_READ_BYTES + 10), encoding="utf-8")
try:
with pytest.raises(ToolError) as exc:
call_tool("read_file", _ctx(), {"vault": "TestVault", "path": "big.txt"})
assert exc.value.code == "file_too_large"
finally:
big.unlink()
def test_read_file_redacts_secrets(self, client):
from backend.indexer import get_vault_data
vault_path = Path(get_vault_data("TestVault")["path"])
secret = vault_path / "secret.md"
fake_jwt = "eyJ" + "a" * 30 + "." + "b" * 30 + "." + "c" * 30
secret.write_text(f"token: {fake_jwt}\n", encoding="utf-8")
try:
result = call_tool("read_file", _ctx(), {"vault": "TestVault", "path": "secret.md"})
assert "[JWT MASQUÉ]" in result.data["content"]
finally:
secret.unlink()
def test_search_fulltext(self, client):
result = call_tool("search_fulltext", _ctx(), {"q": "Python"})
assert result.ok
assert len(result.data) >= 1
assert all(r["vault"] == "TestVault" for r in result.data)
def test_list_tags(self, client):
result = call_tool("list_tags", _ctx(), {})
tags = {t["tag"] for t in result.data}
assert "python" in tags
def test_vault_permission_denied(self, client):
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
with pytest.raises(ToolPermissionError):
call_tool("list_directory", ctx, {"vault": "TestVault"})
# ═══════════════════════════════════════════════════════════════════
# Confirmation gating
# ═══════════════════════════════════════════════════════════════════
class TestConfirmation:
def _register_write_tool(self, monkeypatch):
from backend.tools import registry
spec = ToolSpec(
name="_tmp_write",
description="write tool for tests",
input_model=ListVaultsInput,
handler=lambda ctx, params: {"done": True},
risk=ToolRisk.WRITE,
)
monkeypatch.setitem(registry._REGISTRY, "_tmp_write", spec)
def test_write_tool_requires_confirmation(self, client, monkeypatch):
self._register_write_tool(monkeypatch)
with pytest.raises(ToolConfirmationRequired):
call_tool("_tmp_write", _ctx(), {})
def test_write_tool_runs_when_confirmed(self, client, monkeypatch):
self._register_write_tool(monkeypatch)
result = call_tool("_tmp_write", _ctx(), {}, confirm=True)
assert result.ok
assert result.data == {"done": True}
def test_context_confirmed_flag(self, client, monkeypatch):
self._register_write_tool(monkeypatch)
result = call_tool("_tmp_write", _ctx(confirmed=True), {})
assert result.ok
# ═══════════════════════════════════════════════════════════════════
# Audit
# ═══════════════════════════════════════════════════════════════════
class TestToolAudit:
def test_tool_call_is_audited(self, client, monkeypatch, tmp_path):
from backend import audit as audit_mod
log_file = tmp_path / "audit.log"
monkeypatch.setattr(audit_mod, "AUDIT_LOG_FILE", log_file)
ctx = ToolContext(user={"username": "tester", "vaults": ["*"]}, audit_enabled=True)
call_tool("list_vaults", ctx, {})
assert log_file.exists()
entries = [json.loads(line) for line in log_file.read_text(encoding="utf-8").splitlines() if line.strip()]
assert entries[-1]["action"] == "ai_tool_call"
assert entries[-1]["tool"] == "list_vaults"
assert entries[-1]["ok"] is True
def test_sanitize_arguments_masks_content(self):
from backend.tools.audit import _sanitize_arguments
safe = _sanitize_arguments({"vault": "V", "path": "a.md", "content": "x" * 500})
assert safe["vault"] == "V"
assert safe["path"] == "a.md"
assert safe["content"] == "<500 chars>"
# ═══════════════════════════════════════════════════════════════════
# Phase C — read/search catalogue
# ═══════════════════════════════════════════════════════════════════
def _vault_path() -> Path:
from backend.indexer import get_vault_data
return Path(get_vault_data("TestVault")["path"])
def _make_backup(rel_path: str, ts: int, content: str) -> Path:
from backend.services.backups import get_backup_dir
backup_dir = get_backup_dir("TestVault", rel_path)
backup_dir.mkdir(parents=True, exist_ok=True)
backup = backup_dir / f"{Path(rel_path).name}.{ts}.bak"
backup.write_text(content, encoding="utf-8")
return backup
class TestPhaseCTools:
def test_list_all_files(self, client):
result = call_tool("list_all_files", _ctx(), {"vault": "TestVault"})
assert result.ok
paths = {f["path"] for f in result.data["files"]}
assert "note1.md" in paths
assert "Projets/projet.md" in paths
def test_list_all_files_non_recursive(self, client):
result = call_tool(
"list_all_files", _ctx(), {"vault": "TestVault", "recursive": False}
)
paths = {f["path"] for f in result.data["files"]}
assert "note1.md" in paths
assert "Projets/projet.md" not in paths
def test_list_all_files_permission_denied(self, client):
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
with pytest.raises(ToolPermissionError):
call_tool("list_all_files", ctx, {"vault": "TestVault"})
def test_read_file_raw(self, client):
result = call_tool("read_file_raw", _ctx(), {"vault": "TestVault", "path": "note1.md"})
assert result.ok
assert "Python" in result.data["raw"]
assert result.data["path"] == "note1.md"
def test_get_backlinks(self, client):
result = call_tool("get_backlinks", _ctx(), {"vault": "TestVault", "path": "note1.md"})
assert result.ok
assert all(b["vault"] == "TestVault" for b in result.data)
def test_get_backlinks_filters_inaccessible_vaults(self, client):
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
with pytest.raises(ToolPermissionError):
call_tool("get_backlinks", ctx, {"vault": "TestVault", "path": "note1.md"})
def test_list_backups_empty(self, client):
result = call_tool("list_backups", _ctx(), {"vault": "TestVault", "path": "note1.md"})
assert result.ok
assert result.data["backups"] == []
def test_list_backups_and_diff(self, client):
backup = _make_backup("note1.md", 1_700_000_000, "old content\n")
try:
listed = call_tool("list_backups", _ctx(), {"vault": "TestVault", "path": "note1.md"})
timestamps = [b["timestamp"] for b in listed.data["backups"]]
assert 1_700_000_000 in timestamps
diff = call_tool(
"diff_backup",
_ctx(),
{"vault": "TestVault", "path": "note1.md", "version": 1_700_000_000},
)
assert diff.ok
assert diff.data["version"] == 1_700_000_000
assert diff.data["left_content"] == "old content\n"
assert "-old content" in diff.data["diff"]
finally:
backup.unlink(missing_ok=True)
def test_diff_backup_missing_version(self, client):
with pytest.raises(ToolNotFoundError):
call_tool(
"diff_backup",
_ctx(),
{"vault": "TestVault", "path": "note1.md", "version": 42},
)
def test_get_graph(self, client):
result = call_tool("get_graph", _ctx(), {"vault": "TestVault", "scope": "full", "depth": 2})
assert result.ok
node_paths = {n["path"] for n in result.data["nodes"]}
assert "note1.md" in node_paths
assert result.data["scope"] == "full"
def test_get_graph_missing_path(self, client):
with pytest.raises(ToolNotFoundError):
call_tool("get_graph", _ctx(), {"vault": "TestVault", "path": "nope"})
def test_search_advanced(self, client):
result = call_tool("search_advanced", _ctx(), {"q": "python"})
assert result.ok
assert len(result.data) >= 1
assert all(r["vault"] == "TestVault" for r in result.data)
def test_search_paths(self, client):
result = call_tool("search_paths", _ctx(), {"q": "note1"})
assert result.ok
assert any(r["path"] == "note1.md" for r in result.data)
def test_suggest_tags(self, client):
result = call_tool("suggest_tags", _ctx(), {"q": "py"})
assert result.ok
assert any(s["tag"] == "python" for s in result.data)
def test_suggest_tags_specific_vault_denied(self, client):
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
with pytest.raises(ToolPermissionError):
call_tool("suggest_tags", ctx, {"q": "py", "vault": "TestVault"})
def test_list_recent_modified(self, client):
result = call_tool("list_recent", _ctx(), {"mode": "modified"})
assert result.ok
assert result.data["mode"] == "modified"
assert any(f["path"] == "note1.md" for f in result.data["files"])
def test_list_recent_filters_by_vault(self, client):
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
result = call_tool("list_recent", ctx, {"mode": "modified"})
assert result.data["files"] == []