test_debug_runtime.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325
  1. import asyncio
  2. import importlib
  3. from collections.abc import AsyncIterator
  4. from typing import Any
  5. import pytest
  6. from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
  7. from agent_lab.application.runtime import DebugRuntime
  8. from agent_lab.application.tools import ToolDefinition, ToolRegistry
  9. from agent_lab.domain.events import ToolCallEvent
  10. from agent_lab.domain.messages import ChatMessage, StreamItem
  11. def _runtime_queues_class():
  12. module = importlib.import_module("agent_lab.application.queues")
  13. return module.RuntimeQueues
  14. class RecordingQueue(asyncio.Queue):
  15. def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
  16. super().__init__()
  17. self.name = name
  18. self.log = log
  19. async def put(self, item: Any) -> None:
  20. self.log.append((self.name, "put", self._describe(item)))
  21. await super().put(item)
  22. async def get(self) -> Any:
  23. item = await super().get()
  24. self.log.append((self.name, "get", self._describe(item)))
  25. return item
  26. def _describe(self, item: Any) -> str:
  27. if isinstance(item, ChatMessage):
  28. if item.role == "tool":
  29. return f"tool:{item.tool_call_id}"
  30. return item.role
  31. if isinstance(item, ToolCallEvent):
  32. return f"event:{item.name}:{item.id}"
  33. if isinstance(item, dict):
  34. return f"output:{item.get('type')}"
  35. return type(item).__name__
  36. class FakeChatClient:
  37. def __init__(self) -> None:
  38. self.calls = 0
  39. async def stream_chat(
  40. self,
  41. messages: list[ChatMessage],
  42. tools: list[dict],
  43. params: AgentParams,
  44. ) -> AsyncIterator[StreamItem]:
  45. self.calls += 1
  46. if self.calls == 1:
  47. yield StreamItem.event(
  48. ToolCallEvent(
  49. id="call_1",
  50. name="handoff_note",
  51. arguments={"message": "need event agent"},
  52. raw_arguments='{"message":"need event agent"}',
  53. )
  54. )
  55. return
  56. assert any(message.role == "tool" for message in messages)
  57. yield StreamItem.message_delta("final answer")
  58. class StrictHistoryChatClient:
  59. def __init__(self) -> None:
  60. self.calls = 0
  61. self.second_call_messages: list[ChatMessage] = []
  62. async def stream_chat(
  63. self,
  64. messages: list[ChatMessage],
  65. tools: list[dict],
  66. params: AgentParams,
  67. ) -> AsyncIterator[StreamItem]:
  68. self.calls += 1
  69. if self.calls == 1:
  70. yield StreamItem.event(
  71. ToolCallEvent(
  72. id="call_1",
  73. name="handoff_note",
  74. arguments={"message": "need event agent"},
  75. raw_arguments='{"message":"need event agent"}',
  76. )
  77. )
  78. return
  79. self.second_call_messages = list(messages)
  80. yield StreamItem.message_delta("final answer")
  81. class ToolCapturingChatClient:
  82. def __init__(self) -> None:
  83. self.tools: list[dict[str, Any]] = []
  84. async def stream_chat(
  85. self,
  86. messages: list[ChatMessage],
  87. tools: list[dict],
  88. params: AgentParams,
  89. ) -> AsyncIterator[StreamItem]:
  90. self.tools = list(tools)
  91. yield StreamItem.message_delta("final answer")
  92. def test_runtime_queues_exposes_input_output_and_events_queues():
  93. RuntimeQueues = _runtime_queues_class()
  94. queues = RuntimeQueues()
  95. assert isinstance(queues.input, asyncio.Queue)
  96. assert isinstance(queues.output, asyncio.Queue)
  97. assert isinstance(queues.events, asyncio.Queue)
  98. assert queues.input is not queues.output
  99. assert queues.input is not queues.events
  100. assert queues.output is not queues.events
  101. @pytest.mark.asyncio
  102. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  103. request = DebugRunRequest(
  104. user_message="debug this",
  105. system_prompts=["You are a debugger."],
  106. pre_messages=[],
  107. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  108. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  109. )
  110. client = FakeChatClient()
  111. runtime = DebugRuntime(client)
  112. outputs = [message async for message in runtime.run(request)]
  113. assert client.calls == 2
  114. assert [message["type"] for message in outputs] == [
  115. "session_started",
  116. "event",
  117. "tool_result",
  118. "message_delta",
  119. "done",
  120. ]
  121. assert outputs[1]["event"]["name"] == "handoff_note"
  122. assert outputs[3]["content"] == "final answer"
  123. @pytest.mark.asyncio
  124. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  125. RuntimeQueues = _runtime_queues_class()
  126. queue_log: list[tuple[str, str, str]] = []
  127. queues = RuntimeQueues(
  128. input=RecordingQueue("input", queue_log),
  129. output=RecordingQueue("output", queue_log),
  130. events=RecordingQueue("events", queue_log),
  131. )
  132. request = DebugRunRequest(
  133. user_message="debug this",
  134. system_prompts=["You are a debugger."],
  135. pre_messages=[],
  136. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  137. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  138. )
  139. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  140. outputs = [message async for message in runtime.run(request)]
  141. assert [message["type"] for message in outputs] == [
  142. "session_started",
  143. "event",
  144. "tool_result",
  145. "message_delta",
  146. "done",
  147. ]
  148. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  149. ("input", "get", "user")
  150. )
  151. assert queue_log.index(("events", "put", "event:handoff_note:call_1")) < queue_log.index(
  152. ("events", "get", "event:handoff_note:call_1")
  153. )
  154. assert queue_log.index(("events", "get", "event:handoff_note:call_1")) < queue_log.index(
  155. ("input", "put", "tool:call_1")
  156. )
  157. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  158. ("input", "get", "tool:call_1")
  159. )
  160. @pytest.mark.asyncio
  161. async def test_runtime_yields_existing_output_order_from_output_queue():
  162. RuntimeQueues = _runtime_queues_class()
  163. queue_log: list[tuple[str, str, str]] = []
  164. queues = RuntimeQueues(
  165. input=RecordingQueue("input", queue_log),
  166. output=RecordingQueue("output", queue_log),
  167. events=RecordingQueue("events", queue_log),
  168. )
  169. request = DebugRunRequest(
  170. user_message="debug this",
  171. system_prompts=["You are a debugger."],
  172. pre_messages=[],
  173. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  174. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  175. )
  176. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  177. outputs = [message async for message in runtime.run(request)]
  178. assert [message["type"] for message in outputs] == [
  179. "session_started",
  180. "event",
  181. "tool_result",
  182. "message_delta",
  183. "done",
  184. ]
  185. assert [
  186. entry
  187. for entry in queue_log
  188. if entry[0] == "output" and entry[1] in {"put", "get"}
  189. ] == [
  190. ("output", "put", "output:session_started"),
  191. ("output", "get", "output:session_started"),
  192. ("output", "put", "output:event"),
  193. ("output", "get", "output:event"),
  194. ("output", "put", "output:tool_result"),
  195. ("output", "get", "output:tool_result"),
  196. ("output", "put", "output:message_delta"),
  197. ("output", "get", "output:message_delta"),
  198. ("output", "put", "output:done"),
  199. ("output", "get", "output:done"),
  200. ]
  201. @pytest.mark.asyncio
  202. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  203. request = DebugRunRequest(
  204. user_message="debug this",
  205. system_prompts=["You are a debugger."],
  206. pre_messages=[],
  207. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  208. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  209. )
  210. client = StrictHistoryChatClient()
  211. runtime = DebugRuntime(client)
  212. outputs = [message async for message in runtime.run(request)]
  213. assert client.calls == 2
  214. assert [message.role for message in client.second_call_messages] == [
  215. "system",
  216. "user",
  217. "assistant",
  218. "tool",
  219. ]
  220. assistant_message = client.second_call_messages[2]
  221. tool_message = client.second_call_messages[3]
  222. assert assistant_message.content == ""
  223. assert assistant_message.tool_calls == [
  224. {
  225. "id": "call_1",
  226. "type": "function",
  227. "function": {
  228. "name": "handoff_note",
  229. "arguments": '{"message":"need event agent"}',
  230. },
  231. }
  232. ]
  233. assert tool_message.tool_call_id == "call_1"
  234. assert outputs[-1] == {"type": "done"}
  235. @pytest.mark.asyncio
  236. async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent():
  237. registry = ToolRegistry(
  238. [
  239. ToolDefinition(
  240. name="handoff_note",
  241. description="Registry-owned handoff tool.",
  242. parameters={
  243. "type": "object",
  244. "properties": {
  245. "message": {"type": "string"},
  246. "priority": {"type": "number"},
  247. },
  248. "required": ["message"],
  249. },
  250. handler=lambda event: {"tool": event.name, "message": "handled"},
  251. )
  252. ]
  253. )
  254. request = DebugRunRequest(
  255. user_message="debug this",
  256. system_prompts=[],
  257. pre_messages=[],
  258. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  259. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  260. )
  261. client = ToolCapturingChatClient()
  262. runtime = DebugRuntime(client, registry=registry)
  263. outputs = [message async for message in runtime.run(request)]
  264. assert outputs[-1] == {"type": "done"}
  265. assert client.tools == [
  266. {
  267. "type": "function",
  268. "function": {
  269. "name": "handoff_note",
  270. "description": "Registry-owned handoff tool.",
  271. "parameters": {
  272. "type": "object",
  273. "properties": {
  274. "message": {"type": "string"},
  275. "priority": {"type": "number"},
  276. },
  277. "required": ["message"],
  278. },
  279. },
  280. }
  281. ]