test_debug_runtime.py 34 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050
  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_agent_request",
  374. "chat_event_detected",
  375. "chat_agent_response",
  376. "event_agent_request",
  377. "event_agent_response",
  378. "event_agent_completed",
  379. "chat_round_finished",
  380. "chat_round_started",
  381. "chat_agent_request",
  382. "chat_agent_response",
  383. "chat_round_finished",
  384. "session_finished",
  385. ]
  386. assert "chat_event_detected" in caplog.text
  387. assert "event_agent_completed" in caplog.text
  388. @pytest.mark.asyncio
  389. async def test_runtime_audit_includes_model_params_prompts_results_and_usage():
  390. request = DebugRunRequest(
  391. user_message="debug this",
  392. system_prompts=["You are a debugger."],
  393. pre_messages=[],
  394. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  395. event_agent=EventAgentParams(
  396. model="event-model",
  397. temperature=0.4,
  398. max_tokens=80,
  399. enabled_tools=["handoff_note"],
  400. max_event_loops=1,
  401. ),
  402. )
  403. runtime = DebugRuntime(EventRoundStatsChatClient())
  404. outputs = [message async for message in runtime.run(request)]
  405. audits = [message for message in outputs if message["type"] == "audit"]
  406. chat_request = next(
  407. message for message in audits if message["event"] == "chat_agent_request"
  408. )
  409. assert chat_request["details"]["agent"] == "chat_agent"
  410. assert chat_request["details"]["params"]["model"] == "chat-model"
  411. assert chat_request["details"]["params"]["temperature"] == 0.1
  412. assert chat_request["details"]["params"]["max_tokens"] == 200
  413. assert chat_request["details"]["tools"] == []
  414. assert chat_request["details"]["messages"][0] == {
  415. "role": "system",
  416. "content": "You are a debugger.",
  417. "name": None,
  418. "tool_call_id": None,
  419. }
  420. chat_response = next(
  421. message for message in audits if message["event"] == "chat_agent_response"
  422. )
  423. assert chat_response["details"]["event_names"] == ["handoff_note"]
  424. assert chat_response["details"]["usage"] == {
  425. "prompt_tokens": 3,
  426. "completion_tokens": 0,
  427. "total_tokens": 3,
  428. "cached_tokens": 0,
  429. }
  430. event_request = next(
  431. message for message in audits if message["event"] == "event_agent_request"
  432. )
  433. assert event_request["details"]["agent"] == "event_agent"
  434. assert event_request["details"]["params"]["model"] == "event-model"
  435. assert event_request["details"]["params"]["temperature"] == 0.4
  436. assert event_request["details"]["events"][0]["name"] == "handoff_note"
  437. assert event_request["details"]["tools"][0]["function"]["name"] == "handoff_note"
  438. assert "Authorization" not in str(event_request["details"])
  439. event_response = next(
  440. message for message in audits if message["event"] == "event_agent_response"
  441. )
  442. assert event_response["details"]["replies"][0]["role"] == "tool"
  443. assert '"tool": "handoff_note"' in event_response["details"]["replies"][0]["content"]
  444. @pytest.mark.asyncio
  445. async def test_runtime_outputs_event_as_soon_as_chat_stream_detects_it():
  446. RuntimeQueues = _runtime_queues_class()
  447. queues = RuntimeQueues()
  448. client = SlowAfterEventChatClient()
  449. runtime = DebugRuntime(client, queues=queues)
  450. request = DebugRunRequest(
  451. user_message="debug this",
  452. system_prompts=[],
  453. pre_messages=[],
  454. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  455. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  456. )
  457. runtime.start(request)
  458. try:
  459. assert await _next_non_audit_from_queue(queues) == {"type": "session_started"}
  460. assert await _next_non_audit_from_queue(queues) == {
  461. "type": "message_delta",
  462. "content": "Need event.",
  463. }
  464. await asyncio.wait_for(client.event_seen.wait(), timeout=1)
  465. event_message = await asyncio.wait_for(
  466. _next_non_audit_from_queue(queues),
  467. timeout=0.2,
  468. )
  469. assert event_message["type"] == "event"
  470. assert event_message["event"]["name"] == "handoff_note"
  471. finally:
  472. client.release_stream.set()
  473. await runtime.aclose()
  474. @pytest.mark.asyncio
  475. async def test_runtime_batches_round_events_before_continuing_chat_agent():
  476. def resolve_from_history(
  477. event: ToolCallEvent,
  478. context: ToolExecutionContext,
  479. ) -> dict[str, Any]:
  480. return {"message": context.history[-1].content, "event": event.name}
  481. registry = ToolRegistry(
  482. [
  483. ToolDefinition(
  484. name="handoff_note",
  485. description="Send a handoff note.",
  486. parameters={
  487. "type": "object",
  488. "properties": {"message": {"type": "string"}},
  489. "required": ["message"],
  490. },
  491. handler=lambda event: {
  492. "tool": event.name,
  493. "message": event.arguments["message"],
  494. },
  495. argument_resolver=resolve_from_history,
  496. ),
  497. ToolDefinition(
  498. name="audit_note",
  499. description="Send an audit note.",
  500. parameters={
  501. "type": "object",
  502. "properties": {"message": {"type": "string"}},
  503. "required": ["message"],
  504. },
  505. handler=lambda event: {
  506. "tool": event.name,
  507. "message": event.arguments["message"],
  508. },
  509. argument_resolver=resolve_from_history,
  510. ),
  511. ]
  512. )
  513. request = DebugRunRequest(
  514. user_message="debug this",
  515. system_prompts=[],
  516. pre_messages=[],
  517. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  518. event_agent=EventAgentParams(
  519. enabled_tools=["handoff_note", "audit_note"],
  520. max_event_loops=2,
  521. ),
  522. )
  523. client = MultiEventChatClient()
  524. runtime = DebugRuntime(client, registry=registry)
  525. outputs = [message async for message in runtime.run(request)]
  526. assert _message_types(outputs) == [
  527. "session_started",
  528. "message_delta",
  529. "event",
  530. "event",
  531. "tool_result",
  532. "tool_result",
  533. "round_stats",
  534. "message_delta",
  535. "round_stats",
  536. "done",
  537. ]
  538. assert client.calls == 2
  539. assert [message.role for message in client.second_call_messages] == [
  540. "system",
  541. "user",
  542. "assistant",
  543. "user",
  544. ]
  545. assistant_message = client.second_call_messages[2]
  546. assert assistant_message.content == "Checking events."
  547. assert not any(message.role == "tool" for message in client.second_call_messages)
  548. assert client.second_call_messages[-1].content == (
  549. "EventAgent results:\n"
  550. '{"tool": "handoff_note", "message": "Checking events."}\n'
  551. '{"tool": "audit_note", "message": "Checking events."}'
  552. )
  553. assert client.second_call_messages[-1].name == "event_agent"
  554. @pytest.mark.asyncio
  555. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  556. request = DebugRunRequest(
  557. user_message="debug this",
  558. system_prompts=[],
  559. pre_messages=[],
  560. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  561. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  562. )
  563. runtime = DebugRuntime(RoundStatsChatClient())
  564. queues = runtime.start(request)
  565. outputs: list[dict[str, Any]] = []
  566. while True:
  567. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  568. outputs.append(message)
  569. if message["type"] == "done":
  570. break
  571. assert _message_types(outputs) == [
  572. "session_started",
  573. "message_delta",
  574. "usage",
  575. "round_stats",
  576. "done",
  577. ]
  578. @pytest.mark.asyncio
  579. async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
  580. request = DebugRunRequest(
  581. user_message="debug this",
  582. system_prompts=[],
  583. pre_messages=[],
  584. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  585. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  586. )
  587. client = EventLoopLimitChatClient()
  588. runtime = DebugRuntime(client)
  589. outputs = [message async for message in runtime.run(request)]
  590. business_outputs = _without_audit(outputs)
  591. assert client.calls == 2
  592. assert client.tools_by_call[0] == []
  593. assert client.tools_by_call[1] == []
  594. assert "Available events:" in client.messages_by_call[0][0].content
  595. assert not any(
  596. "Available events:" in message.content
  597. for message in client.messages_by_call[1]
  598. if message.role == "system"
  599. )
  600. assert _message_types(outputs) == [
  601. "session_started",
  602. "event",
  603. "tool_result",
  604. "round_stats",
  605. "message_delta",
  606. "round_stats",
  607. "done",
  608. ]
  609. assert business_outputs[4]["content"] == "final after event limit"
  610. @pytest.mark.asyncio
  611. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  612. RuntimeQueues = _runtime_queues_class()
  613. queues = RuntimeQueues()
  614. request = DebugRunRequest(
  615. user_message="debug this",
  616. system_prompts=[],
  617. pre_messages=[],
  618. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  619. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  620. )
  621. client = StrictHistoryChatClient()
  622. runtime = DebugRuntime(client, queues=queues)
  623. stream = runtime.run(request)
  624. assert await _next_non_audit(stream) == {"type": "session_started"}
  625. event_message = await _next_non_audit(stream)
  626. assert event_message["type"] == "event"
  627. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  628. remaining = [message async for message in stream]
  629. assert remaining[-1] == {"type": "done"}
  630. assert [message.role for message in client.second_call_messages] == [
  631. "system",
  632. "user",
  633. "assistant",
  634. "user",
  635. "user",
  636. ]
  637. assert not any(message.role == "tool" for message in client.second_call_messages)
  638. assert client.second_call_messages[3].content.startswith("EventAgent results:\n")
  639. assert client.second_call_messages[3].name == "event_agent"
  640. assert client.second_call_messages[4].content == "follow-up while tool runs"
  641. @pytest.mark.asyncio
  642. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  643. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  644. raise RuntimeError("boom")
  645. registry = ToolRegistry(
  646. [
  647. ToolDefinition(
  648. name="handoff_note",
  649. description="Broken handoff tool.",
  650. parameters={"type": "object"},
  651. handler=fail_tool,
  652. )
  653. ]
  654. )
  655. request = DebugRunRequest(
  656. user_message="debug this",
  657. system_prompts=[],
  658. pre_messages=[],
  659. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  660. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  661. )
  662. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  663. outputs = await asyncio.wait_for(
  664. _collect_outputs(runtime.run(request)),
  665. timeout=1,
  666. )
  667. business_outputs = _without_audit(outputs)
  668. assert _message_types(outputs) == [
  669. "session_started",
  670. "event",
  671. "tool_result",
  672. "round_stats",
  673. "message_delta",
  674. "round_stats",
  675. "done",
  676. ]
  677. assert json.loads(business_outputs[2]["message"]["content"]) == {
  678. "tool": "handoff_note",
  679. "error": "tool handler failed: boom",
  680. }
  681. @pytest.mark.asyncio
  682. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  683. RuntimeQueues = _runtime_queues_class()
  684. queue_log: list[tuple[str, str, str]] = []
  685. queues = RuntimeQueues(
  686. input=RecordingQueue("input", queue_log),
  687. output=RecordingQueue("output", queue_log),
  688. events=RecordingQueue("events", queue_log),
  689. )
  690. request = DebugRunRequest(
  691. user_message="debug this",
  692. system_prompts=["You are a debugger."],
  693. pre_messages=[],
  694. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  695. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  696. )
  697. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  698. outputs = [message async for message in runtime.run(request)]
  699. assert _message_types(outputs) == [
  700. "session_started",
  701. "event",
  702. "tool_result",
  703. "round_stats",
  704. "message_delta",
  705. "round_stats",
  706. "done",
  707. ]
  708. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  709. ("input", "get", "user")
  710. )
  711. assert queue_log.index(
  712. ("events", "put", "event_request:handoff_note:call_1")
  713. ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
  714. assert queue_log.index(
  715. ("events", "get", "event_request:handoff_note:call_1")
  716. ) < queue_log.index(("input", "put", "tool:call_1"))
  717. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  718. ("input", "get", "tool:call_1")
  719. )
  720. @pytest.mark.asyncio
  721. async def test_runtime_run_consumes_output_queue_in_stream_order():
  722. RuntimeQueues = _runtime_queues_class()
  723. queue_log: list[tuple[str, str, str]] = []
  724. queues = RuntimeQueues(
  725. input=RecordingQueue("input", queue_log),
  726. output=RecordingQueue("output", queue_log),
  727. events=RecordingQueue("events", queue_log),
  728. )
  729. request = DebugRunRequest(
  730. user_message="debug this",
  731. system_prompts=["You are a debugger."],
  732. pre_messages=[],
  733. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  734. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  735. )
  736. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  737. outputs = [message async for message in runtime.run(request)]
  738. assert _message_types(outputs) == [
  739. "session_started",
  740. "event",
  741. "tool_result",
  742. "round_stats",
  743. "message_delta",
  744. "round_stats",
  745. "done",
  746. ]
  747. output_puts = [
  748. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  749. ]
  750. output_gets = [
  751. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  752. ]
  753. assert [message for message in output_puts if message != "output:audit"] == [
  754. "output:session_started",
  755. "output:event",
  756. "output:tool_result",
  757. "output:round_stats",
  758. "output:message_delta",
  759. "output:round_stats",
  760. "output:done",
  761. ]
  762. assert output_gets == output_puts
  763. @pytest.mark.asyncio
  764. async def test_runtime_continues_with_event_summary_without_tool_call_history():
  765. request = DebugRunRequest(
  766. user_message="debug this",
  767. system_prompts=["You are a debugger."],
  768. pre_messages=[],
  769. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  770. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  771. )
  772. client = StrictHistoryChatClient()
  773. runtime = DebugRuntime(client)
  774. outputs = [message async for message in runtime.run(request)]
  775. assert client.calls == 2
  776. assert [message.role for message in client.second_call_messages] == [
  777. "system",
  778. "system",
  779. "user",
  780. "assistant",
  781. "user",
  782. ]
  783. assistant_message = client.second_call_messages[3]
  784. assert assistant_message.content == ""
  785. assert "handoff_note" in client.second_call_messages[1].content
  786. assert not any(message.role == "tool" for message in client.second_call_messages)
  787. assert client.second_call_messages[4].content.startswith("EventAgent results:\n")
  788. assert client.second_call_messages[4].name == "event_agent"
  789. assert outputs[-1] == {"type": "done"}
  790. @pytest.mark.asyncio
  791. async def test_runtime_passes_event_catalog_system_message_without_chat_tools():
  792. registry = ToolRegistry(
  793. [
  794. ToolDefinition(
  795. name="handoff_note",
  796. description="Registry-owned handoff tool.",
  797. parameters={
  798. "type": "object",
  799. "properties": {
  800. "message": {"type": "string"},
  801. "priority": {"type": "number"},
  802. },
  803. "required": ["message"],
  804. },
  805. handler=lambda event: {"tool": event.name, "message": "handled"},
  806. )
  807. ]
  808. )
  809. request = DebugRunRequest(
  810. user_message="debug this",
  811. system_prompts=[],
  812. pre_messages=[],
  813. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  814. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  815. )
  816. client = ToolCapturingChatClient()
  817. runtime = DebugRuntime(client, registry=registry)
  818. outputs = [message async for message in runtime.run(request)]
  819. assert outputs[-1] == {"type": "done"}
  820. assert client.tools == []
  821. assert client.messages[0].role == "system"
  822. assert "Available events:" in client.messages[0].content
  823. assert "- handoff_note: Registry-owned handoff tool." in client.messages[0].content
  824. assert "message" not in client.messages[0].content
  825. assert "priority" not in client.messages[0].content
  826. @pytest.mark.asyncio
  827. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  828. request = DebugRunRequest(
  829. user_message="debug this",
  830. system_prompts=[],
  831. pre_messages=[],
  832. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  833. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  834. )
  835. ticks = iter([1.0, 1.123, 1.456])
  836. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  837. outputs = [message async for message in runtime.run(request)]
  838. business_outputs = _without_audit(outputs)
  839. assert _message_types(outputs) == [
  840. "session_started",
  841. "message_delta",
  842. "usage",
  843. "round_stats",
  844. "done",
  845. ]
  846. assert business_outputs[3] == {
  847. "type": "round_stats",
  848. "round_index": 1,
  849. "ttft_ms": 123,
  850. "elapsed_ms": 456,
  851. "prompt_tokens": 10,
  852. "completion_tokens": 20,
  853. "total_tokens": 30,
  854. "cached_tokens": 5,
  855. "had_event": False,
  856. }
  857. @pytest.mark.asyncio
  858. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  859. request = DebugRunRequest(
  860. user_message="debug this",
  861. system_prompts=[],
  862. pre_messages=[],
  863. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  864. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  865. )
  866. ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
  867. client = EventRoundStatsChatClient()
  868. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  869. outputs = [message async for message in runtime.run(request)]
  870. stats = [message for message in outputs if message["type"] == "round_stats"]
  871. assert client.calls == 2
  872. assert stats == [
  873. {
  874. "type": "round_stats",
  875. "round_index": 1,
  876. "ttft_ms": None,
  877. "elapsed_ms": 250,
  878. "prompt_tokens": 3,
  879. "completion_tokens": 0,
  880. "total_tokens": 3,
  881. "cached_tokens": 0,
  882. "had_event": True,
  883. },
  884. {
  885. "type": "round_stats",
  886. "round_index": 2,
  887. "ttft_ms": 50,
  888. "elapsed_ms": 200,
  889. "prompt_tokens": 4,
  890. "completion_tokens": 6,
  891. "total_tokens": 10,
  892. "cached_tokens": 0,
  893. "had_event": False,
  894. },
  895. ]
  896. @pytest.mark.asyncio
  897. async def test_runtime_session_resets_event_budget_for_each_user_turn():
  898. request = DebugRunRequest(
  899. user_message="first turn",
  900. system_prompts=[],
  901. pre_messages=[],
  902. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  903. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  904. )
  905. client = TwoTurnSessionChatClient()
  906. runtime = DebugRuntime(client)
  907. queues = runtime.start_session(request)
  908. outputs: list[dict[str, Any]] = []
  909. while len([message for message in outputs if message["type"] == "turn_completed"]) < 1:
  910. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  911. await queues.input.put(ChatMessage(role="user", content="second turn"))
  912. while len([message for message in outputs if message["type"] == "turn_completed"]) < 2:
  913. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  914. await runtime.aclose()
  915. business_types = [message["type"] for message in _without_audit(outputs)]
  916. assert business_types.count("turn_started") == 2
  917. assert business_types.count("turn_completed") == 2
  918. assert client.calls == 4
  919. assert "Available events:" in client.messages_by_call[0][0].content
  920. assert "Available events:" in client.messages_by_call[2][0].content