266 lines
9.7 KiB
Python
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"]
|