CI / lint (push) Successful in 57s
CI / security (push) Successful in 39s
CI / test (push) Successful in 1m12s
CI / build (push) Successful in 36s
CI / e2e (push) Successful in 10m22s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
Le .gitignore exclut _*.py, donc backend/tools/__init__.py n'etait pas versionne et la CI echouait a l'import (unknown location). Remplacement par une facade api.py qui enregistre les outils et reexporte l'API publique. Tests adaptes.
281 lines
11 KiB
Python
281 lines
11 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 = {"list_vaults", "list_directory", "read_file", "search_fulltext", "list_tags"}
|
|
|
|
|
|
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>"
|