|
@@ -0,0 +1,229 @@
|
|
|
|
|
+"""AgentLoop shared-path behavior tests."""
|
|
|
|
|
+
|
|
|
|
|
+from __future__ import annotations
|
|
|
|
|
+
|
|
|
|
|
+from typing import Any
|
|
|
|
|
+
|
|
|
|
|
+import pytest
|
|
|
|
|
+
|
|
|
|
|
+from supply_agent.agent.loop import AgentLoop
|
|
|
|
|
+from supply_agent.tools.base import tool
|
|
|
|
|
+from supply_agent.tools.registry import ToolRegistry
|
|
|
|
|
+from supply_agent.types import (
|
|
|
|
|
+ AgentEventType,
|
|
|
|
|
+ Message,
|
|
|
|
|
+ Role,
|
|
|
|
|
+ ToolCall,
|
|
|
|
|
+)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+class _FakeLLM:
|
|
|
|
|
+ """Deterministic LLM that returns scripted assistant messages."""
|
|
|
|
|
+
|
|
|
|
|
+ def __init__(self, responses: list[Message]) -> None:
|
|
|
|
|
+ self._responses = list(responses)
|
|
|
|
|
+ self.calls: list[dict[str, Any]] = []
|
|
|
|
|
+
|
|
|
|
|
+ def chat(
|
|
|
|
|
+ self,
|
|
|
|
|
+ messages: list[Message],
|
|
|
|
|
+ tools: list[Any] | None = None,
|
|
|
|
|
+ temperature: float | None = None,
|
|
|
|
|
+ *,
|
|
|
|
|
+ iteration: int = 0,
|
|
|
|
|
+ ) -> Message:
|
|
|
|
|
+ self.calls.append(
|
|
|
|
|
+ {
|
|
|
|
|
+ "messages": list(messages),
|
|
|
|
|
+ "tools": tools,
|
|
|
|
|
+ "iteration": iteration,
|
|
|
|
|
+ }
|
|
|
|
|
+ )
|
|
|
|
|
+ if not self._responses:
|
|
|
|
|
+ raise AssertionError("FakeLLM: no more scripted responses")
|
|
|
|
|
+ return self._responses.pop(0)
|
|
|
|
|
+
|
|
|
|
|
+ async def achat(
|
|
|
|
|
+ self,
|
|
|
|
|
+ messages: list[Message],
|
|
|
|
|
+ tools: list[Any] | None = None,
|
|
|
|
|
+ temperature: float | None = None,
|
|
|
|
|
+ *,
|
|
|
|
|
+ iteration: int = 0,
|
|
|
|
|
+ ) -> Message:
|
|
|
|
|
+ return self.chat(
|
|
|
|
|
+ messages,
|
|
|
|
|
+ tools=tools,
|
|
|
|
|
+ temperature=temperature,
|
|
|
|
|
+ iteration=iteration,
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _system() -> Message:
|
|
|
|
|
+ return Message(role=Role.SYSTEM, content="base system")
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _assistant_text(content: str) -> Message:
|
|
|
|
|
+ return Message(role=Role.ASSISTANT, content=content)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _assistant_tools(*calls: ToolCall) -> Message:
|
|
|
|
|
+ return Message(role=Role.ASSISTANT, content=None, tool_calls=list(calls))
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+@pytest.fixture
|
|
|
|
|
+def skill_tools() -> ToolRegistry:
|
|
|
|
|
+ registry = ToolRegistry()
|
|
|
|
|
+
|
|
|
|
|
+ @tool(name="load_skill")
|
|
|
|
|
+ def load_skill(name: str) -> str:
|
|
|
|
|
+ return f"skill-body:{name}"
|
|
|
|
|
+
|
|
|
|
|
+ registry.register(load_skill, name="load_skill")
|
|
|
|
|
+ return registry
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def test_stream_loads_skill_into_system_message(skill_tools: ToolRegistry) -> None:
|
|
|
|
|
+ active_skills: list[str] = []
|
|
|
|
|
+ system_parts = {"value": "base system"}
|
|
|
|
|
+
|
|
|
|
|
+ def builder() -> Message:
|
|
|
|
|
+ parts = [system_parts["value"]]
|
|
|
|
|
+ for name in active_skills:
|
|
|
|
|
+ parts.append(f"LOADED:{name}")
|
|
|
|
|
+ return Message(role=Role.SYSTEM, content="\n\n".join(parts))
|
|
|
|
|
+
|
|
|
|
|
+ llm = _FakeLLM(
|
|
|
|
|
+ [
|
|
|
|
|
+ _assistant_tools(
|
|
|
|
|
+ ToolCall(
|
|
|
|
|
+ id="c1",
|
|
|
|
|
+ name="load_skill",
|
|
|
|
|
+ arguments='{"name": "demo"}',
|
|
|
|
|
+ )
|
|
|
|
|
+ ),
|
|
|
|
|
+ _assistant_text("done"),
|
|
|
|
|
+ ]
|
|
|
|
|
+ )
|
|
|
|
|
+ loop = AgentLoop(
|
|
|
|
|
+ llm=llm, # type: ignore[arg-type]
|
|
|
|
|
+ tools=skill_tools,
|
|
|
|
|
+ system_message=builder(),
|
|
|
|
|
+ system_message_builder=builder,
|
|
|
|
|
+ messages=[Message(role=Role.USER, content="hi")],
|
|
|
|
|
+ max_iterations=5,
|
|
|
|
|
+ active_skills=active_skills,
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ events = list(loop.stream())
|
|
|
|
|
+ assert any(e.type == AgentEventType.DONE for e in events)
|
|
|
|
|
+ assert "demo" in active_skills
|
|
|
|
|
+ # Second LLM call must see refreshed system prompt with skill content.
|
|
|
|
|
+ assert "LOADED:demo" in (llm.calls[1]["messages"][0].content or "")
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def test_stream_nudges_best_answer_on_max_iterations() -> None:
|
|
|
|
|
+ registry = ToolRegistry()
|
|
|
|
|
+
|
|
|
|
|
+ @tool(name="noop")
|
|
|
|
|
+ def noop() -> str:
|
|
|
|
|
+ return '{"ok": true}'
|
|
|
|
|
+
|
|
|
|
|
+ registry.register(noop, name="noop")
|
|
|
|
|
+
|
|
|
|
|
+ llm = _FakeLLM(
|
|
|
|
|
+ [
|
|
|
|
|
+ _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
|
|
|
|
|
+ _assistant_text("final after nudge"),
|
|
|
|
|
+ ]
|
|
|
|
|
+ )
|
|
|
|
|
+ loop = AgentLoop(
|
|
|
|
|
+ llm=llm, # type: ignore[arg-type]
|
|
|
|
|
+ tools=registry,
|
|
|
|
|
+ system_message=_system(),
|
|
|
|
|
+ messages=[Message(role=Role.USER, content="hi")],
|
|
|
|
|
+ max_iterations=1,
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ events = list(loop.stream())
|
|
|
|
|
+ done = next(e for e in events if e.type == AgentEventType.DONE)
|
|
|
|
|
+ assert done.data["content"] == "final after nudge"
|
|
|
|
|
+ assert len(llm.calls) == 2
|
|
|
|
|
+ assert "best answer" in (llm.calls[1]["messages"][-1].content or "").lower()
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def test_run_and_stream_share_max_iteration_nudge() -> None:
|
|
|
|
|
+ registry = ToolRegistry()
|
|
|
|
|
+
|
|
|
|
|
+ @tool(name="noop")
|
|
|
|
|
+ def noop() -> str:
|
|
|
|
|
+ return '{"ok": true}'
|
|
|
|
|
+
|
|
|
|
|
+ registry.register(noop, name="noop")
|
|
|
|
|
+
|
|
|
|
|
+ def make_loop(responses: list[Message]) -> AgentLoop:
|
|
|
|
|
+ return AgentLoop(
|
|
|
|
|
+ llm=_FakeLLM(responses), # type: ignore[arg-type]
|
|
|
|
|
+ tools=registry,
|
|
|
|
|
+ system_message=_system(),
|
|
|
|
|
+ messages=[Message(role=Role.USER, content="hi")],
|
|
|
|
|
+ max_iterations=1,
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ run_result = make_loop(
|
|
|
|
|
+ [
|
|
|
|
|
+ _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
|
|
|
|
|
+ _assistant_text("from-run"),
|
|
|
|
|
+ ]
|
|
|
|
|
+ ).run()
|
|
|
|
|
+ stream_events = list(
|
|
|
|
|
+ make_loop(
|
|
|
|
|
+ [
|
|
|
|
|
+ _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
|
|
|
|
|
+ _assistant_text("from-stream"),
|
|
|
|
|
+ ]
|
|
|
|
|
+ ).stream()
|
|
|
|
|
+ )
|
|
|
|
|
+ stream_done = next(e for e in stream_events if e.type == AgentEventType.DONE)
|
|
|
|
|
+
|
|
|
|
|
+ assert run_result.content == "from-run"
|
|
|
|
|
+ assert stream_done.data["content"] == "from-stream"
|
|
|
|
|
+ assert run_result.iterations == stream_done.data["iterations"] == 1
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+@pytest.mark.asyncio
|
|
|
|
|
+async def test_astream_loads_skill(skill_tools: ToolRegistry) -> None:
|
|
|
|
|
+ active_skills: list[str] = []
|
|
|
|
|
+
|
|
|
|
|
+ def builder() -> Message:
|
|
|
|
|
+ parts = ["base"]
|
|
|
|
|
+ for name in active_skills:
|
|
|
|
|
+ parts.append(f"LOADED:{name}")
|
|
|
|
|
+ return Message(role=Role.SYSTEM, content="\n\n".join(parts))
|
|
|
|
|
+
|
|
|
|
|
+ llm = _FakeLLM(
|
|
|
|
|
+ [
|
|
|
|
|
+ _assistant_tools(
|
|
|
|
|
+ ToolCall(
|
|
|
|
|
+ id="c1",
|
|
|
|
|
+ name="load_skill",
|
|
|
|
|
+ arguments='{"name": "async-demo"}',
|
|
|
|
|
+ )
|
|
|
|
|
+ ),
|
|
|
|
|
+ _assistant_text("ok"),
|
|
|
|
|
+ ]
|
|
|
|
|
+ )
|
|
|
|
|
+ loop = AgentLoop(
|
|
|
|
|
+ llm=llm, # type: ignore[arg-type]
|
|
|
|
|
+ tools=skill_tools,
|
|
|
|
|
+ system_message=builder(),
|
|
|
|
|
+ system_message_builder=builder,
|
|
|
|
|
+ messages=[Message(role=Role.USER, content="hi")],
|
|
|
|
|
+ max_iterations=5,
|
|
|
|
|
+ active_skills=active_skills,
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ events = [event async for event in loop.astream()]
|
|
|
|
|
+ assert any(e.type == AgentEventType.DONE for e in events)
|
|
|
|
|
+ assert "async-demo" in active_skills
|
|
|
|
|
+ assert "LOADED:async-demo" in (llm.calls[1]["messages"][0].content or "")
|