test_debug_runtime.py 8.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259
  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.domain.events import ToolCallEvent
  9. from agent_lab.domain.messages import ChatMessage, StreamItem
  10. def _runtime_queues_class():
  11. module = importlib.import_module("agent_lab.application.queues")
  12. return module.RuntimeQueues
  13. class RecordingQueue(asyncio.Queue):
  14. def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
  15. super().__init__()
  16. self.name = name
  17. self.log = log
  18. async def put(self, item: Any) -> None:
  19. self.log.append((self.name, "put", self._describe(item)))
  20. await super().put(item)
  21. async def get(self) -> Any:
  22. item = await super().get()
  23. self.log.append((self.name, "get", self._describe(item)))
  24. return item
  25. def _describe(self, item: Any) -> str:
  26. if isinstance(item, ChatMessage):
  27. if item.role == "tool":
  28. return f"tool:{item.tool_call_id}"
  29. return item.role
  30. if isinstance(item, ToolCallEvent):
  31. return f"event:{item.name}:{item.id}"
  32. if isinstance(item, dict):
  33. return f"output:{item.get('type')}"
  34. return type(item).__name__
  35. class FakeChatClient:
  36. def __init__(self) -> None:
  37. self.calls = 0
  38. async def stream_chat(
  39. self,
  40. messages: list[ChatMessage],
  41. tools: list[dict],
  42. params: AgentParams,
  43. ) -> AsyncIterator[StreamItem]:
  44. self.calls += 1
  45. if self.calls == 1:
  46. yield StreamItem.event(
  47. ToolCallEvent(
  48. id="call_1",
  49. name="handoff_note",
  50. arguments={"message": "need event agent"},
  51. raw_arguments='{"message":"need event agent"}',
  52. )
  53. )
  54. return
  55. assert any(message.role == "tool" for message in messages)
  56. yield StreamItem.message_delta("final answer")
  57. class StrictHistoryChatClient:
  58. def __init__(self) -> None:
  59. self.calls = 0
  60. self.second_call_messages: list[ChatMessage] = []
  61. async def stream_chat(
  62. self,
  63. messages: list[ChatMessage],
  64. tools: list[dict],
  65. params: AgentParams,
  66. ) -> AsyncIterator[StreamItem]:
  67. self.calls += 1
  68. if self.calls == 1:
  69. yield StreamItem.event(
  70. ToolCallEvent(
  71. id="call_1",
  72. name="handoff_note",
  73. arguments={"message": "need event agent"},
  74. raw_arguments='{"message":"need event agent"}',
  75. )
  76. )
  77. return
  78. self.second_call_messages = list(messages)
  79. yield StreamItem.message_delta("final answer")
  80. def test_runtime_queues_exposes_input_output_and_events_queues():
  81. RuntimeQueues = _runtime_queues_class()
  82. queues = RuntimeQueues()
  83. assert isinstance(queues.input, asyncio.Queue)
  84. assert isinstance(queues.output, asyncio.Queue)
  85. assert isinstance(queues.events, asyncio.Queue)
  86. assert queues.input is not queues.output
  87. assert queues.input is not queues.events
  88. assert queues.output is not queues.events
  89. @pytest.mark.asyncio
  90. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  91. request = DebugRunRequest(
  92. user_message="debug this",
  93. system_prompts=["You are a debugger."],
  94. pre_messages=[],
  95. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  96. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  97. )
  98. client = FakeChatClient()
  99. runtime = DebugRuntime(client)
  100. outputs = [message async for message in runtime.run(request)]
  101. assert client.calls == 2
  102. assert [message["type"] for message in outputs] == [
  103. "session_started",
  104. "event",
  105. "tool_result",
  106. "message_delta",
  107. "done",
  108. ]
  109. assert outputs[1]["event"]["name"] == "handoff_note"
  110. assert outputs[3]["content"] == "final answer"
  111. @pytest.mark.asyncio
  112. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  113. RuntimeQueues = _runtime_queues_class()
  114. queue_log: list[tuple[str, str, str]] = []
  115. queues = RuntimeQueues(
  116. input=RecordingQueue("input", queue_log),
  117. output=RecordingQueue("output", queue_log),
  118. events=RecordingQueue("events", queue_log),
  119. )
  120. request = DebugRunRequest(
  121. user_message="debug this",
  122. system_prompts=["You are a debugger."],
  123. pre_messages=[],
  124. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  125. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  126. )
  127. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  128. outputs = [message async for message in runtime.run(request)]
  129. assert [message["type"] for message in outputs] == [
  130. "session_started",
  131. "event",
  132. "tool_result",
  133. "message_delta",
  134. "done",
  135. ]
  136. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  137. ("input", "get", "user")
  138. )
  139. assert queue_log.index(("events", "put", "event:handoff_note:call_1")) < queue_log.index(
  140. ("events", "get", "event:handoff_note:call_1")
  141. )
  142. assert queue_log.index(("events", "get", "event:handoff_note:call_1")) < queue_log.index(
  143. ("input", "put", "tool:call_1")
  144. )
  145. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  146. ("input", "get", "tool:call_1")
  147. )
  148. @pytest.mark.asyncio
  149. async def test_runtime_yields_existing_output_order_from_output_queue():
  150. RuntimeQueues = _runtime_queues_class()
  151. queue_log: list[tuple[str, str, str]] = []
  152. queues = RuntimeQueues(
  153. input=RecordingQueue("input", queue_log),
  154. output=RecordingQueue("output", queue_log),
  155. events=RecordingQueue("events", queue_log),
  156. )
  157. request = DebugRunRequest(
  158. user_message="debug this",
  159. system_prompts=["You are a debugger."],
  160. pre_messages=[],
  161. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  162. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  163. )
  164. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  165. outputs = [message async for message in runtime.run(request)]
  166. assert [message["type"] for message in outputs] == [
  167. "session_started",
  168. "event",
  169. "tool_result",
  170. "message_delta",
  171. "done",
  172. ]
  173. assert [
  174. entry
  175. for entry in queue_log
  176. if entry[0] == "output" and entry[1] in {"put", "get"}
  177. ] == [
  178. ("output", "put", "output:session_started"),
  179. ("output", "get", "output:session_started"),
  180. ("output", "put", "output:event"),
  181. ("output", "get", "output:event"),
  182. ("output", "put", "output:tool_result"),
  183. ("output", "get", "output:tool_result"),
  184. ("output", "put", "output:message_delta"),
  185. ("output", "get", "output:message_delta"),
  186. ("output", "put", "output:done"),
  187. ("output", "get", "output:done"),
  188. ]
  189. @pytest.mark.asyncio
  190. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  191. request = DebugRunRequest(
  192. user_message="debug this",
  193. system_prompts=["You are a debugger."],
  194. pre_messages=[],
  195. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  196. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  197. )
  198. client = StrictHistoryChatClient()
  199. runtime = DebugRuntime(client)
  200. outputs = [message async for message in runtime.run(request)]
  201. assert client.calls == 2
  202. assert [message.role for message in client.second_call_messages] == [
  203. "system",
  204. "user",
  205. "assistant",
  206. "tool",
  207. ]
  208. assistant_message = client.second_call_messages[2]
  209. tool_message = client.second_call_messages[3]
  210. assert assistant_message.content == ""
  211. assert assistant_message.tool_calls == [
  212. {
  213. "id": "call_1",
  214. "type": "function",
  215. "function": {
  216. "name": "handoff_note",
  217. "arguments": '{"message":"need event agent"}',
  218. },
  219. }
  220. ]
  221. assert tool_message.tool_call_id == "call_1"
  222. assert outputs[-1] == {"type": "done"}