test_debug_runtime.py 28 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884
  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. 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. history: list[ChatMessage],
  338. ) -> dict[str, Any]:
  339. return {"message": 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. ]
  404. assistant_message = client.second_call_messages[1]
  405. assert assistant_message.content == "Checking events."
  406. assert assistant_message.tool_calls == [
  407. {
  408. "id": "call_1",
  409. "type": "function",
  410. "function": {"name": "handoff_note", "arguments": "{}"},
  411. },
  412. {
  413. "id": "call_2",
  414. "type": "function",
  415. "function": {"name": "audit_note", "arguments": "{}"},
  416. },
  417. ]
  418. assert [
  419. json.loads(message.content)
  420. for message in client.second_call_messages
  421. if message.role == "tool"
  422. ] == [
  423. {"tool": "handoff_note", "message": "Checking events."},
  424. {"tool": "audit_note", "message": "Checking events."},
  425. ]
  426. @pytest.mark.asyncio
  427. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  428. request = DebugRunRequest(
  429. user_message="debug this",
  430. system_prompts=[],
  431. pre_messages=[],
  432. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  433. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  434. )
  435. runtime = DebugRuntime(RoundStatsChatClient())
  436. queues = runtime.start(request)
  437. outputs: list[dict[str, Any]] = []
  438. while True:
  439. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  440. outputs.append(message)
  441. if message["type"] == "done":
  442. break
  443. assert _message_types(outputs) == [
  444. "session_started",
  445. "message_delta",
  446. "usage",
  447. "round_stats",
  448. "done",
  449. ]
  450. @pytest.mark.asyncio
  451. async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
  452. request = DebugRunRequest(
  453. user_message="debug this",
  454. system_prompts=[],
  455. pre_messages=[],
  456. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  457. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  458. )
  459. client = EventLoopLimitChatClient()
  460. runtime = DebugRuntime(client)
  461. outputs = [message async for message in runtime.run(request)]
  462. business_outputs = _without_audit(outputs)
  463. assert client.calls == 2
  464. assert client.tools_by_call[0][0]["function"]["name"] == "handoff_note"
  465. assert client.tools_by_call[1] == []
  466. assert _message_types(outputs) == [
  467. "session_started",
  468. "event",
  469. "tool_result",
  470. "round_stats",
  471. "message_delta",
  472. "round_stats",
  473. "done",
  474. ]
  475. assert business_outputs[4]["content"] == "final after event limit"
  476. @pytest.mark.asyncio
  477. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  478. RuntimeQueues = _runtime_queues_class()
  479. queues = RuntimeQueues()
  480. request = DebugRunRequest(
  481. user_message="debug this",
  482. system_prompts=[],
  483. pre_messages=[],
  484. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  485. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  486. )
  487. client = StrictHistoryChatClient()
  488. runtime = DebugRuntime(client, queues=queues)
  489. stream = runtime.run(request)
  490. assert await _next_non_audit(stream) == {"type": "session_started"}
  491. event_message = await _next_non_audit(stream)
  492. assert event_message["type"] == "event"
  493. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  494. remaining = [message async for message in stream]
  495. assert remaining[-1] == {"type": "done"}
  496. assert [message.role for message in client.second_call_messages] == [
  497. "user",
  498. "assistant",
  499. "tool",
  500. "user",
  501. ]
  502. assert client.second_call_messages[2].tool_call_id == "call_1"
  503. assert client.second_call_messages[3].content == "follow-up while tool runs"
  504. @pytest.mark.asyncio
  505. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  506. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  507. raise RuntimeError("boom")
  508. registry = ToolRegistry(
  509. [
  510. ToolDefinition(
  511. name="handoff_note",
  512. description="Broken handoff tool.",
  513. parameters={"type": "object"},
  514. handler=fail_tool,
  515. )
  516. ]
  517. )
  518. request = DebugRunRequest(
  519. user_message="debug this",
  520. system_prompts=[],
  521. pre_messages=[],
  522. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  523. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  524. )
  525. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  526. outputs = await asyncio.wait_for(
  527. _collect_outputs(runtime.run(request)),
  528. timeout=1,
  529. )
  530. business_outputs = _without_audit(outputs)
  531. assert _message_types(outputs) == [
  532. "session_started",
  533. "event",
  534. "tool_result",
  535. "round_stats",
  536. "message_delta",
  537. "round_stats",
  538. "done",
  539. ]
  540. assert json.loads(business_outputs[2]["message"]["content"]) == {
  541. "tool": "handoff_note",
  542. "error": "tool handler failed: boom",
  543. }
  544. @pytest.mark.asyncio
  545. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  546. RuntimeQueues = _runtime_queues_class()
  547. queue_log: list[tuple[str, str, str]] = []
  548. queues = RuntimeQueues(
  549. input=RecordingQueue("input", queue_log),
  550. output=RecordingQueue("output", queue_log),
  551. events=RecordingQueue("events", queue_log),
  552. )
  553. request = DebugRunRequest(
  554. user_message="debug this",
  555. system_prompts=["You are a debugger."],
  556. pre_messages=[],
  557. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  558. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  559. )
  560. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  561. outputs = [message async for message in runtime.run(request)]
  562. assert _message_types(outputs) == [
  563. "session_started",
  564. "event",
  565. "tool_result",
  566. "round_stats",
  567. "message_delta",
  568. "round_stats",
  569. "done",
  570. ]
  571. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  572. ("input", "get", "user")
  573. )
  574. assert queue_log.index(
  575. ("events", "put", "event_request:handoff_note:call_1")
  576. ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
  577. assert queue_log.index(
  578. ("events", "get", "event_request:handoff_note:call_1")
  579. ) < queue_log.index(("input", "put", "tool:call_1"))
  580. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  581. ("input", "get", "tool:call_1")
  582. )
  583. @pytest.mark.asyncio
  584. async def test_runtime_run_consumes_output_queue_in_stream_order():
  585. RuntimeQueues = _runtime_queues_class()
  586. queue_log: list[tuple[str, str, str]] = []
  587. queues = RuntimeQueues(
  588. input=RecordingQueue("input", queue_log),
  589. output=RecordingQueue("output", queue_log),
  590. events=RecordingQueue("events", queue_log),
  591. )
  592. request = DebugRunRequest(
  593. user_message="debug this",
  594. system_prompts=["You are a debugger."],
  595. pre_messages=[],
  596. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  597. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  598. )
  599. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  600. outputs = [message async for message in runtime.run(request)]
  601. assert _message_types(outputs) == [
  602. "session_started",
  603. "event",
  604. "tool_result",
  605. "round_stats",
  606. "message_delta",
  607. "round_stats",
  608. "done",
  609. ]
  610. output_puts = [
  611. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  612. ]
  613. output_gets = [
  614. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  615. ]
  616. assert [message for message in output_puts if message != "output:audit"] == [
  617. "output:session_started",
  618. "output:event",
  619. "output:tool_result",
  620. "output:round_stats",
  621. "output:message_delta",
  622. "output:round_stats",
  623. "output:done",
  624. ]
  625. assert output_gets == output_puts
  626. @pytest.mark.asyncio
  627. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  628. request = DebugRunRequest(
  629. user_message="debug this",
  630. system_prompts=["You are a debugger."],
  631. pre_messages=[],
  632. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  633. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  634. )
  635. client = StrictHistoryChatClient()
  636. runtime = DebugRuntime(client)
  637. outputs = [message async for message in runtime.run(request)]
  638. assert client.calls == 2
  639. assert [message.role for message in client.second_call_messages] == [
  640. "system",
  641. "user",
  642. "assistant",
  643. "tool",
  644. ]
  645. assistant_message = client.second_call_messages[2]
  646. tool_message = client.second_call_messages[3]
  647. assert assistant_message.content == ""
  648. assert assistant_message.tool_calls == [
  649. {
  650. "id": "call_1",
  651. "type": "function",
  652. "function": {
  653. "name": "handoff_note",
  654. "arguments": "{}",
  655. },
  656. }
  657. ]
  658. assert tool_message.tool_call_id == "call_1"
  659. assert outputs[-1] == {"type": "done"}
  660. @pytest.mark.asyncio
  661. async def test_runtime_passes_event_names_without_tool_parameters_to_chat_agent():
  662. registry = ToolRegistry(
  663. [
  664. ToolDefinition(
  665. name="handoff_note",
  666. description="Registry-owned handoff tool.",
  667. parameters={
  668. "type": "object",
  669. "properties": {
  670. "message": {"type": "string"},
  671. "priority": {"type": "number"},
  672. },
  673. "required": ["message"],
  674. },
  675. handler=lambda event: {"tool": event.name, "message": "handled"},
  676. )
  677. ]
  678. )
  679. request = DebugRunRequest(
  680. user_message="debug this",
  681. system_prompts=[],
  682. pre_messages=[],
  683. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  684. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  685. )
  686. client = ToolCapturingChatClient()
  687. runtime = DebugRuntime(client, registry=registry)
  688. outputs = [message async for message in runtime.run(request)]
  689. assert outputs[-1] == {"type": "done"}
  690. assert client.tools == [
  691. {
  692. "type": "function",
  693. "function": {
  694. "name": "handoff_note",
  695. "description": "Registry-owned handoff tool.",
  696. "parameters": {
  697. "type": "object",
  698. "properties": {},
  699. "additionalProperties": False,
  700. },
  701. },
  702. }
  703. ]
  704. @pytest.mark.asyncio
  705. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  706. request = DebugRunRequest(
  707. user_message="debug this",
  708. system_prompts=[],
  709. pre_messages=[],
  710. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  711. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  712. )
  713. ticks = iter([1.0, 1.123, 1.456])
  714. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  715. outputs = [message async for message in runtime.run(request)]
  716. business_outputs = _without_audit(outputs)
  717. assert _message_types(outputs) == [
  718. "session_started",
  719. "message_delta",
  720. "usage",
  721. "round_stats",
  722. "done",
  723. ]
  724. assert business_outputs[3] == {
  725. "type": "round_stats",
  726. "round_index": 1,
  727. "ttft_ms": 123,
  728. "elapsed_ms": 456,
  729. "prompt_tokens": 10,
  730. "completion_tokens": 20,
  731. "total_tokens": 30,
  732. "cached_tokens": 5,
  733. "had_event": False,
  734. }
  735. @pytest.mark.asyncio
  736. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  737. request = DebugRunRequest(
  738. user_message="debug this",
  739. system_prompts=[],
  740. pre_messages=[],
  741. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  742. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  743. )
  744. ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
  745. client = EventRoundStatsChatClient()
  746. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  747. outputs = [message async for message in runtime.run(request)]
  748. stats = [message for message in outputs if message["type"] == "round_stats"]
  749. assert client.calls == 2
  750. assert stats == [
  751. {
  752. "type": "round_stats",
  753. "round_index": 1,
  754. "ttft_ms": None,
  755. "elapsed_ms": 250,
  756. "prompt_tokens": 3,
  757. "completion_tokens": 0,
  758. "total_tokens": 3,
  759. "cached_tokens": 0,
  760. "had_event": True,
  761. },
  762. {
  763. "type": "round_stats",
  764. "round_index": 2,
  765. "ttft_ms": 50,
  766. "elapsed_ms": 200,
  767. "prompt_tokens": 4,
  768. "completion_tokens": 6,
  769. "total_tokens": 10,
  770. "cached_tokens": 0,
  771. "had_event": False,
  772. },
  773. ]