CI / lint (push) Successful in 59s
CI / security (push) Successful in 45s
CI / test (push) Successful in 1m22s
CI / build (push) Successful in 37s
CI / e2e (push) Successful in 10m37s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
445 lines
18 KiB
Python
445 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_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"] == []
|