test_debug_runtime.py 24 KB

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