test_debug_runtime.py 25 KB

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