test_debug_runtime.py 112 KB

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