CI / lint (push) Successful in 58s
CI / security (push) Successful in 40s
CI / test (push) Successful in 1m15s
CI / build (push) Successful in 37s
CI / e2e (push) Successful in 10m15s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
259 lines
12 KiB
Python
259 lines
12 KiB
Python
# tests/test_agent_loop.py — Unit tests for the in-app agent loop (Phase B)
|
|
"""Tests for backend.agent.loop.run_agent using a scripted (mocked) LLM."""
|
|
|
|
import pytest
|
|
|
|
from backend.agent.loop import (
|
|
STOP_CONFIRMATION_REQUIRED,
|
|
STOP_DONE,
|
|
STOP_MAX_ITERATIONS,
|
|
run_agent,
|
|
)
|
|
from backend.ai_chat import LLMResponse, ToolCall
|
|
from backend.tools.api import ToolContext, ToolRisk
|
|
from backend.tools.registry import ToolSpec
|
|
from backend.tools.schemas import ListVaultsInput
|
|
|
|
|
|
def _ctx(vaults=None) -> ToolContext:
|
|
return ToolContext(
|
|
user={"username": "tester", "role": "admin", "vaults": vaults or ["*"]},
|
|
audit_enabled=False,
|
|
)
|
|
|
|
|
|
def _register(monkeypatch, name, handler, risk=ToolRisk.READ):
|
|
from backend.tools import registry
|
|
|
|
spec = ToolSpec(
|
|
name=name,
|
|
description=f"test tool {name}",
|
|
input_model=ListVaultsInput,
|
|
handler=handler,
|
|
risk=risk,
|
|
)
|
|
monkeypatch.setitem(registry._REGISTRY, name, spec)
|
|
return spec
|
|
|
|
|
|
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)
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# Plain completion
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestPlainCompletion:
|
|
@pytest.mark.asyncio
|
|
async def test_no_tool_calls_returns_content(self):
|
|
llm = ScriptedLLM([LLMResponse(content="bonjour")])
|
|
result = await run_agent([{"role": "user", "content": "hi"}], ctx=_ctx(), llm=llm)
|
|
assert result.stopped == STOP_DONE
|
|
assert result.content == "bonjour"
|
|
assert result.iterations == 1
|
|
assert result.tool_calls == []
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_tools_empty_disables_tool_calling(self):
|
|
llm = ScriptedLLM([LLMResponse(content="ok")])
|
|
await run_agent([{"role": "user", "content": "hi"}], ctx=_ctx(), llm=llm, tools=[])
|
|
assert llm.calls[0]["tools"] == []
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_default_tools_expose_in_app_tools(self):
|
|
llm = ScriptedLLM([LLMResponse(content="ok")])
|
|
await run_agent([{"role": "user", "content": "hi"}], ctx=_ctx(), llm=llm)
|
|
names = {t["function"]["name"] for t in llm.calls[0]["tools"]}
|
|
assert "list_vaults" in names
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# Tool calling
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestToolCalling:
|
|
@pytest.mark.asyncio
|
|
async def test_executes_tool_then_answers(self, monkeypatch):
|
|
_register(monkeypatch, "_echo", lambda ctx, params: {"echoed": True})
|
|
llm = ScriptedLLM([
|
|
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
|
|
LLMResponse(content="final"),
|
|
])
|
|
result = await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm)
|
|
assert result.stopped == STOP_DONE
|
|
assert result.content == "final"
|
|
assert result.iterations == 2
|
|
assert len(result.tool_calls) == 1
|
|
assert result.tool_calls[0].ok is True
|
|
assert result.tool_calls[0].result == {"echoed": True}
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_tool_result_is_fed_back(self, monkeypatch):
|
|
_register(monkeypatch, "_echo", lambda ctx, params: {"value": 42})
|
|
llm = ScriptedLLM([
|
|
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
|
|
LLMResponse(content="done"),
|
|
])
|
|
await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm)
|
|
# Second LLM call must contain a tool message with the result.
|
|
second_messages = llm.calls[1]["messages"]
|
|
tool_msgs = [m for m in second_messages if m.get("role") == "tool"]
|
|
assert len(tool_msgs) == 1
|
|
assert '"value": 42' in tool_msgs[0]["content"]
|
|
assert tool_msgs[0]["tool_call_id"] == "1"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_unknown_tool_recorded_as_error(self):
|
|
llm = ScriptedLLM([
|
|
LLMResponse(tool_calls=[ToolCall(id="1", name="does_not_exist", arguments={})]),
|
|
LLMResponse(content="recovered"),
|
|
])
|
|
result = await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm)
|
|
assert result.stopped == STOP_DONE
|
|
assert len(result.tool_calls) == 1
|
|
assert result.tool_calls[0].ok is False
|
|
assert result.tool_calls[0].result["error"]["code"] == "not_found"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_on_tool_call_callback(self, monkeypatch):
|
|
_register(monkeypatch, "_echo", lambda ctx, params: "x")
|
|
seen = []
|
|
llm = ScriptedLLM([
|
|
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
|
|
LLMResponse(content="done"),
|
|
])
|
|
await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm, on_tool_call=seen.append)
|
|
assert len(seen) == 1
|
|
assert seen[0].name == "_echo"
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# Confirmation & limits
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestConfirmationAndLimits:
|
|
@pytest.mark.asyncio
|
|
async def test_write_tool_pauses_for_confirmation(self, monkeypatch):
|
|
_register(monkeypatch, "_write", lambda ctx, params: {"done": True}, risk=ToolRisk.WRITE)
|
|
llm = ScriptedLLM([LLMResponse(tool_calls=[ToolCall(id="1", name="_write", arguments={})])])
|
|
result = await run_agent([{"role": "user", "content": "write"}], ctx=_ctx(), llm=llm)
|
|
assert result.stopped == STOP_CONFIRMATION_REQUIRED
|
|
assert result.pending is not None
|
|
assert result.pending["error"]["tool"] == "_write"
|
|
assert result.tool_calls == []
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_write_tool_runs_when_confirmed(self, monkeypatch):
|
|
_register(monkeypatch, "_write", lambda ctx, params: {"done": True}, risk=ToolRisk.WRITE)
|
|
llm = ScriptedLLM([
|
|
LLMResponse(tool_calls=[ToolCall(id="1", name="_write", arguments={})]),
|
|
LLMResponse(content="done"),
|
|
])
|
|
ctx = _ctx()
|
|
ctx.confirmed = True
|
|
result = await run_agent([{"role": "user", "content": "write"}], ctx=ctx, llm=llm)
|
|
assert result.stopped == STOP_DONE
|
|
assert result.tool_calls[0].ok is True
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_max_iterations_stops_loop(self, monkeypatch):
|
|
_register(monkeypatch, "_echo", lambda ctx, params: "x")
|
|
llm = ScriptedLLM([
|
|
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
|
|
LLMResponse(tool_calls=[ToolCall(id="2", name="_echo", arguments={})]),
|
|
LLMResponse(content="never reached"),
|
|
])
|
|
result = await run_agent([{"role": "user", "content": "loop"}], ctx=_ctx(), llm=llm, max_iterations=2)
|
|
assert result.stopped == STOP_MAX_ITERATIONS
|
|
assert result.iterations == 2
|
|
assert len(result.tool_calls) == 2
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# Permissions (index-backed)
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestConfirmationResume:
|
|
@pytest.mark.asyncio
|
|
async def test_resume_applies_pending_and_continues(self, monkeypatch):
|
|
seen = {}
|
|
|
|
def handler(ctx, params):
|
|
seen["ctx_confirmed"] = ctx.confirmed
|
|
return {"done": True}
|
|
|
|
_register(monkeypatch, "_write", handler, risk=ToolRisk.WRITE)
|
|
|
|
# First run pauses on the mutating tool.
|
|
llm1 = ScriptedLLM([LLMResponse(tool_calls=[ToolCall(id="1", name="_write", arguments={"x": 1})])])
|
|
paused = await run_agent([{"role": "user", "content": "write"}], ctx=_ctx(), llm=llm1)
|
|
assert paused.stopped == STOP_CONFIRMATION_REQUIRED
|
|
assert paused.pending["error"]["id"] == "1"
|
|
|
|
# Resume: the pending call is applied (one-shot confirm), then the loop
|
|
# continues and produces the final answer.
|
|
llm2 = ScriptedLLM([LLMResponse(content="applied")])
|
|
resumed = await run_agent(
|
|
[{"role": "user", "content": "write"}],
|
|
ctx=_ctx(),
|
|
llm=llm2,
|
|
resume_messages=paused.messages,
|
|
confirm_pending=paused.pending,
|
|
)
|
|
assert resumed.stopped == STOP_DONE
|
|
assert resumed.content == "applied"
|
|
assert len(resumed.tool_calls) == 1
|
|
assert resumed.tool_calls[0].ok is True
|
|
# The one-shot confirmation must not leak into the context.
|
|
assert seen["ctx_confirmed"] is False
|
|
# The tool result is fed back to the model on the resumed turn.
|
|
tool_msgs = [m for m in llm2.calls[0]["messages"] if m.get("role") == "tool"]
|
|
assert len(tool_msgs) == 1
|
|
assert tool_msgs[0]["tool_call_id"] == "1"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resume_without_assistant_message_reconstructs_it(self, monkeypatch):
|
|
_register(monkeypatch, "_write", lambda ctx, params: {"done": True}, risk=ToolRisk.WRITE)
|
|
llm = ScriptedLLM([LLMResponse(content="ok")])
|
|
pending = {"error": {"tool": "_write", "arguments": {}, "id": "call_9"}}
|
|
result = await run_agent(
|
|
[{"role": "user", "content": "write"}],
|
|
ctx=_ctx(),
|
|
llm=llm,
|
|
resume_messages=[{"role": "user", "content": "write"}],
|
|
confirm_pending=pending,
|
|
)
|
|
assert result.stopped == STOP_DONE
|
|
assistant_tool_msgs = [
|
|
m for m in llm.calls[0]["messages"]
|
|
if m.get("role") == "assistant" and m.get("tool_calls")
|
|
]
|
|
assert len(assistant_tool_msgs) == 1
|
|
assert assistant_tool_msgs[0]["tool_calls"][0]["id"] == "call_9"
|
|
|
|
|
|
class TestAgentPermissions:
|
|
@pytest.mark.asyncio
|
|
async def test_permission_denied_recorded(self, client):
|
|
llm = ScriptedLLM([
|
|
LLMResponse(tool_calls=[ToolCall(id="1", name="list_directory", arguments={"vault": "TestVault"})]),
|
|
LLMResponse(content="denied"),
|
|
])
|
|
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
|
|
result = await run_agent([{"role": "user", "content": "list"}], ctx=ctx, llm=llm)
|
|
assert result.tool_calls[0].ok is False
|
|
assert result.tool_calls[0].result["error"]["code"] == "vault_access_denied"
|