test_debug_runtime.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566
  1. import asyncio
  2. import importlib
  3. import json
  4. from collections.abc import AsyncIterator
  5. from typing import Any
  6. import pytest
  7. from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
  8. from agent_lab.application.runtime import DebugRuntime
  9. from agent_lab.application.tools import ToolDefinition, ToolRegistry
  10. from agent_lab.domain.events import ToolCallEvent
  11. from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
  12. def _runtime_queues_class():
  13. module = importlib.import_module("agent_lab.application.queues")
  14. return module.RuntimeQueues
  15. async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[str, Any]]:
  16. return [message async for message in stream]
  17. class RecordingQueue(asyncio.Queue):
  18. def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
  19. super().__init__()
  20. self.name = name
  21. self.log = log
  22. async def put(self, item: Any) -> None:
  23. self.log.append((self.name, "put", self._describe(item)))
  24. await super().put(item)
  25. async def get(self) -> Any:
  26. item = await super().get()
  27. self.log.append((self.name, "get", self._describe(item)))
  28. return item
  29. def _describe(self, item: Any) -> str:
  30. if isinstance(item, ChatMessage):
  31. if item.role == "tool":
  32. return f"tool:{item.tool_call_id}"
  33. return item.role
  34. if isinstance(item, ToolCallEvent):
  35. return f"event:{item.name}:{item.id}"
  36. if isinstance(item, dict):
  37. return f"output:{item.get('type')}"
  38. return type(item).__name__
  39. class FakeChatClient:
  40. def __init__(self) -> None:
  41. self.calls = 0
  42. async def stream_chat(
  43. self,
  44. messages: list[ChatMessage],
  45. tools: list[dict],
  46. params: AgentParams,
  47. ) -> AsyncIterator[StreamItem]:
  48. self.calls += 1
  49. if self.calls == 1:
  50. yield StreamItem.event(
  51. ToolCallEvent(
  52. id="call_1",
  53. name="handoff_note",
  54. arguments={"message": "need event agent"},
  55. raw_arguments='{"message":"need event agent"}',
  56. )
  57. )
  58. return
  59. assert any(message.role == "tool" for message in messages)
  60. yield StreamItem.message_delta("final answer")
  61. class StrictHistoryChatClient:
  62. def __init__(self) -> None:
  63. self.calls = 0
  64. self.second_call_messages: list[ChatMessage] = []
  65. async def stream_chat(
  66. self,
  67. messages: list[ChatMessage],
  68. tools: list[dict],
  69. params: AgentParams,
  70. ) -> AsyncIterator[StreamItem]:
  71. self.calls += 1
  72. if self.calls == 1:
  73. yield StreamItem.event(
  74. ToolCallEvent(
  75. id="call_1",
  76. name="handoff_note",
  77. arguments={"message": "need event agent"},
  78. raw_arguments='{"message":"need event agent"}',
  79. )
  80. )
  81. return
  82. self.second_call_messages = list(messages)
  83. yield StreamItem.message_delta("final answer")
  84. class ToolCapturingChatClient:
  85. def __init__(self) -> None:
  86. self.tools: list[dict[str, Any]] = []
  87. async def stream_chat(
  88. self,
  89. messages: list[ChatMessage],
  90. tools: list[dict],
  91. params: AgentParams,
  92. ) -> AsyncIterator[StreamItem]:
  93. self.tools = list(tools)
  94. yield StreamItem.message_delta("final answer")
  95. class RoundStatsChatClient:
  96. async def stream_chat(
  97. self,
  98. messages: list[ChatMessage],
  99. tools: list[dict],
  100. params: AgentParams,
  101. ) -> AsyncIterator[StreamItem]:
  102. yield StreamItem.message_delta("hello")
  103. yield StreamItem.usage_item(
  104. TokenUsage(
  105. prompt_tokens=10,
  106. completion_tokens=20,
  107. total_tokens=30,
  108. cached_tokens=5,
  109. )
  110. )
  111. class EventRoundStatsChatClient:
  112. def __init__(self) -> None:
  113. self.calls = 0
  114. async def stream_chat(
  115. self,
  116. messages: list[ChatMessage],
  117. tools: list[dict],
  118. params: AgentParams,
  119. ) -> AsyncIterator[StreamItem]:
  120. self.calls += 1
  121. if self.calls == 1:
  122. yield StreamItem.event(
  123. ToolCallEvent(
  124. id="call_1",
  125. name="handoff_note",
  126. arguments={"message": "need event agent"},
  127. raw_arguments='{"message":"need event agent"}',
  128. )
  129. )
  130. yield StreamItem.usage_item(
  131. TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3)
  132. )
  133. return
  134. yield StreamItem.message_delta("final answer")
  135. yield StreamItem.usage_item(
  136. TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
  137. )
  138. def test_runtime_queues_exposes_input_output_and_events_queues():
  139. RuntimeQueues = _runtime_queues_class()
  140. queues = RuntimeQueues()
  141. assert isinstance(queues.input, asyncio.Queue)
  142. assert isinstance(queues.output, asyncio.Queue)
  143. assert isinstance(queues.events, asyncio.Queue)
  144. assert queues.input is not queues.output
  145. assert queues.input is not queues.events
  146. assert queues.output is not queues.events
  147. @pytest.mark.asyncio
  148. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  149. request = DebugRunRequest(
  150. user_message="debug this",
  151. system_prompts=["You are a debugger."],
  152. pre_messages=[],
  153. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  154. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  155. )
  156. client = FakeChatClient()
  157. runtime = DebugRuntime(client)
  158. outputs = [message async for message in runtime.run(request)]
  159. assert client.calls == 2
  160. assert [message["type"] for message in outputs] == [
  161. "session_started",
  162. "event",
  163. "tool_result",
  164. "round_stats",
  165. "message_delta",
  166. "round_stats",
  167. "done",
  168. ]
  169. assert outputs[1]["event"]["name"] == "handoff_note"
  170. assert outputs[4]["content"] == "final answer"
  171. @pytest.mark.asyncio
  172. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  173. request = DebugRunRequest(
  174. user_message="debug this",
  175. system_prompts=[],
  176. pre_messages=[],
  177. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  178. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  179. )
  180. runtime = DebugRuntime(RoundStatsChatClient())
  181. queues = runtime.start(request)
  182. outputs: list[dict[str, Any]] = []
  183. while True:
  184. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  185. outputs.append(message)
  186. if message["type"] == "done":
  187. break
  188. assert [message["type"] for message in outputs] == [
  189. "session_started",
  190. "message_delta",
  191. "usage",
  192. "round_stats",
  193. "done",
  194. ]
  195. @pytest.mark.asyncio
  196. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  197. RuntimeQueues = _runtime_queues_class()
  198. queues = RuntimeQueues()
  199. request = DebugRunRequest(
  200. user_message="debug this",
  201. system_prompts=[],
  202. pre_messages=[],
  203. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  204. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  205. )
  206. client = StrictHistoryChatClient()
  207. runtime = DebugRuntime(client, queues=queues)
  208. stream = runtime.run(request)
  209. assert await anext(stream) == {"type": "session_started"}
  210. event_message = await anext(stream)
  211. assert event_message["type"] == "event"
  212. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  213. remaining = [message async for message in stream]
  214. assert remaining[-1] == {"type": "done"}
  215. assert [message.role for message in client.second_call_messages] == [
  216. "user",
  217. "assistant",
  218. "tool",
  219. "user",
  220. ]
  221. assert client.second_call_messages[2].tool_call_id == "call_1"
  222. assert client.second_call_messages[3].content == "follow-up while tool runs"
  223. @pytest.mark.asyncio
  224. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  225. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  226. raise RuntimeError("boom")
  227. registry = ToolRegistry(
  228. [
  229. ToolDefinition(
  230. name="handoff_note",
  231. description="Broken handoff tool.",
  232. parameters={"type": "object"},
  233. handler=fail_tool,
  234. )
  235. ]
  236. )
  237. request = DebugRunRequest(
  238. user_message="debug this",
  239. system_prompts=[],
  240. pre_messages=[],
  241. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  242. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  243. )
  244. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  245. outputs = await asyncio.wait_for(
  246. _collect_outputs(runtime.run(request)),
  247. timeout=1,
  248. )
  249. assert [message["type"] for message in outputs] == [
  250. "session_started",
  251. "event",
  252. "tool_result",
  253. "round_stats",
  254. "message_delta",
  255. "round_stats",
  256. "done",
  257. ]
  258. assert json.loads(outputs[2]["message"]["content"]) == {
  259. "tool": "handoff_note",
  260. "error": "tool handler failed: boom",
  261. }
  262. @pytest.mark.asyncio
  263. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  264. RuntimeQueues = _runtime_queues_class()
  265. queue_log: list[tuple[str, str, str]] = []
  266. queues = RuntimeQueues(
  267. input=RecordingQueue("input", queue_log),
  268. output=RecordingQueue("output", queue_log),
  269. events=RecordingQueue("events", queue_log),
  270. )
  271. request = DebugRunRequest(
  272. user_message="debug this",
  273. system_prompts=["You are a debugger."],
  274. pre_messages=[],
  275. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  276. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  277. )
  278. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  279. outputs = [message async for message in runtime.run(request)]
  280. assert [message["type"] for message in outputs] == [
  281. "session_started",
  282. "event",
  283. "tool_result",
  284. "round_stats",
  285. "message_delta",
  286. "round_stats",
  287. "done",
  288. ]
  289. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  290. ("input", "get", "user")
  291. )
  292. assert queue_log.index(("events", "put", "event:handoff_note:call_1")) < queue_log.index(
  293. ("events", "get", "event:handoff_note:call_1")
  294. )
  295. assert queue_log.index(("events", "get", "event:handoff_note:call_1")) < queue_log.index(
  296. ("input", "put", "tool:call_1")
  297. )
  298. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  299. ("input", "get", "tool:call_1")
  300. )
  301. @pytest.mark.asyncio
  302. async def test_runtime_run_consumes_output_queue_in_stream_order():
  303. RuntimeQueues = _runtime_queues_class()
  304. queue_log: list[tuple[str, str, str]] = []
  305. queues = RuntimeQueues(
  306. input=RecordingQueue("input", queue_log),
  307. output=RecordingQueue("output", queue_log),
  308. events=RecordingQueue("events", queue_log),
  309. )
  310. request = DebugRunRequest(
  311. user_message="debug this",
  312. system_prompts=["You are a debugger."],
  313. pre_messages=[],
  314. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  315. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  316. )
  317. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  318. outputs = [message async for message in runtime.run(request)]
  319. assert [message["type"] for message in outputs] == [
  320. "session_started",
  321. "event",
  322. "tool_result",
  323. "round_stats",
  324. "message_delta",
  325. "round_stats",
  326. "done",
  327. ]
  328. output_puts = [
  329. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  330. ]
  331. output_gets = [
  332. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  333. ]
  334. assert output_puts == [
  335. "output:session_started",
  336. "output:event",
  337. "output:tool_result",
  338. "output:round_stats",
  339. "output:message_delta",
  340. "output:round_stats",
  341. "output:done",
  342. ]
  343. assert output_gets == output_puts
  344. @pytest.mark.asyncio
  345. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  346. request = DebugRunRequest(
  347. user_message="debug this",
  348. system_prompts=["You are a debugger."],
  349. pre_messages=[],
  350. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  351. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  352. )
  353. client = StrictHistoryChatClient()
  354. runtime = DebugRuntime(client)
  355. outputs = [message async for message in runtime.run(request)]
  356. assert client.calls == 2
  357. assert [message.role for message in client.second_call_messages] == [
  358. "system",
  359. "user",
  360. "assistant",
  361. "tool",
  362. ]
  363. assistant_message = client.second_call_messages[2]
  364. tool_message = client.second_call_messages[3]
  365. assert assistant_message.content == ""
  366. assert assistant_message.tool_calls == [
  367. {
  368. "id": "call_1",
  369. "type": "function",
  370. "function": {
  371. "name": "handoff_note",
  372. "arguments": '{"message":"need event agent"}',
  373. },
  374. }
  375. ]
  376. assert tool_message.tool_call_id == "call_1"
  377. assert outputs[-1] == {"type": "done"}
  378. @pytest.mark.asyncio
  379. async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent():
  380. registry = ToolRegistry(
  381. [
  382. ToolDefinition(
  383. name="handoff_note",
  384. description="Registry-owned handoff tool.",
  385. parameters={
  386. "type": "object",
  387. "properties": {
  388. "message": {"type": "string"},
  389. "priority": {"type": "number"},
  390. },
  391. "required": ["message"],
  392. },
  393. handler=lambda event: {"tool": event.name, "message": "handled"},
  394. )
  395. ]
  396. )
  397. request = DebugRunRequest(
  398. user_message="debug this",
  399. system_prompts=[],
  400. pre_messages=[],
  401. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  402. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  403. )
  404. client = ToolCapturingChatClient()
  405. runtime = DebugRuntime(client, registry=registry)
  406. outputs = [message async for message in runtime.run(request)]
  407. assert outputs[-1] == {"type": "done"}
  408. assert client.tools == [
  409. {
  410. "type": "function",
  411. "function": {
  412. "name": "handoff_note",
  413. "description": "Registry-owned handoff tool.",
  414. "parameters": {
  415. "type": "object",
  416. "properties": {
  417. "message": {"type": "string"},
  418. "priority": {"type": "number"},
  419. },
  420. "required": ["message"],
  421. },
  422. },
  423. }
  424. ]
  425. @pytest.mark.asyncio
  426. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  427. request = DebugRunRequest(
  428. user_message="debug this",
  429. system_prompts=[],
  430. pre_messages=[],
  431. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  432. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  433. )
  434. ticks = iter([1.0, 1.123, 1.456])
  435. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  436. outputs = [message async for message in runtime.run(request)]
  437. assert [message["type"] for message in outputs] == [
  438. "session_started",
  439. "message_delta",
  440. "usage",
  441. "round_stats",
  442. "done",
  443. ]
  444. assert outputs[3] == {
  445. "type": "round_stats",
  446. "round_index": 1,
  447. "ttft_ms": 123,
  448. "elapsed_ms": 456,
  449. "prompt_tokens": 10,
  450. "completion_tokens": 20,
  451. "total_tokens": 30,
  452. "cached_tokens": 5,
  453. "had_event": False,
  454. }
  455. @pytest.mark.asyncio
  456. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  457. request = DebugRunRequest(
  458. user_message="debug this",
  459. system_prompts=[],
  460. pre_messages=[],
  461. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  462. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  463. )
  464. ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
  465. client = EventRoundStatsChatClient()
  466. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  467. outputs = [message async for message in runtime.run(request)]
  468. stats = [message for message in outputs if message["type"] == "round_stats"]
  469. assert client.calls == 2
  470. assert stats == [
  471. {
  472. "type": "round_stats",
  473. "round_index": 1,
  474. "ttft_ms": None,
  475. "elapsed_ms": 250,
  476. "prompt_tokens": 3,
  477. "completion_tokens": 0,
  478. "total_tokens": 3,
  479. "cached_tokens": 0,
  480. "had_event": True,
  481. },
  482. {
  483. "type": "round_stats",
  484. "round_index": 2,
  485. "ttft_ms": 50,
  486. "elapsed_ms": 200,
  487. "prompt_tokens": 4,
  488. "completion_tokens": 6,
  489. "total_tokens": 10,
  490. "cached_tokens": 0,
  491. "had_event": False,
  492. },
  493. ]