Files
ObsiGate/tests/test_ai_e2e.py
T

266 lines
9.7 KiB
Python

# tests/test_ai_e2e.py — End-to-end tests for the AI tool layer (Phase F)
"""End-to-end coverage of the AI assistant, combining the shared tool layer,
the in-app agent loop and the MCP server.
The LLM is scripted (no network) so the tests are deterministic; everything
else (permissions, services, confirmations, redaction, rate limiting, audit) is
the real production code path.
"""
import json
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
from backend.agent.loop import (
STOP_CONFIRMATION_REQUIRED,
STOP_DONE,
STOP_QUOTA_EXCEEDED,
run_agent,
)
from backend.ai_chat import LLMResponse, ToolCall
from backend.tools.api import (
ToolContext,
ToolRateLimitError,
ToolRisk,
call_tool,
)
def _ctx(vaults=None) -> ToolContext:
return ToolContext(
user={"username": "e2e", "role": "admin", "vaults": vaults or ["*"]},
audit_enabled=False,
)
def _vault_path() -> Path:
from backend.indexer import get_vault_data
return Path(get_vault_data("TestVault")["path"])
class ScriptedLLM:
"""Async LLM returning pre-scripted responses and recording invocations."""
def __init__(self, responses):
self.responses = list(responses)
self.calls = []
async def __call__(self, messages, tools):
self.calls.append({"messages": [dict(m) for m in messages], "tools": tools})
return self.responses.pop(0)
# ═══════════════════════════════════════════════════════════════════
# In-app agent — read then confirmed write, end to end
# ═══════════════════════════════════════════════════════════════════
class TestAgentEndToEnd:
@pytest.mark.asyncio
async def test_read_then_confirmed_create_file(self, client):
target = _vault_path() / "agent-note.md"
assert not target.exists()
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(
id="1", name="read_file", arguments={"vault": "TestVault", "path": "note1.md"},
)]),
LLMResponse(tool_calls=[ToolCall(
id="2",
name="create_file",
arguments={"vault": "TestVault", "path": "agent-note.md", "content": "# Agent\n"},
)]),
])
paused = await run_agent([{"role": "user", "content": "read then create"}], ctx=_ctx(), llm=llm)
assert paused.stopped == STOP_CONFIRMATION_REQUIRED
# The read ran, the write is pending confirmation.
assert len(paused.tool_calls) == 1
assert paused.tool_calls[0].name == "read_file"
assert paused.pending["error"]["tool"] == "create_file"
assert not target.exists()
llm2 = ScriptedLLM([LLMResponse(content="Note créée.")])
resumed = await run_agent(
[{"role": "user", "content": "read then create"}],
ctx=_ctx(),
llm=llm2,
resume_messages=paused.messages,
confirm_pending=paused.pending,
)
assert resumed.stopped == STOP_DONE
assert resumed.content == "Note créée."
assert target.read_text(encoding="utf-8") == "# Agent\n"
assert any(r.name == "create_file" and r.ok for r in resumed.tool_calls)
@pytest.mark.asyncio
async def test_tool_call_quota_stops_run(self, client):
llm = ScriptedLLM([
LLMResponse(tool_calls=[
ToolCall(id="1", name="list_vaults", arguments={}),
ToolCall(id="2", name="list_vaults", arguments={}),
]),
])
result = await run_agent(
[{"role": "user", "content": "loop"}],
ctx=_ctx(),
llm=llm,
max_tool_calls=1,
)
assert result.stopped == STOP_QUOTA_EXCEEDED
assert len(result.tool_calls) == 1
# ═══════════════════════════════════════════════════════════════════
# Hardening — rate limiting & redaction
# ═══════════════════════════════════════════════════════════════════
class TestHardening:
def test_rate_limit_blocks_excess_calls(self, client, monkeypatch):
monkeypatch.setenv("OBSIGATE_TOOL_RATE_LIMIT", "1")
monkeypatch.setenv("OBSIGATE_TOOL_RATE_LIMIT_PER_TOOL", "1")
first = call_tool("list_vaults", _ctx(), {})
assert first.ok
with pytest.raises(ToolRateLimitError) as exc:
call_tool("list_vaults", _ctx(), {})
assert exc.value.code == "rate_limited"
assert exc.value.retry_after >= 1
def test_rate_limit_is_per_tool(self, client, monkeypatch):
monkeypatch.setenv("OBSIGATE_TOOL_RATE_LIMIT", "100")
monkeypatch.setenv("OBSIGATE_TOOL_RATE_LIMIT_PER_TOOL", "1")
call_tool("list_vaults", _ctx(), {})
# A different tool still has budget.
assert call_tool("list_tags", _ctx(), {}).ok
with pytest.raises(ToolRateLimitError):
call_tool("list_vaults", _ctx(), {})
def test_tool_results_are_redacted(self, client, monkeypatch):
from backend.tools import registry
from backend.tools.registry import ToolSpec
from backend.tools.schemas import ListVaultsInput
fake_jwt = "eyJ" + "a" * 30 + "." + "b" * 30 + "." + "c" * 30
spec = ToolSpec(
name="_leak",
description="returns a secret for tests",
input_model=ListVaultsInput,
handler=lambda ctx, params: {"content": f"token: {fake_jwt}", "nested": [fake_jwt]},
risk=ToolRisk.READ,
)
monkeypatch.setitem(registry._REGISTRY, "_leak", spec)
result = call_tool("_leak", _ctx(), {})
assert fake_jwt not in json.dumps(result.data)
assert "[JWT MASQUÉ]" in result.data["content"]
assert "[JWT MASQUÉ]" in result.data["nested"][0]
# ═══════════════════════════════════════════════════════════════════
# MCP — full propose/apply flow over Streamable HTTP
# ═══════════════════════════════════════════════════════════════════
ACCEPT = "application/json, text/event-stream"
PROTOCOL = "2025-03-26"
@pytest.fixture
def mcp_client(app_with_vault, monkeypatch):
from backend import main
async def _noop_build(*args, **kwargs):
return None
monkeypatch.setattr(main, "build_index", _noop_build)
monkeypatch.setattr(main, "init_inverted_index", lambda: None)
main.mcp_app._manager = None
main.mcp_app._run_task = None
main.mcp_app._start_lock = None
with TestClient(main.app) as client:
yield client
def _mcp_post(client, payload, session=None):
headers = {"Accept": ACCEPT, "Content-Type": "application/json"}
if session:
headers["Mcp-Session-Id"] = session
return client.post("/mcp", content=json.dumps(payload), headers=headers)
def _mcp_call(client, session, method, params=None, req_id=2):
resp = _mcp_post(
client,
{"jsonrpc": "2.0", "id": req_id, "method": method, "params": params or {}},
session=session,
)
assert resp.status_code == 200, resp.text
return resp.json()
def _text_json(result):
return json.loads(result["content"][0]["text"])
class TestMcpEndToEnd:
def test_read_propose_apply(self, mcp_client):
init = _mcp_post(
mcp_client,
{
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": PROTOCOL,
"capabilities": {},
"clientInfo": {"name": "pytest-e2e", "version": "1.0"},
},
},
)
assert init.status_code == 200
session = init.headers["mcp-session-id"]
_mcp_post(mcp_client, {"jsonrpc": "2.0", "method": "notifications/initialized"}, session)
# 1. Read the current content.
read = _mcp_call(
mcp_client, session, "tools/call",
{"name": "read_file", "arguments": {"vault": "TestVault", "path": "note1.md"}},
)
assert "Python" in _text_json(read["result"])["data"]["content"]
# 2. Propose an edit (no change yet).
proposed = _mcp_call(
mcp_client, session, "tools/call",
{
"name": "propose_edit_file",
"arguments": {"vault": "TestVault", "path": "note1.md", "content": "# E2E\n"},
},
)
payload = _text_json(proposed["result"])
token = payload["confirmation_token"]
assert payload["tool"] == "edit_file"
assert "diff" in payload
# 3. Apply it.
applied = _mcp_call(
mcp_client, session, "tools/call",
{"name": "apply_edit_file", "arguments": {"confirmation_token": token}},
req_id=3,
)
assert _text_json(applied["result"])["ok"] is True
assert (_vault_path() / "note1.md").read_text(encoding="utf-8") == "# E2E\n"
# 4. Resource read reflects the new content.
resource = _mcp_call(
mcp_client, session, "resources/read", {"uri": "vault://TestVault/note1.md"}, req_id=4
)
assert "E2E" in resource["result"]["contents"][0]["text"]