# 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"] == []