test_debug_runtime.py 28 KB

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