feat(ai): durcissement phase F (#79) - rate limit, redaction, OpenAPI/MCP, E2E
This commit is contained in:
@@ -12,6 +12,16 @@ from fastapi.testclient import TestClient
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_tool_ratelimit():
|
||||
"""Isolate the per-identity tool rate limiter between tests."""
|
||||
from backend.tools import ratelimit
|
||||
|
||||
ratelimit.reset()
|
||||
yield
|
||||
ratelimit.reset()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_env():
|
||||
"""Ensure no vault env vars leak between tests — but preserve test vault config."""
|
||||
|
||||
@@ -0,0 +1,265 @@
|
||||
# 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"]
|
||||
Reference in New Issue
Block a user