test_debug_runtime.py 101 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117111811191120112111221123112411251126112711281129113011311132113311341135113611371138113911401141114211431144114511461147114811491150115111521153115411551156115711581159116011611162116311641165116611671168116911701171117211731174117511761177117811791180118111821183118411851186118711881189119011911192119311941195119611971198119912001201120212031204120512061207120812091210121112121213121412151216121712181219122012211222122312241225122612271228122912301231123212331234123512361237123812391240124112421243124412451246124712481249125012511252125312541255125612571258125912601261126212631264126512661267126812691270127112721273127412751276127712781279128012811282128312841285128612871288128912901291129212931294129512961297129812991300130113021303130413051306130713081309131013111312131313141315131613171318131913201321132213231324132513261327132813291330133113321333133413351336133713381339134013411342134313441345134613471348134913501351135213531354135513561357135813591360136113621363136413651366136713681369137013711372137313741375137613771378137913801381138213831384138513861387138813891390139113921393139413951396139713981399140014011402140314041405140614071408140914101411141214131414141514161417141814191420142114221423142414251426142714281429143014311432143314341435143614371438143914401441144214431444144514461447144814491450145114521453145414551456145714581459146014611462146314641465146614671468146914701471147214731474147514761477147814791480148114821483148414851486148714881489149014911492149314941495149614971498149915001501150215031504150515061507150815091510151115121513151415151516151715181519152015211522152315241525152615271528152915301531153215331534153515361537153815391540154115421543154415451546154715481549155015511552155315541555155615571558155915601561156215631564156515661567156815691570157115721573157415751576157715781579158015811582158315841585158615871588158915901591159215931594159515961597159815991600160116021603160416051606160716081609161016111612161316141615161616171618161916201621162216231624162516261627162816291630163116321633163416351636163716381639164016411642164316441645164616471648164916501651165216531654165516561657165816591660166116621663166416651666166716681669167016711672167316741675167616771678167916801681168216831684168516861687168816891690169116921693169416951696169716981699170017011702170317041705170617071708170917101711171217131714171517161717171817191720172117221723172417251726172717281729173017311732173317341735173617371738173917401741174217431744174517461747174817491750175117521753175417551756175717581759176017611762176317641765176617671768176917701771177217731774177517761777177817791780178117821783178417851786178717881789179017911792179317941795179617971798179918001801180218031804180518061807180818091810181118121813181418151816181718181819182018211822182318241825182618271828182918301831183218331834183518361837183818391840184118421843184418451846184718481849185018511852185318541855185618571858185918601861186218631864186518661867186818691870187118721873187418751876187718781879188018811882188318841885188618871888188918901891189218931894189518961897189818991900190119021903190419051906190719081909191019111912191319141915191619171918191919201921192219231924192519261927192819291930193119321933193419351936193719381939194019411942194319441945194619471948194919501951195219531954195519561957195819591960196119621963196419651966196719681969197019711972197319741975197619771978197919801981198219831984198519861987198819891990199119921993199419951996199719981999200020012002200320042005200620072008200920102011201220132014201520162017201820192020202120222023202420252026202720282029203020312032203320342035203620372038203920402041204220432044204520462047204820492050205120522053205420552056205720582059206020612062206320642065206620672068206920702071207220732074207520762077207820792080208120822083208420852086208720882089209020912092209320942095209620972098209921002101210221032104210521062107210821092110211121122113211421152116211721182119212021212122212321242125212621272128212921302131213221332134213521362137213821392140214121422143214421452146214721482149215021512152215321542155215621572158215921602161216221632164216521662167216821692170217121722173217421752176217721782179218021812182218321842185218621872188218921902191219221932194219521962197219821992200220122022203220422052206220722082209221022112212221322142215221622172218221922202221222222232224222522262227222822292230223122322233223422352236223722382239224022412242224322442245224622472248224922502251225222532254225522562257225822592260226122622263226422652266226722682269227022712272227322742275227622772278227922802281228222832284228522862287228822892290229122922293229422952296229722982299230023012302230323042305230623072308230923102311231223132314231523162317231823192320232123222323232423252326232723282329233023312332233323342335233623372338233923402341234223432344234523462347234823492350235123522353235423552356235723582359236023612362236323642365236623672368236923702371237223732374237523762377237823792380238123822383238423852386238723882389239023912392239323942395239623972398239924002401240224032404240524062407240824092410241124122413241424152416241724182419242024212422242324242425242624272428242924302431243224332434243524362437243824392440244124422443244424452446244724482449245024512452245324542455245624572458245924602461246224632464246524662467246824692470247124722473247424752476247724782479248024812482248324842485248624872488248924902491249224932494249524962497249824992500250125022503250425052506250725082509251025112512251325142515251625172518251925202521252225232524252525262527252825292530253125322533253425352536253725382539254025412542254325442545254625472548254925502551255225532554255525562557255825592560256125622563256425652566256725682569257025712572257325742575257625772578257925802581258225832584258525862587258825892590259125922593259425952596259725982599260026012602260326042605260626072608260926102611261226132614261526162617261826192620262126222623262426252626262726282629263026312632263326342635263626372638263926402641264226432644264526462647264826492650265126522653265426552656265726582659266026612662266326642665266626672668266926702671267226732674267526762677267826792680268126822683268426852686268726882689269026912692269326942695269626972698269927002701270227032704270527062707270827092710271127122713271427152716271727182719272027212722272327242725272627272728272927302731273227332734273527362737273827392740274127422743274427452746274727482749275027512752275327542755275627572758275927602761276227632764276527662767276827692770277127722773277427752776277727782779278027812782278327842785278627872788278927902791279227932794279527962797279827992800280128022803280428052806280728082809281028112812281328142815281628172818281928202821282228232824282528262827282828292830283128322833283428352836283728382839284028412842284328442845284628472848284928502851285228532854285528562857285828592860286128622863286428652866286728682869287028712872287328742875287628772878287928802881288228832884288528862887288828892890289128922893289428952896289728982899290029012902290329042905290629072908290929102911291229132914291529162917291829192920292129222923292429252926292729282929293029312932293329342935293629372938293929402941294229432944294529462947294829492950295129522953295429552956295729582959296029612962296329642965296629672968296929702971297229732974297529762977297829792980298129822983298429852986298729882989299029912992299329942995299629972998299930003001300230033004300530063007300830093010301130123013301430153016301730183019302030213022302330243025302630273028302930303031303230333034303530363037303830393040304130423043304430453046304730483049305030513052305330543055305630573058305930603061306230633064306530663067306830693070307130723073307430753076307730783079308030813082308330843085308630873088308930903091309230933094309530963097309830993100310131023103310431053106310731083109311031113112311331143115311631173118311931203121312231233124312531263127312831293130
  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.events import EventBatchExecutor, ResultPolicy
  11. from agent_lab.application.runtime import DebugRuntime
  12. from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
  13. from agent_lab.domain.events import ToolCallEvent
  14. from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
  15. from agent_lab.infrastructure.chat_client import OpenAICompatibleChatClient
  16. from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
  17. def _runtime_queues_class():
  18. module = importlib.import_module("agent_lab.application.queues")
  19. return module.RuntimeQueues
  20. async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[str, Any]]:
  21. return [message async for message in stream]
  22. def _without_audit(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
  23. return [message for message in outputs if message["type"] != "audit"]
  24. def _message_types(outputs: list[dict[str, Any]]) -> list[str]:
  25. return [message["type"] for message in _without_audit(outputs)]
  26. def _event_tool_call_from_tools(
  27. tools: list[dict],
  28. messages: list[ChatMessage],
  29. ) -> StreamItem:
  30. tool_name = tools[0]["function"]["name"]
  31. content = ""
  32. for role in ("assistant", "user"):
  33. content = next(
  34. (
  35. message.content
  36. for message in reversed(messages)
  37. if message.role == role and message.content.strip()
  38. ),
  39. "",
  40. )
  41. if content:
  42. break
  43. arguments = {
  44. "message": content,
  45. "query": content,
  46. "title": content,
  47. }
  48. return StreamItem.provider_tool_call(
  49. ToolCallEvent(
  50. id="event_agent_call_1",
  51. name=tool_name,
  52. arguments=arguments,
  53. raw_arguments=json.dumps(arguments),
  54. )
  55. )
  56. async def _next_non_audit(
  57. stream: AsyncIterator[dict[str, Any]],
  58. ) -> dict[str, Any]:
  59. while True:
  60. message = await anext(stream)
  61. if message["type"] != "audit":
  62. return message
  63. async def _next_non_audit_from_queue(queues: Any) -> dict[str, Any]:
  64. while True:
  65. message = await queues.output.get()
  66. if message["type"] != "audit":
  67. return message
  68. class RecordingQueue(asyncio.Queue):
  69. def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
  70. super().__init__()
  71. self.name = name
  72. self.log = log
  73. async def put(self, item: Any) -> None:
  74. self.log.append((self.name, "put", self._describe(item)))
  75. await super().put(item)
  76. async def get(self) -> Any:
  77. item = await super().get()
  78. self.log.append((self.name, "get", self._describe(item)))
  79. return item
  80. def _describe(self, item: Any) -> str:
  81. if isinstance(item, ChatMessage):
  82. if item.role == "tool":
  83. return f"tool:{item.tool_call_id}"
  84. return item.role
  85. if isinstance(item, ToolCallEvent):
  86. return f"event:{item.name}:{item.id}"
  87. if isinstance(item, EventAgentRequest):
  88. events = ",".join(f"{event.name}:{event.id}" for event in item.events)
  89. return f"event_request:{events}"
  90. if isinstance(item, dict):
  91. return f"output:{item.get('type')}"
  92. return type(item).__name__
  93. class FakeChatClient:
  94. def __init__(self) -> None:
  95. self.calls = 0
  96. async def stream_chat(
  97. self,
  98. messages: list[ChatMessage],
  99. tools: list[dict],
  100. params: AgentParams,
  101. tool_choice: dict[str, Any] | None = None,
  102. ) -> AsyncIterator[StreamItem]:
  103. if tools:
  104. yield StreamItem.raw_response_chunk(
  105. {
  106. "choices": [
  107. {
  108. "delta": {
  109. "tool_calls": [
  110. {
  111. "index": 0,
  112. "function": {"name": tools[0]["function"]["name"]},
  113. }
  114. ]
  115. },
  116. "finish_reason": None,
  117. }
  118. ]
  119. }
  120. )
  121. yield _event_tool_call_from_tools(tools, messages)
  122. return
  123. self.calls += 1
  124. if self.calls == 1:
  125. yield StreamItem.raw_response_chunk(
  126. {
  127. "choices": [
  128. {
  129. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  130. "finish_reason": None,
  131. }
  132. ]
  133. }
  134. )
  135. yield StreamItem.text_event(
  136. ToolCallEvent(
  137. id="call_1",
  138. name="handoff_note",
  139. arguments={},
  140. raw_arguments="{}",
  141. )
  142. )
  143. return
  144. assert not any(message.role == "tool" for message in messages)
  145. assert any(message.name == "event_agent" for message in messages)
  146. yield StreamItem.message_delta("final answer")
  147. class WrongSourceEventChatClient:
  148. def __init__(self, item: StreamItem) -> None:
  149. self.item = item
  150. self.calls = 0
  151. async def stream_chat(
  152. self,
  153. messages: list[ChatMessage],
  154. tools: list[dict],
  155. params: AgentParams,
  156. tool_choice: dict[str, Any] | None = None,
  157. ) -> AsyncIterator[StreamItem]:
  158. if tools:
  159. yield _event_tool_call_from_tools(tools, messages)
  160. return
  161. self.calls += 1
  162. if self.calls == 1:
  163. yield self.item
  164. return
  165. yield StreamItem.message_delta("unexpected continuation")
  166. class IncrementingClock:
  167. def __init__(self, current: float = 100.0, step: float = 0.01) -> None:
  168. self.current = current
  169. self.step = step
  170. def __call__(self) -> float:
  171. self.current += self.step
  172. return self.current
  173. class StrictHistoryChatClient:
  174. def __init__(self) -> None:
  175. self.calls = 0
  176. self.second_call_messages: list[ChatMessage] = []
  177. async def stream_chat(
  178. self,
  179. messages: list[ChatMessage],
  180. tools: list[dict],
  181. params: AgentParams,
  182. tool_choice: dict[str, Any] | None = None,
  183. ) -> AsyncIterator[StreamItem]:
  184. if tools:
  185. yield StreamItem.raw_response_chunk(
  186. {
  187. "choices": [
  188. {
  189. "delta": {
  190. "tool_calls": [
  191. {
  192. "index": 0,
  193. "function": {"name": tools[0]["function"]["name"]},
  194. }
  195. ]
  196. },
  197. "finish_reason": None,
  198. }
  199. ]
  200. }
  201. )
  202. yield _event_tool_call_from_tools(tools, messages)
  203. return
  204. self.calls += 1
  205. if self.calls == 1:
  206. yield StreamItem.raw_response_chunk(
  207. {
  208. "choices": [
  209. {
  210. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  211. "finish_reason": None,
  212. }
  213. ]
  214. }
  215. )
  216. yield StreamItem.text_event(
  217. ToolCallEvent(
  218. id="call_1",
  219. name="handoff_note",
  220. arguments={},
  221. raw_arguments="{}",
  222. )
  223. )
  224. return
  225. self.second_call_messages = list(messages)
  226. yield StreamItem.message_delta("final answer")
  227. class MultiEventChatClient:
  228. def __init__(self) -> None:
  229. self.calls = 0
  230. self.second_call_messages: list[ChatMessage] = []
  231. async def stream_chat(
  232. self,
  233. messages: list[ChatMessage],
  234. tools: list[dict],
  235. params: AgentParams,
  236. tool_choice: dict[str, Any] | None = None,
  237. ) -> AsyncIterator[StreamItem]:
  238. if tools:
  239. yield _event_tool_call_from_tools(tools, messages)
  240. return
  241. self.calls += 1
  242. if self.calls == 1:
  243. yield StreamItem.message_delta("Checking events.")
  244. yield StreamItem.text_event(
  245. ToolCallEvent(
  246. id="call_1",
  247. name="handoff_note",
  248. arguments={"message": "ignored chat argument"},
  249. raw_arguments='{"message":"ignored chat argument"}',
  250. )
  251. )
  252. yield StreamItem.text_event(
  253. ToolCallEvent(
  254. id="call_2",
  255. name="audit_note",
  256. arguments={"message": "ignored chat argument"},
  257. raw_arguments='{"message":"ignored chat argument"}',
  258. )
  259. )
  260. return
  261. self.second_call_messages = list(messages)
  262. yield StreamItem.message_delta("Final answer.")
  263. class ToolCapturingChatClient:
  264. def __init__(self) -> None:
  265. self.tools: list[dict[str, Any]] = []
  266. self.messages: list[ChatMessage] = []
  267. async def stream_chat(
  268. self,
  269. messages: list[ChatMessage],
  270. tools: list[dict],
  271. params: AgentParams,
  272. tool_choice: dict[str, Any] | None = None,
  273. ) -> AsyncIterator[StreamItem]:
  274. self.messages = list(messages)
  275. self.tools = list(tools)
  276. yield StreamItem.message_delta("final answer")
  277. class EventLoopLimitChatClient:
  278. def __init__(self) -> None:
  279. self.calls = 0
  280. self.tools_by_call: list[list[dict[str, Any]]] = []
  281. self.messages_by_call: list[list[ChatMessage]] = []
  282. async def stream_chat(
  283. self,
  284. messages: list[ChatMessage],
  285. tools: list[dict],
  286. params: AgentParams,
  287. tool_choice: dict[str, Any] | None = None,
  288. ) -> AsyncIterator[StreamItem]:
  289. if tools:
  290. yield _event_tool_call_from_tools(tools, messages)
  291. return
  292. self.calls += 1
  293. self.tools_by_call.append(list(tools))
  294. self.messages_by_call.append(list(messages))
  295. if self.calls == 1:
  296. yield StreamItem.text_event(
  297. ToolCallEvent(
  298. id="call_1",
  299. name="handoff_note",
  300. arguments={},
  301. raw_arguments="{}",
  302. )
  303. )
  304. return
  305. yield StreamItem.message_delta("final after event limit")
  306. class RoundStatsChatClient:
  307. async def stream_chat(
  308. self,
  309. messages: list[ChatMessage],
  310. tools: list[dict],
  311. params: AgentParams,
  312. tool_choice: dict[str, Any] | None = None,
  313. ) -> AsyncIterator[StreamItem]:
  314. yield StreamItem.message_delta("hello")
  315. yield StreamItem.usage_item(
  316. TokenUsage(
  317. prompt_tokens=10,
  318. completion_tokens=20,
  319. total_tokens=30,
  320. cached_tokens=5,
  321. )
  322. )
  323. class EventRoundStatsChatClient:
  324. def __init__(self) -> None:
  325. self.calls = 0
  326. async def stream_chat(
  327. self,
  328. messages: list[ChatMessage],
  329. tools: list[dict],
  330. params: AgentParams,
  331. tool_choice: dict[str, Any] | None = None,
  332. ) -> AsyncIterator[StreamItem]:
  333. if tools:
  334. yield StreamItem.raw_response_chunk(
  335. {
  336. "choices": [
  337. {
  338. "delta": {
  339. "tool_calls": [
  340. {
  341. "index": 0,
  342. "function": {
  343. "name": tools[0]["function"]["name"],
  344. },
  345. }
  346. ]
  347. },
  348. "finish_reason": None,
  349. }
  350. ]
  351. }
  352. )
  353. yield _event_tool_call_from_tools(tools, messages)
  354. return
  355. self.calls += 1
  356. if self.calls == 1:
  357. yield StreamItem.raw_response_chunk(
  358. {
  359. "choices": [
  360. {
  361. "delta": {
  362. "content": "<agent_events>handoff_note</agent_events>",
  363. },
  364. "finish_reason": None,
  365. }
  366. ]
  367. }
  368. )
  369. yield StreamItem.text_event(
  370. ToolCallEvent(
  371. id="call_1",
  372. name="handoff_note",
  373. arguments={},
  374. raw_arguments="{}",
  375. )
  376. )
  377. yield StreamItem.usage_item(
  378. TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3)
  379. )
  380. return
  381. yield StreamItem.raw_response_chunk(
  382. {
  383. "choices": [
  384. {
  385. "delta": {"content": "final answer"},
  386. "finish_reason": None,
  387. }
  388. ]
  389. }
  390. )
  391. yield StreamItem.message_delta("final answer")
  392. yield StreamItem.usage_item(
  393. TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
  394. )
  395. class ContextBoundaryChatClient:
  396. def __init__(self) -> None:
  397. self.calls = 0
  398. self.event_agent_messages: list[list[ChatMessage]] = []
  399. async def stream_chat(
  400. self,
  401. messages: list[ChatMessage],
  402. tools: list[dict],
  403. params: AgentParams,
  404. tool_choice: dict[str, Any] | None = None,
  405. ) -> AsyncIterator[StreamItem]:
  406. if tools:
  407. self.event_agent_messages.append(list(messages))
  408. yield _event_tool_call_from_tools(tools, messages)
  409. return
  410. self.calls += 1
  411. if self.calls == 1:
  412. yield StreamItem.message_delta("I will check.")
  413. yield StreamItem.text_event(
  414. ToolCallEvent(
  415. id="call_1",
  416. name="handoff_note",
  417. arguments={},
  418. raw_arguments="{}",
  419. )
  420. )
  421. return
  422. yield StreamItem.message_delta("final answer")
  423. class TwoTurnSessionChatClient:
  424. def __init__(self) -> None:
  425. self.calls = 0
  426. self.messages_by_call: list[list[ChatMessage]] = []
  427. async def stream_chat(
  428. self,
  429. messages: list[ChatMessage],
  430. tools: list[dict],
  431. params: AgentParams,
  432. tool_choice: dict[str, Any] | None = None,
  433. ) -> AsyncIterator[StreamItem]:
  434. if tools:
  435. yield _event_tool_call_from_tools(tools, messages)
  436. return
  437. self.calls += 1
  438. self.messages_by_call.append(list(messages))
  439. if self.calls in {1, 3}:
  440. yield StreamItem.text_event(
  441. ToolCallEvent(
  442. id=f"call_{self.calls}",
  443. name="handoff_note",
  444. arguments={},
  445. raw_arguments="{}",
  446. )
  447. )
  448. return
  449. yield StreamItem.message_delta(f"final answer {self.calls}")
  450. class SlowAfterEventChatClient:
  451. def __init__(self) -> None:
  452. self.calls = 0
  453. self.event_seen = asyncio.Event()
  454. self.release_stream = asyncio.Event()
  455. async def stream_chat(
  456. self,
  457. messages: list[ChatMessage],
  458. tools: list[dict],
  459. params: AgentParams,
  460. tool_choice: dict[str, Any] | None = None,
  461. ) -> AsyncIterator[StreamItem]:
  462. if tools:
  463. yield _event_tool_call_from_tools(tools, messages)
  464. return
  465. self.calls += 1
  466. if self.calls == 1:
  467. yield StreamItem.message_delta("Need event.")
  468. yield StreamItem.text_event(
  469. ToolCallEvent(
  470. id="call_1",
  471. name="handoff_note",
  472. arguments={},
  473. raw_arguments="{}",
  474. )
  475. )
  476. self.event_seen.set()
  477. await self.release_stream.wait()
  478. return
  479. yield StreamItem.message_delta("final answer")
  480. def test_runtime_queues_exposes_input_output_and_events_queues():
  481. RuntimeQueues = _runtime_queues_class()
  482. queues = RuntimeQueues()
  483. assert isinstance(queues.input, asyncio.Queue)
  484. assert isinstance(queues.output, asyncio.Queue)
  485. assert isinstance(queues.events, asyncio.Queue)
  486. assert queues.input is not queues.output
  487. assert queues.input is not queues.events
  488. assert queues.output is not queues.events
  489. @pytest.mark.asyncio
  490. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  491. request = DebugRunRequest(
  492. user_message="debug this",
  493. system_prompts=["You are a debugger."],
  494. pre_messages=[],
  495. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  496. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  497. )
  498. client = FakeChatClient()
  499. runtime = DebugRuntime(client)
  500. outputs = [message async for message in runtime.run(request)]
  501. business_outputs = _without_audit(outputs)
  502. assert client.calls == 2
  503. assert _message_types(outputs) == [
  504. "session_started",
  505. "event",
  506. "tool_result",
  507. "round_stats",
  508. "message_delta",
  509. "round_stats",
  510. "done",
  511. ]
  512. assert business_outputs[1]["event"]["name"] == "handoff_note"
  513. assert business_outputs[4]["content"] == "final answer"
  514. @pytest.mark.asyncio
  515. @pytest.mark.parametrize(
  516. "item",
  517. [
  518. StreamItem.provider_tool_call(
  519. ToolCallEvent(
  520. id="provider_call_1",
  521. name="handoff_note",
  522. arguments={"message": "wrong source"},
  523. raw_arguments='{"message":"wrong source"}',
  524. )
  525. ),
  526. StreamItem.event(
  527. ToolCallEvent(
  528. id="legacy_event_1",
  529. name="handoff_note",
  530. arguments={},
  531. raw_arguments="{}",
  532. )
  533. ),
  534. ],
  535. ids=["provider_tool_call", "legacy_event"],
  536. )
  537. async def test_runtime_ignores_non_text_event_sources(item: StreamItem):
  538. request = DebugRunRequest(
  539. user_message="debug this",
  540. system_prompts=[],
  541. pre_messages=[],
  542. chat_agent=AgentParams(model="fake-model"),
  543. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  544. )
  545. client = WrongSourceEventChatClient(item)
  546. outputs = [message async for message in DebugRuntime(client).run(request)]
  547. assert client.calls == 1
  548. assert "event" not in _message_types(outputs)
  549. assert "tool_result" not in _message_types(outputs)
  550. @pytest.mark.asyncio
  551. async def test_runtime_emits_audit_events_and_backend_logs(caplog):
  552. request = DebugRunRequest(
  553. user_message="debug this",
  554. system_prompts=[],
  555. pre_messages=[],
  556. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  557. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  558. )
  559. runtime = DebugRuntime(FakeChatClient())
  560. with caplog.at_level(logging.INFO, logger="agent_lab.application.runtime"):
  561. outputs = [message async for message in runtime.run(request)]
  562. audit_events = [
  563. message["event"]
  564. for message in outputs
  565. if message["type"] == "audit"
  566. ]
  567. assert audit_events == [
  568. "session_started",
  569. "chat_round_started",
  570. "chat_agent_request",
  571. "chat_event_detected",
  572. "chat_agent_response",
  573. "event_batch_started",
  574. "event_agent_request",
  575. "event_agent_response",
  576. "event_agent_completed",
  577. "event_batch_results",
  578. "event_policy_decision",
  579. "chat_round_finished",
  580. "chat_round_started",
  581. "chat_agent_request",
  582. "chat_message_stream_started",
  583. "chat_message_stream_finished",
  584. "chat_agent_response",
  585. "chat_round_finished",
  586. "session_finished",
  587. ]
  588. assert "chat_event_detected" in caplog.text
  589. assert "event_agent_completed" in caplog.text
  590. @pytest.mark.asyncio
  591. async def test_runtime_audits_chat_message_stream_boundaries_in_output_order():
  592. request = DebugRunRequest(
  593. user_message="debug this",
  594. system_prompts=[],
  595. pre_messages=[],
  596. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  597. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  598. )
  599. runtime = DebugRuntime(RoundStatsChatClient())
  600. outputs = [message async for message in runtime.run(request)]
  601. ordered_labels = [
  602. message["event"] if message["type"] == "audit" else message["type"]
  603. for message in outputs
  604. ]
  605. assert ordered_labels.index("chat_message_stream_started") < ordered_labels.index(
  606. "message_delta"
  607. )
  608. assert ordered_labels.index("message_delta") < ordered_labels.index(
  609. "chat_message_stream_finished"
  610. )
  611. assert ordered_labels.index("chat_message_stream_finished") < ordered_labels.index(
  612. "chat_agent_response"
  613. )
  614. stream_started = next(
  615. message
  616. for message in outputs
  617. if message.get("event") == "chat_message_stream_started"
  618. )
  619. stream_finished = next(
  620. message
  621. for message in outputs
  622. if message.get("event") == "chat_message_stream_finished"
  623. )
  624. assert stream_started["details"]["agent"] == "chat_agent"
  625. assert stream_started["details"]["round_index"] == 1
  626. assert stream_finished["details"]["delta_count"] == 1
  627. assert stream_finished["details"]["content_length"] == len("hello")
  628. @pytest.mark.asyncio
  629. async def test_runtime_audit_events_include_turn_relative_elapsed_time():
  630. request = DebugRunRequest(
  631. user_message="debug this",
  632. system_prompts=[],
  633. pre_messages=[],
  634. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  635. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  636. )
  637. runtime = DebugRuntime(FakeChatClient(), clock=IncrementingClock())
  638. queues = runtime.start_session(request)
  639. outputs: list[dict[str, Any]] = []
  640. while not any(message["type"] == "turn_completed" for message in outputs):
  641. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  642. await runtime.aclose()
  643. turn_audits = [
  644. message
  645. for message in outputs
  646. if message["type"] == "audit"
  647. ]
  648. elapsed_values = [
  649. message["details"].get("turn_elapsed_ms") for message in turn_audits
  650. ]
  651. assert turn_audits
  652. assert all(isinstance(value, int) for value in elapsed_values)
  653. assert all(value >= 0 for value in elapsed_values)
  654. assert elapsed_values == sorted(elapsed_values)
  655. assert next(
  656. message
  657. for message in turn_audits
  658. if message["event"] == "event_agent_request"
  659. )["details"]["turn_elapsed_ms"] >= 0
  660. @pytest.mark.asyncio
  661. async def test_runtime_audit_includes_model_params_prompts_results_and_usage():
  662. request = DebugRunRequest(
  663. user_message="debug this",
  664. system_prompts=["You are a debugger."],
  665. pre_messages=[],
  666. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  667. event_agent=EventAgentParams(
  668. model="event-model",
  669. temperature=0.4,
  670. max_tokens=80,
  671. enabled_tools=["handoff_note"],
  672. max_event_loops=1,
  673. ),
  674. )
  675. runtime = DebugRuntime(EventRoundStatsChatClient())
  676. outputs = [message async for message in runtime.run(request)]
  677. audits = [message for message in outputs if message["type"] == "audit"]
  678. chat_request = next(
  679. message for message in audits if message["event"] == "chat_agent_request"
  680. )
  681. assert chat_request["details"]["agent"] == "chat_agent"
  682. assert chat_request["details"]["params"]["model"] == "chat-model"
  683. assert chat_request["details"]["params"]["temperature"] == 0.1
  684. assert chat_request["details"]["params"]["max_tokens"] == 200
  685. assert chat_request["details"]["tools"] == []
  686. assert chat_request["details"]["messages"][0] == {
  687. "role": "system",
  688. "content": "You are a debugger.",
  689. "name": None,
  690. "tool_call_id": None,
  691. "tool_calls": [],
  692. }
  693. chat_response = next(
  694. message for message in audits if message["event"] == "chat_agent_response"
  695. )
  696. assert chat_response["details"]["event_names"] == ["handoff_note"]
  697. assert chat_response["details"]["raw_chunks"] == [
  698. {
  699. "choices": [
  700. {
  701. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  702. "finish_reason": None,
  703. }
  704. ]
  705. }
  706. ]
  707. assert chat_response["details"]["usage"] == {
  708. "prompt_tokens": 3,
  709. "completion_tokens": 0,
  710. "total_tokens": 3,
  711. "cached_tokens": 0,
  712. }
  713. event_request = next(
  714. message for message in audits if message["event"] == "event_agent_request"
  715. )
  716. assert event_request["details"]["agent"] == "event_agent"
  717. assert event_request["details"]["params"]["model"] == "event-model"
  718. assert event_request["details"]["params"]["temperature"] == 0.4
  719. assert event_request["details"]["events"][0]["name"] == "handoff_note"
  720. assert event_request["details"]["tools"][0]["function"]["name"] == "handoff_note"
  721. assert "Authorization" not in str(event_request["details"])
  722. event_response = next(
  723. message for message in audits if message["event"] == "event_agent_response"
  724. )
  725. assert event_response["details"]["replies"][0]["role"] == "tool"
  726. assert '"tool": "handoff_note"' in event_response["details"]["replies"][0]["content"]
  727. assert event_response["details"]["raw_model_chunks"][0]["event_name"] == "handoff_note"
  728. assert event_response["details"]["raw_model_chunks"][0]["chunks"] == []
  729. @pytest.mark.asyncio
  730. async def test_runtime_round_started_separates_available_events_from_round_budget():
  731. request = DebugRunRequest(
  732. user_message="debug this",
  733. system_prompts=[],
  734. pre_messages=[],
  735. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  736. event_agent=EventAgentParams(
  737. enabled_tools=["handoff_note"],
  738. max_event_loops=1,
  739. ),
  740. )
  741. runtime = DebugRuntime(EventRoundStatsChatClient())
  742. outputs = [message async for message in runtime.run(request)]
  743. round_starts = [
  744. message
  745. for message in outputs
  746. if message.get("event") == "chat_round_started"
  747. ]
  748. assert round_starts[0]["details"]["events_enabled"] == ["handoff_note"]
  749. assert round_starts[0]["details"]["configured_events"] == ["handoff_note"]
  750. assert round_starts[0]["details"]["event_generation_enabled"] is True
  751. assert round_starts[0]["details"]["event_prompt_events"] == ["handoff_note"]
  752. assert round_starts[1]["details"]["events_enabled"] == []
  753. assert round_starts[1]["details"]["configured_events"] == ["handoff_note"]
  754. assert round_starts[1]["details"]["event_generation_enabled"] is False
  755. assert round_starts[1]["details"]["event_prompt_events"] == []
  756. @pytest.mark.asyncio
  757. async def test_runtime_event_agent_history_excludes_chat_agent_system_context():
  758. request = DebugRunRequest(
  759. user_message="debug this",
  760. system_prompts=["ChatAgent root prompt."],
  761. pre_messages=[
  762. ChatMessage(role="system", content="ChatAgent pre system."),
  763. ChatMessage(role="user", content="earlier user"),
  764. ChatMessage(role="assistant", content="earlier assistant"),
  765. ],
  766. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  767. event_agent=EventAgentParams(
  768. model="event-model",
  769. enabled_tools=["handoff_note"],
  770. max_event_loops=1,
  771. ),
  772. )
  773. client = ContextBoundaryChatClient()
  774. runtime = DebugRuntime(client)
  775. outputs = [message async for message in runtime.run(request)]
  776. event_request = next(
  777. message for message in outputs if message.get("event") == "event_agent_request"
  778. )
  779. assert [
  780. (message["role"], message["content"])
  781. for message in event_request["details"]["history"]
  782. ] == [
  783. ("user", "earlier user"),
  784. ("assistant", "earlier assistant"),
  785. ("user", "debug this"),
  786. ("assistant", "I will check."),
  787. ]
  788. assert not any(
  789. message["role"] == "system"
  790. for message in event_request["details"]["history"]
  791. )
  792. assert client.event_agent_messages == []
  793. @pytest.mark.asyncio
  794. async def test_runtime_session_event_agent_history_excludes_internal_replies():
  795. request = DebugRunRequest(
  796. user_message="debug this",
  797. system_prompts=["ChatAgent root prompt."],
  798. pre_messages=[ChatMessage(role="assistant", content="prior answer")],
  799. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  800. event_agent=EventAgentParams(
  801. model="event-model",
  802. enabled_tools=["handoff_note"],
  803. max_event_loops=1,
  804. ),
  805. )
  806. runtime = DebugRuntime(ContextBoundaryChatClient())
  807. queues = runtime.start_session(request)
  808. outputs: list[dict[str, Any]] = []
  809. while not any(message["type"] == "turn_completed" for message in outputs):
  810. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  811. await runtime.aclose()
  812. event_request = next(
  813. message for message in outputs if message.get("event") == "event_agent_request"
  814. )
  815. assert [
  816. (message["role"], message["content"], message["name"])
  817. for message in event_request["details"]["history"]
  818. ] == [
  819. ("assistant", "prior answer", None),
  820. ("user", "debug this", None),
  821. ("assistant", "I will check.", None),
  822. ]
  823. assert not any(
  824. message["name"] == "event_agent"
  825. for message in event_request["details"]["history"]
  826. )
  827. @pytest.mark.asyncio
  828. async def test_runtime_outputs_event_as_soon_as_chat_stream_detects_it():
  829. RuntimeQueues = _runtime_queues_class()
  830. queues = RuntimeQueues()
  831. client = SlowAfterEventChatClient()
  832. runtime = DebugRuntime(client, queues=queues)
  833. request = DebugRunRequest(
  834. user_message="debug this",
  835. system_prompts=[],
  836. pre_messages=[],
  837. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  838. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  839. )
  840. runtime.start(request)
  841. try:
  842. assert await _next_non_audit_from_queue(queues) == {"type": "session_started"}
  843. assert await _next_non_audit_from_queue(queues) == {
  844. "type": "message_delta",
  845. "content": "Need event.",
  846. }
  847. await asyncio.wait_for(client.event_seen.wait(), timeout=1)
  848. event_message = await asyncio.wait_for(
  849. _next_non_audit_from_queue(queues),
  850. timeout=0.2,
  851. )
  852. assert event_message["type"] == "event"
  853. assert event_message["event"]["name"] == "handoff_note"
  854. finally:
  855. client.release_stream.set()
  856. await runtime.aclose()
  857. @pytest.mark.asyncio
  858. async def test_runtime_batches_round_events_before_continuing_chat_agent():
  859. def resolve_from_history(
  860. event: ToolCallEvent,
  861. context: ToolExecutionContext,
  862. ) -> dict[str, Any]:
  863. return {"message": context.history[-1].content, "event": event.name}
  864. registry = ToolRegistry(
  865. [
  866. ToolDefinition(
  867. name="handoff_note",
  868. description="Send a handoff note.",
  869. parameters={
  870. "type": "object",
  871. "properties": {"message": {"type": "string"}},
  872. "required": ["message"],
  873. },
  874. handler=lambda event: {
  875. "tool": event.name,
  876. "message": event.arguments["message"],
  877. },
  878. argument_resolver=resolve_from_history,
  879. ),
  880. ToolDefinition(
  881. name="audit_note",
  882. description="Send an audit note.",
  883. parameters={
  884. "type": "object",
  885. "properties": {"message": {"type": "string"}},
  886. "required": ["message"],
  887. },
  888. handler=lambda event: {
  889. "tool": event.name,
  890. "message": event.arguments["message"],
  891. },
  892. argument_resolver=resolve_from_history,
  893. ),
  894. ]
  895. )
  896. request = DebugRunRequest(
  897. user_message="debug this",
  898. system_prompts=[],
  899. pre_messages=[],
  900. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  901. event_agent=EventAgentParams(
  902. enabled_tools=["handoff_note", "audit_note"],
  903. max_event_loops=2,
  904. ),
  905. )
  906. client = MultiEventChatClient()
  907. runtime = DebugRuntime(client, registry=registry)
  908. outputs = [message async for message in runtime.run(request)]
  909. assert _message_types(outputs) == [
  910. "session_started",
  911. "message_delta",
  912. "event",
  913. "event",
  914. "tool_result",
  915. "tool_result",
  916. "round_stats",
  917. "message_delta",
  918. "round_stats",
  919. "done",
  920. ]
  921. assert client.calls == 2
  922. assert [message.role for message in client.second_call_messages] == [
  923. "system",
  924. "user",
  925. "assistant",
  926. "user",
  927. ]
  928. assistant_message = client.second_call_messages[2]
  929. assert assistant_message.content == "Checking events."
  930. assert not any(message.role == "tool" for message in client.second_call_messages)
  931. assert client.second_call_messages[-1].content == (
  932. "EventAgent results:\n"
  933. '{"tool": "handoff_note", "message": "Checking events."}\n'
  934. '{"tool": "audit_note", "message": "Checking events."}'
  935. )
  936. assert client.second_call_messages[-1].name == "event_agent"
  937. @pytest.mark.asyncio
  938. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  939. request = DebugRunRequest(
  940. user_message="debug this",
  941. system_prompts=[],
  942. pre_messages=[],
  943. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  944. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  945. )
  946. runtime = DebugRuntime(RoundStatsChatClient())
  947. queues = runtime.start(request)
  948. outputs: list[dict[str, Any]] = []
  949. while True:
  950. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  951. outputs.append(message)
  952. if message["type"] == "done":
  953. break
  954. assert _message_types(outputs) == [
  955. "session_started",
  956. "message_delta",
  957. "usage",
  958. "round_stats",
  959. "done",
  960. ]
  961. @pytest.mark.asyncio
  962. async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
  963. request = DebugRunRequest(
  964. user_message="debug this",
  965. system_prompts=[],
  966. pre_messages=[],
  967. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  968. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  969. )
  970. client = EventLoopLimitChatClient()
  971. runtime = DebugRuntime(client)
  972. outputs = [message async for message in runtime.run(request)]
  973. business_outputs = _without_audit(outputs)
  974. assert client.calls == 2
  975. assert client.tools_by_call[0] == []
  976. assert client.tools_by_call[1] == []
  977. assert "Available events:" in client.messages_by_call[0][0].content
  978. assert not any(
  979. "Available events:" in message.content
  980. for message in client.messages_by_call[1]
  981. if message.role == "system"
  982. )
  983. assert _message_types(outputs) == [
  984. "session_started",
  985. "event",
  986. "tool_result",
  987. "round_stats",
  988. "message_delta",
  989. "round_stats",
  990. "done",
  991. ]
  992. assert business_outputs[4]["content"] == "final after event limit"
  993. @pytest.mark.asyncio
  994. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  995. RuntimeQueues = _runtime_queues_class()
  996. queues = RuntimeQueues()
  997. request = DebugRunRequest(
  998. user_message="debug this",
  999. system_prompts=[],
  1000. pre_messages=[],
  1001. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1002. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1003. )
  1004. client = StrictHistoryChatClient()
  1005. runtime = DebugRuntime(client, queues=queues)
  1006. stream = runtime.run(request)
  1007. assert await _next_non_audit(stream) == {"type": "session_started"}
  1008. event_message = await _next_non_audit(stream)
  1009. assert event_message["type"] == "event"
  1010. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  1011. remaining = [message async for message in stream]
  1012. assert remaining[-1] == {"type": "done"}
  1013. assert [message.role for message in client.second_call_messages] == [
  1014. "system",
  1015. "user",
  1016. "assistant",
  1017. "user",
  1018. "user",
  1019. ]
  1020. assert not any(message.role == "tool" for message in client.second_call_messages)
  1021. assert client.second_call_messages[3].content.startswith("EventAgent results:\n")
  1022. assert client.second_call_messages[3].name == "event_agent"
  1023. assert client.second_call_messages[4].content == "follow-up while tool runs"
  1024. @pytest.mark.asyncio
  1025. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  1026. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  1027. raise RuntimeError("boom")
  1028. registry = ToolRegistry(
  1029. [
  1030. ToolDefinition(
  1031. name="handoff_note",
  1032. description="Broken handoff tool.",
  1033. parameters={"type": "object"},
  1034. handler=fail_tool,
  1035. )
  1036. ]
  1037. )
  1038. request = DebugRunRequest(
  1039. user_message="debug this",
  1040. system_prompts=[],
  1041. pre_messages=[],
  1042. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1043. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1044. )
  1045. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  1046. outputs = await asyncio.wait_for(
  1047. _collect_outputs(runtime.run(request)),
  1048. timeout=1,
  1049. )
  1050. business_outputs = _without_audit(outputs)
  1051. assert _message_types(outputs) == [
  1052. "session_started",
  1053. "event",
  1054. "tool_result",
  1055. "round_stats",
  1056. "message_delta",
  1057. "round_stats",
  1058. "done",
  1059. ]
  1060. assert json.loads(business_outputs[2]["message"]["content"]) == {
  1061. "tool": "handoff_note",
  1062. "error": "tool handler failed: boom",
  1063. }
  1064. @pytest.mark.asyncio
  1065. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  1066. RuntimeQueues = _runtime_queues_class()
  1067. queue_log: list[tuple[str, str, str]] = []
  1068. queues = RuntimeQueues(
  1069. input=RecordingQueue("input", queue_log),
  1070. output=RecordingQueue("output", queue_log),
  1071. events=RecordingQueue("events", queue_log),
  1072. )
  1073. request = DebugRunRequest(
  1074. user_message="debug this",
  1075. system_prompts=["You are a debugger."],
  1076. pre_messages=[],
  1077. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1078. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  1079. )
  1080. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  1081. outputs = [message async for message in runtime.run(request)]
  1082. assert _message_types(outputs) == [
  1083. "session_started",
  1084. "event",
  1085. "tool_result",
  1086. "round_stats",
  1087. "message_delta",
  1088. "round_stats",
  1089. "done",
  1090. ]
  1091. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  1092. ("input", "get", "user")
  1093. )
  1094. assert queue_log.index(
  1095. ("events", "put", "event_request:handoff_note:call_1")
  1096. ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
  1097. assert queue_log.index(
  1098. ("events", "get", "event_request:handoff_note:call_1")
  1099. ) < queue_log.index(("input", "put", "tool:call_1"))
  1100. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  1101. ("input", "get", "tool:call_1")
  1102. )
  1103. @pytest.mark.asyncio
  1104. async def test_runtime_run_consumes_output_queue_in_stream_order():
  1105. RuntimeQueues = _runtime_queues_class()
  1106. queue_log: list[tuple[str, str, str]] = []
  1107. queues = RuntimeQueues(
  1108. input=RecordingQueue("input", queue_log),
  1109. output=RecordingQueue("output", queue_log),
  1110. events=RecordingQueue("events", queue_log),
  1111. )
  1112. request = DebugRunRequest(
  1113. user_message="debug this",
  1114. system_prompts=["You are a debugger."],
  1115. pre_messages=[],
  1116. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1117. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  1118. )
  1119. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  1120. outputs = [message async for message in runtime.run(request)]
  1121. assert _message_types(outputs) == [
  1122. "session_started",
  1123. "event",
  1124. "tool_result",
  1125. "round_stats",
  1126. "message_delta",
  1127. "round_stats",
  1128. "done",
  1129. ]
  1130. output_puts = [
  1131. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  1132. ]
  1133. output_gets = [
  1134. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  1135. ]
  1136. assert [message for message in output_puts if message != "output:audit"] == [
  1137. "output:session_started",
  1138. "output:event",
  1139. "output:tool_result",
  1140. "output:round_stats",
  1141. "output:message_delta",
  1142. "output:round_stats",
  1143. "output:done",
  1144. ]
  1145. assert output_gets == output_puts
  1146. @pytest.mark.asyncio
  1147. async def test_runtime_continues_with_event_summary_without_tool_call_history():
  1148. request = DebugRunRequest(
  1149. user_message="debug this",
  1150. system_prompts=["You are a debugger."],
  1151. pre_messages=[],
  1152. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1153. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1154. )
  1155. client = StrictHistoryChatClient()
  1156. runtime = DebugRuntime(client)
  1157. outputs = [message async for message in runtime.run(request)]
  1158. assert client.calls == 2
  1159. assert [message.role for message in client.second_call_messages] == [
  1160. "system",
  1161. "system",
  1162. "user",
  1163. "assistant",
  1164. "user",
  1165. ]
  1166. assistant_message = client.second_call_messages[3]
  1167. assert assistant_message.content == ""
  1168. assert "handoff_note" in client.second_call_messages[1].content
  1169. assert not any(message.role == "tool" for message in client.second_call_messages)
  1170. assert client.second_call_messages[4].content.startswith("EventAgent results:\n")
  1171. assert client.second_call_messages[4].name == "event_agent"
  1172. assert outputs[-1] == {"type": "done"}
  1173. @pytest.mark.asyncio
  1174. async def test_runtime_passes_event_catalog_system_message_without_chat_tools():
  1175. registry = ToolRegistry(
  1176. [
  1177. ToolDefinition(
  1178. name="handoff_note",
  1179. description="Registry-owned handoff tool.",
  1180. parameters={
  1181. "type": "object",
  1182. "properties": {
  1183. "message": {"type": "string"},
  1184. "priority": {"type": "number"},
  1185. },
  1186. "required": ["message"],
  1187. },
  1188. handler=lambda event: {"tool": event.name, "message": "handled"},
  1189. )
  1190. ]
  1191. )
  1192. request = DebugRunRequest(
  1193. user_message="debug this",
  1194. system_prompts=[],
  1195. pre_messages=[],
  1196. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1197. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1198. )
  1199. client = ToolCapturingChatClient()
  1200. runtime = DebugRuntime(client, registry=registry)
  1201. outputs = [message async for message in runtime.run(request)]
  1202. assert outputs[-1] == {"type": "done"}
  1203. assert client.tools == []
  1204. assert client.messages[0].role == "system"
  1205. assert "Available events:" in client.messages[0].content
  1206. assert "- handoff_note: Registry-owned handoff tool." in client.messages[0].content
  1207. assert "message" not in client.messages[0].content
  1208. assert "priority" not in client.messages[0].content
  1209. @pytest.mark.asyncio
  1210. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  1211. request = DebugRunRequest(
  1212. user_message="debug this",
  1213. system_prompts=[],
  1214. pre_messages=[],
  1215. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1216. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1217. )
  1218. ticks = iter(
  1219. [
  1220. 0.9,
  1221. 0.91,
  1222. 1.0,
  1223. 1.01,
  1224. 1.02,
  1225. 1.123,
  1226. 1.2,
  1227. 1.25,
  1228. 1.3,
  1229. 1.456,
  1230. 1.7,
  1231. 1.8,
  1232. ]
  1233. )
  1234. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  1235. outputs = [message async for message in runtime.run(request)]
  1236. business_outputs = _without_audit(outputs)
  1237. assert _message_types(outputs) == [
  1238. "session_started",
  1239. "message_delta",
  1240. "usage",
  1241. "round_stats",
  1242. "done",
  1243. ]
  1244. assert business_outputs[3] == {
  1245. "type": "round_stats",
  1246. "round_index": 1,
  1247. "ttft_ms": 123,
  1248. "elapsed_ms": 456,
  1249. "prompt_tokens": 10,
  1250. "completion_tokens": 20,
  1251. "total_tokens": 30,
  1252. "cached_tokens": 5,
  1253. "had_event": False,
  1254. }
  1255. @pytest.mark.asyncio
  1256. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  1257. request = DebugRunRequest(
  1258. user_message="debug this",
  1259. system_prompts=[],
  1260. pre_messages=[],
  1261. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1262. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1263. )
  1264. ticks = iter(
  1265. [
  1266. 1.9,
  1267. 1.91,
  1268. 2.0,
  1269. 2.01,
  1270. 2.02,
  1271. 2.03,
  1272. 2.04,
  1273. 2.05,
  1274. 2.06,
  1275. 2.07,
  1276. 2.08,
  1277. 2.09,
  1278. 2.1,
  1279. 2.25,
  1280. 2.26,
  1281. 3.0,
  1282. 3.01,
  1283. 3.02,
  1284. 3.05,
  1285. 3.1,
  1286. 3.15,
  1287. 3.18,
  1288. 3.2,
  1289. 3.23,
  1290. 3.24,
  1291. 3.25,
  1292. 3.26,
  1293. ]
  1294. )
  1295. client = EventRoundStatsChatClient()
  1296. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  1297. outputs = [message async for message in runtime.run(request)]
  1298. stats = [message for message in outputs if message["type"] == "round_stats"]
  1299. assert client.calls == 2
  1300. assert stats == [
  1301. {
  1302. "type": "round_stats",
  1303. "round_index": 1,
  1304. "ttft_ms": None,
  1305. "elapsed_ms": 250,
  1306. "prompt_tokens": 3,
  1307. "completion_tokens": 0,
  1308. "total_tokens": 3,
  1309. "cached_tokens": 0,
  1310. "had_event": True,
  1311. },
  1312. {
  1313. "type": "round_stats",
  1314. "round_index": 2,
  1315. "ttft_ms": 50,
  1316. "elapsed_ms": 200,
  1317. "prompt_tokens": 4,
  1318. "completion_tokens": 6,
  1319. "total_tokens": 10,
  1320. "cached_tokens": 0,
  1321. "had_event": False,
  1322. },
  1323. ]
  1324. @pytest.mark.asyncio
  1325. async def test_runtime_session_resets_event_budget_for_each_user_turn():
  1326. request = DebugRunRequest(
  1327. user_message="first turn",
  1328. system_prompts=[],
  1329. pre_messages=[],
  1330. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1331. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  1332. )
  1333. client = TwoTurnSessionChatClient()
  1334. runtime = DebugRuntime(client)
  1335. queues = runtime.start_session(request)
  1336. outputs: list[dict[str, Any]] = []
  1337. while len([message for message in outputs if message["type"] == "turn_completed"]) < 1:
  1338. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1339. await queues.input.put(ChatMessage(role="user", content="second turn"))
  1340. while len([message for message in outputs if message["type"] == "turn_completed"]) < 2:
  1341. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1342. await runtime.aclose()
  1343. business_types = [message["type"] for message in _without_audit(outputs)]
  1344. assert business_types.count("turn_started") == 2
  1345. assert business_types.count("turn_completed") == 2
  1346. assert client.calls == 4
  1347. assert "Available events:" in client.messages_by_call[0][0].content
  1348. assert "Available events:" in client.messages_by_call[2][0].content
  1349. @pytest.mark.asyncio
  1350. async def test_runtime_persists_session_turn_messages_audit_and_usage(tmp_path):
  1351. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1352. request = DebugRunRequest(
  1353. session_id="session-1",
  1354. user_message="persist this",
  1355. system_prompts=["You are a debugger."],
  1356. pre_messages=[],
  1357. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  1358. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1359. )
  1360. runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
  1361. queues = runtime.start_session(request)
  1362. outputs: list[dict[str, Any]] = []
  1363. while not any(message["type"] == "turn_completed" for message in outputs):
  1364. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1365. await runtime.aclose()
  1366. assert _without_audit(outputs)[0] == {
  1367. "type": "session_started",
  1368. "session_id": "session-1",
  1369. }
  1370. assert store.get_session("session-1")["config"]["chat_agent"]["model"] == "chat-model"
  1371. assert [
  1372. (message["turn_index"], message["role"], message["content"])
  1373. for message in store.list_messages("session-1")
  1374. ] == [
  1375. (1, "user", "persist this"),
  1376. (1, "assistant", "hello"),
  1377. ]
  1378. audit_events = [
  1379. audit["event"]
  1380. for audit in store.list_audit_logs("session-1")
  1381. ]
  1382. assert "session_started" in audit_events
  1383. assert "chat_agent_request" in audit_events
  1384. assert "turn_completed" in audit_events
  1385. usage = store.usage_summary("session-1")
  1386. assert usage["calls"][0]["total_tokens"] == 30
  1387. assert usage["turns"][0]["turn_index"] == 1
  1388. assert usage["session"]["total_tokens"] == 30
  1389. @pytest.mark.asyncio
  1390. async def test_runtime_continues_persisted_turn_indexes_for_existing_session(tmp_path):
  1391. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1392. session_id = store.create_session(title="existing", config={})
  1393. store.start_turn(session_id, turn_index=1, user_message="old")
  1394. request = DebugRunRequest(
  1395. session_id=session_id,
  1396. user_message="new",
  1397. system_prompts=[],
  1398. pre_messages=[],
  1399. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  1400. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1401. )
  1402. runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
  1403. queues = runtime.start_session(request)
  1404. outputs: list[dict[str, Any]] = []
  1405. while not any(message["type"] == "turn_completed" for message in outputs):
  1406. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1407. await runtime.aclose()
  1408. assert [
  1409. (message["turn_index"], message["content"])
  1410. for message in store.list_messages(session_id)
  1411. ] == [
  1412. (2, "new"),
  1413. (2, "hello"),
  1414. ]
  1415. class ScriptedChatClient:
  1416. def __init__(self, rounds: list[list[StreamItem]]) -> None:
  1417. self.rounds = rounds
  1418. self.calls = 0
  1419. self.messages_by_call: list[list[ChatMessage]] = []
  1420. self.tools_by_call: list[list[dict[str, Any]]] = []
  1421. async def stream_chat(
  1422. self,
  1423. messages: list[ChatMessage],
  1424. tools: list[dict],
  1425. params: AgentParams,
  1426. tool_choice: dict[str, Any] | None = None,
  1427. ) -> AsyncIterator[StreamItem]:
  1428. self.messages_by_call.append(list(messages))
  1429. self.tools_by_call.append(list(tools))
  1430. round_items = self.rounds[self.calls]
  1431. self.calls += 1
  1432. for item in round_items:
  1433. yield item
  1434. class TranscriptValidatingScriptedChatClient(ScriptedChatClient):
  1435. async def stream_chat(
  1436. self,
  1437. messages: list[ChatMessage],
  1438. tools: list[dict],
  1439. params: AgentParams,
  1440. tool_choice: dict[str, Any] | None = None,
  1441. ) -> AsyncIterator[StreamItem]:
  1442. OpenAICompatibleChatClient._validate_tool_transcript(self, messages)
  1443. async for item in super().stream_chat(
  1444. messages,
  1445. tools,
  1446. params,
  1447. tool_choice,
  1448. ):
  1449. yield item
  1450. def _direct_request(
  1451. *,
  1452. enabled_tools: list[str],
  1453. max_event_loops: int = 1,
  1454. session_id: str | None = None,
  1455. ) -> DebugRunRequest:
  1456. return DebugRunRequest(
  1457. session_id=session_id,
  1458. user_message="debug direct tools",
  1459. system_prompts=["You are a debugger."],
  1460. pre_messages=[],
  1461. chat_agent=AgentParams(model="chat-model"),
  1462. event_agent=EventAgentParams(
  1463. enabled_tools=enabled_tools,
  1464. max_event_loops=max_event_loops,
  1465. ),
  1466. tool_invocation_mode="chat_agent_tools",
  1467. )
  1468. @pytest.mark.asyncio
  1469. async def test_direct_mode_uses_enabled_schemas_and_valid_ordered_provider_transcript():
  1470. executed: list[tuple[str, dict[str, Any], str]] = []
  1471. def handle(event: ToolCallEvent) -> dict[str, Any]:
  1472. executed.append((event.name, event.arguments, event.raw_arguments))
  1473. return {"tool": event.name, "value": event.arguments["value"]}
  1474. registry = ToolRegistry(
  1475. [
  1476. ToolDefinition(
  1477. name="first_tool",
  1478. description="First direct tool.",
  1479. parameters={
  1480. "type": "object",
  1481. "properties": {"value": {"type": "string"}},
  1482. "required": ["value"],
  1483. },
  1484. handler=handle,
  1485. ),
  1486. ToolDefinition(
  1487. name="second_tool",
  1488. description="Second direct tool.",
  1489. parameters={
  1490. "type": "object",
  1491. "properties": {"value": {"type": "string"}},
  1492. "required": ["value"],
  1493. },
  1494. handler=handle,
  1495. ),
  1496. ]
  1497. )
  1498. calls = [
  1499. ToolCallEvent(
  1500. id="provider-1",
  1501. name="first_tool",
  1502. arguments={"value": "one"},
  1503. raw_arguments='{"value":"one"}',
  1504. ),
  1505. ToolCallEvent(
  1506. id="provider-2",
  1507. name="second_tool",
  1508. arguments={"value": "two"},
  1509. raw_arguments='{"value":"two"}',
  1510. ),
  1511. ]
  1512. ignored_text_event = ToolCallEvent(
  1513. id="text-ignored",
  1514. name="first_tool",
  1515. arguments={},
  1516. raw_arguments="{}",
  1517. )
  1518. client = ScriptedChatClient(
  1519. [
  1520. [
  1521. StreamItem.message_delta("Visible before tools."),
  1522. StreamItem.provider_tool_call(calls[0]),
  1523. StreamItem.text_event(ignored_text_event),
  1524. StreamItem.provider_tool_call(calls[1]),
  1525. ],
  1526. [StreamItem.message_delta("Final answer.")],
  1527. ]
  1528. )
  1529. outputs = await _collect_outputs(
  1530. DebugRuntime(client, registry=registry).run(
  1531. _direct_request(enabled_tools=["first_tool", "second_tool"])
  1532. )
  1533. )
  1534. assert [tool["function"]["name"] for tool in client.tools_by_call[0]] == [
  1535. "first_tool",
  1536. "second_tool",
  1537. ]
  1538. assert client.tools_by_call[1] == []
  1539. assert not any(
  1540. "Available events:" in message.content
  1541. for message in client.messages_by_call[0]
  1542. if message.role == "system"
  1543. )
  1544. transcript = client.messages_by_call[1]
  1545. assert [message.role for message in transcript[-4:]] == [
  1546. "user",
  1547. "assistant",
  1548. "tool",
  1549. "tool",
  1550. ]
  1551. assert transcript[-3].content == "Visible before tools."
  1552. assert transcript[-3].tool_calls == calls
  1553. assert [message.tool_call_id for message in transcript[-2:]] == [
  1554. "provider-1",
  1555. "provider-2",
  1556. ]
  1557. assert [message.name for message in transcript[-2:]] == [
  1558. "first_tool",
  1559. "second_tool",
  1560. ]
  1561. assert executed == [
  1562. ("first_tool", {"value": "one"}, '{"value":"one"}'),
  1563. ("second_tool", {"value": "two"}, '{"value":"two"}'),
  1564. ]
  1565. emitted_events = [
  1566. message["event"] for message in outputs if message["type"] == "event"
  1567. ]
  1568. assert emitted_events == [call.model_dump() for call in calls]
  1569. assert "event_agent_request" not in [
  1570. message.get("event") for message in outputs if message["type"] == "audit"
  1571. ]
  1572. assert not any(
  1573. message.get("event", {}).get("id") == "text-ignored"
  1574. for message in outputs
  1575. if message["type"] == "event"
  1576. )
  1577. assert all(
  1578. message["details"]["tool_invocation_mode"] == "chat_agent_tools"
  1579. for message in outputs
  1580. if message["type"] == "audit"
  1581. and message["event"] in {"chat_round_started", "chat_agent_request"}
  1582. )
  1583. @pytest.mark.asyncio
  1584. async def test_direct_mode_does_not_execute_unknown_disabled_or_failed_handlers():
  1585. executed: list[str] = []
  1586. def disabled_handler(event: ToolCallEvent) -> dict[str, Any]:
  1587. executed.append(event.name)
  1588. return {"tool": event.name}
  1589. def failing_handler(event: ToolCallEvent) -> dict[str, Any]:
  1590. executed.append(event.name)
  1591. raise RuntimeError("direct boom")
  1592. registry = ToolRegistry(
  1593. [
  1594. ToolDefinition(
  1595. name="enabled_tool",
  1596. description="Enabled.",
  1597. parameters={"type": "object"},
  1598. handler=lambda event: {"tool": event.name},
  1599. ),
  1600. ToolDefinition(
  1601. name="disabled_tool",
  1602. description="Disabled.",
  1603. parameters={"type": "object"},
  1604. handler=disabled_handler,
  1605. ),
  1606. ToolDefinition(
  1607. name="failing_tool",
  1608. description="Fails.",
  1609. parameters={"type": "object"},
  1610. handler=failing_handler,
  1611. ),
  1612. ]
  1613. )
  1614. client = ScriptedChatClient(
  1615. [
  1616. [
  1617. StreamItem.provider_tool_call(
  1618. ToolCallEvent(
  1619. id="unknown-1",
  1620. name="unknown_tool",
  1621. arguments={"kept": True},
  1622. raw_arguments='{"kept":true}',
  1623. )
  1624. ),
  1625. StreamItem.provider_tool_call(
  1626. ToolCallEvent(
  1627. id="disabled-1",
  1628. name="disabled_tool",
  1629. arguments={},
  1630. raw_arguments="{}",
  1631. )
  1632. ),
  1633. StreamItem.provider_tool_call(
  1634. ToolCallEvent(
  1635. id="failed-1",
  1636. name="failing_tool",
  1637. arguments={},
  1638. raw_arguments="{}",
  1639. )
  1640. ),
  1641. ],
  1642. [StreamItem.message_delta("continued")],
  1643. ]
  1644. )
  1645. outputs = await _collect_outputs(
  1646. DebugRuntime(client, registry=registry).run(
  1647. _direct_request(enabled_tools=["enabled_tool", "failing_tool"])
  1648. )
  1649. )
  1650. assert [tool["function"]["name"] for tool in client.tools_by_call[0]] == [
  1651. "enabled_tool",
  1652. "failing_tool",
  1653. ]
  1654. assert executed == ["failing_tool"]
  1655. results = [
  1656. json.loads(message["message"]["content"])
  1657. for message in outputs
  1658. if message["type"] == "tool_result"
  1659. ]
  1660. assert results == [
  1661. {"tool": "unknown_tool", "error": "unknown tool"},
  1662. {"tool": "disabled_tool", "error": "tool disabled"},
  1663. {"tool": "failing_tool", "error": "tool handler failed: direct boom"},
  1664. ]
  1665. @pytest.mark.asyncio
  1666. async def test_direct_mode_ignores_provider_calls_after_event_budget_exhaustion():
  1667. executed: list[str] = []
  1668. registry = ToolRegistry(
  1669. [
  1670. ToolDefinition(
  1671. name="once_tool",
  1672. description="Run once.",
  1673. parameters={"type": "object"},
  1674. handler=lambda event: executed.append(event.id) or {"tool": event.name},
  1675. )
  1676. ]
  1677. )
  1678. client = ScriptedChatClient(
  1679. [
  1680. [
  1681. StreamItem.provider_tool_call(
  1682. ToolCallEvent(
  1683. id="accepted",
  1684. name="once_tool",
  1685. arguments={},
  1686. raw_arguments="{}",
  1687. )
  1688. )
  1689. ],
  1690. [
  1691. StreamItem.provider_tool_call(
  1692. ToolCallEvent(
  1693. id="ignored",
  1694. name="once_tool",
  1695. arguments={},
  1696. raw_arguments="{}",
  1697. )
  1698. ),
  1699. StreamItem.message_delta("budget exhausted"),
  1700. ],
  1701. ]
  1702. )
  1703. outputs = await _collect_outputs(
  1704. DebugRuntime(client, registry=registry).run(
  1705. _direct_request(enabled_tools=["once_tool"], max_event_loops=1)
  1706. )
  1707. )
  1708. assert client.tools_by_call == [
  1709. [registry.tool_schema("once_tool")],
  1710. [],
  1711. ]
  1712. assert executed == ["accepted"]
  1713. assert [
  1714. message["event"]["id"]
  1715. for message in outputs
  1716. if message["type"] == "event"
  1717. ] == ["accepted"]
  1718. assert outputs[-1] == {"type": "done"}
  1719. @pytest.mark.asyncio
  1720. async def test_direct_mode_session_turns_reset_budget_and_snapshot_mode(tmp_path):
  1721. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1722. registry = ToolRegistry(
  1723. [
  1724. ToolDefinition(
  1725. name="turn_tool",
  1726. description="Per-turn tool.",
  1727. parameters={"type": "object"},
  1728. handler=lambda event: {"tool": event.name, "call_id": event.id},
  1729. )
  1730. ]
  1731. )
  1732. client = ScriptedChatClient(
  1733. [
  1734. [
  1735. StreamItem.provider_tool_call(
  1736. ToolCallEvent(
  1737. id="turn-1-call",
  1738. name="turn_tool",
  1739. arguments={},
  1740. raw_arguments="{}",
  1741. )
  1742. )
  1743. ],
  1744. [StreamItem.message_delta("turn one done")],
  1745. [
  1746. StreamItem.provider_tool_call(
  1747. ToolCallEvent(
  1748. id="turn-2-call",
  1749. name="turn_tool",
  1750. arguments={},
  1751. raw_arguments="{}",
  1752. )
  1753. )
  1754. ],
  1755. [StreamItem.message_delta("turn two done")],
  1756. ]
  1757. )
  1758. runtime = DebugRuntime(client, registry=registry, session_store=store)
  1759. request = _direct_request(
  1760. enabled_tools=["turn_tool"],
  1761. max_event_loops=1,
  1762. session_id="direct-session",
  1763. )
  1764. queues = runtime.start_session(request)
  1765. outputs: list[dict[str, Any]] = []
  1766. while sum(message["type"] == "turn_completed" for message in outputs) < 1:
  1767. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1768. await queues.input.put(ChatMessage(role="user", content="second direct turn"))
  1769. while sum(message["type"] == "turn_completed" for message in outputs) < 2:
  1770. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1771. await runtime.aclose()
  1772. assert [bool(tools) for tools in client.tools_by_call] == [True, False, True, False]
  1773. assert [message.role for message in client.messages_by_call[1][-3:]] == [
  1774. "user",
  1775. "assistant",
  1776. "tool",
  1777. ]
  1778. second_turn_transcript = client.messages_by_call[3]
  1779. assert not any(message.name == "event_agent" for message in second_turn_transcript)
  1780. assert [
  1781. message.tool_call_id
  1782. for message in second_turn_transcript
  1783. if message.role == "tool"
  1784. ] == ["turn-1-call", "turn-2-call"]
  1785. session = store.get_session("direct-session")
  1786. assert session is not None
  1787. assert session["config"]["tool_invocation_mode"] == "chat_agent_tools"
  1788. persisted_messages = store.list_messages("direct-session")
  1789. assert [message["role"] for message in persisted_messages] == [
  1790. "user",
  1791. "assistant",
  1792. "tool",
  1793. "assistant",
  1794. "user",
  1795. "assistant",
  1796. "tool",
  1797. "assistant",
  1798. ]
  1799. assert [
  1800. call["id"]
  1801. for message in persisted_messages
  1802. for call in message["tool_calls"]
  1803. ] == ["turn-1-call", "turn-2-call"]
  1804. assert [
  1805. message["tool_call_id"]
  1806. for message in persisted_messages
  1807. if message["role"] == "tool"
  1808. ] == ["turn-1-call", "turn-2-call"]
  1809. request_audits = [
  1810. audit
  1811. for audit in store.list_audit_logs("direct-session")
  1812. if audit["event"] == "chat_agent_request"
  1813. ]
  1814. assert request_audits
  1815. assert all(
  1816. audit["details"]["tool_invocation_mode"] == "chat_agent_tools"
  1817. for audit in request_audits
  1818. )
  1819. @pytest.mark.asyncio
  1820. async def test_direct_mode_reconnect_restores_provider_valid_session_transcript(tmp_path):
  1821. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1822. session_id = store.create_session(
  1823. title="restored direct",
  1824. config={"tool_invocation_mode": "chat_agent_tools"},
  1825. session_id="restored-direct",
  1826. )
  1827. store.start_turn(session_id, turn_index=1, user_message="past user")
  1828. previous_call = ToolCallEvent(
  1829. id="past-call",
  1830. name="safe_tool",
  1831. arguments={"value": 1},
  1832. raw_arguments='{"value":1}',
  1833. )
  1834. for message in [
  1835. ChatMessage(role="user", content="past user"),
  1836. ChatMessage(role="assistant", content="", tool_calls=[previous_call]),
  1837. ChatMessage(
  1838. role="tool",
  1839. content='{"ok":true}',
  1840. name="safe_tool",
  1841. tool_call_id="past-call",
  1842. ),
  1843. ChatMessage(role="assistant", content="past answer"),
  1844. ]:
  1845. store.append_message(session_id, turn_index=1, message=message)
  1846. store.complete_turn(session_id, turn_index=1)
  1847. client = TranscriptValidatingScriptedChatClient(
  1848. [[StreamItem.message_delta("current answer")]]
  1849. )
  1850. runtime = DebugRuntime(
  1851. client,
  1852. registry=_single_tool_registry(lambda event: {"ok": True}),
  1853. session_store=store,
  1854. )
  1855. request = _direct_request(
  1856. enabled_tools=["safe_tool"],
  1857. session_id=session_id,
  1858. ).model_copy(
  1859. update={
  1860. "user_message": "current user",
  1861. "pre_messages": [ChatMessage(role="user", content="configured pre")],
  1862. }
  1863. )
  1864. queues = runtime.start_session(request)
  1865. outputs: list[dict[str, Any]] = []
  1866. while not any(message["type"] == "turn_completed" for message in outputs):
  1867. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1868. await runtime.aclose()
  1869. transcript = client.messages_by_call[0]
  1870. assert [(message.role, message.content) for message in transcript] == [
  1871. ("system", "You are a debugger."),
  1872. ("user", "configured pre"),
  1873. ("user", "past user"),
  1874. ("assistant", ""),
  1875. ("tool", '{"ok":true}'),
  1876. ("assistant", "past answer"),
  1877. ("user", "current user"),
  1878. ]
  1879. assert transcript[3].tool_calls == [previous_call]
  1880. assert sum(
  1881. message.role == "user" and message.content == "current user"
  1882. for message in transcript
  1883. ) == 1
  1884. @pytest.mark.asyncio
  1885. async def test_cancelled_direct_batch_persists_no_partial_protocol_and_reconnects(tmp_path):
  1886. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1887. handler_started = asyncio.Event()
  1888. release_handler = asyncio.Event()
  1889. async def blocking_handler(event: ToolCallEvent) -> dict[str, Any]:
  1890. handler_started.set()
  1891. await release_handler.wait()
  1892. return {"ok": True, "call_id": event.id}
  1893. registry = _single_tool_registry(blocking_handler)
  1894. first_client = ScriptedChatClient(
  1895. [[StreamItem.provider_tool_call(_tool_call("cancelled-call"))]]
  1896. )
  1897. first_runtime = DebugRuntime(
  1898. first_client,
  1899. registry=registry,
  1900. session_store=store,
  1901. )
  1902. request = _direct_request(
  1903. enabled_tools=["safe_tool"],
  1904. session_id="cancelled-direct",
  1905. )
  1906. first_runtime.start_session(request)
  1907. await asyncio.wait_for(handler_started.wait(), timeout=1)
  1908. await first_runtime.aclose()
  1909. persisted_after_cancel = store.list_messages("cancelled-direct")
  1910. assert [message["role"] for message in persisted_after_cancel] == ["user"]
  1911. assert not any(message["tool_calls"] for message in persisted_after_cancel)
  1912. reconnect_client = TranscriptValidatingScriptedChatClient(
  1913. [[StreamItem.message_delta("reconnected")]]
  1914. )
  1915. reconnect_runtime = DebugRuntime(
  1916. reconnect_client,
  1917. registry=registry,
  1918. session_store=store,
  1919. )
  1920. reconnect_request = request.model_copy(
  1921. update={"user_message": "after cancellation"}
  1922. )
  1923. queues = reconnect_runtime.start_session(reconnect_request)
  1924. outputs: list[dict[str, Any]] = []
  1925. while not any(
  1926. message["type"] in {"turn_completed", "error"} for message in outputs
  1927. ):
  1928. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1929. await reconnect_runtime.aclose()
  1930. assert not any(message["type"] == "error" for message in outputs)
  1931. assert [
  1932. (message.role, message.content)
  1933. for message in reconnect_client.messages_by_call[0]
  1934. if message.role != "system"
  1935. ] == [
  1936. ("user", "debug direct tools"),
  1937. ("user", "after cancellation"),
  1938. ]
  1939. @pytest.mark.asyncio
  1940. async def test_direct_restore_drops_invalid_protocol_fragments_and_keeps_visible_history(
  1941. tmp_path,
  1942. ):
  1943. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1944. session_id = store.create_session(
  1945. title="repair transcript",
  1946. config={"tool_invocation_mode": "chat_agent_tools"},
  1947. session_id="repair-transcript",
  1948. )
  1949. store.start_turn(session_id, turn_index=1, user_message="visible user")
  1950. complete_calls = [_tool_call("complete-1"), _tool_call("complete-2")]
  1951. store.append_messages(
  1952. session_id,
  1953. turn_index=1,
  1954. messages=[
  1955. ChatMessage(role="user", content="visible user"),
  1956. ChatMessage(
  1957. role="assistant",
  1958. content="using tools",
  1959. tool_calls=complete_calls,
  1960. ),
  1961. ChatMessage(
  1962. role="tool",
  1963. content='{"id":1}',
  1964. name="safe_tool",
  1965. tool_call_id="complete-1",
  1966. ),
  1967. ChatMessage(
  1968. role="tool",
  1969. content='{"id":2}',
  1970. name="safe_tool",
  1971. tool_call_id="complete-2",
  1972. ),
  1973. ChatMessage(role="assistant", content="visible answer"),
  1974. ChatMessage(
  1975. role="tool",
  1976. content='{"orphan":true}',
  1977. name="safe_tool",
  1978. tool_call_id="orphan-call",
  1979. ),
  1980. ChatMessage(role="assistant", content="visible after orphan"),
  1981. ChatMessage(
  1982. role="assistant",
  1983. content="visible incomplete tool preface",
  1984. tool_calls=[
  1985. _tool_call("incomplete-visible-1"),
  1986. _tool_call("incomplete-visible-2"),
  1987. ],
  1988. ),
  1989. ChatMessage(
  1990. role="tool",
  1991. content='{"partial":true}',
  1992. name="safe_tool",
  1993. tool_call_id="incomplete-visible-1",
  1994. ),
  1995. ChatMessage(
  1996. role="assistant",
  1997. content=" \n\t",
  1998. tool_calls=[_tool_call("incomplete-whitespace")],
  1999. ),
  2000. ],
  2001. )
  2002. store.complete_turn(session_id, turn_index=1)
  2003. client = TranscriptValidatingScriptedChatClient(
  2004. [[StreamItem.message_delta("repaired")]]
  2005. )
  2006. runtime = DebugRuntime(
  2007. client,
  2008. registry=_single_tool_registry(lambda event: {"ok": True}),
  2009. session_store=store,
  2010. )
  2011. request = _direct_request(
  2012. enabled_tools=["safe_tool"],
  2013. session_id=session_id,
  2014. ).model_copy(update={"user_message": "current user"})
  2015. queues = runtime.start_session(request)
  2016. outputs: list[dict[str, Any]] = []
  2017. while not any(
  2018. message["type"] in {"turn_completed", "error"} for message in outputs
  2019. ):
  2020. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  2021. await runtime.aclose()
  2022. assert not any(message["type"] == "error" for message in outputs)
  2023. transcript = client.messages_by_call[0]
  2024. assert [(message.role, message.content) for message in transcript] == [
  2025. ("system", "You are a debugger."),
  2026. ("user", "visible user"),
  2027. ("assistant", "using tools"),
  2028. ("tool", '{"id":1}'),
  2029. ("tool", '{"id":2}'),
  2030. ("assistant", "visible answer"),
  2031. ("assistant", "visible after orphan"),
  2032. ("assistant", "visible incomplete tool preface"),
  2033. ("user", "current user"),
  2034. ]
  2035. assert transcript[2].tool_calls == complete_calls
  2036. assert transcript[-2].tool_calls == []
  2037. @pytest.mark.asyncio
  2038. async def test_dual_mode_reconnect_restores_visible_history_without_tool_protocol(tmp_path):
  2039. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  2040. session_id = store.create_session(
  2041. title="restored dual",
  2042. config={"tool_invocation_mode": "dual_agent"},
  2043. session_id="restored-dual",
  2044. )
  2045. store.start_turn(session_id, turn_index=1, user_message="past user")
  2046. store.append_message(
  2047. session_id,
  2048. turn_index=1,
  2049. message=ChatMessage(role="user", content="past user"),
  2050. )
  2051. store.append_message(
  2052. session_id,
  2053. turn_index=1,
  2054. message=ChatMessage(role="assistant", content="past answer"),
  2055. )
  2056. store.complete_turn(session_id, turn_index=1)
  2057. registry = _single_tool_registry(lambda event: {"ok": True})
  2058. client = ScriptedChatClient(
  2059. [
  2060. [StreamItem.text_event(_tool_call("dual-call"))],
  2061. [StreamItem.message_delta("current answer")],
  2062. ]
  2063. )
  2064. runtime = DebugRuntime(client, registry=registry, session_store=store)
  2065. request = DebugRunRequest(
  2066. session_id=session_id,
  2067. user_message="current user",
  2068. system_prompts=["system context"],
  2069. pre_messages=[ChatMessage(role="assistant", content="configured pre")],
  2070. chat_agent=AgentParams(model="chat-model"),
  2071. event_agent=EventAgentParams(
  2072. enabled_tools=["safe_tool"],
  2073. max_event_loops=1,
  2074. ),
  2075. tool_invocation_mode="dual_agent",
  2076. )
  2077. queues = runtime.start_session(request)
  2078. outputs: list[dict[str, Any]] = []
  2079. while not any(message["type"] == "turn_completed" for message in outputs):
  2080. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  2081. await runtime.aclose()
  2082. first_request = client.messages_by_call[0]
  2083. assert [(message.role, message.content) for message in first_request if message.role != "system"] == [
  2084. ("assistant", "configured pre"),
  2085. ("user", "past user"),
  2086. ("assistant", "past answer"),
  2087. ("user", "current user"),
  2088. ]
  2089. assert sum(
  2090. message.role == "user" and message.content == "current user"
  2091. for message in first_request
  2092. ) == 1
  2093. persisted_messages = store.list_messages(session_id)
  2094. assert not any(message["role"] == "tool" for message in persisted_messages)
  2095. assert not any(message["tool_calls"] for message in persisted_messages)
  2096. assert not any(message["content"] == "" for message in persisted_messages)
  2097. def _single_tool_registry(
  2098. handler: Any,
  2099. *,
  2100. name: str = "safe_tool",
  2101. ) -> ToolRegistry:
  2102. return ToolRegistry(
  2103. [
  2104. ToolDefinition(
  2105. name=name,
  2106. description="Tool used by execution-safety tests.",
  2107. parameters={"type": "object"},
  2108. handler=handler,
  2109. )
  2110. ]
  2111. )
  2112. def _tool_call(call_id: str, *, name: str = "safe_tool") -> ToolCallEvent:
  2113. return ToolCallEvent(
  2114. id=call_id,
  2115. name=name,
  2116. arguments={},
  2117. raw_arguments="{}",
  2118. )
  2119. @pytest.mark.asyncio
  2120. async def test_direct_mode_rejects_provider_call_id_reused_across_one_shot_rounds():
  2121. executed: list[str] = []
  2122. registry = _single_tool_registry(
  2123. lambda event: executed.append(event.id) or {"tool": event.name}
  2124. )
  2125. client = ScriptedChatClient(
  2126. [
  2127. [StreamItem.provider_tool_call(_tool_call("duplicate-call"))],
  2128. [StreamItem.provider_tool_call(_tool_call("duplicate-call"))],
  2129. ]
  2130. )
  2131. outputs = await _collect_outputs(
  2132. DebugRuntime(client, registry=registry).run(
  2133. _direct_request(enabled_tools=["safe_tool"], max_event_loops=2)
  2134. )
  2135. )
  2136. assert executed == ["duplicate-call"]
  2137. assert [
  2138. message["message"]["tool_call_id"]
  2139. for message in outputs
  2140. if message["type"] == "tool_result"
  2141. ] == ["duplicate-call"]
  2142. assert outputs[-1]["type"] == "error"
  2143. assert "duplicate provider tool-call ID: duplicate-call" in outputs[-1]["message"]
  2144. @pytest.mark.asyncio
  2145. async def test_direct_mode_rejects_provider_call_id_reused_in_session_rounds():
  2146. executed: list[str] = []
  2147. registry = _single_tool_registry(
  2148. lambda event: executed.append(event.id) or {"tool": event.name}
  2149. )
  2150. client = ScriptedChatClient(
  2151. [
  2152. [StreamItem.provider_tool_call(_tool_call("session-duplicate"))],
  2153. [StreamItem.provider_tool_call(_tool_call("session-duplicate"))],
  2154. ]
  2155. )
  2156. runtime = DebugRuntime(client, registry=registry)
  2157. queues = runtime.start_session(
  2158. _direct_request(enabled_tools=["safe_tool"], max_event_loops=2)
  2159. )
  2160. outputs: list[dict[str, Any]] = []
  2161. while not any(message["type"] == "error" for message in outputs):
  2162. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  2163. await runtime.aclose()
  2164. assert executed == ["session-duplicate"]
  2165. assert [
  2166. message["message"]["tool_call_id"]
  2167. for message in outputs
  2168. if message["type"] == "tool_result"
  2169. ] == ["session-duplicate"]
  2170. assert "duplicate provider tool-call ID: session-duplicate" in outputs[-1][
  2171. "message"
  2172. ]
  2173. @pytest.mark.asyncio
  2174. async def test_direct_mode_rejects_duplicate_provider_call_ids_within_batch():
  2175. executed: list[str] = []
  2176. registry = _single_tool_registry(
  2177. lambda event: executed.append(event.id) or {"tool": event.name}
  2178. )
  2179. client = ScriptedChatClient(
  2180. [
  2181. [
  2182. StreamItem.provider_tool_call(_tool_call("same-batch")),
  2183. StreamItem.provider_tool_call(_tool_call("same-batch")),
  2184. ]
  2185. ]
  2186. )
  2187. outputs = await _collect_outputs(
  2188. DebugRuntime(client, registry=registry).run(
  2189. _direct_request(enabled_tools=["safe_tool"])
  2190. )
  2191. )
  2192. assert executed == []
  2193. assert outputs[-1]["type"] == "error"
  2194. assert "duplicate provider tool-call ID: same-batch" in outputs[-1]["message"]
  2195. @pytest.mark.asyncio
  2196. async def test_direct_mode_executes_batch_concurrently_but_replies_in_provider_order():
  2197. both_started = asyncio.Event()
  2198. first_release = asyncio.Event()
  2199. second_release = asyncio.Event()
  2200. second_finished = asyncio.Event()
  2201. started: list[str] = []
  2202. async def handler(event: ToolCallEvent) -> dict[str, Any]:
  2203. started.append(event.id)
  2204. if len(started) == 2:
  2205. both_started.set()
  2206. if event.id == "first-call":
  2207. await first_release.wait()
  2208. else:
  2209. await second_release.wait()
  2210. second_finished.set()
  2211. return {"tool": event.name, "call_id": event.id}
  2212. registry = _single_tool_registry(handler)
  2213. client = ScriptedChatClient(
  2214. [
  2215. [
  2216. StreamItem.provider_tool_call(
  2217. _tool_call("first-call").model_copy(
  2218. update={"arguments": {"order": 1}, "raw_arguments": '{"order":1}'}
  2219. )
  2220. ),
  2221. StreamItem.provider_tool_call(
  2222. _tool_call("second-call").model_copy(
  2223. update={"arguments": {"order": 2}, "raw_arguments": '{"order":2}'}
  2224. )
  2225. ),
  2226. ],
  2227. [StreamItem.message_delta("finished")],
  2228. ]
  2229. )
  2230. runtime = DebugRuntime(client, registry=registry)
  2231. run_task = asyncio.create_task(
  2232. _collect_outputs(
  2233. runtime.run(_direct_request(enabled_tools=["safe_tool"]))
  2234. )
  2235. )
  2236. overlapped = False
  2237. try:
  2238. try:
  2239. await asyncio.wait_for(both_started.wait(), timeout=1)
  2240. except TimeoutError:
  2241. pass
  2242. else:
  2243. overlapped = True
  2244. second_release.set()
  2245. await asyncio.wait_for(second_finished.wait(), timeout=1)
  2246. assert not run_task.done()
  2247. finally:
  2248. first_release.set()
  2249. second_release.set()
  2250. outputs = await asyncio.wait_for(run_task, timeout=1)
  2251. assert overlapped is True
  2252. assert started == ["first-call", "second-call"]
  2253. assert [
  2254. message["message"]["tool_call_id"]
  2255. for message in outputs
  2256. if message["type"] == "tool_result"
  2257. ] == ["first-call", "second-call"]
  2258. def _dual_budget_request() -> DebugRunRequest:
  2259. return DebugRunRequest(
  2260. user_message="dual budget",
  2261. system_prompts=[],
  2262. pre_messages=[],
  2263. chat_agent=AgentParams(model="chat-model"),
  2264. event_agent=EventAgentParams(
  2265. enabled_tools=["safe_tool"],
  2266. max_event_loops=1,
  2267. ),
  2268. )
  2269. @pytest.mark.asyncio
  2270. async def test_dual_mode_ignores_text_event_after_one_shot_budget_exhaustion():
  2271. executed: list[str] = []
  2272. registry = _single_tool_registry(
  2273. lambda event: executed.append(event.id) or {"tool": event.name}
  2274. )
  2275. client = ScriptedChatClient(
  2276. [
  2277. [StreamItem.text_event(_tool_call("accepted-text-event"))],
  2278. [
  2279. StreamItem.text_event(_tool_call("ignored-text-event")),
  2280. StreamItem.message_delta("final dual answer"),
  2281. ],
  2282. ]
  2283. )
  2284. outputs = await _collect_outputs(
  2285. DebugRuntime(client, registry=registry).run(_dual_budget_request())
  2286. )
  2287. assert executed == ["accepted-text-event"]
  2288. assert [
  2289. message["event"]["id"]
  2290. for message in outputs
  2291. if message["type"] == "event"
  2292. ] == ["accepted-text-event"]
  2293. assert outputs[-1] == {"type": "done"}
  2294. @pytest.mark.asyncio
  2295. async def test_dual_mode_session_main_path_ignores_text_event_after_budget_exhaustion():
  2296. executed: list[str] = []
  2297. registry = _single_tool_registry(
  2298. lambda event: executed.append(event.id) or {"tool": event.name}
  2299. )
  2300. client = ScriptedChatClient(
  2301. [
  2302. [StreamItem.text_event(_tool_call("session-accepted"))],
  2303. [
  2304. StreamItem.text_event(_tool_call("session-ignored")),
  2305. StreamItem.message_delta("session final"),
  2306. ],
  2307. ]
  2308. )
  2309. runtime = DebugRuntime(client, registry=registry)
  2310. queues = runtime.start_session(_dual_budget_request())
  2311. outputs: list[dict[str, Any]] = []
  2312. while not any(
  2313. message["type"] in {"turn_completed", "error"} for message in outputs
  2314. ):
  2315. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  2316. await runtime.aclose()
  2317. assert executed == ["session-accepted"]
  2318. assert [
  2319. message["event"]["id"]
  2320. for message in outputs
  2321. if message["type"] == "event"
  2322. ] == ["session-accepted"]
  2323. assert outputs[-1]["type"] == "turn_completed"
  2324. def _policy_request(
  2325. *,
  2326. enabled_tools: list[str],
  2327. mode: str = "dual_agent",
  2328. session_id: str | None = None,
  2329. ) -> DebugRunRequest:
  2330. return DebugRunRequest(
  2331. session_id=session_id,
  2332. user_message="exercise event result policy",
  2333. system_prompts=[],
  2334. pre_messages=[],
  2335. chat_agent=AgentParams(model="chat-model"),
  2336. event_agent=EventAgentParams(
  2337. enabled_tools=enabled_tools,
  2338. max_event_loops=1,
  2339. ),
  2340. tool_invocation_mode=mode,
  2341. )
  2342. @pytest.mark.asyncio
  2343. async def test_dual_silent_success_does_not_enqueue_legacy_summary_or_call_chat_again():
  2344. registry = ToolRegistry(
  2345. [
  2346. ToolDefinition(
  2347. name="example.silent",
  2348. description="Complete silently.",
  2349. parameters={"type": "object"},
  2350. handler=lambda event: {"tool": event.name, "status": "applied"},
  2351. result_policy=ResultPolicy.SILENT_SUCCESS,
  2352. )
  2353. ]
  2354. )
  2355. client = ScriptedChatClient(
  2356. [
  2357. [
  2358. StreamItem.message_delta("Applying it now."),
  2359. StreamItem.text_event(_tool_call("silent-1", name="example.silent")),
  2360. ]
  2361. ]
  2362. )
  2363. outputs = await _collect_outputs(
  2364. DebugRuntime(client, registry=registry).run(
  2365. _policy_request(enabled_tools=["example.silent"])
  2366. )
  2367. )
  2368. assert client.calls == 1
  2369. assert outputs[-1] == {"type": "done"}
  2370. assert not any(
  2371. message.get("details", {}).get("result_summary")
  2372. for message in outputs
  2373. if message.get("event") == "event_agent_completed"
  2374. )
  2375. @pytest.mark.asyncio
  2376. async def test_direct_template_success_emits_plugin_confirmation_without_model_call():
  2377. schedule = ToolCallEvent(
  2378. id="schedule-1",
  2379. name="calendar.schedule.create",
  2380. arguments={
  2381. "title": "Design review",
  2382. "start_at": "2026-07-14T09:30:00+08:00",
  2383. "timezone": "Asia/Shanghai",
  2384. },
  2385. raw_arguments=(
  2386. '{"title":"Design review","start_at":"2026-07-14T09:30:00+08:00",'
  2387. '"timezone":"Asia/Shanghai"}'
  2388. ),
  2389. )
  2390. client = ScriptedChatClient(
  2391. [
  2392. [
  2393. StreamItem.message_delta("I will add it."),
  2394. StreamItem.provider_tool_call(schedule),
  2395. ]
  2396. ]
  2397. )
  2398. outputs = await _collect_outputs(
  2399. DebugRuntime(client).run(
  2400. _policy_request(
  2401. enabled_tools=["calendar.schedule.create"],
  2402. mode="chat_agent_tools",
  2403. )
  2404. )
  2405. )
  2406. assert client.calls == 1
  2407. assert [
  2408. message["content"]
  2409. for message in outputs
  2410. if message["type"] == "message_delta"
  2411. ] == [
  2412. "I will add it.",
  2413. "Scheduled Design review for 2026-07-14T09:30:00+08:00 (Asia/Shanghai).",
  2414. ]
  2415. assert outputs[-1] == {"type": "done"}
  2416. @pytest.mark.asyncio
  2417. async def test_direct_silent_failure_requests_exactly_one_corrective_follow_up():
  2418. def fail(event: ToolCallEvent) -> dict[str, Any]:
  2419. raise RuntimeError("device offline")
  2420. registry = ToolRegistry(
  2421. [
  2422. ToolDefinition(
  2423. name="example.silent",
  2424. description="Complete silently on success.",
  2425. parameters={"type": "object"},
  2426. handler=fail,
  2427. result_policy=ResultPolicy.SILENT_SUCCESS,
  2428. )
  2429. ]
  2430. )
  2431. client = ScriptedChatClient(
  2432. [
  2433. [
  2434. StreamItem.message_delta("Trying it."),
  2435. StreamItem.provider_tool_call(
  2436. _tool_call("silent-failure", name="example.silent")
  2437. ),
  2438. ],
  2439. [StreamItem.message_delta("I could not apply that change.")],
  2440. ]
  2441. )
  2442. outputs = await _collect_outputs(
  2443. DebugRuntime(client, registry=registry).run(
  2444. _policy_request(
  2445. enabled_tools=["example.silent"],
  2446. mode="chat_agent_tools",
  2447. )
  2448. )
  2449. )
  2450. assert client.calls == 2
  2451. assert [
  2452. message["content"]
  2453. for message in outputs
  2454. if message["type"] == "message_delta"
  2455. ] == ["Trying it.", "I could not apply that change."]
  2456. @pytest.mark.asyncio
  2457. async def test_mixed_builtin_batch_emits_template_then_terminates_without_llm():
  2458. calls = [
  2459. ToolCallEvent(
  2460. id="terminate-1",
  2461. name="session.terminate",
  2462. arguments={},
  2463. raw_arguments="{}",
  2464. ),
  2465. ToolCallEvent(
  2466. id="schedule-1",
  2467. name="calendar.schedule.create",
  2468. arguments={
  2469. "title": "Design review",
  2470. "start_at": "2026-07-14T09:30:00+08:00",
  2471. "timezone": "Asia/Shanghai",
  2472. },
  2473. raw_arguments=(
  2474. '{"title":"Design review","start_at":"2026-07-14T09:30:00+08:00",'
  2475. '"timezone":"Asia/Shanghai"}'
  2476. ),
  2477. ),
  2478. ToolCallEvent(
  2479. id="search-1",
  2480. name="knowledge.web.search",
  2481. arguments={"query": "event batch executors"},
  2482. raw_arguments='{"query":"event batch executors"}',
  2483. ),
  2484. ToolCallEvent(
  2485. id="volume-1",
  2486. name="device.volume.adjust",
  2487. arguments={"mode": "absolute", "value": 30},
  2488. raw_arguments='{"mode":"absolute","value":30}',
  2489. ),
  2490. ]
  2491. client = ScriptedChatClient(
  2492. [
  2493. [
  2494. StreamItem.message_delta("Goodbye, I will finish those first."),
  2495. *[StreamItem.provider_tool_call(call) for call in calls],
  2496. ]
  2497. ]
  2498. )
  2499. outputs = await _collect_outputs(
  2500. DebugRuntime(client).run(
  2501. _policy_request(
  2502. enabled_tools=[call.name for call in calls],
  2503. mode="chat_agent_tools",
  2504. )
  2505. )
  2506. )
  2507. assert client.calls == 1
  2508. assert [
  2509. message["message"]["tool_call_id"]
  2510. for message in outputs
  2511. if message["type"] == "tool_result"
  2512. ] == [call.id for call in calls]
  2513. assert [
  2514. message["content"]
  2515. for message in outputs
  2516. if message["type"] == "message_delta"
  2517. ][-1] == (
  2518. "Scheduled Design review for 2026-07-14T09:30:00+08:00 "
  2519. "(Asia/Shanghai)."
  2520. )
  2521. decision = next(
  2522. message
  2523. for message in outputs
  2524. if message.get("event") == "event_policy_decision"
  2525. )
  2526. assert decision["details"]["terminate"] is True
  2527. assert decision["details"]["llm_follow_up"] is False
  2528. assert any(
  2529. message.get("event") == "terminal_completed" for message in outputs
  2530. )
  2531. assert outputs[-1] == {"type": "done"}
  2532. @pytest.mark.asyncio
  2533. async def test_reusable_session_terminate_completes_turn_and_session_task():
  2534. client = ScriptedChatClient(
  2535. [
  2536. [
  2537. StreamItem.message_delta("Goodbye."),
  2538. StreamItem.provider_tool_call(
  2539. ToolCallEvent(
  2540. id="terminate-session",
  2541. name="session.terminate",
  2542. arguments={},
  2543. raw_arguments="{}",
  2544. )
  2545. ),
  2546. ]
  2547. ]
  2548. )
  2549. runtime = DebugRuntime(client)
  2550. queues = runtime.start_session(
  2551. _policy_request(
  2552. enabled_tools=["session.terminate"],
  2553. mode="chat_agent_tools",
  2554. session_id="reusable-session",
  2555. )
  2556. )
  2557. outputs: list[dict[str, Any]] = []
  2558. while not any(message["type"] == "done" for message in outputs):
  2559. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  2560. await asyncio.wait_for(runtime._wait_for_tasks(), timeout=1)
  2561. business_types = [message["type"] for message in _without_audit(outputs)]
  2562. assert business_types[-2:] == ["turn_completed", "done"]
  2563. @pytest.mark.asyncio
  2564. async def test_dual_text_event_scope_uses_round_and_does_not_replay_stale_history():
  2565. handled_messages: list[str] = []
  2566. def resolve_from_history(
  2567. event: ToolCallEvent,
  2568. context: ToolExecutionContext,
  2569. ) -> dict[str, Any]:
  2570. del event
  2571. return {
  2572. "message": next(
  2573. message.content
  2574. for message in reversed(context.history)
  2575. if message.role == "assistant" and message.content
  2576. )
  2577. }
  2578. registry = ToolRegistry(
  2579. [
  2580. ToolDefinition(
  2581. name="example.history",
  2582. description="Capture the current assistant history.",
  2583. parameters={
  2584. "type": "object",
  2585. "properties": {"message": {"type": "string"}},
  2586. "required": ["message"],
  2587. },
  2588. handler=lambda event: handled_messages.append(
  2589. event.arguments["message"]
  2590. )
  2591. or {"message": event.arguments["message"]},
  2592. argument_resolver=resolve_from_history,
  2593. )
  2594. ]
  2595. )
  2596. repeated = _tool_call("event_1", name="example.history")
  2597. client = ScriptedChatClient(
  2598. [
  2599. [StreamItem.message_delta("first history"), StreamItem.text_event(repeated)],
  2600. [StreamItem.message_delta("second history"), StreamItem.text_event(repeated)],
  2601. [StreamItem.message_delta("final answer")],
  2602. ]
  2603. )
  2604. request = DebugRunRequest(
  2605. user_message="capture each round",
  2606. system_prompts=[],
  2607. pre_messages=[],
  2608. chat_agent=AgentParams(model="chat-model"),
  2609. event_agent=EventAgentParams(
  2610. enabled_tools=["example.history"],
  2611. max_event_loops=2,
  2612. ),
  2613. )
  2614. outputs = await _collect_outputs(
  2615. DebugRuntime(client, registry=registry).run(request)
  2616. )
  2617. assert handled_messages == ["first history", "second history"]
  2618. scopes = [
  2619. message["details"]["scope"]
  2620. for message in outputs
  2621. if message.get("event") == "event_batch_started"
  2622. ]
  2623. assert scopes[0].endswith(":round-1")
  2624. assert scopes[1].endswith(":round-2")
  2625. @pytest.mark.asyncio
  2626. @pytest.mark.parametrize("mode", ["dual_agent", "chat_agent_tools"])
  2627. async def test_runtime_releases_batch_scope_after_results(monkeypatch, mode: str):
  2628. released: list[Any] = []
  2629. original_release = EventBatchExecutor.release_scope
  2630. def record_release(self: EventBatchExecutor, scope: Any) -> None:
  2631. released.append(scope)
  2632. original_release(self, scope)
  2633. monkeypatch.setattr(EventBatchExecutor, "release_scope", record_release)
  2634. registry = ToolRegistry(
  2635. [
  2636. ToolDefinition(
  2637. name="example.silent",
  2638. description="Complete without another model call.",
  2639. parameters={"type": "object"},
  2640. handler=lambda event: {"event_id": event.id},
  2641. result_policy=ResultPolicy.SILENT_SUCCESS,
  2642. )
  2643. ]
  2644. )
  2645. event = _tool_call("release-1", name="example.silent")
  2646. item = (
  2647. StreamItem.text_event(event)
  2648. if mode == "dual_agent"
  2649. else StreamItem.provider_tool_call(event)
  2650. )
  2651. client = ScriptedChatClient([[item]])
  2652. await _collect_outputs(
  2653. DebugRuntime(client, registry=registry).run(
  2654. _policy_request(enabled_tools=["example.silent"], mode=mode)
  2655. )
  2656. )
  2657. assert len(released) == 1
  2658. assert str(released[0]).endswith(":round-1")
  2659. @pytest.mark.asyncio
  2660. @pytest.mark.parametrize("mode", ["dual_agent", "chat_agent_tools"])
  2661. async def test_tool_only_terminate_emits_plugin_farewell_once(mode: str):
  2662. event = ToolCallEvent(
  2663. id="terminate-only",
  2664. name="session.terminate",
  2665. arguments={},
  2666. raw_arguments="{}",
  2667. )
  2668. item = (
  2669. StreamItem.text_event(event)
  2670. if mode == "dual_agent"
  2671. else StreamItem.provider_tool_call(event)
  2672. )
  2673. client = ScriptedChatClient([[item]])
  2674. outputs = await _collect_outputs(
  2675. DebugRuntime(client).run(
  2676. _policy_request(enabled_tools=["session.terminate"], mode=mode)
  2677. )
  2678. )
  2679. assert client.calls == 1
  2680. assert [
  2681. message["content"]
  2682. for message in outputs
  2683. if message["type"] == "message_delta"
  2684. ] == ["Goodbye."]
  2685. assert outputs[-1] == {"type": "done"}
  2686. @pytest.mark.asyncio
  2687. @pytest.mark.parametrize("mode", ["dual_agent", "chat_agent_tools"])
  2688. async def test_coalesced_terminate_has_one_farewell_and_one_audit_result(mode: str):
  2689. events = [
  2690. ToolCallEvent(
  2691. id=event_id,
  2692. name="session.terminate",
  2693. arguments={},
  2694. raw_arguments="{}",
  2695. )
  2696. for event_id in ("terminate-primary", "terminate-duplicate")
  2697. ]
  2698. items = [
  2699. (
  2700. StreamItem.text_event(event)
  2701. if mode == "dual_agent"
  2702. else StreamItem.provider_tool_call(event)
  2703. )
  2704. for event in events
  2705. ]
  2706. client = ScriptedChatClient([items])
  2707. outputs = await _collect_outputs(
  2708. DebugRuntime(client).run(
  2709. _policy_request(enabled_tools=["session.terminate"], mode=mode)
  2710. )
  2711. )
  2712. assert [
  2713. message["content"]
  2714. for message in outputs
  2715. if message["type"] == "message_delta"
  2716. ] == ["Goodbye."]
  2717. assert [
  2718. message["message"]["tool_call_id"]
  2719. for message in outputs
  2720. if message["type"] == "tool_result"
  2721. ] == [event.id for event in events]
  2722. batch_audit = next(
  2723. message
  2724. for message in outputs
  2725. if message.get("event") == "event_batch_results"
  2726. )
  2727. terminal_audit = next(
  2728. message
  2729. for message in outputs
  2730. if message.get("event") == "terminal_completed"
  2731. )
  2732. completion_audit = next(
  2733. message
  2734. for message in outputs
  2735. if message.get("event")
  2736. == (
  2737. "provider_tools_completed"
  2738. if mode == "chat_agent_tools"
  2739. else "event_agent_completed"
  2740. )
  2741. )
  2742. assert [
  2743. result["event_id"] for result in batch_audit["details"]["results"]
  2744. ] == ["terminate-primary"]
  2745. assert completion_audit["details"]["result_count"] == 1
  2746. assert terminal_audit["details"]["event_ids"] == ["terminate-primary"]
  2747. @pytest.mark.asyncio
  2748. async def test_tool_only_terminate_ends_reusable_session_after_plugin_farewell():
  2749. client = ScriptedChatClient(
  2750. [
  2751. [
  2752. StreamItem.provider_tool_call(
  2753. ToolCallEvent(
  2754. id="terminate-session-only",
  2755. name="session.terminate",
  2756. arguments={},
  2757. raw_arguments="{}",
  2758. )
  2759. )
  2760. ]
  2761. ]
  2762. )
  2763. runtime = DebugRuntime(client)
  2764. queues = runtime.start_session(
  2765. _policy_request(
  2766. enabled_tools=["session.terminate"],
  2767. mode="chat_agent_tools",
  2768. session_id="terminal-only-session",
  2769. )
  2770. )
  2771. outputs: list[dict[str, Any]] = []
  2772. while not any(message["type"] == "done" for message in outputs):
  2773. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  2774. await asyncio.wait_for(runtime._wait_for_tasks(), timeout=1)
  2775. assert [
  2776. message["content"]
  2777. for message in outputs
  2778. if message["type"] == "message_delta"
  2779. ] == ["Goodbye."]
  2780. assert [message["type"] for message in _without_audit(outputs)][-2:] == [
  2781. "turn_completed",
  2782. "done",
  2783. ]
  2784. @pytest.mark.asyncio
  2785. @pytest.mark.parametrize("mode", ["dual_agent", "chat_agent_tools"])
  2786. async def test_web_search_has_first_answer_tool_result_and_one_second_answer(mode: str):
  2787. search = ToolCallEvent(
  2788. id="search-two-answer",
  2789. name="knowledge.web.search",
  2790. arguments={"query": "event batch safety"},
  2791. raw_arguments='{"query":"event batch safety"}',
  2792. )
  2793. item = (
  2794. StreamItem.text_event(search)
  2795. if mode == "dual_agent"
  2796. else StreamItem.provider_tool_call(search)
  2797. )
  2798. client = ScriptedChatClient(
  2799. [
  2800. [StreamItem.message_delta("I will check."), item],
  2801. [StreamItem.message_delta("Grounded search update.")],
  2802. ]
  2803. )
  2804. outputs = await _collect_outputs(
  2805. DebugRuntime(client).run(
  2806. _policy_request(enabled_tools=["knowledge.web.search"], mode=mode)
  2807. )
  2808. )
  2809. business = _without_audit(outputs)
  2810. visible = [
  2811. (index, message["content"])
  2812. for index, message in enumerate(business)
  2813. if message["type"] == "message_delta"
  2814. ]
  2815. tool_index = next(
  2816. index for index, message in enumerate(business) if message["type"] == "tool_result"
  2817. )
  2818. assert client.calls == 2
  2819. assert [content for _, content in visible] == [
  2820. "I will check.",
  2821. "Grounded search update.",
  2822. ]
  2823. assert visible[0][0] < tool_index < visible[1][0]
  2824. assert "sources" in business[tool_index]["message"]["content"]