test_agent_loop.py 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229
  1. """AgentLoop shared-path behavior tests."""
  2. from __future__ import annotations
  3. from typing import Any
  4. import pytest
  5. from supply_agent.agent.loop import AgentLoop
  6. from supply_agent.tools.base import tool
  7. from supply_agent.tools.registry import ToolRegistry
  8. from supply_agent.types import (
  9. AgentEventType,
  10. Message,
  11. Role,
  12. ToolCall,
  13. )
  14. class _FakeLLM:
  15. """Deterministic LLM that returns scripted assistant messages."""
  16. def __init__(self, responses: list[Message]) -> None:
  17. self._responses = list(responses)
  18. self.calls: list[dict[str, Any]] = []
  19. def chat(
  20. self,
  21. messages: list[Message],
  22. tools: list[Any] | None = None,
  23. temperature: float | None = None,
  24. *,
  25. iteration: int = 0,
  26. ) -> Message:
  27. self.calls.append(
  28. {
  29. "messages": list(messages),
  30. "tools": tools,
  31. "iteration": iteration,
  32. }
  33. )
  34. if not self._responses:
  35. raise AssertionError("FakeLLM: no more scripted responses")
  36. return self._responses.pop(0)
  37. async def achat(
  38. self,
  39. messages: list[Message],
  40. tools: list[Any] | None = None,
  41. temperature: float | None = None,
  42. *,
  43. iteration: int = 0,
  44. ) -> Message:
  45. return self.chat(
  46. messages,
  47. tools=tools,
  48. temperature=temperature,
  49. iteration=iteration,
  50. )
  51. def _system() -> Message:
  52. return Message(role=Role.SYSTEM, content="base system")
  53. def _assistant_text(content: str) -> Message:
  54. return Message(role=Role.ASSISTANT, content=content)
  55. def _assistant_tools(*calls: ToolCall) -> Message:
  56. return Message(role=Role.ASSISTANT, content=None, tool_calls=list(calls))
  57. @pytest.fixture
  58. def skill_tools() -> ToolRegistry:
  59. registry = ToolRegistry()
  60. @tool(name="load_skill")
  61. def load_skill(name: str) -> str:
  62. return f"skill-body:{name}"
  63. registry.register(load_skill, name="load_skill")
  64. return registry
  65. def test_stream_loads_skill_into_system_message(skill_tools: ToolRegistry) -> None:
  66. active_skills: list[str] = []
  67. system_parts = {"value": "base system"}
  68. def builder() -> Message:
  69. parts = [system_parts["value"]]
  70. for name in active_skills:
  71. parts.append(f"LOADED:{name}")
  72. return Message(role=Role.SYSTEM, content="\n\n".join(parts))
  73. llm = _FakeLLM(
  74. [
  75. _assistant_tools(
  76. ToolCall(
  77. id="c1",
  78. name="load_skill",
  79. arguments='{"name": "demo"}',
  80. )
  81. ),
  82. _assistant_text("done"),
  83. ]
  84. )
  85. loop = AgentLoop(
  86. llm=llm, # type: ignore[arg-type]
  87. tools=skill_tools,
  88. system_message=builder(),
  89. system_message_builder=builder,
  90. messages=[Message(role=Role.USER, content="hi")],
  91. max_iterations=5,
  92. active_skills=active_skills,
  93. )
  94. events = list(loop.stream())
  95. assert any(e.type == AgentEventType.DONE for e in events)
  96. assert "demo" in active_skills
  97. # Second LLM call must see refreshed system prompt with skill content.
  98. assert "LOADED:demo" in (llm.calls[1]["messages"][0].content or "")
  99. def test_stream_nudges_best_answer_on_max_iterations() -> None:
  100. registry = ToolRegistry()
  101. @tool(name="noop")
  102. def noop() -> str:
  103. return '{"ok": true}'
  104. registry.register(noop, name="noop")
  105. llm = _FakeLLM(
  106. [
  107. _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
  108. _assistant_text("final after nudge"),
  109. ]
  110. )
  111. loop = AgentLoop(
  112. llm=llm, # type: ignore[arg-type]
  113. tools=registry,
  114. system_message=_system(),
  115. messages=[Message(role=Role.USER, content="hi")],
  116. max_iterations=1,
  117. )
  118. events = list(loop.stream())
  119. done = next(e for e in events if e.type == AgentEventType.DONE)
  120. assert done.data["content"] == "final after nudge"
  121. assert len(llm.calls) == 2
  122. assert "best answer" in (llm.calls[1]["messages"][-1].content or "").lower()
  123. def test_run_and_stream_share_max_iteration_nudge() -> None:
  124. registry = ToolRegistry()
  125. @tool(name="noop")
  126. def noop() -> str:
  127. return '{"ok": true}'
  128. registry.register(noop, name="noop")
  129. def make_loop(responses: list[Message]) -> AgentLoop:
  130. return AgentLoop(
  131. llm=_FakeLLM(responses), # type: ignore[arg-type]
  132. tools=registry,
  133. system_message=_system(),
  134. messages=[Message(role=Role.USER, content="hi")],
  135. max_iterations=1,
  136. )
  137. run_result = make_loop(
  138. [
  139. _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
  140. _assistant_text("from-run"),
  141. ]
  142. ).run()
  143. stream_events = list(
  144. make_loop(
  145. [
  146. _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
  147. _assistant_text("from-stream"),
  148. ]
  149. ).stream()
  150. )
  151. stream_done = next(e for e in stream_events if e.type == AgentEventType.DONE)
  152. assert run_result.content == "from-run"
  153. assert stream_done.data["content"] == "from-stream"
  154. assert run_result.iterations == stream_done.data["iterations"] == 1
  155. @pytest.mark.asyncio
  156. async def test_astream_loads_skill(skill_tools: ToolRegistry) -> None:
  157. active_skills: list[str] = []
  158. def builder() -> Message:
  159. parts = ["base"]
  160. for name in active_skills:
  161. parts.append(f"LOADED:{name}")
  162. return Message(role=Role.SYSTEM, content="\n\n".join(parts))
  163. llm = _FakeLLM(
  164. [
  165. _assistant_tools(
  166. ToolCall(
  167. id="c1",
  168. name="load_skill",
  169. arguments='{"name": "async-demo"}',
  170. )
  171. ),
  172. _assistant_text("ok"),
  173. ]
  174. )
  175. loop = AgentLoop(
  176. llm=llm, # type: ignore[arg-type]
  177. tools=skill_tools,
  178. system_message=builder(),
  179. system_message_builder=builder,
  180. messages=[Message(role=Role.USER, content="hi")],
  181. max_iterations=5,
  182. active_skills=active_skills,
  183. )
  184. events = [event async for event in loop.astream()]
  185. assert any(e.type == AgentEventType.DONE for e in events)
  186. assert "async-demo" in active_skills
  187. assert "LOADED:async-demo" in (llm.calls[1]["messages"][0].content or "")