test_debug_runtime.py 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517
  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_start_returns_queues_for_downstream_output_consumer():
  170. request = DebugRunRequest(
  171. user_message="debug this",
  172. system_prompts=[],
  173. pre_messages=[],
  174. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  175. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  176. )
  177. runtime = DebugRuntime(RoundStatsChatClient())
  178. queues = runtime.start(request)
  179. outputs: list[dict[str, Any]] = []
  180. while True:
  181. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  182. outputs.append(message)
  183. if message["type"] == "done":
  184. break
  185. assert [message["type"] for message in outputs] == [
  186. "session_started",
  187. "message_delta",
  188. "usage",
  189. "round_stats",
  190. "done",
  191. ]
  192. @pytest.mark.asyncio
  193. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  194. RuntimeQueues = _runtime_queues_class()
  195. queues = RuntimeQueues()
  196. request = DebugRunRequest(
  197. user_message="debug this",
  198. system_prompts=[],
  199. pre_messages=[],
  200. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  201. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  202. )
  203. client = StrictHistoryChatClient()
  204. runtime = DebugRuntime(client, queues=queues)
  205. stream = runtime.run(request)
  206. assert await anext(stream) == {"type": "session_started"}
  207. event_message = await anext(stream)
  208. assert event_message["type"] == "event"
  209. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  210. remaining = [message async for message in stream]
  211. assert remaining[-1] == {"type": "done"}
  212. assert [message.role for message in client.second_call_messages] == [
  213. "user",
  214. "assistant",
  215. "tool",
  216. "user",
  217. ]
  218. assert client.second_call_messages[2].tool_call_id == "call_1"
  219. assert client.second_call_messages[3].content == "follow-up while tool runs"
  220. @pytest.mark.asyncio
  221. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  222. RuntimeQueues = _runtime_queues_class()
  223. queue_log: list[tuple[str, str, str]] = []
  224. queues = RuntimeQueues(
  225. input=RecordingQueue("input", queue_log),
  226. output=RecordingQueue("output", queue_log),
  227. events=RecordingQueue("events", queue_log),
  228. )
  229. request = DebugRunRequest(
  230. user_message="debug this",
  231. system_prompts=["You are a debugger."],
  232. pre_messages=[],
  233. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  234. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  235. )
  236. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  237. outputs = [message async for message in runtime.run(request)]
  238. assert [message["type"] for message in outputs] == [
  239. "session_started",
  240. "event",
  241. "tool_result",
  242. "round_stats",
  243. "message_delta",
  244. "round_stats",
  245. "done",
  246. ]
  247. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  248. ("input", "get", "user")
  249. )
  250. assert queue_log.index(("events", "put", "event:handoff_note:call_1")) < queue_log.index(
  251. ("events", "get", "event:handoff_note:call_1")
  252. )
  253. assert queue_log.index(("events", "get", "event:handoff_note:call_1")) < queue_log.index(
  254. ("input", "put", "tool:call_1")
  255. )
  256. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  257. ("input", "get", "tool:call_1")
  258. )
  259. @pytest.mark.asyncio
  260. async def test_runtime_run_consumes_output_queue_in_stream_order():
  261. RuntimeQueues = _runtime_queues_class()
  262. queue_log: list[tuple[str, str, str]] = []
  263. queues = RuntimeQueues(
  264. input=RecordingQueue("input", queue_log),
  265. output=RecordingQueue("output", queue_log),
  266. events=RecordingQueue("events", queue_log),
  267. )
  268. request = DebugRunRequest(
  269. user_message="debug this",
  270. system_prompts=["You are a debugger."],
  271. pre_messages=[],
  272. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  273. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  274. )
  275. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  276. outputs = [message async for message in runtime.run(request)]
  277. assert [message["type"] for message in outputs] == [
  278. "session_started",
  279. "event",
  280. "tool_result",
  281. "round_stats",
  282. "message_delta",
  283. "round_stats",
  284. "done",
  285. ]
  286. output_puts = [
  287. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  288. ]
  289. output_gets = [
  290. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  291. ]
  292. assert output_puts == [
  293. "output:session_started",
  294. "output:event",
  295. "output:tool_result",
  296. "output:round_stats",
  297. "output:message_delta",
  298. "output:round_stats",
  299. "output:done",
  300. ]
  301. assert output_gets == output_puts
  302. @pytest.mark.asyncio
  303. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  304. request = DebugRunRequest(
  305. user_message="debug this",
  306. system_prompts=["You are a debugger."],
  307. pre_messages=[],
  308. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  309. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  310. )
  311. client = StrictHistoryChatClient()
  312. runtime = DebugRuntime(client)
  313. outputs = [message async for message in runtime.run(request)]
  314. assert client.calls == 2
  315. assert [message.role for message in client.second_call_messages] == [
  316. "system",
  317. "user",
  318. "assistant",
  319. "tool",
  320. ]
  321. assistant_message = client.second_call_messages[2]
  322. tool_message = client.second_call_messages[3]
  323. assert assistant_message.content == ""
  324. assert assistant_message.tool_calls == [
  325. {
  326. "id": "call_1",
  327. "type": "function",
  328. "function": {
  329. "name": "handoff_note",
  330. "arguments": '{"message":"need event agent"}',
  331. },
  332. }
  333. ]
  334. assert tool_message.tool_call_id == "call_1"
  335. assert outputs[-1] == {"type": "done"}
  336. @pytest.mark.asyncio
  337. async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent():
  338. registry = ToolRegistry(
  339. [
  340. ToolDefinition(
  341. name="handoff_note",
  342. description="Registry-owned handoff tool.",
  343. parameters={
  344. "type": "object",
  345. "properties": {
  346. "message": {"type": "string"},
  347. "priority": {"type": "number"},
  348. },
  349. "required": ["message"],
  350. },
  351. handler=lambda event: {"tool": event.name, "message": "handled"},
  352. )
  353. ]
  354. )
  355. request = DebugRunRequest(
  356. user_message="debug this",
  357. system_prompts=[],
  358. pre_messages=[],
  359. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  360. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  361. )
  362. client = ToolCapturingChatClient()
  363. runtime = DebugRuntime(client, registry=registry)
  364. outputs = [message async for message in runtime.run(request)]
  365. assert outputs[-1] == {"type": "done"}
  366. assert client.tools == [
  367. {
  368. "type": "function",
  369. "function": {
  370. "name": "handoff_note",
  371. "description": "Registry-owned handoff tool.",
  372. "parameters": {
  373. "type": "object",
  374. "properties": {
  375. "message": {"type": "string"},
  376. "priority": {"type": "number"},
  377. },
  378. "required": ["message"],
  379. },
  380. },
  381. }
  382. ]
  383. @pytest.mark.asyncio
  384. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  385. request = DebugRunRequest(
  386. user_message="debug this",
  387. system_prompts=[],
  388. pre_messages=[],
  389. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  390. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  391. )
  392. ticks = iter([1.0, 1.123, 1.456])
  393. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  394. outputs = [message async for message in runtime.run(request)]
  395. assert [message["type"] for message in outputs] == [
  396. "session_started",
  397. "message_delta",
  398. "usage",
  399. "round_stats",
  400. "done",
  401. ]
  402. assert outputs[3] == {
  403. "type": "round_stats",
  404. "round_index": 1,
  405. "ttft_ms": 123,
  406. "elapsed_ms": 456,
  407. "prompt_tokens": 10,
  408. "completion_tokens": 20,
  409. "total_tokens": 30,
  410. "cached_tokens": 5,
  411. "had_event": False,
  412. }
  413. @pytest.mark.asyncio
  414. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  415. request = DebugRunRequest(
  416. user_message="debug this",
  417. system_prompts=[],
  418. pre_messages=[],
  419. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  420. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  421. )
  422. ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
  423. client = EventRoundStatsChatClient()
  424. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  425. outputs = [message async for message in runtime.run(request)]
  426. stats = [message for message in outputs if message["type"] == "round_stats"]
  427. assert client.calls == 2
  428. assert stats == [
  429. {
  430. "type": "round_stats",
  431. "round_index": 1,
  432. "ttft_ms": None,
  433. "elapsed_ms": 250,
  434. "prompt_tokens": 3,
  435. "completion_tokens": 0,
  436. "total_tokens": 3,
  437. "cached_tokens": 0,
  438. "had_event": True,
  439. },
  440. {
  441. "type": "round_stats",
  442. "round_index": 2,
  443. "ttft_ms": 50,
  444. "elapsed_ms": 200,
  445. "prompt_tokens": 4,
  446. "completion_tokens": 6,
  447. "total_tokens": 10,
  448. "cached_tokens": 0,
  449. "had_event": False,
  450. },
  451. ]