test_debug_runtime.py 29 KB

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