test_debug_runtime.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461
  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, TokenUsage
  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. class RoundStatsChatClient:
  93. async def stream_chat(
  94. self,
  95. messages: list[ChatMessage],
  96. tools: list[dict],
  97. params: AgentParams,
  98. ) -> AsyncIterator[StreamItem]:
  99. yield StreamItem.message_delta("hello")
  100. yield StreamItem.usage_item(
  101. TokenUsage(
  102. prompt_tokens=10,
  103. completion_tokens=20,
  104. total_tokens=30,
  105. cached_tokens=5,
  106. )
  107. )
  108. class EventRoundStatsChatClient:
  109. def __init__(self) -> None:
  110. self.calls = 0
  111. async def stream_chat(
  112. self,
  113. messages: list[ChatMessage],
  114. tools: list[dict],
  115. params: AgentParams,
  116. ) -> AsyncIterator[StreamItem]:
  117. self.calls += 1
  118. if self.calls == 1:
  119. yield StreamItem.event(
  120. ToolCallEvent(
  121. id="call_1",
  122. name="handoff_note",
  123. arguments={"message": "need event agent"},
  124. raw_arguments='{"message":"need event agent"}',
  125. )
  126. )
  127. yield StreamItem.usage_item(
  128. TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3)
  129. )
  130. return
  131. yield StreamItem.message_delta("final answer")
  132. yield StreamItem.usage_item(
  133. TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
  134. )
  135. def test_runtime_queues_exposes_input_output_and_events_queues():
  136. RuntimeQueues = _runtime_queues_class()
  137. queues = RuntimeQueues()
  138. assert isinstance(queues.input, asyncio.Queue)
  139. assert isinstance(queues.output, asyncio.Queue)
  140. assert isinstance(queues.events, asyncio.Queue)
  141. assert queues.input is not queues.output
  142. assert queues.input is not queues.events
  143. assert queues.output is not queues.events
  144. @pytest.mark.asyncio
  145. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  146. request = DebugRunRequest(
  147. user_message="debug this",
  148. system_prompts=["You are a debugger."],
  149. pre_messages=[],
  150. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  151. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  152. )
  153. client = FakeChatClient()
  154. runtime = DebugRuntime(client)
  155. outputs = [message async for message in runtime.run(request)]
  156. assert client.calls == 2
  157. assert [message["type"] for message in outputs] == [
  158. "session_started",
  159. "event",
  160. "tool_result",
  161. "round_stats",
  162. "message_delta",
  163. "round_stats",
  164. "done",
  165. ]
  166. assert outputs[1]["event"]["name"] == "handoff_note"
  167. assert outputs[4]["content"] == "final answer"
  168. @pytest.mark.asyncio
  169. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  170. RuntimeQueues = _runtime_queues_class()
  171. queue_log: list[tuple[str, str, str]] = []
  172. queues = RuntimeQueues(
  173. input=RecordingQueue("input", queue_log),
  174. output=RecordingQueue("output", queue_log),
  175. events=RecordingQueue("events", queue_log),
  176. )
  177. request = DebugRunRequest(
  178. user_message="debug this",
  179. system_prompts=["You are a debugger."],
  180. pre_messages=[],
  181. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  182. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  183. )
  184. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  185. outputs = [message async for message in runtime.run(request)]
  186. assert [message["type"] for message in outputs] == [
  187. "session_started",
  188. "event",
  189. "tool_result",
  190. "round_stats",
  191. "message_delta",
  192. "round_stats",
  193. "done",
  194. ]
  195. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  196. ("input", "get", "user")
  197. )
  198. assert queue_log.index(("events", "put", "event:handoff_note:call_1")) < queue_log.index(
  199. ("events", "get", "event:handoff_note:call_1")
  200. )
  201. assert queue_log.index(("events", "get", "event:handoff_note:call_1")) < queue_log.index(
  202. ("input", "put", "tool:call_1")
  203. )
  204. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  205. ("input", "get", "tool:call_1")
  206. )
  207. @pytest.mark.asyncio
  208. async def test_runtime_yields_existing_output_order_from_output_queue():
  209. RuntimeQueues = _runtime_queues_class()
  210. queue_log: list[tuple[str, str, str]] = []
  211. queues = RuntimeQueues(
  212. input=RecordingQueue("input", queue_log),
  213. output=RecordingQueue("output", queue_log),
  214. events=RecordingQueue("events", queue_log),
  215. )
  216. request = DebugRunRequest(
  217. user_message="debug this",
  218. system_prompts=["You are a debugger."],
  219. pre_messages=[],
  220. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  221. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  222. )
  223. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  224. outputs = [message async for message in runtime.run(request)]
  225. assert [message["type"] for message in outputs] == [
  226. "session_started",
  227. "event",
  228. "tool_result",
  229. "round_stats",
  230. "message_delta",
  231. "round_stats",
  232. "done",
  233. ]
  234. assert [
  235. entry
  236. for entry in queue_log
  237. if entry[0] == "output" and entry[1] in {"put", "get"}
  238. ] == [
  239. ("output", "put", "output:session_started"),
  240. ("output", "get", "output:session_started"),
  241. ("output", "put", "output:event"),
  242. ("output", "get", "output:event"),
  243. ("output", "put", "output:tool_result"),
  244. ("output", "get", "output:tool_result"),
  245. ("output", "put", "output:round_stats"),
  246. ("output", "get", "output:round_stats"),
  247. ("output", "put", "output:message_delta"),
  248. ("output", "get", "output:message_delta"),
  249. ("output", "put", "output:round_stats"),
  250. ("output", "get", "output:round_stats"),
  251. ("output", "put", "output:done"),
  252. ("output", "get", "output:done"),
  253. ]
  254. @pytest.mark.asyncio
  255. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  256. request = DebugRunRequest(
  257. user_message="debug this",
  258. system_prompts=["You are a debugger."],
  259. pre_messages=[],
  260. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  261. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  262. )
  263. client = StrictHistoryChatClient()
  264. runtime = DebugRuntime(client)
  265. outputs = [message async for message in runtime.run(request)]
  266. assert client.calls == 2
  267. assert [message.role for message in client.second_call_messages] == [
  268. "system",
  269. "user",
  270. "assistant",
  271. "tool",
  272. ]
  273. assistant_message = client.second_call_messages[2]
  274. tool_message = client.second_call_messages[3]
  275. assert assistant_message.content == ""
  276. assert assistant_message.tool_calls == [
  277. {
  278. "id": "call_1",
  279. "type": "function",
  280. "function": {
  281. "name": "handoff_note",
  282. "arguments": '{"message":"need event agent"}',
  283. },
  284. }
  285. ]
  286. assert tool_message.tool_call_id == "call_1"
  287. assert outputs[-1] == {"type": "done"}
  288. @pytest.mark.asyncio
  289. async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent():
  290. registry = ToolRegistry(
  291. [
  292. ToolDefinition(
  293. name="handoff_note",
  294. description="Registry-owned handoff tool.",
  295. parameters={
  296. "type": "object",
  297. "properties": {
  298. "message": {"type": "string"},
  299. "priority": {"type": "number"},
  300. },
  301. "required": ["message"],
  302. },
  303. handler=lambda event: {"tool": event.name, "message": "handled"},
  304. )
  305. ]
  306. )
  307. request = DebugRunRequest(
  308. user_message="debug this",
  309. system_prompts=[],
  310. pre_messages=[],
  311. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  312. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  313. )
  314. client = ToolCapturingChatClient()
  315. runtime = DebugRuntime(client, registry=registry)
  316. outputs = [message async for message in runtime.run(request)]
  317. assert outputs[-1] == {"type": "done"}
  318. assert client.tools == [
  319. {
  320. "type": "function",
  321. "function": {
  322. "name": "handoff_note",
  323. "description": "Registry-owned handoff tool.",
  324. "parameters": {
  325. "type": "object",
  326. "properties": {
  327. "message": {"type": "string"},
  328. "priority": {"type": "number"},
  329. },
  330. "required": ["message"],
  331. },
  332. },
  333. }
  334. ]
  335. @pytest.mark.asyncio
  336. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  337. request = DebugRunRequest(
  338. user_message="debug this",
  339. system_prompts=[],
  340. pre_messages=[],
  341. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  342. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  343. )
  344. ticks = iter([1.0, 1.123, 1.456])
  345. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  346. outputs = [message async for message in runtime.run(request)]
  347. assert [message["type"] for message in outputs] == [
  348. "session_started",
  349. "message_delta",
  350. "usage",
  351. "round_stats",
  352. "done",
  353. ]
  354. assert outputs[3] == {
  355. "type": "round_stats",
  356. "round_index": 1,
  357. "ttft_ms": 123,
  358. "elapsed_ms": 456,
  359. "prompt_tokens": 10,
  360. "completion_tokens": 20,
  361. "total_tokens": 30,
  362. "cached_tokens": 5,
  363. "had_event": False,
  364. }
  365. @pytest.mark.asyncio
  366. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  367. request = DebugRunRequest(
  368. user_message="debug this",
  369. system_prompts=[],
  370. pre_messages=[],
  371. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  372. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  373. )
  374. ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
  375. client = EventRoundStatsChatClient()
  376. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  377. outputs = [message async for message in runtime.run(request)]
  378. stats = [message for message in outputs if message["type"] == "round_stats"]
  379. assert client.calls == 2
  380. assert stats == [
  381. {
  382. "type": "round_stats",
  383. "round_index": 1,
  384. "ttft_ms": None,
  385. "elapsed_ms": 250,
  386. "prompt_tokens": 3,
  387. "completion_tokens": 0,
  388. "total_tokens": 3,
  389. "cached_tokens": 0,
  390. "had_event": True,
  391. },
  392. {
  393. "type": "round_stats",
  394. "round_index": 2,
  395. "ttft_ms": 50,
  396. "elapsed_ms": 200,
  397. "prompt_tokens": 4,
  398. "completion_tokens": 6,
  399. "total_tokens": 10,
  400. "cached_tokens": 0,
  401. "had_event": False,
  402. },
  403. ]