test_debug_runtime.py 31 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981
  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 TwoTurnSessionChatClient:
  265. def __init__(self) -> None:
  266. self.calls = 0
  267. self.messages_by_call: list[list[ChatMessage]] = []
  268. async def stream_chat(
  269. self,
  270. messages: list[ChatMessage],
  271. tools: list[dict],
  272. params: AgentParams,
  273. ) -> AsyncIterator[StreamItem]:
  274. if tools:
  275. yield _event_tool_call_from_tools(tools, messages)
  276. return
  277. self.calls += 1
  278. self.messages_by_call.append(list(messages))
  279. if self.calls in {1, 3}:
  280. yield StreamItem.event(
  281. ToolCallEvent(
  282. id=f"call_{self.calls}",
  283. name="handoff_note",
  284. arguments={},
  285. raw_arguments="{}",
  286. )
  287. )
  288. return
  289. yield StreamItem.message_delta(f"final answer {self.calls}")
  290. class SlowAfterEventChatClient:
  291. def __init__(self) -> None:
  292. self.calls = 0
  293. self.event_seen = asyncio.Event()
  294. self.release_stream = asyncio.Event()
  295. async def stream_chat(
  296. self,
  297. messages: list[ChatMessage],
  298. tools: list[dict],
  299. params: AgentParams,
  300. ) -> AsyncIterator[StreamItem]:
  301. if tools:
  302. yield _event_tool_call_from_tools(tools, messages)
  303. return
  304. self.calls += 1
  305. if self.calls == 1:
  306. yield StreamItem.message_delta("Need event.")
  307. yield StreamItem.event(
  308. ToolCallEvent(
  309. id="call_1",
  310. name="handoff_note",
  311. arguments={},
  312. raw_arguments="{}",
  313. )
  314. )
  315. self.event_seen.set()
  316. await self.release_stream.wait()
  317. return
  318. yield StreamItem.message_delta("final answer")
  319. def test_runtime_queues_exposes_input_output_and_events_queues():
  320. RuntimeQueues = _runtime_queues_class()
  321. queues = RuntimeQueues()
  322. assert isinstance(queues.input, asyncio.Queue)
  323. assert isinstance(queues.output, asyncio.Queue)
  324. assert isinstance(queues.events, asyncio.Queue)
  325. assert queues.input is not queues.output
  326. assert queues.input is not queues.events
  327. assert queues.output is not queues.events
  328. @pytest.mark.asyncio
  329. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  330. request = DebugRunRequest(
  331. user_message="debug this",
  332. system_prompts=["You are a debugger."],
  333. pre_messages=[],
  334. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  335. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  336. )
  337. client = FakeChatClient()
  338. runtime = DebugRuntime(client)
  339. outputs = [message async for message in runtime.run(request)]
  340. business_outputs = _without_audit(outputs)
  341. assert client.calls == 2
  342. assert _message_types(outputs) == [
  343. "session_started",
  344. "event",
  345. "tool_result",
  346. "round_stats",
  347. "message_delta",
  348. "round_stats",
  349. "done",
  350. ]
  351. assert business_outputs[1]["event"]["name"] == "handoff_note"
  352. assert business_outputs[4]["content"] == "final answer"
  353. @pytest.mark.asyncio
  354. async def test_runtime_emits_audit_events_and_backend_logs(caplog):
  355. request = DebugRunRequest(
  356. user_message="debug this",
  357. system_prompts=[],
  358. pre_messages=[],
  359. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  360. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  361. )
  362. runtime = DebugRuntime(FakeChatClient())
  363. with caplog.at_level(logging.INFO, logger="agent_lab.application.runtime"):
  364. outputs = [message async for message in runtime.run(request)]
  365. audit_events = [
  366. message["event"]
  367. for message in outputs
  368. if message["type"] == "audit"
  369. ]
  370. assert audit_events == [
  371. "session_started",
  372. "chat_round_started",
  373. "chat_event_detected",
  374. "event_agent_completed",
  375. "chat_round_finished",
  376. "chat_round_started",
  377. "chat_round_finished",
  378. "session_finished",
  379. ]
  380. assert "chat_event_detected" in caplog.text
  381. assert "event_agent_completed" in caplog.text
  382. @pytest.mark.asyncio
  383. async def test_runtime_outputs_event_as_soon_as_chat_stream_detects_it():
  384. RuntimeQueues = _runtime_queues_class()
  385. queues = RuntimeQueues()
  386. client = SlowAfterEventChatClient()
  387. runtime = DebugRuntime(client, queues=queues)
  388. request = DebugRunRequest(
  389. user_message="debug this",
  390. system_prompts=[],
  391. pre_messages=[],
  392. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  393. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  394. )
  395. runtime.start(request)
  396. try:
  397. assert await _next_non_audit_from_queue(queues) == {"type": "session_started"}
  398. assert await _next_non_audit_from_queue(queues) == {
  399. "type": "message_delta",
  400. "content": "Need event.",
  401. }
  402. await asyncio.wait_for(client.event_seen.wait(), timeout=1)
  403. event_message = await asyncio.wait_for(
  404. _next_non_audit_from_queue(queues),
  405. timeout=0.2,
  406. )
  407. assert event_message["type"] == "event"
  408. assert event_message["event"]["name"] == "handoff_note"
  409. finally:
  410. client.release_stream.set()
  411. await runtime.aclose()
  412. @pytest.mark.asyncio
  413. async def test_runtime_batches_round_events_before_continuing_chat_agent():
  414. def resolve_from_history(
  415. event: ToolCallEvent,
  416. context: ToolExecutionContext,
  417. ) -> dict[str, Any]:
  418. return {"message": context.history[-1].content, "event": event.name}
  419. registry = ToolRegistry(
  420. [
  421. ToolDefinition(
  422. name="handoff_note",
  423. description="Send a handoff note.",
  424. parameters={
  425. "type": "object",
  426. "properties": {"message": {"type": "string"}},
  427. "required": ["message"],
  428. },
  429. handler=lambda event: {
  430. "tool": event.name,
  431. "message": event.arguments["message"],
  432. },
  433. argument_resolver=resolve_from_history,
  434. ),
  435. ToolDefinition(
  436. name="audit_note",
  437. description="Send an audit note.",
  438. parameters={
  439. "type": "object",
  440. "properties": {"message": {"type": "string"}},
  441. "required": ["message"],
  442. },
  443. handler=lambda event: {
  444. "tool": event.name,
  445. "message": event.arguments["message"],
  446. },
  447. argument_resolver=resolve_from_history,
  448. ),
  449. ]
  450. )
  451. request = DebugRunRequest(
  452. user_message="debug this",
  453. system_prompts=[],
  454. pre_messages=[],
  455. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  456. event_agent=EventAgentParams(
  457. enabled_tools=["handoff_note", "audit_note"],
  458. max_event_loops=2,
  459. ),
  460. )
  461. client = MultiEventChatClient()
  462. runtime = DebugRuntime(client, registry=registry)
  463. outputs = [message async for message in runtime.run(request)]
  464. assert _message_types(outputs) == [
  465. "session_started",
  466. "message_delta",
  467. "event",
  468. "event",
  469. "tool_result",
  470. "tool_result",
  471. "round_stats",
  472. "message_delta",
  473. "round_stats",
  474. "done",
  475. ]
  476. assert client.calls == 2
  477. assert [message.role for message in client.second_call_messages] == [
  478. "system",
  479. "user",
  480. "assistant",
  481. "user",
  482. ]
  483. assistant_message = client.second_call_messages[2]
  484. assert assistant_message.content == "Checking events."
  485. assert not any(message.role == "tool" for message in client.second_call_messages)
  486. assert client.second_call_messages[-1].content == (
  487. "EventAgent results:\n"
  488. '{"tool": "handoff_note", "message": "Checking events."}\n'
  489. '{"tool": "audit_note", "message": "Checking events."}'
  490. )
  491. assert client.second_call_messages[-1].name == "event_agent"
  492. @pytest.mark.asyncio
  493. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  494. request = DebugRunRequest(
  495. user_message="debug this",
  496. system_prompts=[],
  497. pre_messages=[],
  498. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  499. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  500. )
  501. runtime = DebugRuntime(RoundStatsChatClient())
  502. queues = runtime.start(request)
  503. outputs: list[dict[str, Any]] = []
  504. while True:
  505. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  506. outputs.append(message)
  507. if message["type"] == "done":
  508. break
  509. assert _message_types(outputs) == [
  510. "session_started",
  511. "message_delta",
  512. "usage",
  513. "round_stats",
  514. "done",
  515. ]
  516. @pytest.mark.asyncio
  517. async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
  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=1),
  524. )
  525. client = EventLoopLimitChatClient()
  526. runtime = DebugRuntime(client)
  527. outputs = [message async for message in runtime.run(request)]
  528. business_outputs = _without_audit(outputs)
  529. assert client.calls == 2
  530. assert client.tools_by_call[0] == []
  531. assert client.tools_by_call[1] == []
  532. assert "Available events:" in client.messages_by_call[0][0].content
  533. assert not any(
  534. "Available events:" in message.content
  535. for message in client.messages_by_call[1]
  536. if message.role == "system"
  537. )
  538. assert _message_types(outputs) == [
  539. "session_started",
  540. "event",
  541. "tool_result",
  542. "round_stats",
  543. "message_delta",
  544. "round_stats",
  545. "done",
  546. ]
  547. assert business_outputs[4]["content"] == "final after event limit"
  548. @pytest.mark.asyncio
  549. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  550. RuntimeQueues = _runtime_queues_class()
  551. queues = RuntimeQueues()
  552. request = DebugRunRequest(
  553. user_message="debug this",
  554. system_prompts=[],
  555. pre_messages=[],
  556. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  557. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  558. )
  559. client = StrictHistoryChatClient()
  560. runtime = DebugRuntime(client, queues=queues)
  561. stream = runtime.run(request)
  562. assert await _next_non_audit(stream) == {"type": "session_started"}
  563. event_message = await _next_non_audit(stream)
  564. assert event_message["type"] == "event"
  565. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  566. remaining = [message async for message in stream]
  567. assert remaining[-1] == {"type": "done"}
  568. assert [message.role for message in client.second_call_messages] == [
  569. "system",
  570. "user",
  571. "assistant",
  572. "user",
  573. "user",
  574. ]
  575. assert not any(message.role == "tool" for message in client.second_call_messages)
  576. assert client.second_call_messages[3].content.startswith("EventAgent results:\n")
  577. assert client.second_call_messages[3].name == "event_agent"
  578. assert client.second_call_messages[4].content == "follow-up while tool runs"
  579. @pytest.mark.asyncio
  580. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  581. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  582. raise RuntimeError("boom")
  583. registry = ToolRegistry(
  584. [
  585. ToolDefinition(
  586. name="handoff_note",
  587. description="Broken handoff tool.",
  588. parameters={"type": "object"},
  589. handler=fail_tool,
  590. )
  591. ]
  592. )
  593. request = DebugRunRequest(
  594. user_message="debug this",
  595. system_prompts=[],
  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=2),
  599. )
  600. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  601. outputs = await asyncio.wait_for(
  602. _collect_outputs(runtime.run(request)),
  603. timeout=1,
  604. )
  605. business_outputs = _without_audit(outputs)
  606. assert _message_types(outputs) == [
  607. "session_started",
  608. "event",
  609. "tool_result",
  610. "round_stats",
  611. "message_delta",
  612. "round_stats",
  613. "done",
  614. ]
  615. assert json.loads(business_outputs[2]["message"]["content"]) == {
  616. "tool": "handoff_note",
  617. "error": "tool handler failed: boom",
  618. }
  619. @pytest.mark.asyncio
  620. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  621. RuntimeQueues = _runtime_queues_class()
  622. queue_log: list[tuple[str, str, str]] = []
  623. queues = RuntimeQueues(
  624. input=RecordingQueue("input", queue_log),
  625. output=RecordingQueue("output", queue_log),
  626. events=RecordingQueue("events", queue_log),
  627. )
  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=3),
  634. )
  635. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  636. outputs = [message async for message in runtime.run(request)]
  637. assert _message_types(outputs) == [
  638. "session_started",
  639. "event",
  640. "tool_result",
  641. "round_stats",
  642. "message_delta",
  643. "round_stats",
  644. "done",
  645. ]
  646. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  647. ("input", "get", "user")
  648. )
  649. assert queue_log.index(
  650. ("events", "put", "event_request:handoff_note:call_1")
  651. ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
  652. assert queue_log.index(
  653. ("events", "get", "event_request:handoff_note:call_1")
  654. ) < queue_log.index(("input", "put", "tool:call_1"))
  655. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  656. ("input", "get", "tool:call_1")
  657. )
  658. @pytest.mark.asyncio
  659. async def test_runtime_run_consumes_output_queue_in_stream_order():
  660. RuntimeQueues = _runtime_queues_class()
  661. queue_log: list[tuple[str, str, str]] = []
  662. queues = RuntimeQueues(
  663. input=RecordingQueue("input", queue_log),
  664. output=RecordingQueue("output", queue_log),
  665. events=RecordingQueue("events", queue_log),
  666. )
  667. request = DebugRunRequest(
  668. user_message="debug this",
  669. system_prompts=["You are a debugger."],
  670. pre_messages=[],
  671. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  672. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  673. )
  674. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  675. outputs = [message async for message in runtime.run(request)]
  676. assert _message_types(outputs) == [
  677. "session_started",
  678. "event",
  679. "tool_result",
  680. "round_stats",
  681. "message_delta",
  682. "round_stats",
  683. "done",
  684. ]
  685. output_puts = [
  686. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  687. ]
  688. output_gets = [
  689. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  690. ]
  691. assert [message for message in output_puts if message != "output:audit"] == [
  692. "output:session_started",
  693. "output:event",
  694. "output:tool_result",
  695. "output:round_stats",
  696. "output:message_delta",
  697. "output:round_stats",
  698. "output:done",
  699. ]
  700. assert output_gets == output_puts
  701. @pytest.mark.asyncio
  702. async def test_runtime_continues_with_event_summary_without_tool_call_history():
  703. request = DebugRunRequest(
  704. user_message="debug this",
  705. system_prompts=["You are a debugger."],
  706. pre_messages=[],
  707. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  708. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  709. )
  710. client = StrictHistoryChatClient()
  711. runtime = DebugRuntime(client)
  712. outputs = [message async for message in runtime.run(request)]
  713. assert client.calls == 2
  714. assert [message.role for message in client.second_call_messages] == [
  715. "system",
  716. "system",
  717. "user",
  718. "assistant",
  719. "user",
  720. ]
  721. assistant_message = client.second_call_messages[3]
  722. assert assistant_message.content == ""
  723. assert "handoff_note" in client.second_call_messages[1].content
  724. assert not any(message.role == "tool" for message in client.second_call_messages)
  725. assert client.second_call_messages[4].content.startswith("EventAgent results:\n")
  726. assert client.second_call_messages[4].name == "event_agent"
  727. assert outputs[-1] == {"type": "done"}
  728. @pytest.mark.asyncio
  729. async def test_runtime_passes_event_catalog_system_message_without_chat_tools():
  730. registry = ToolRegistry(
  731. [
  732. ToolDefinition(
  733. name="handoff_note",
  734. description="Registry-owned handoff tool.",
  735. parameters={
  736. "type": "object",
  737. "properties": {
  738. "message": {"type": "string"},
  739. "priority": {"type": "number"},
  740. },
  741. "required": ["message"],
  742. },
  743. handler=lambda event: {"tool": event.name, "message": "handled"},
  744. )
  745. ]
  746. )
  747. request = DebugRunRequest(
  748. user_message="debug this",
  749. system_prompts=[],
  750. pre_messages=[],
  751. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  752. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  753. )
  754. client = ToolCapturingChatClient()
  755. runtime = DebugRuntime(client, registry=registry)
  756. outputs = [message async for message in runtime.run(request)]
  757. assert outputs[-1] == {"type": "done"}
  758. assert client.tools == []
  759. assert client.messages[0].role == "system"
  760. assert "Available events:" in client.messages[0].content
  761. assert "- handoff_note: Registry-owned handoff tool." in client.messages[0].content
  762. assert "message" not in client.messages[0].content
  763. assert "priority" not in client.messages[0].content
  764. @pytest.mark.asyncio
  765. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  766. request = DebugRunRequest(
  767. user_message="debug this",
  768. system_prompts=[],
  769. pre_messages=[],
  770. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  771. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  772. )
  773. ticks = iter([1.0, 1.123, 1.456])
  774. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  775. outputs = [message async for message in runtime.run(request)]
  776. business_outputs = _without_audit(outputs)
  777. assert _message_types(outputs) == [
  778. "session_started",
  779. "message_delta",
  780. "usage",
  781. "round_stats",
  782. "done",
  783. ]
  784. assert business_outputs[3] == {
  785. "type": "round_stats",
  786. "round_index": 1,
  787. "ttft_ms": 123,
  788. "elapsed_ms": 456,
  789. "prompt_tokens": 10,
  790. "completion_tokens": 20,
  791. "total_tokens": 30,
  792. "cached_tokens": 5,
  793. "had_event": False,
  794. }
  795. @pytest.mark.asyncio
  796. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  797. request = DebugRunRequest(
  798. user_message="debug this",
  799. system_prompts=[],
  800. pre_messages=[],
  801. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  802. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  803. )
  804. ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
  805. client = EventRoundStatsChatClient()
  806. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  807. outputs = [message async for message in runtime.run(request)]
  808. stats = [message for message in outputs if message["type"] == "round_stats"]
  809. assert client.calls == 2
  810. assert stats == [
  811. {
  812. "type": "round_stats",
  813. "round_index": 1,
  814. "ttft_ms": None,
  815. "elapsed_ms": 250,
  816. "prompt_tokens": 3,
  817. "completion_tokens": 0,
  818. "total_tokens": 3,
  819. "cached_tokens": 0,
  820. "had_event": True,
  821. },
  822. {
  823. "type": "round_stats",
  824. "round_index": 2,
  825. "ttft_ms": 50,
  826. "elapsed_ms": 200,
  827. "prompt_tokens": 4,
  828. "completion_tokens": 6,
  829. "total_tokens": 10,
  830. "cached_tokens": 0,
  831. "had_event": False,
  832. },
  833. ]
  834. @pytest.mark.asyncio
  835. async def test_runtime_session_resets_event_budget_for_each_user_turn():
  836. request = DebugRunRequest(
  837. user_message="first turn",
  838. system_prompts=[],
  839. pre_messages=[],
  840. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  841. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  842. )
  843. client = TwoTurnSessionChatClient()
  844. runtime = DebugRuntime(client)
  845. queues = runtime.start_session(request)
  846. outputs: list[dict[str, Any]] = []
  847. while len([message for message in outputs if message["type"] == "turn_completed"]) < 1:
  848. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  849. await queues.input.put(ChatMessage(role="user", content="second turn"))
  850. while len([message for message in outputs if message["type"] == "turn_completed"]) < 2:
  851. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  852. await runtime.aclose()
  853. business_types = [message["type"] for message in _without_audit(outputs)]
  854. assert business_types.count("turn_started") == 2
  855. assert business_types.count("turn_completed") == 2
  856. assert client.calls == 4
  857. assert "Available events:" in client.messages_by_call[0][0].content
  858. assert "Available events:" in client.messages_by_call[2][0].content