test_debug_runtime.py 52 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053105410551056105710581059106010611062106310641065106610671068106910701071107210731074107510761077107810791080108110821083108410851086108710881089109010911092109310941095109610971098109911001101110211031104110511061107110811091110111111121113111411151116111711181119112011211122112311241125112611271128112911301131113211331134113511361137113811391140114111421143114411451146114711481149115011511152115311541155115611571158115911601161116211631164116511661167116811691170117111721173117411751176117711781179118011811182118311841185118611871188118911901191119211931194119511961197119811991200120112021203120412051206120712081209121012111212121312141215121612171218121912201221122212231224122512261227122812291230123112321233123412351236123712381239124012411242124312441245124612471248124912501251125212531254125512561257125812591260126112621263126412651266126712681269127012711272127312741275127612771278127912801281128212831284128512861287128812891290129112921293129412951296129712981299130013011302130313041305130613071308130913101311131213131314131513161317131813191320132113221323132413251326132713281329133013311332133313341335133613371338133913401341134213431344134513461347134813491350135113521353135413551356135713581359136013611362136313641365136613671368136913701371137213731374137513761377137813791380138113821383138413851386138713881389139013911392139313941395139613971398139914001401140214031404140514061407140814091410141114121413141414151416141714181419142014211422142314241425142614271428142914301431143214331434143514361437143814391440144114421443144414451446144714481449145014511452145314541455145614571458145914601461146214631464146514661467146814691470147114721473147414751476147714781479148014811482148314841485148614871488148914901491149214931494149514961497149814991500150115021503150415051506150715081509151015111512151315141515151615171518151915201521152215231524152515261527152815291530153115321533153415351536153715381539154015411542154315441545154615471548154915501551155215531554155515561557155815591560156115621563156415651566156715681569157015711572157315741575157615771578157915801581
  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. from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
  15. def _runtime_queues_class():
  16. module = importlib.import_module("agent_lab.application.queues")
  17. return module.RuntimeQueues
  18. async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[str, Any]]:
  19. return [message async for message in stream]
  20. def _without_audit(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
  21. return [message for message in outputs if message["type"] != "audit"]
  22. def _message_types(outputs: list[dict[str, Any]]) -> list[str]:
  23. return [message["type"] for message in _without_audit(outputs)]
  24. def _event_tool_call_from_tools(
  25. tools: list[dict],
  26. messages: list[ChatMessage],
  27. ) -> StreamItem:
  28. tool_name = tools[0]["function"]["name"]
  29. content = ""
  30. for role in ("assistant", "user"):
  31. content = next(
  32. (
  33. message.content
  34. for message in reversed(messages)
  35. if message.role == role and message.content.strip()
  36. ),
  37. "",
  38. )
  39. if content:
  40. break
  41. arguments = {
  42. "message": content,
  43. "query": content,
  44. "title": content,
  45. }
  46. return StreamItem.provider_tool_call(
  47. ToolCallEvent(
  48. id="event_agent_call_1",
  49. name=tool_name,
  50. arguments=arguments,
  51. raw_arguments=json.dumps(arguments),
  52. )
  53. )
  54. async def _next_non_audit(
  55. stream: AsyncIterator[dict[str, Any]],
  56. ) -> dict[str, Any]:
  57. while True:
  58. message = await anext(stream)
  59. if message["type"] != "audit":
  60. return message
  61. async def _next_non_audit_from_queue(queues: Any) -> dict[str, Any]:
  62. while True:
  63. message = await queues.output.get()
  64. if message["type"] != "audit":
  65. return message
  66. class RecordingQueue(asyncio.Queue):
  67. def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
  68. super().__init__()
  69. self.name = name
  70. self.log = log
  71. async def put(self, item: Any) -> None:
  72. self.log.append((self.name, "put", self._describe(item)))
  73. await super().put(item)
  74. async def get(self) -> Any:
  75. item = await super().get()
  76. self.log.append((self.name, "get", self._describe(item)))
  77. return item
  78. def _describe(self, item: Any) -> str:
  79. if isinstance(item, ChatMessage):
  80. if item.role == "tool":
  81. return f"tool:{item.tool_call_id}"
  82. return item.role
  83. if isinstance(item, ToolCallEvent):
  84. return f"event:{item.name}:{item.id}"
  85. if isinstance(item, EventAgentRequest):
  86. events = ",".join(f"{event.name}:{event.id}" for event in item.events)
  87. return f"event_request:{events}"
  88. if isinstance(item, dict):
  89. return f"output:{item.get('type')}"
  90. return type(item).__name__
  91. class FakeChatClient:
  92. def __init__(self) -> None:
  93. self.calls = 0
  94. async def stream_chat(
  95. self,
  96. messages: list[ChatMessage],
  97. tools: list[dict],
  98. params: AgentParams,
  99. ) -> AsyncIterator[StreamItem]:
  100. if tools:
  101. yield StreamItem.raw_response_chunk(
  102. {
  103. "choices": [
  104. {
  105. "delta": {
  106. "tool_calls": [
  107. {
  108. "index": 0,
  109. "function": {"name": tools[0]["function"]["name"]},
  110. }
  111. ]
  112. },
  113. "finish_reason": None,
  114. }
  115. ]
  116. }
  117. )
  118. yield _event_tool_call_from_tools(tools, messages)
  119. return
  120. self.calls += 1
  121. if self.calls == 1:
  122. yield StreamItem.raw_response_chunk(
  123. {
  124. "choices": [
  125. {
  126. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  127. "finish_reason": None,
  128. }
  129. ]
  130. }
  131. )
  132. yield StreamItem.text_event(
  133. ToolCallEvent(
  134. id="call_1",
  135. name="handoff_note",
  136. arguments={},
  137. raw_arguments="{}",
  138. )
  139. )
  140. return
  141. assert not any(message.role == "tool" for message in messages)
  142. assert any(message.name == "event_agent" for message in messages)
  143. yield StreamItem.message_delta("final answer")
  144. class WrongSourceEventChatClient:
  145. def __init__(self, item: StreamItem) -> None:
  146. self.item = item
  147. self.calls = 0
  148. async def stream_chat(
  149. self,
  150. messages: list[ChatMessage],
  151. tools: list[dict],
  152. params: AgentParams,
  153. ) -> AsyncIterator[StreamItem]:
  154. if tools:
  155. yield _event_tool_call_from_tools(tools, messages)
  156. return
  157. self.calls += 1
  158. if self.calls == 1:
  159. yield self.item
  160. return
  161. yield StreamItem.message_delta("unexpected continuation")
  162. class IncrementingClock:
  163. def __init__(self, current: float = 100.0, step: float = 0.01) -> None:
  164. self.current = current
  165. self.step = step
  166. def __call__(self) -> float:
  167. self.current += self.step
  168. return self.current
  169. class StrictHistoryChatClient:
  170. def __init__(self) -> None:
  171. self.calls = 0
  172. self.second_call_messages: list[ChatMessage] = []
  173. async def stream_chat(
  174. self,
  175. messages: list[ChatMessage],
  176. tools: list[dict],
  177. params: AgentParams,
  178. ) -> AsyncIterator[StreamItem]:
  179. if tools:
  180. yield StreamItem.raw_response_chunk(
  181. {
  182. "choices": [
  183. {
  184. "delta": {
  185. "tool_calls": [
  186. {
  187. "index": 0,
  188. "function": {"name": tools[0]["function"]["name"]},
  189. }
  190. ]
  191. },
  192. "finish_reason": None,
  193. }
  194. ]
  195. }
  196. )
  197. yield _event_tool_call_from_tools(tools, messages)
  198. return
  199. self.calls += 1
  200. if self.calls == 1:
  201. yield StreamItem.raw_response_chunk(
  202. {
  203. "choices": [
  204. {
  205. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  206. "finish_reason": None,
  207. }
  208. ]
  209. }
  210. )
  211. yield StreamItem.text_event(
  212. ToolCallEvent(
  213. id="call_1",
  214. name="handoff_note",
  215. arguments={},
  216. raw_arguments="{}",
  217. )
  218. )
  219. return
  220. self.second_call_messages = list(messages)
  221. yield StreamItem.message_delta("final answer")
  222. class MultiEventChatClient:
  223. def __init__(self) -> None:
  224. self.calls = 0
  225. self.second_call_messages: list[ChatMessage] = []
  226. async def stream_chat(
  227. self,
  228. messages: list[ChatMessage],
  229. tools: list[dict],
  230. params: AgentParams,
  231. ) -> AsyncIterator[StreamItem]:
  232. if tools:
  233. yield _event_tool_call_from_tools(tools, messages)
  234. return
  235. self.calls += 1
  236. if self.calls == 1:
  237. yield StreamItem.message_delta("Checking events.")
  238. yield StreamItem.text_event(
  239. ToolCallEvent(
  240. id="call_1",
  241. name="handoff_note",
  242. arguments={"message": "ignored chat argument"},
  243. raw_arguments='{"message":"ignored chat argument"}',
  244. )
  245. )
  246. yield StreamItem.text_event(
  247. ToolCallEvent(
  248. id="call_2",
  249. name="audit_note",
  250. arguments={"message": "ignored chat argument"},
  251. raw_arguments='{"message":"ignored chat argument"}',
  252. )
  253. )
  254. return
  255. self.second_call_messages = list(messages)
  256. yield StreamItem.message_delta("Final answer.")
  257. class ToolCapturingChatClient:
  258. def __init__(self) -> None:
  259. self.tools: list[dict[str, Any]] = []
  260. self.messages: list[ChatMessage] = []
  261. async def stream_chat(
  262. self,
  263. messages: list[ChatMessage],
  264. tools: list[dict],
  265. params: AgentParams,
  266. ) -> AsyncIterator[StreamItem]:
  267. self.messages = list(messages)
  268. self.tools = list(tools)
  269. yield StreamItem.message_delta("final answer")
  270. class EventLoopLimitChatClient:
  271. def __init__(self) -> None:
  272. self.calls = 0
  273. self.tools_by_call: list[list[dict[str, Any]]] = []
  274. self.messages_by_call: list[list[ChatMessage]] = []
  275. async def stream_chat(
  276. self,
  277. messages: list[ChatMessage],
  278. tools: list[dict],
  279. params: AgentParams,
  280. ) -> AsyncIterator[StreamItem]:
  281. if tools:
  282. yield _event_tool_call_from_tools(tools, messages)
  283. return
  284. self.calls += 1
  285. self.tools_by_call.append(list(tools))
  286. self.messages_by_call.append(list(messages))
  287. if self.calls == 1:
  288. yield StreamItem.text_event(
  289. ToolCallEvent(
  290. id="call_1",
  291. name="handoff_note",
  292. arguments={},
  293. raw_arguments="{}",
  294. )
  295. )
  296. return
  297. yield StreamItem.message_delta("final after event limit")
  298. class RoundStatsChatClient:
  299. async def stream_chat(
  300. self,
  301. messages: list[ChatMessage],
  302. tools: list[dict],
  303. params: AgentParams,
  304. ) -> AsyncIterator[StreamItem]:
  305. yield StreamItem.message_delta("hello")
  306. yield StreamItem.usage_item(
  307. TokenUsage(
  308. prompt_tokens=10,
  309. completion_tokens=20,
  310. total_tokens=30,
  311. cached_tokens=5,
  312. )
  313. )
  314. class EventRoundStatsChatClient:
  315. def __init__(self) -> None:
  316. self.calls = 0
  317. async def stream_chat(
  318. self,
  319. messages: list[ChatMessage],
  320. tools: list[dict],
  321. params: AgentParams,
  322. ) -> AsyncIterator[StreamItem]:
  323. if tools:
  324. yield StreamItem.raw_response_chunk(
  325. {
  326. "choices": [
  327. {
  328. "delta": {
  329. "tool_calls": [
  330. {
  331. "index": 0,
  332. "function": {
  333. "name": tools[0]["function"]["name"],
  334. },
  335. }
  336. ]
  337. },
  338. "finish_reason": None,
  339. }
  340. ]
  341. }
  342. )
  343. yield _event_tool_call_from_tools(tools, messages)
  344. return
  345. self.calls += 1
  346. if self.calls == 1:
  347. yield StreamItem.raw_response_chunk(
  348. {
  349. "choices": [
  350. {
  351. "delta": {
  352. "content": "<agent_events>handoff_note</agent_events>",
  353. },
  354. "finish_reason": None,
  355. }
  356. ]
  357. }
  358. )
  359. yield StreamItem.text_event(
  360. ToolCallEvent(
  361. id="call_1",
  362. name="handoff_note",
  363. arguments={},
  364. raw_arguments="{}",
  365. )
  366. )
  367. yield StreamItem.usage_item(
  368. TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3)
  369. )
  370. return
  371. yield StreamItem.raw_response_chunk(
  372. {
  373. "choices": [
  374. {
  375. "delta": {"content": "final answer"},
  376. "finish_reason": None,
  377. }
  378. ]
  379. }
  380. )
  381. yield StreamItem.message_delta("final answer")
  382. yield StreamItem.usage_item(
  383. TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
  384. )
  385. class ContextBoundaryChatClient:
  386. def __init__(self) -> None:
  387. self.calls = 0
  388. self.event_agent_messages: list[list[ChatMessage]] = []
  389. async def stream_chat(
  390. self,
  391. messages: list[ChatMessage],
  392. tools: list[dict],
  393. params: AgentParams,
  394. ) -> AsyncIterator[StreamItem]:
  395. if tools:
  396. self.event_agent_messages.append(list(messages))
  397. yield _event_tool_call_from_tools(tools, messages)
  398. return
  399. self.calls += 1
  400. if self.calls == 1:
  401. yield StreamItem.message_delta("I will check.")
  402. yield StreamItem.text_event(
  403. ToolCallEvent(
  404. id="call_1",
  405. name="handoff_note",
  406. arguments={},
  407. raw_arguments="{}",
  408. )
  409. )
  410. return
  411. yield StreamItem.message_delta("final answer")
  412. class TwoTurnSessionChatClient:
  413. def __init__(self) -> None:
  414. self.calls = 0
  415. self.messages_by_call: list[list[ChatMessage]] = []
  416. async def stream_chat(
  417. self,
  418. messages: list[ChatMessage],
  419. tools: list[dict],
  420. params: AgentParams,
  421. ) -> AsyncIterator[StreamItem]:
  422. if tools:
  423. yield _event_tool_call_from_tools(tools, messages)
  424. return
  425. self.calls += 1
  426. self.messages_by_call.append(list(messages))
  427. if self.calls in {1, 3}:
  428. yield StreamItem.text_event(
  429. ToolCallEvent(
  430. id=f"call_{self.calls}",
  431. name="handoff_note",
  432. arguments={},
  433. raw_arguments="{}",
  434. )
  435. )
  436. return
  437. yield StreamItem.message_delta(f"final answer {self.calls}")
  438. class SlowAfterEventChatClient:
  439. def __init__(self) -> None:
  440. self.calls = 0
  441. self.event_seen = asyncio.Event()
  442. self.release_stream = asyncio.Event()
  443. async def stream_chat(
  444. self,
  445. messages: list[ChatMessage],
  446. tools: list[dict],
  447. params: AgentParams,
  448. ) -> AsyncIterator[StreamItem]:
  449. if tools:
  450. yield _event_tool_call_from_tools(tools, messages)
  451. return
  452. self.calls += 1
  453. if self.calls == 1:
  454. yield StreamItem.message_delta("Need event.")
  455. yield StreamItem.text_event(
  456. ToolCallEvent(
  457. id="call_1",
  458. name="handoff_note",
  459. arguments={},
  460. raw_arguments="{}",
  461. )
  462. )
  463. self.event_seen.set()
  464. await self.release_stream.wait()
  465. return
  466. yield StreamItem.message_delta("final answer")
  467. def test_runtime_queues_exposes_input_output_and_events_queues():
  468. RuntimeQueues = _runtime_queues_class()
  469. queues = RuntimeQueues()
  470. assert isinstance(queues.input, asyncio.Queue)
  471. assert isinstance(queues.output, asyncio.Queue)
  472. assert isinstance(queues.events, asyncio.Queue)
  473. assert queues.input is not queues.output
  474. assert queues.input is not queues.events
  475. assert queues.output is not queues.events
  476. @pytest.mark.asyncio
  477. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  478. request = DebugRunRequest(
  479. user_message="debug this",
  480. system_prompts=["You are a debugger."],
  481. pre_messages=[],
  482. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  483. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  484. )
  485. client = FakeChatClient()
  486. runtime = DebugRuntime(client)
  487. outputs = [message async for message in runtime.run(request)]
  488. business_outputs = _without_audit(outputs)
  489. assert client.calls == 2
  490. assert _message_types(outputs) == [
  491. "session_started",
  492. "event",
  493. "tool_result",
  494. "round_stats",
  495. "message_delta",
  496. "round_stats",
  497. "done",
  498. ]
  499. assert business_outputs[1]["event"]["name"] == "handoff_note"
  500. assert business_outputs[4]["content"] == "final answer"
  501. @pytest.mark.asyncio
  502. @pytest.mark.parametrize(
  503. "item",
  504. [
  505. StreamItem.provider_tool_call(
  506. ToolCallEvent(
  507. id="provider_call_1",
  508. name="handoff_note",
  509. arguments={"message": "wrong source"},
  510. raw_arguments='{"message":"wrong source"}',
  511. )
  512. ),
  513. StreamItem.event(
  514. ToolCallEvent(
  515. id="legacy_event_1",
  516. name="handoff_note",
  517. arguments={},
  518. raw_arguments="{}",
  519. )
  520. ),
  521. ],
  522. ids=["provider_tool_call", "legacy_event"],
  523. )
  524. async def test_runtime_ignores_non_text_event_sources(item: StreamItem):
  525. request = DebugRunRequest(
  526. user_message="debug this",
  527. system_prompts=[],
  528. pre_messages=[],
  529. chat_agent=AgentParams(model="fake-model"),
  530. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  531. )
  532. client = WrongSourceEventChatClient(item)
  533. outputs = [message async for message in DebugRuntime(client).run(request)]
  534. assert client.calls == 1
  535. assert "event" not in _message_types(outputs)
  536. assert "tool_result" not in _message_types(outputs)
  537. @pytest.mark.asyncio
  538. async def test_runtime_emits_audit_events_and_backend_logs(caplog):
  539. request = DebugRunRequest(
  540. user_message="debug this",
  541. system_prompts=[],
  542. pre_messages=[],
  543. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  544. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  545. )
  546. runtime = DebugRuntime(FakeChatClient())
  547. with caplog.at_level(logging.INFO, logger="agent_lab.application.runtime"):
  548. outputs = [message async for message in runtime.run(request)]
  549. audit_events = [
  550. message["event"]
  551. for message in outputs
  552. if message["type"] == "audit"
  553. ]
  554. assert audit_events == [
  555. "session_started",
  556. "chat_round_started",
  557. "chat_agent_request",
  558. "chat_event_detected",
  559. "chat_agent_response",
  560. "event_agent_request",
  561. "event_agent_response",
  562. "event_agent_completed",
  563. "chat_round_finished",
  564. "chat_round_started",
  565. "chat_agent_request",
  566. "chat_message_stream_started",
  567. "chat_message_stream_finished",
  568. "chat_agent_response",
  569. "chat_round_finished",
  570. "session_finished",
  571. ]
  572. assert "chat_event_detected" in caplog.text
  573. assert "event_agent_completed" in caplog.text
  574. @pytest.mark.asyncio
  575. async def test_runtime_audits_chat_message_stream_boundaries_in_output_order():
  576. request = DebugRunRequest(
  577. user_message="debug this",
  578. system_prompts=[],
  579. pre_messages=[],
  580. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  581. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  582. )
  583. runtime = DebugRuntime(RoundStatsChatClient())
  584. outputs = [message async for message in runtime.run(request)]
  585. ordered_labels = [
  586. message["event"] if message["type"] == "audit" else message["type"]
  587. for message in outputs
  588. ]
  589. assert ordered_labels.index("chat_message_stream_started") < ordered_labels.index(
  590. "message_delta"
  591. )
  592. assert ordered_labels.index("message_delta") < ordered_labels.index(
  593. "chat_message_stream_finished"
  594. )
  595. assert ordered_labels.index("chat_message_stream_finished") < ordered_labels.index(
  596. "chat_agent_response"
  597. )
  598. stream_started = next(
  599. message
  600. for message in outputs
  601. if message.get("event") == "chat_message_stream_started"
  602. )
  603. stream_finished = next(
  604. message
  605. for message in outputs
  606. if message.get("event") == "chat_message_stream_finished"
  607. )
  608. assert stream_started["details"]["agent"] == "chat_agent"
  609. assert stream_started["details"]["round_index"] == 1
  610. assert stream_finished["details"]["delta_count"] == 1
  611. assert stream_finished["details"]["content_length"] == len("hello")
  612. @pytest.mark.asyncio
  613. async def test_runtime_audit_events_include_turn_relative_elapsed_time():
  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=1),
  620. )
  621. runtime = DebugRuntime(FakeChatClient(), clock=IncrementingClock())
  622. queues = runtime.start_session(request)
  623. outputs: list[dict[str, Any]] = []
  624. while not any(message["type"] == "turn_completed" for message in outputs):
  625. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  626. await runtime.aclose()
  627. turn_audits = [
  628. message
  629. for message in outputs
  630. if message["type"] == "audit"
  631. ]
  632. elapsed_values = [
  633. message["details"].get("turn_elapsed_ms") for message in turn_audits
  634. ]
  635. assert turn_audits
  636. assert all(isinstance(value, int) for value in elapsed_values)
  637. assert all(value >= 0 for value in elapsed_values)
  638. assert elapsed_values == sorted(elapsed_values)
  639. assert next(
  640. message
  641. for message in turn_audits
  642. if message["event"] == "event_agent_request"
  643. )["details"]["turn_elapsed_ms"] >= 0
  644. @pytest.mark.asyncio
  645. async def test_runtime_audit_includes_model_params_prompts_results_and_usage():
  646. request = DebugRunRequest(
  647. user_message="debug this",
  648. system_prompts=["You are a debugger."],
  649. pre_messages=[],
  650. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  651. event_agent=EventAgentParams(
  652. model="event-model",
  653. temperature=0.4,
  654. max_tokens=80,
  655. enabled_tools=["handoff_note"],
  656. max_event_loops=1,
  657. ),
  658. )
  659. runtime = DebugRuntime(EventRoundStatsChatClient())
  660. outputs = [message async for message in runtime.run(request)]
  661. audits = [message for message in outputs if message["type"] == "audit"]
  662. chat_request = next(
  663. message for message in audits if message["event"] == "chat_agent_request"
  664. )
  665. assert chat_request["details"]["agent"] == "chat_agent"
  666. assert chat_request["details"]["params"]["model"] == "chat-model"
  667. assert chat_request["details"]["params"]["temperature"] == 0.1
  668. assert chat_request["details"]["params"]["max_tokens"] == 200
  669. assert chat_request["details"]["tools"] == []
  670. assert chat_request["details"]["messages"][0] == {
  671. "role": "system",
  672. "content": "You are a debugger.",
  673. "name": None,
  674. "tool_call_id": None,
  675. }
  676. chat_response = next(
  677. message for message in audits if message["event"] == "chat_agent_response"
  678. )
  679. assert chat_response["details"]["event_names"] == ["handoff_note"]
  680. assert chat_response["details"]["raw_chunks"] == [
  681. {
  682. "choices": [
  683. {
  684. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  685. "finish_reason": None,
  686. }
  687. ]
  688. }
  689. ]
  690. assert chat_response["details"]["usage"] == {
  691. "prompt_tokens": 3,
  692. "completion_tokens": 0,
  693. "total_tokens": 3,
  694. "cached_tokens": 0,
  695. }
  696. event_request = next(
  697. message for message in audits if message["event"] == "event_agent_request"
  698. )
  699. assert event_request["details"]["agent"] == "event_agent"
  700. assert event_request["details"]["params"]["model"] == "event-model"
  701. assert event_request["details"]["params"]["temperature"] == 0.4
  702. assert event_request["details"]["events"][0]["name"] == "handoff_note"
  703. assert event_request["details"]["tools"][0]["function"]["name"] == "handoff_note"
  704. assert "Authorization" not in str(event_request["details"])
  705. event_response = next(
  706. message for message in audits if message["event"] == "event_agent_response"
  707. )
  708. assert event_response["details"]["replies"][0]["role"] == "tool"
  709. assert '"tool": "handoff_note"' in event_response["details"]["replies"][0]["content"]
  710. assert event_response["details"]["raw_model_chunks"][0]["event_name"] == "handoff_note"
  711. assert event_response["details"]["raw_model_chunks"][0]["chunks"][0]["choices"][0]["delta"] == {
  712. "tool_calls": [
  713. {
  714. "index": 0,
  715. "function": {"name": "handoff_note"},
  716. }
  717. ]
  718. }
  719. @pytest.mark.asyncio
  720. async def test_runtime_round_started_separates_available_events_from_round_budget():
  721. request = DebugRunRequest(
  722. user_message="debug this",
  723. system_prompts=[],
  724. pre_messages=[],
  725. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  726. event_agent=EventAgentParams(
  727. enabled_tools=["handoff_note"],
  728. max_event_loops=1,
  729. ),
  730. )
  731. runtime = DebugRuntime(EventRoundStatsChatClient())
  732. outputs = [message async for message in runtime.run(request)]
  733. round_starts = [
  734. message
  735. for message in outputs
  736. if message.get("event") == "chat_round_started"
  737. ]
  738. assert round_starts[0]["details"]["events_enabled"] == ["handoff_note"]
  739. assert round_starts[0]["details"]["configured_events"] == ["handoff_note"]
  740. assert round_starts[0]["details"]["event_generation_enabled"] is True
  741. assert round_starts[0]["details"]["event_prompt_events"] == ["handoff_note"]
  742. assert round_starts[1]["details"]["events_enabled"] == []
  743. assert round_starts[1]["details"]["configured_events"] == ["handoff_note"]
  744. assert round_starts[1]["details"]["event_generation_enabled"] is False
  745. assert round_starts[1]["details"]["event_prompt_events"] == []
  746. @pytest.mark.asyncio
  747. async def test_runtime_event_agent_history_excludes_chat_agent_system_context():
  748. request = DebugRunRequest(
  749. user_message="debug this",
  750. system_prompts=["ChatAgent root prompt."],
  751. pre_messages=[
  752. ChatMessage(role="system", content="ChatAgent pre system."),
  753. ChatMessage(role="user", content="earlier user"),
  754. ChatMessage(role="assistant", content="earlier assistant"),
  755. ],
  756. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  757. event_agent=EventAgentParams(
  758. model="event-model",
  759. enabled_tools=["handoff_note"],
  760. max_event_loops=1,
  761. ),
  762. )
  763. client = ContextBoundaryChatClient()
  764. runtime = DebugRuntime(client)
  765. outputs = [message async for message in runtime.run(request)]
  766. event_request = next(
  767. message for message in outputs if message.get("event") == "event_agent_request"
  768. )
  769. assert [
  770. (message["role"], message["content"])
  771. for message in event_request["details"]["history"]
  772. ] == [
  773. ("user", "earlier user"),
  774. ("assistant", "earlier assistant"),
  775. ("user", "debug this"),
  776. ("assistant", "I will check."),
  777. ]
  778. assert not any(
  779. message["role"] == "system"
  780. for message in event_request["details"]["history"]
  781. )
  782. assert client.event_agent_messages
  783. assert [
  784. (message.role, message.content)
  785. for message in client.event_agent_messages[0]
  786. if message.content in {"ChatAgent root prompt.", "ChatAgent pre system."}
  787. ] == []
  788. @pytest.mark.asyncio
  789. async def test_runtime_session_event_agent_history_excludes_internal_replies():
  790. request = DebugRunRequest(
  791. user_message="debug this",
  792. system_prompts=["ChatAgent root prompt."],
  793. pre_messages=[ChatMessage(role="assistant", content="prior answer")],
  794. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  795. event_agent=EventAgentParams(
  796. model="event-model",
  797. enabled_tools=["handoff_note"],
  798. max_event_loops=1,
  799. ),
  800. )
  801. runtime = DebugRuntime(ContextBoundaryChatClient())
  802. queues = runtime.start_session(request)
  803. outputs: list[dict[str, Any]] = []
  804. while not any(message["type"] == "turn_completed" for message in outputs):
  805. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  806. await runtime.aclose()
  807. event_request = next(
  808. message for message in outputs if message.get("event") == "event_agent_request"
  809. )
  810. assert [
  811. (message["role"], message["content"], message["name"])
  812. for message in event_request["details"]["history"]
  813. ] == [
  814. ("assistant", "prior answer", None),
  815. ("user", "debug this", None),
  816. ("assistant", "I will check.", None),
  817. ]
  818. assert not any(
  819. message["name"] == "event_agent"
  820. for message in event_request["details"]["history"]
  821. )
  822. @pytest.mark.asyncio
  823. async def test_runtime_outputs_event_as_soon_as_chat_stream_detects_it():
  824. RuntimeQueues = _runtime_queues_class()
  825. queues = RuntimeQueues()
  826. client = SlowAfterEventChatClient()
  827. runtime = DebugRuntime(client, queues=queues)
  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=["handoff_note"], max_event_loops=1),
  834. )
  835. runtime.start(request)
  836. try:
  837. assert await _next_non_audit_from_queue(queues) == {"type": "session_started"}
  838. assert await _next_non_audit_from_queue(queues) == {
  839. "type": "message_delta",
  840. "content": "Need event.",
  841. }
  842. await asyncio.wait_for(client.event_seen.wait(), timeout=1)
  843. event_message = await asyncio.wait_for(
  844. _next_non_audit_from_queue(queues),
  845. timeout=0.2,
  846. )
  847. assert event_message["type"] == "event"
  848. assert event_message["event"]["name"] == "handoff_note"
  849. finally:
  850. client.release_stream.set()
  851. await runtime.aclose()
  852. @pytest.mark.asyncio
  853. async def test_runtime_batches_round_events_before_continuing_chat_agent():
  854. def resolve_from_history(
  855. event: ToolCallEvent,
  856. context: ToolExecutionContext,
  857. ) -> dict[str, Any]:
  858. return {"message": context.history[-1].content, "event": event.name}
  859. registry = ToolRegistry(
  860. [
  861. ToolDefinition(
  862. name="handoff_note",
  863. description="Send a handoff note.",
  864. parameters={
  865. "type": "object",
  866. "properties": {"message": {"type": "string"}},
  867. "required": ["message"],
  868. },
  869. handler=lambda event: {
  870. "tool": event.name,
  871. "message": event.arguments["message"],
  872. },
  873. argument_resolver=resolve_from_history,
  874. ),
  875. ToolDefinition(
  876. name="audit_note",
  877. description="Send an audit note.",
  878. parameters={
  879. "type": "object",
  880. "properties": {"message": {"type": "string"}},
  881. "required": ["message"],
  882. },
  883. handler=lambda event: {
  884. "tool": event.name,
  885. "message": event.arguments["message"],
  886. },
  887. argument_resolver=resolve_from_history,
  888. ),
  889. ]
  890. )
  891. request = DebugRunRequest(
  892. user_message="debug this",
  893. system_prompts=[],
  894. pre_messages=[],
  895. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  896. event_agent=EventAgentParams(
  897. enabled_tools=["handoff_note", "audit_note"],
  898. max_event_loops=2,
  899. ),
  900. )
  901. client = MultiEventChatClient()
  902. runtime = DebugRuntime(client, registry=registry)
  903. outputs = [message async for message in runtime.run(request)]
  904. assert _message_types(outputs) == [
  905. "session_started",
  906. "message_delta",
  907. "event",
  908. "event",
  909. "tool_result",
  910. "tool_result",
  911. "round_stats",
  912. "message_delta",
  913. "round_stats",
  914. "done",
  915. ]
  916. assert client.calls == 2
  917. assert [message.role for message in client.second_call_messages] == [
  918. "system",
  919. "user",
  920. "assistant",
  921. "user",
  922. ]
  923. assistant_message = client.second_call_messages[2]
  924. assert assistant_message.content == "Checking events."
  925. assert not any(message.role == "tool" for message in client.second_call_messages)
  926. assert client.second_call_messages[-1].content == (
  927. "EventAgent results:\n"
  928. '{"tool": "handoff_note", "message": "Checking events."}\n'
  929. '{"tool": "audit_note", "message": "Checking events."}'
  930. )
  931. assert client.second_call_messages[-1].name == "event_agent"
  932. @pytest.mark.asyncio
  933. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  934. request = DebugRunRequest(
  935. user_message="debug this",
  936. system_prompts=[],
  937. pre_messages=[],
  938. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  939. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  940. )
  941. runtime = DebugRuntime(RoundStatsChatClient())
  942. queues = runtime.start(request)
  943. outputs: list[dict[str, Any]] = []
  944. while True:
  945. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  946. outputs.append(message)
  947. if message["type"] == "done":
  948. break
  949. assert _message_types(outputs) == [
  950. "session_started",
  951. "message_delta",
  952. "usage",
  953. "round_stats",
  954. "done",
  955. ]
  956. @pytest.mark.asyncio
  957. async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
  958. request = DebugRunRequest(
  959. user_message="debug this",
  960. system_prompts=[],
  961. pre_messages=[],
  962. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  963. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  964. )
  965. client = EventLoopLimitChatClient()
  966. runtime = DebugRuntime(client)
  967. outputs = [message async for message in runtime.run(request)]
  968. business_outputs = _without_audit(outputs)
  969. assert client.calls == 2
  970. assert client.tools_by_call[0] == []
  971. assert client.tools_by_call[1] == []
  972. assert "Available events:" in client.messages_by_call[0][0].content
  973. assert not any(
  974. "Available events:" in message.content
  975. for message in client.messages_by_call[1]
  976. if message.role == "system"
  977. )
  978. assert _message_types(outputs) == [
  979. "session_started",
  980. "event",
  981. "tool_result",
  982. "round_stats",
  983. "message_delta",
  984. "round_stats",
  985. "done",
  986. ]
  987. assert business_outputs[4]["content"] == "final after event limit"
  988. @pytest.mark.asyncio
  989. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  990. RuntimeQueues = _runtime_queues_class()
  991. queues = RuntimeQueues()
  992. request = DebugRunRequest(
  993. user_message="debug this",
  994. system_prompts=[],
  995. pre_messages=[],
  996. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  997. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  998. )
  999. client = StrictHistoryChatClient()
  1000. runtime = DebugRuntime(client, queues=queues)
  1001. stream = runtime.run(request)
  1002. assert await _next_non_audit(stream) == {"type": "session_started"}
  1003. event_message = await _next_non_audit(stream)
  1004. assert event_message["type"] == "event"
  1005. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  1006. remaining = [message async for message in stream]
  1007. assert remaining[-1] == {"type": "done"}
  1008. assert [message.role for message in client.second_call_messages] == [
  1009. "system",
  1010. "user",
  1011. "assistant",
  1012. "user",
  1013. "user",
  1014. ]
  1015. assert not any(message.role == "tool" for message in client.second_call_messages)
  1016. assert client.second_call_messages[3].content.startswith("EventAgent results:\n")
  1017. assert client.second_call_messages[3].name == "event_agent"
  1018. assert client.second_call_messages[4].content == "follow-up while tool runs"
  1019. @pytest.mark.asyncio
  1020. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  1021. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  1022. raise RuntimeError("boom")
  1023. registry = ToolRegistry(
  1024. [
  1025. ToolDefinition(
  1026. name="handoff_note",
  1027. description="Broken handoff tool.",
  1028. parameters={"type": "object"},
  1029. handler=fail_tool,
  1030. )
  1031. ]
  1032. )
  1033. request = DebugRunRequest(
  1034. user_message="debug this",
  1035. system_prompts=[],
  1036. pre_messages=[],
  1037. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1038. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1039. )
  1040. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  1041. outputs = await asyncio.wait_for(
  1042. _collect_outputs(runtime.run(request)),
  1043. timeout=1,
  1044. )
  1045. business_outputs = _without_audit(outputs)
  1046. assert _message_types(outputs) == [
  1047. "session_started",
  1048. "event",
  1049. "tool_result",
  1050. "round_stats",
  1051. "message_delta",
  1052. "round_stats",
  1053. "done",
  1054. ]
  1055. assert json.loads(business_outputs[2]["message"]["content"]) == {
  1056. "tool": "handoff_note",
  1057. "error": "tool handler failed: boom",
  1058. }
  1059. @pytest.mark.asyncio
  1060. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  1061. RuntimeQueues = _runtime_queues_class()
  1062. queue_log: list[tuple[str, str, str]] = []
  1063. queues = RuntimeQueues(
  1064. input=RecordingQueue("input", queue_log),
  1065. output=RecordingQueue("output", queue_log),
  1066. events=RecordingQueue("events", queue_log),
  1067. )
  1068. request = DebugRunRequest(
  1069. user_message="debug this",
  1070. system_prompts=["You are a debugger."],
  1071. pre_messages=[],
  1072. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1073. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  1074. )
  1075. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  1076. outputs = [message async for message in runtime.run(request)]
  1077. assert _message_types(outputs) == [
  1078. "session_started",
  1079. "event",
  1080. "tool_result",
  1081. "round_stats",
  1082. "message_delta",
  1083. "round_stats",
  1084. "done",
  1085. ]
  1086. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  1087. ("input", "get", "user")
  1088. )
  1089. assert queue_log.index(
  1090. ("events", "put", "event_request:handoff_note:call_1")
  1091. ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
  1092. assert queue_log.index(
  1093. ("events", "get", "event_request:handoff_note:call_1")
  1094. ) < queue_log.index(("input", "put", "tool:call_1"))
  1095. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  1096. ("input", "get", "tool:call_1")
  1097. )
  1098. @pytest.mark.asyncio
  1099. async def test_runtime_run_consumes_output_queue_in_stream_order():
  1100. RuntimeQueues = _runtime_queues_class()
  1101. queue_log: list[tuple[str, str, str]] = []
  1102. queues = RuntimeQueues(
  1103. input=RecordingQueue("input", queue_log),
  1104. output=RecordingQueue("output", queue_log),
  1105. events=RecordingQueue("events", queue_log),
  1106. )
  1107. request = DebugRunRequest(
  1108. user_message="debug this",
  1109. system_prompts=["You are a debugger."],
  1110. pre_messages=[],
  1111. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1112. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  1113. )
  1114. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  1115. outputs = [message async for message in runtime.run(request)]
  1116. assert _message_types(outputs) == [
  1117. "session_started",
  1118. "event",
  1119. "tool_result",
  1120. "round_stats",
  1121. "message_delta",
  1122. "round_stats",
  1123. "done",
  1124. ]
  1125. output_puts = [
  1126. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  1127. ]
  1128. output_gets = [
  1129. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  1130. ]
  1131. assert [message for message in output_puts if message != "output:audit"] == [
  1132. "output:session_started",
  1133. "output:event",
  1134. "output:tool_result",
  1135. "output:round_stats",
  1136. "output:message_delta",
  1137. "output:round_stats",
  1138. "output:done",
  1139. ]
  1140. assert output_gets == output_puts
  1141. @pytest.mark.asyncio
  1142. async def test_runtime_continues_with_event_summary_without_tool_call_history():
  1143. request = DebugRunRequest(
  1144. user_message="debug this",
  1145. system_prompts=["You are a debugger."],
  1146. pre_messages=[],
  1147. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1148. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1149. )
  1150. client = StrictHistoryChatClient()
  1151. runtime = DebugRuntime(client)
  1152. outputs = [message async for message in runtime.run(request)]
  1153. assert client.calls == 2
  1154. assert [message.role for message in client.second_call_messages] == [
  1155. "system",
  1156. "system",
  1157. "user",
  1158. "assistant",
  1159. "user",
  1160. ]
  1161. assistant_message = client.second_call_messages[3]
  1162. assert assistant_message.content == ""
  1163. assert "handoff_note" in client.second_call_messages[1].content
  1164. assert not any(message.role == "tool" for message in client.second_call_messages)
  1165. assert client.second_call_messages[4].content.startswith("EventAgent results:\n")
  1166. assert client.second_call_messages[4].name == "event_agent"
  1167. assert outputs[-1] == {"type": "done"}
  1168. @pytest.mark.asyncio
  1169. async def test_runtime_passes_event_catalog_system_message_without_chat_tools():
  1170. registry = ToolRegistry(
  1171. [
  1172. ToolDefinition(
  1173. name="handoff_note",
  1174. description="Registry-owned handoff tool.",
  1175. parameters={
  1176. "type": "object",
  1177. "properties": {
  1178. "message": {"type": "string"},
  1179. "priority": {"type": "number"},
  1180. },
  1181. "required": ["message"],
  1182. },
  1183. handler=lambda event: {"tool": event.name, "message": "handled"},
  1184. )
  1185. ]
  1186. )
  1187. request = DebugRunRequest(
  1188. user_message="debug this",
  1189. system_prompts=[],
  1190. pre_messages=[],
  1191. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1192. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1193. )
  1194. client = ToolCapturingChatClient()
  1195. runtime = DebugRuntime(client, registry=registry)
  1196. outputs = [message async for message in runtime.run(request)]
  1197. assert outputs[-1] == {"type": "done"}
  1198. assert client.tools == []
  1199. assert client.messages[0].role == "system"
  1200. assert "Available events:" in client.messages[0].content
  1201. assert "- handoff_note: Registry-owned handoff tool." in client.messages[0].content
  1202. assert "message" not in client.messages[0].content
  1203. assert "priority" not in client.messages[0].content
  1204. @pytest.mark.asyncio
  1205. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  1206. request = DebugRunRequest(
  1207. user_message="debug this",
  1208. system_prompts=[],
  1209. pre_messages=[],
  1210. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1211. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1212. )
  1213. ticks = iter(
  1214. [
  1215. 0.9,
  1216. 0.91,
  1217. 1.0,
  1218. 1.01,
  1219. 1.02,
  1220. 1.123,
  1221. 1.2,
  1222. 1.25,
  1223. 1.3,
  1224. 1.456,
  1225. 1.7,
  1226. 1.8,
  1227. ]
  1228. )
  1229. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  1230. outputs = [message async for message in runtime.run(request)]
  1231. business_outputs = _without_audit(outputs)
  1232. assert _message_types(outputs) == [
  1233. "session_started",
  1234. "message_delta",
  1235. "usage",
  1236. "round_stats",
  1237. "done",
  1238. ]
  1239. assert business_outputs[3] == {
  1240. "type": "round_stats",
  1241. "round_index": 1,
  1242. "ttft_ms": 123,
  1243. "elapsed_ms": 456,
  1244. "prompt_tokens": 10,
  1245. "completion_tokens": 20,
  1246. "total_tokens": 30,
  1247. "cached_tokens": 5,
  1248. "had_event": False,
  1249. }
  1250. @pytest.mark.asyncio
  1251. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  1252. request = DebugRunRequest(
  1253. user_message="debug this",
  1254. system_prompts=[],
  1255. pre_messages=[],
  1256. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1257. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1258. )
  1259. ticks = iter(
  1260. [
  1261. 1.9,
  1262. 1.91,
  1263. 2.0,
  1264. 2.01,
  1265. 2.02,
  1266. 2.03,
  1267. 2.04,
  1268. 2.05,
  1269. 2.06,
  1270. 2.07,
  1271. 2.25,
  1272. 2.26,
  1273. 3.0,
  1274. 3.01,
  1275. 3.02,
  1276. 3.05,
  1277. 3.1,
  1278. 3.15,
  1279. 3.18,
  1280. 3.2,
  1281. 3.23,
  1282. 3.24,
  1283. 3.25,
  1284. 3.26,
  1285. ]
  1286. )
  1287. client = EventRoundStatsChatClient()
  1288. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  1289. outputs = [message async for message in runtime.run(request)]
  1290. stats = [message for message in outputs if message["type"] == "round_stats"]
  1291. assert client.calls == 2
  1292. assert stats == [
  1293. {
  1294. "type": "round_stats",
  1295. "round_index": 1,
  1296. "ttft_ms": None,
  1297. "elapsed_ms": 250,
  1298. "prompt_tokens": 3,
  1299. "completion_tokens": 0,
  1300. "total_tokens": 3,
  1301. "cached_tokens": 0,
  1302. "had_event": True,
  1303. },
  1304. {
  1305. "type": "round_stats",
  1306. "round_index": 2,
  1307. "ttft_ms": 50,
  1308. "elapsed_ms": 200,
  1309. "prompt_tokens": 4,
  1310. "completion_tokens": 6,
  1311. "total_tokens": 10,
  1312. "cached_tokens": 0,
  1313. "had_event": False,
  1314. },
  1315. ]
  1316. @pytest.mark.asyncio
  1317. async def test_runtime_session_resets_event_budget_for_each_user_turn():
  1318. request = DebugRunRequest(
  1319. user_message="first turn",
  1320. system_prompts=[],
  1321. pre_messages=[],
  1322. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1323. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  1324. )
  1325. client = TwoTurnSessionChatClient()
  1326. runtime = DebugRuntime(client)
  1327. queues = runtime.start_session(request)
  1328. outputs: list[dict[str, Any]] = []
  1329. while len([message for message in outputs if message["type"] == "turn_completed"]) < 1:
  1330. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1331. await queues.input.put(ChatMessage(role="user", content="second turn"))
  1332. while len([message for message in outputs if message["type"] == "turn_completed"]) < 2:
  1333. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1334. await runtime.aclose()
  1335. business_types = [message["type"] for message in _without_audit(outputs)]
  1336. assert business_types.count("turn_started") == 2
  1337. assert business_types.count("turn_completed") == 2
  1338. assert client.calls == 4
  1339. assert "Available events:" in client.messages_by_call[0][0].content
  1340. assert "Available events:" in client.messages_by_call[2][0].content
  1341. @pytest.mark.asyncio
  1342. async def test_runtime_persists_session_turn_messages_audit_and_usage(tmp_path):
  1343. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1344. request = DebugRunRequest(
  1345. session_id="session-1",
  1346. user_message="persist this",
  1347. system_prompts=["You are a debugger."],
  1348. pre_messages=[],
  1349. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  1350. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1351. )
  1352. runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
  1353. queues = runtime.start_session(request)
  1354. outputs: list[dict[str, Any]] = []
  1355. while not any(message["type"] == "turn_completed" for message in outputs):
  1356. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1357. await runtime.aclose()
  1358. assert _without_audit(outputs)[0] == {
  1359. "type": "session_started",
  1360. "session_id": "session-1",
  1361. }
  1362. assert store.get_session("session-1")["config"]["chat_agent"]["model"] == "chat-model"
  1363. assert [
  1364. (message["turn_index"], message["role"], message["content"])
  1365. for message in store.list_messages("session-1")
  1366. ] == [
  1367. (1, "user", "persist this"),
  1368. (1, "assistant", "hello"),
  1369. ]
  1370. audit_events = [
  1371. audit["event"]
  1372. for audit in store.list_audit_logs("session-1")
  1373. ]
  1374. assert "session_started" in audit_events
  1375. assert "chat_agent_request" in audit_events
  1376. assert "turn_completed" in audit_events
  1377. usage = store.usage_summary("session-1")
  1378. assert usage["calls"][0]["total_tokens"] == 30
  1379. assert usage["turns"][0]["turn_index"] == 1
  1380. assert usage["session"]["total_tokens"] == 30
  1381. @pytest.mark.asyncio
  1382. async def test_runtime_continues_persisted_turn_indexes_for_existing_session(tmp_path):
  1383. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1384. session_id = store.create_session(title="existing", config={})
  1385. store.start_turn(session_id, turn_index=1, user_message="old")
  1386. request = DebugRunRequest(
  1387. session_id=session_id,
  1388. user_message="new",
  1389. system_prompts=[],
  1390. pre_messages=[],
  1391. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  1392. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1393. )
  1394. runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
  1395. queues = runtime.start_session(request)
  1396. outputs: list[dict[str, Any]] = []
  1397. while not any(message["type"] == "turn_completed" for message in outputs):
  1398. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1399. await runtime.aclose()
  1400. assert [
  1401. (message["turn_index"], message["content"])
  1402. for message in store.list_messages(session_id)
  1403. ] == [
  1404. (2, "new"),
  1405. (2, "hello"),
  1406. ]