test_debug_runtime.py 28 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894
  1. import asyncio
  2. import importlib
  3. import json
  4. import logging
  5. from collections.abc import AsyncIterator
  6. from typing import Any
  7. import pytest
  8. from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
  9. from agent_lab.application.event_agent import EventAgentRequest
  10. from agent_lab.application.runtime import DebugRuntime
  11. from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
  12. from agent_lab.domain.events import ToolCallEvent
  13. from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
  14. def _runtime_queues_class():
  15. module = importlib.import_module("agent_lab.application.queues")
  16. return module.RuntimeQueues
  17. async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[str, Any]]:
  18. return [message async for message in stream]
  19. def _without_audit(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
  20. return [message for message in outputs if message["type"] != "audit"]
  21. def _message_types(outputs: list[dict[str, Any]]) -> list[str]:
  22. return [message["type"] for message in _without_audit(outputs)]
  23. async def _next_non_audit(
  24. stream: AsyncIterator[dict[str, Any]],
  25. ) -> dict[str, Any]:
  26. while True:
  27. message = await anext(stream)
  28. if message["type"] != "audit":
  29. return message
  30. async def _next_non_audit_from_queue(queues: Any) -> dict[str, Any]:
  31. while True:
  32. message = await queues.output.get()
  33. if message["type"] != "audit":
  34. return message
  35. class RecordingQueue(asyncio.Queue):
  36. def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
  37. super().__init__()
  38. self.name = name
  39. self.log = log
  40. async def put(self, item: Any) -> None:
  41. self.log.append((self.name, "put", self._describe(item)))
  42. await super().put(item)
  43. async def get(self) -> Any:
  44. item = await super().get()
  45. self.log.append((self.name, "get", self._describe(item)))
  46. return item
  47. def _describe(self, item: Any) -> str:
  48. if isinstance(item, ChatMessage):
  49. if item.role == "tool":
  50. return f"tool:{item.tool_call_id}"
  51. return item.role
  52. if isinstance(item, ToolCallEvent):
  53. return f"event:{item.name}:{item.id}"
  54. if isinstance(item, EventAgentRequest):
  55. events = ",".join(f"{event.name}:{event.id}" for event in item.events)
  56. return f"event_request:{events}"
  57. if isinstance(item, dict):
  58. return f"output:{item.get('type')}"
  59. return type(item).__name__
  60. class FakeChatClient:
  61. def __init__(self) -> None:
  62. self.calls = 0
  63. async def stream_chat(
  64. self,
  65. messages: list[ChatMessage],
  66. tools: list[dict],
  67. params: AgentParams,
  68. ) -> AsyncIterator[StreamItem]:
  69. self.calls += 1
  70. if self.calls == 1:
  71. yield StreamItem.event(
  72. ToolCallEvent(
  73. id="call_1",
  74. name="handoff_note",
  75. arguments={},
  76. raw_arguments="{}",
  77. )
  78. )
  79. return
  80. assert any(message.role == "tool" for message in messages)
  81. yield StreamItem.message_delta("final answer")
  82. class StrictHistoryChatClient:
  83. def __init__(self) -> None:
  84. self.calls = 0
  85. self.second_call_messages: list[ChatMessage] = []
  86. async def stream_chat(
  87. self,
  88. messages: list[ChatMessage],
  89. tools: list[dict],
  90. params: AgentParams,
  91. ) -> AsyncIterator[StreamItem]:
  92. self.calls += 1
  93. if self.calls == 1:
  94. yield StreamItem.event(
  95. ToolCallEvent(
  96. id="call_1",
  97. name="handoff_note",
  98. arguments={},
  99. raw_arguments="{}",
  100. )
  101. )
  102. return
  103. self.second_call_messages = list(messages)
  104. yield StreamItem.message_delta("final answer")
  105. class MultiEventChatClient:
  106. def __init__(self) -> None:
  107. self.calls = 0
  108. self.second_call_messages: list[ChatMessage] = []
  109. async def stream_chat(
  110. self,
  111. messages: list[ChatMessage],
  112. tools: list[dict],
  113. params: AgentParams,
  114. ) -> AsyncIterator[StreamItem]:
  115. self.calls += 1
  116. if self.calls == 1:
  117. yield StreamItem.message_delta("Checking events.")
  118. yield StreamItem.event(
  119. ToolCallEvent(
  120. id="call_1",
  121. name="handoff_note",
  122. arguments={"message": "ignored chat argument"},
  123. raw_arguments='{"message":"ignored chat argument"}',
  124. )
  125. )
  126. yield StreamItem.event(
  127. ToolCallEvent(
  128. id="call_2",
  129. name="audit_note",
  130. arguments={"message": "ignored chat argument"},
  131. raw_arguments='{"message":"ignored chat argument"}',
  132. )
  133. )
  134. return
  135. self.second_call_messages = list(messages)
  136. yield StreamItem.message_delta("Final answer.")
  137. class ToolCapturingChatClient:
  138. def __init__(self) -> None:
  139. self.tools: list[dict[str, Any]] = []
  140. async def stream_chat(
  141. self,
  142. messages: list[ChatMessage],
  143. tools: list[dict],
  144. params: AgentParams,
  145. ) -> AsyncIterator[StreamItem]:
  146. self.tools = list(tools)
  147. yield StreamItem.message_delta("final answer")
  148. class EventLoopLimitChatClient:
  149. def __init__(self) -> None:
  150. self.calls = 0
  151. self.tools_by_call: list[list[dict[str, Any]]] = []
  152. async def stream_chat(
  153. self,
  154. messages: list[ChatMessage],
  155. tools: list[dict],
  156. params: AgentParams,
  157. ) -> AsyncIterator[StreamItem]:
  158. self.calls += 1
  159. self.tools_by_call.append(list(tools))
  160. if self.calls == 1:
  161. yield StreamItem.event(
  162. ToolCallEvent(
  163. id="call_1",
  164. name="handoff_note",
  165. arguments={},
  166. raw_arguments="{}",
  167. )
  168. )
  169. return
  170. yield StreamItem.message_delta("final after event limit")
  171. class RoundStatsChatClient:
  172. async def stream_chat(
  173. self,
  174. messages: list[ChatMessage],
  175. tools: list[dict],
  176. params: AgentParams,
  177. ) -> AsyncIterator[StreamItem]:
  178. yield StreamItem.message_delta("hello")
  179. yield StreamItem.usage_item(
  180. TokenUsage(
  181. prompt_tokens=10,
  182. completion_tokens=20,
  183. total_tokens=30,
  184. cached_tokens=5,
  185. )
  186. )
  187. class EventRoundStatsChatClient:
  188. def __init__(self) -> None:
  189. self.calls = 0
  190. async def stream_chat(
  191. self,
  192. messages: list[ChatMessage],
  193. tools: list[dict],
  194. params: AgentParams,
  195. ) -> AsyncIterator[StreamItem]:
  196. self.calls += 1
  197. if self.calls == 1:
  198. yield StreamItem.event(
  199. ToolCallEvent(
  200. id="call_1",
  201. name="handoff_note",
  202. arguments={},
  203. raw_arguments="{}",
  204. )
  205. )
  206. yield StreamItem.usage_item(
  207. TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3)
  208. )
  209. return
  210. yield StreamItem.message_delta("final answer")
  211. yield StreamItem.usage_item(
  212. TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
  213. )
  214. class SlowAfterEventChatClient:
  215. def __init__(self) -> None:
  216. self.calls = 0
  217. self.event_seen = asyncio.Event()
  218. self.release_stream = asyncio.Event()
  219. async def stream_chat(
  220. self,
  221. messages: list[ChatMessage],
  222. tools: list[dict],
  223. params: AgentParams,
  224. ) -> AsyncIterator[StreamItem]:
  225. self.calls += 1
  226. if self.calls == 1:
  227. yield StreamItem.message_delta("Need event.")
  228. yield StreamItem.event(
  229. ToolCallEvent(
  230. id="call_1",
  231. name="handoff_note",
  232. arguments={},
  233. raw_arguments="{}",
  234. )
  235. )
  236. self.event_seen.set()
  237. await self.release_stream.wait()
  238. return
  239. yield StreamItem.message_delta("final answer")
  240. def test_runtime_queues_exposes_input_output_and_events_queues():
  241. RuntimeQueues = _runtime_queues_class()
  242. queues = RuntimeQueues()
  243. assert isinstance(queues.input, asyncio.Queue)
  244. assert isinstance(queues.output, asyncio.Queue)
  245. assert isinstance(queues.events, asyncio.Queue)
  246. assert queues.input is not queues.output
  247. assert queues.input is not queues.events
  248. assert queues.output is not queues.events
  249. @pytest.mark.asyncio
  250. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  251. request = DebugRunRequest(
  252. user_message="debug this",
  253. system_prompts=["You are a debugger."],
  254. pre_messages=[],
  255. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  256. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  257. )
  258. client = FakeChatClient()
  259. runtime = DebugRuntime(client)
  260. outputs = [message async for message in runtime.run(request)]
  261. business_outputs = _without_audit(outputs)
  262. assert client.calls == 2
  263. assert _message_types(outputs) == [
  264. "session_started",
  265. "event",
  266. "tool_result",
  267. "round_stats",
  268. "message_delta",
  269. "round_stats",
  270. "done",
  271. ]
  272. assert business_outputs[1]["event"]["name"] == "handoff_note"
  273. assert business_outputs[4]["content"] == "final answer"
  274. @pytest.mark.asyncio
  275. async def test_runtime_emits_audit_events_and_backend_logs(caplog):
  276. request = DebugRunRequest(
  277. user_message="debug this",
  278. system_prompts=[],
  279. pre_messages=[],
  280. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  281. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  282. )
  283. runtime = DebugRuntime(FakeChatClient())
  284. with caplog.at_level(logging.INFO, logger="agent_lab.application.runtime"):
  285. outputs = [message async for message in runtime.run(request)]
  286. audit_events = [
  287. message["event"]
  288. for message in outputs
  289. if message["type"] == "audit"
  290. ]
  291. assert audit_events == [
  292. "session_started",
  293. "chat_round_started",
  294. "chat_event_detected",
  295. "event_agent_completed",
  296. "chat_round_finished",
  297. "chat_round_started",
  298. "chat_round_finished",
  299. "session_finished",
  300. ]
  301. assert "chat_event_detected" in caplog.text
  302. assert "event_agent_completed" in caplog.text
  303. @pytest.mark.asyncio
  304. async def test_runtime_outputs_event_as_soon_as_chat_stream_detects_it():
  305. RuntimeQueues = _runtime_queues_class()
  306. queues = RuntimeQueues()
  307. client = SlowAfterEventChatClient()
  308. runtime = DebugRuntime(client, queues=queues)
  309. request = DebugRunRequest(
  310. user_message="debug this",
  311. system_prompts=[],
  312. pre_messages=[],
  313. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  314. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  315. )
  316. runtime.start(request)
  317. try:
  318. assert await _next_non_audit_from_queue(queues) == {"type": "session_started"}
  319. assert await _next_non_audit_from_queue(queues) == {
  320. "type": "message_delta",
  321. "content": "Need event.",
  322. }
  323. await asyncio.wait_for(client.event_seen.wait(), timeout=1)
  324. event_message = await asyncio.wait_for(
  325. _next_non_audit_from_queue(queues),
  326. timeout=0.2,
  327. )
  328. assert event_message["type"] == "event"
  329. assert event_message["event"]["name"] == "handoff_note"
  330. finally:
  331. client.release_stream.set()
  332. await runtime.aclose()
  333. @pytest.mark.asyncio
  334. async def test_runtime_batches_round_events_before_continuing_chat_agent():
  335. def resolve_from_history(
  336. event: ToolCallEvent,
  337. context: ToolExecutionContext,
  338. ) -> dict[str, Any]:
  339. return {"message": context.history[-1].content, "event": event.name}
  340. registry = ToolRegistry(
  341. [
  342. ToolDefinition(
  343. name="handoff_note",
  344. description="Send a handoff note.",
  345. parameters={
  346. "type": "object",
  347. "properties": {"message": {"type": "string"}},
  348. "required": ["message"],
  349. },
  350. handler=lambda event: {
  351. "tool": event.name,
  352. "message": event.arguments["message"],
  353. },
  354. argument_resolver=resolve_from_history,
  355. ),
  356. ToolDefinition(
  357. name="audit_note",
  358. description="Send an audit note.",
  359. parameters={
  360. "type": "object",
  361. "properties": {"message": {"type": "string"}},
  362. "required": ["message"],
  363. },
  364. handler=lambda event: {
  365. "tool": event.name,
  366. "message": event.arguments["message"],
  367. },
  368. argument_resolver=resolve_from_history,
  369. ),
  370. ]
  371. )
  372. request = DebugRunRequest(
  373. user_message="debug this",
  374. system_prompts=[],
  375. pre_messages=[],
  376. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  377. event_agent=EventAgentParams(
  378. enabled_tools=["handoff_note", "audit_note"],
  379. max_event_loops=2,
  380. ),
  381. )
  382. client = MultiEventChatClient()
  383. runtime = DebugRuntime(client, registry=registry)
  384. outputs = [message async for message in runtime.run(request)]
  385. assert _message_types(outputs) == [
  386. "session_started",
  387. "message_delta",
  388. "event",
  389. "event",
  390. "tool_result",
  391. "tool_result",
  392. "round_stats",
  393. "message_delta",
  394. "round_stats",
  395. "done",
  396. ]
  397. assert client.calls == 2
  398. assert [message.role for message in client.second_call_messages] == [
  399. "user",
  400. "assistant",
  401. "tool",
  402. "tool",
  403. "user",
  404. ]
  405. assistant_message = client.second_call_messages[1]
  406. assert assistant_message.content == "Checking events."
  407. assert assistant_message.tool_calls == [
  408. {
  409. "id": "call_1",
  410. "type": "function",
  411. "function": {"name": "handoff_note", "arguments": "{}"},
  412. },
  413. {
  414. "id": "call_2",
  415. "type": "function",
  416. "function": {"name": "audit_note", "arguments": "{}"},
  417. },
  418. ]
  419. assert [
  420. json.loads(message.content)
  421. for message in client.second_call_messages
  422. if message.role == "tool"
  423. ] == [
  424. {"tool": "handoff_note", "message": "Checking events."},
  425. {"tool": "audit_note", "message": "Checking events."},
  426. ]
  427. assert client.second_call_messages[-1].content == (
  428. "EventAgent results:\n"
  429. '{"tool": "handoff_note", "message": "Checking events."}\n'
  430. '{"tool": "audit_note", "message": "Checking events."}'
  431. )
  432. @pytest.mark.asyncio
  433. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  434. request = DebugRunRequest(
  435. user_message="debug this",
  436. system_prompts=[],
  437. pre_messages=[],
  438. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  439. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  440. )
  441. runtime = DebugRuntime(RoundStatsChatClient())
  442. queues = runtime.start(request)
  443. outputs: list[dict[str, Any]] = []
  444. while True:
  445. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  446. outputs.append(message)
  447. if message["type"] == "done":
  448. break
  449. assert _message_types(outputs) == [
  450. "session_started",
  451. "message_delta",
  452. "usage",
  453. "round_stats",
  454. "done",
  455. ]
  456. @pytest.mark.asyncio
  457. async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
  458. request = DebugRunRequest(
  459. user_message="debug this",
  460. system_prompts=[],
  461. pre_messages=[],
  462. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  463. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  464. )
  465. client = EventLoopLimitChatClient()
  466. runtime = DebugRuntime(client)
  467. outputs = [message async for message in runtime.run(request)]
  468. business_outputs = _without_audit(outputs)
  469. assert client.calls == 2
  470. assert client.tools_by_call[0][0]["function"]["name"] == "handoff_note"
  471. assert client.tools_by_call[1] == []
  472. assert _message_types(outputs) == [
  473. "session_started",
  474. "event",
  475. "tool_result",
  476. "round_stats",
  477. "message_delta",
  478. "round_stats",
  479. "done",
  480. ]
  481. assert business_outputs[4]["content"] == "final after event limit"
  482. @pytest.mark.asyncio
  483. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  484. RuntimeQueues = _runtime_queues_class()
  485. queues = RuntimeQueues()
  486. request = DebugRunRequest(
  487. user_message="debug this",
  488. system_prompts=[],
  489. pre_messages=[],
  490. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  491. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  492. )
  493. client = StrictHistoryChatClient()
  494. runtime = DebugRuntime(client, queues=queues)
  495. stream = runtime.run(request)
  496. assert await _next_non_audit(stream) == {"type": "session_started"}
  497. event_message = await _next_non_audit(stream)
  498. assert event_message["type"] == "event"
  499. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  500. remaining = [message async for message in stream]
  501. assert remaining[-1] == {"type": "done"}
  502. assert [message.role for message in client.second_call_messages] == [
  503. "user",
  504. "assistant",
  505. "tool",
  506. "user",
  507. "user",
  508. ]
  509. assert client.second_call_messages[2].tool_call_id == "call_1"
  510. assert client.second_call_messages[3].content.startswith("EventAgent results:\n")
  511. assert client.second_call_messages[4].content == "follow-up while tool runs"
  512. @pytest.mark.asyncio
  513. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  514. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  515. raise RuntimeError("boom")
  516. registry = ToolRegistry(
  517. [
  518. ToolDefinition(
  519. name="handoff_note",
  520. description="Broken handoff tool.",
  521. parameters={"type": "object"},
  522. handler=fail_tool,
  523. )
  524. ]
  525. )
  526. request = DebugRunRequest(
  527. user_message="debug this",
  528. system_prompts=[],
  529. pre_messages=[],
  530. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  531. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  532. )
  533. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  534. outputs = await asyncio.wait_for(
  535. _collect_outputs(runtime.run(request)),
  536. timeout=1,
  537. )
  538. business_outputs = _without_audit(outputs)
  539. assert _message_types(outputs) == [
  540. "session_started",
  541. "event",
  542. "tool_result",
  543. "round_stats",
  544. "message_delta",
  545. "round_stats",
  546. "done",
  547. ]
  548. assert json.loads(business_outputs[2]["message"]["content"]) == {
  549. "tool": "handoff_note",
  550. "error": "tool handler failed: boom",
  551. }
  552. @pytest.mark.asyncio
  553. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  554. RuntimeQueues = _runtime_queues_class()
  555. queue_log: list[tuple[str, str, str]] = []
  556. queues = RuntimeQueues(
  557. input=RecordingQueue("input", queue_log),
  558. output=RecordingQueue("output", queue_log),
  559. events=RecordingQueue("events", queue_log),
  560. )
  561. request = DebugRunRequest(
  562. user_message="debug this",
  563. system_prompts=["You are a debugger."],
  564. pre_messages=[],
  565. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  566. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  567. )
  568. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  569. outputs = [message async for message in runtime.run(request)]
  570. assert _message_types(outputs) == [
  571. "session_started",
  572. "event",
  573. "tool_result",
  574. "round_stats",
  575. "message_delta",
  576. "round_stats",
  577. "done",
  578. ]
  579. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  580. ("input", "get", "user")
  581. )
  582. assert queue_log.index(
  583. ("events", "put", "event_request:handoff_note:call_1")
  584. ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
  585. assert queue_log.index(
  586. ("events", "get", "event_request:handoff_note:call_1")
  587. ) < queue_log.index(("input", "put", "tool:call_1"))
  588. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  589. ("input", "get", "tool:call_1")
  590. )
  591. @pytest.mark.asyncio
  592. async def test_runtime_run_consumes_output_queue_in_stream_order():
  593. RuntimeQueues = _runtime_queues_class()
  594. queue_log: list[tuple[str, str, str]] = []
  595. queues = RuntimeQueues(
  596. input=RecordingQueue("input", queue_log),
  597. output=RecordingQueue("output", queue_log),
  598. events=RecordingQueue("events", queue_log),
  599. )
  600. request = DebugRunRequest(
  601. user_message="debug this",
  602. system_prompts=["You are a debugger."],
  603. pre_messages=[],
  604. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  605. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  606. )
  607. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  608. outputs = [message async for message in runtime.run(request)]
  609. assert _message_types(outputs) == [
  610. "session_started",
  611. "event",
  612. "tool_result",
  613. "round_stats",
  614. "message_delta",
  615. "round_stats",
  616. "done",
  617. ]
  618. output_puts = [
  619. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  620. ]
  621. output_gets = [
  622. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  623. ]
  624. assert [message for message in output_puts if message != "output:audit"] == [
  625. "output:session_started",
  626. "output:event",
  627. "output:tool_result",
  628. "output:round_stats",
  629. "output:message_delta",
  630. "output:round_stats",
  631. "output:done",
  632. ]
  633. assert output_gets == output_puts
  634. @pytest.mark.asyncio
  635. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  636. request = DebugRunRequest(
  637. user_message="debug this",
  638. system_prompts=["You are a debugger."],
  639. pre_messages=[],
  640. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  641. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  642. )
  643. client = StrictHistoryChatClient()
  644. runtime = DebugRuntime(client)
  645. outputs = [message async for message in runtime.run(request)]
  646. assert client.calls == 2
  647. assert [message.role for message in client.second_call_messages] == [
  648. "system",
  649. "user",
  650. "assistant",
  651. "tool",
  652. "user",
  653. ]
  654. assistant_message = client.second_call_messages[2]
  655. tool_message = client.second_call_messages[3]
  656. assert assistant_message.content == ""
  657. assert assistant_message.tool_calls == [
  658. {
  659. "id": "call_1",
  660. "type": "function",
  661. "function": {
  662. "name": "handoff_note",
  663. "arguments": "{}",
  664. },
  665. }
  666. ]
  667. assert tool_message.tool_call_id == "call_1"
  668. assert client.second_call_messages[4].content.startswith("EventAgent results:\n")
  669. assert outputs[-1] == {"type": "done"}
  670. @pytest.mark.asyncio
  671. async def test_runtime_passes_event_names_without_tool_parameters_to_chat_agent():
  672. registry = ToolRegistry(
  673. [
  674. ToolDefinition(
  675. name="handoff_note",
  676. description="Registry-owned handoff tool.",
  677. parameters={
  678. "type": "object",
  679. "properties": {
  680. "message": {"type": "string"},
  681. "priority": {"type": "number"},
  682. },
  683. "required": ["message"],
  684. },
  685. handler=lambda event: {"tool": event.name, "message": "handled"},
  686. )
  687. ]
  688. )
  689. request = DebugRunRequest(
  690. user_message="debug this",
  691. system_prompts=[],
  692. pre_messages=[],
  693. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  694. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  695. )
  696. client = ToolCapturingChatClient()
  697. runtime = DebugRuntime(client, registry=registry)
  698. outputs = [message async for message in runtime.run(request)]
  699. assert outputs[-1] == {"type": "done"}
  700. assert client.tools == [
  701. {
  702. "type": "function",
  703. "function": {
  704. "name": "handoff_note",
  705. "description": "Registry-owned handoff tool.",
  706. "parameters": {
  707. "type": "object",
  708. "properties": {},
  709. "additionalProperties": False,
  710. },
  711. },
  712. }
  713. ]
  714. @pytest.mark.asyncio
  715. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  716. request = DebugRunRequest(
  717. user_message="debug this",
  718. system_prompts=[],
  719. pre_messages=[],
  720. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  721. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  722. )
  723. ticks = iter([1.0, 1.123, 1.456])
  724. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  725. outputs = [message async for message in runtime.run(request)]
  726. business_outputs = _without_audit(outputs)
  727. assert _message_types(outputs) == [
  728. "session_started",
  729. "message_delta",
  730. "usage",
  731. "round_stats",
  732. "done",
  733. ]
  734. assert business_outputs[3] == {
  735. "type": "round_stats",
  736. "round_index": 1,
  737. "ttft_ms": 123,
  738. "elapsed_ms": 456,
  739. "prompt_tokens": 10,
  740. "completion_tokens": 20,
  741. "total_tokens": 30,
  742. "cached_tokens": 5,
  743. "had_event": False,
  744. }
  745. @pytest.mark.asyncio
  746. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  747. request = DebugRunRequest(
  748. user_message="debug this",
  749. system_prompts=[],
  750. pre_messages=[],
  751. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  752. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  753. )
  754. ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
  755. client = EventRoundStatsChatClient()
  756. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  757. outputs = [message async for message in runtime.run(request)]
  758. stats = [message for message in outputs if message["type"] == "round_stats"]
  759. assert client.calls == 2
  760. assert stats == [
  761. {
  762. "type": "round_stats",
  763. "round_index": 1,
  764. "ttft_ms": None,
  765. "elapsed_ms": 250,
  766. "prompt_tokens": 3,
  767. "completion_tokens": 0,
  768. "total_tokens": 3,
  769. "cached_tokens": 0,
  770. "had_event": True,
  771. },
  772. {
  773. "type": "round_stats",
  774. "round_index": 2,
  775. "ttft_ms": 50,
  776. "elapsed_ms": 200,
  777. "prompt_tokens": 4,
  778. "completion_tokens": 6,
  779. "total_tokens": 10,
  780. "cached_tokens": 0,
  781. "had_event": False,
  782. },
  783. ]