test_debug_runtime.py 72 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117111811191120112111221123112411251126112711281129113011311132113311341135113611371138113911401141114211431144114511461147114811491150115111521153115411551156115711581159116011611162116311641165116611671168116911701171117211731174117511761177117811791180118111821183118411851186118711881189119011911192119311941195119611971198119912001201120212031204120512061207120812091210121112121213121412151216121712181219122012211222122312241225122612271228122912301231123212331234123512361237123812391240124112421243124412451246124712481249125012511252125312541255125612571258125912601261126212631264126512661267126812691270127112721273127412751276127712781279128012811282128312841285128612871288128912901291129212931294129512961297129812991300130113021303130413051306130713081309131013111312131313141315131613171318131913201321132213231324132513261327132813291330133113321333133413351336133713381339134013411342134313441345134613471348134913501351135213531354135513561357135813591360136113621363136413651366136713681369137013711372137313741375137613771378137913801381138213831384138513861387138813891390139113921393139413951396139713981399140014011402140314041405140614071408140914101411141214131414141514161417141814191420142114221423142414251426142714281429143014311432143314341435143614371438143914401441144214431444144514461447144814491450145114521453145414551456145714581459146014611462146314641465146614671468146914701471147214731474147514761477147814791480148114821483148414851486148714881489149014911492149314941495149614971498149915001501150215031504150515061507150815091510151115121513151415151516151715181519152015211522152315241525152615271528152915301531153215331534153515361537153815391540154115421543154415451546154715481549155015511552155315541555155615571558155915601561156215631564156515661567156815691570157115721573157415751576157715781579158015811582158315841585158615871588158915901591159215931594159515961597159815991600160116021603160416051606160716081609161016111612161316141615161616171618161916201621162216231624162516261627162816291630163116321633163416351636163716381639164016411642164316441645164616471648164916501651165216531654165516561657165816591660166116621663166416651666166716681669167016711672167316741675167616771678167916801681168216831684168516861687168816891690169116921693169416951696169716981699170017011702170317041705170617071708170917101711171217131714171517161717171817191720172117221723172417251726172717281729173017311732173317341735173617371738173917401741174217431744174517461747174817491750175117521753175417551756175717581759176017611762176317641765176617671768176917701771177217731774177517761777177817791780178117821783178417851786178717881789179017911792179317941795179617971798179918001801180218031804180518061807180818091810181118121813181418151816181718181819182018211822182318241825182618271828182918301831183218331834183518361837183818391840184118421843184418451846184718481849185018511852185318541855185618571858185918601861186218631864186518661867186818691870187118721873187418751876187718781879188018811882188318841885188618871888188918901891189218931894189518961897189818991900190119021903190419051906190719081909191019111912191319141915191619171918191919201921192219231924192519261927192819291930193119321933193419351936193719381939194019411942194319441945194619471948194919501951195219531954195519561957195819591960196119621963196419651966196719681969197019711972197319741975197619771978197919801981198219831984198519861987198819891990199119921993199419951996199719981999200020012002200320042005200620072008200920102011201220132014201520162017201820192020202120222023202420252026202720282029203020312032203320342035203620372038203920402041204220432044204520462047204820492050205120522053205420552056205720582059206020612062206320642065206620672068206920702071207220732074207520762077207820792080208120822083208420852086208720882089209020912092209320942095209620972098209921002101210221032104210521062107210821092110211121122113211421152116211721182119212021212122212321242125212621272128212921302131213221332134213521362137213821392140214121422143214421452146214721482149215021512152215321542155215621572158215921602161216221632164216521662167216821692170217121722173217421752176217721782179218021812182218321842185218621872188218921902191219221932194219521962197219821992200220122022203220422052206220722082209221022112212221322142215221622172218221922202221222222232224
  1. import asyncio
  2. import importlib
  3. import json
  4. import logging
  5. from collections.abc import AsyncIterator
  6. from typing import Any
  7. import pytest
  8. from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
  9. from agent_lab.application.event_agent import EventAgentRequest
  10. from agent_lab.application.runtime import DebugRuntime
  11. from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
  12. from agent_lab.domain.events import ToolCallEvent
  13. from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
  14. from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
  15. def _runtime_queues_class():
  16. module = importlib.import_module("agent_lab.application.queues")
  17. return module.RuntimeQueues
  18. async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[str, Any]]:
  19. return [message async for message in stream]
  20. def _without_audit(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
  21. return [message for message in outputs if message["type"] != "audit"]
  22. def _message_types(outputs: list[dict[str, Any]]) -> list[str]:
  23. return [message["type"] for message in _without_audit(outputs)]
  24. def _event_tool_call_from_tools(
  25. tools: list[dict],
  26. messages: list[ChatMessage],
  27. ) -> StreamItem:
  28. tool_name = tools[0]["function"]["name"]
  29. content = ""
  30. for role in ("assistant", "user"):
  31. content = next(
  32. (
  33. message.content
  34. for message in reversed(messages)
  35. if message.role == role and message.content.strip()
  36. ),
  37. "",
  38. )
  39. if content:
  40. break
  41. arguments = {
  42. "message": content,
  43. "query": content,
  44. "title": content,
  45. }
  46. return StreamItem.provider_tool_call(
  47. ToolCallEvent(
  48. id="event_agent_call_1",
  49. name=tool_name,
  50. arguments=arguments,
  51. raw_arguments=json.dumps(arguments),
  52. )
  53. )
  54. async def _next_non_audit(
  55. stream: AsyncIterator[dict[str, Any]],
  56. ) -> dict[str, Any]:
  57. while True:
  58. message = await anext(stream)
  59. if message["type"] != "audit":
  60. return message
  61. async def _next_non_audit_from_queue(queues: Any) -> dict[str, Any]:
  62. while True:
  63. message = await queues.output.get()
  64. if message["type"] != "audit":
  65. return message
  66. class RecordingQueue(asyncio.Queue):
  67. def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
  68. super().__init__()
  69. self.name = name
  70. self.log = log
  71. async def put(self, item: Any) -> None:
  72. self.log.append((self.name, "put", self._describe(item)))
  73. await super().put(item)
  74. async def get(self) -> Any:
  75. item = await super().get()
  76. self.log.append((self.name, "get", self._describe(item)))
  77. return item
  78. def _describe(self, item: Any) -> str:
  79. if isinstance(item, ChatMessage):
  80. if item.role == "tool":
  81. return f"tool:{item.tool_call_id}"
  82. return item.role
  83. if isinstance(item, ToolCallEvent):
  84. return f"event:{item.name}:{item.id}"
  85. if isinstance(item, EventAgentRequest):
  86. events = ",".join(f"{event.name}:{event.id}" for event in item.events)
  87. return f"event_request:{events}"
  88. if isinstance(item, dict):
  89. return f"output:{item.get('type')}"
  90. return type(item).__name__
  91. class FakeChatClient:
  92. def __init__(self) -> None:
  93. self.calls = 0
  94. async def stream_chat(
  95. self,
  96. messages: list[ChatMessage],
  97. tools: list[dict],
  98. params: AgentParams,
  99. tool_choice: dict[str, Any] | None = None,
  100. ) -> AsyncIterator[StreamItem]:
  101. if tools:
  102. yield StreamItem.raw_response_chunk(
  103. {
  104. "choices": [
  105. {
  106. "delta": {
  107. "tool_calls": [
  108. {
  109. "index": 0,
  110. "function": {"name": tools[0]["function"]["name"]},
  111. }
  112. ]
  113. },
  114. "finish_reason": None,
  115. }
  116. ]
  117. }
  118. )
  119. yield _event_tool_call_from_tools(tools, messages)
  120. return
  121. self.calls += 1
  122. if self.calls == 1:
  123. yield StreamItem.raw_response_chunk(
  124. {
  125. "choices": [
  126. {
  127. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  128. "finish_reason": None,
  129. }
  130. ]
  131. }
  132. )
  133. yield StreamItem.text_event(
  134. ToolCallEvent(
  135. id="call_1",
  136. name="handoff_note",
  137. arguments={},
  138. raw_arguments="{}",
  139. )
  140. )
  141. return
  142. assert not any(message.role == "tool" for message in messages)
  143. assert any(message.name == "event_agent" for message in messages)
  144. yield StreamItem.message_delta("final answer")
  145. class WrongSourceEventChatClient:
  146. def __init__(self, item: StreamItem) -> None:
  147. self.item = item
  148. self.calls = 0
  149. async def stream_chat(
  150. self,
  151. messages: list[ChatMessage],
  152. tools: list[dict],
  153. params: AgentParams,
  154. tool_choice: dict[str, Any] | None = None,
  155. ) -> AsyncIterator[StreamItem]:
  156. if tools:
  157. yield _event_tool_call_from_tools(tools, messages)
  158. return
  159. self.calls += 1
  160. if self.calls == 1:
  161. yield self.item
  162. return
  163. yield StreamItem.message_delta("unexpected continuation")
  164. class IncrementingClock:
  165. def __init__(self, current: float = 100.0, step: float = 0.01) -> None:
  166. self.current = current
  167. self.step = step
  168. def __call__(self) -> float:
  169. self.current += self.step
  170. return self.current
  171. class StrictHistoryChatClient:
  172. def __init__(self) -> None:
  173. self.calls = 0
  174. self.second_call_messages: list[ChatMessage] = []
  175. async def stream_chat(
  176. self,
  177. messages: list[ChatMessage],
  178. tools: list[dict],
  179. params: AgentParams,
  180. tool_choice: dict[str, Any] | None = None,
  181. ) -> AsyncIterator[StreamItem]:
  182. if tools:
  183. yield StreamItem.raw_response_chunk(
  184. {
  185. "choices": [
  186. {
  187. "delta": {
  188. "tool_calls": [
  189. {
  190. "index": 0,
  191. "function": {"name": tools[0]["function"]["name"]},
  192. }
  193. ]
  194. },
  195. "finish_reason": None,
  196. }
  197. ]
  198. }
  199. )
  200. yield _event_tool_call_from_tools(tools, messages)
  201. return
  202. self.calls += 1
  203. if self.calls == 1:
  204. yield StreamItem.raw_response_chunk(
  205. {
  206. "choices": [
  207. {
  208. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  209. "finish_reason": None,
  210. }
  211. ]
  212. }
  213. )
  214. yield StreamItem.text_event(
  215. ToolCallEvent(
  216. id="call_1",
  217. name="handoff_note",
  218. arguments={},
  219. raw_arguments="{}",
  220. )
  221. )
  222. return
  223. self.second_call_messages = list(messages)
  224. yield StreamItem.message_delta("final answer")
  225. class MultiEventChatClient:
  226. def __init__(self) -> None:
  227. self.calls = 0
  228. self.second_call_messages: list[ChatMessage] = []
  229. async def stream_chat(
  230. self,
  231. messages: list[ChatMessage],
  232. tools: list[dict],
  233. params: AgentParams,
  234. tool_choice: dict[str, Any] | None = None,
  235. ) -> AsyncIterator[StreamItem]:
  236. if tools:
  237. yield _event_tool_call_from_tools(tools, messages)
  238. return
  239. self.calls += 1
  240. if self.calls == 1:
  241. yield StreamItem.message_delta("Checking events.")
  242. yield StreamItem.text_event(
  243. ToolCallEvent(
  244. id="call_1",
  245. name="handoff_note",
  246. arguments={"message": "ignored chat argument"},
  247. raw_arguments='{"message":"ignored chat argument"}',
  248. )
  249. )
  250. yield StreamItem.text_event(
  251. ToolCallEvent(
  252. id="call_2",
  253. name="audit_note",
  254. arguments={"message": "ignored chat argument"},
  255. raw_arguments='{"message":"ignored chat argument"}',
  256. )
  257. )
  258. return
  259. self.second_call_messages = list(messages)
  260. yield StreamItem.message_delta("Final answer.")
  261. class ToolCapturingChatClient:
  262. def __init__(self) -> None:
  263. self.tools: list[dict[str, Any]] = []
  264. self.messages: list[ChatMessage] = []
  265. async def stream_chat(
  266. self,
  267. messages: list[ChatMessage],
  268. tools: list[dict],
  269. params: AgentParams,
  270. tool_choice: dict[str, Any] | None = None,
  271. ) -> AsyncIterator[StreamItem]:
  272. self.messages = list(messages)
  273. self.tools = list(tools)
  274. yield StreamItem.message_delta("final answer")
  275. class EventLoopLimitChatClient:
  276. def __init__(self) -> None:
  277. self.calls = 0
  278. self.tools_by_call: list[list[dict[str, Any]]] = []
  279. self.messages_by_call: list[list[ChatMessage]] = []
  280. async def stream_chat(
  281. self,
  282. messages: list[ChatMessage],
  283. tools: list[dict],
  284. params: AgentParams,
  285. tool_choice: dict[str, Any] | None = None,
  286. ) -> AsyncIterator[StreamItem]:
  287. if tools:
  288. yield _event_tool_call_from_tools(tools, messages)
  289. return
  290. self.calls += 1
  291. self.tools_by_call.append(list(tools))
  292. self.messages_by_call.append(list(messages))
  293. if self.calls == 1:
  294. yield StreamItem.text_event(
  295. ToolCallEvent(
  296. id="call_1",
  297. name="handoff_note",
  298. arguments={},
  299. raw_arguments="{}",
  300. )
  301. )
  302. return
  303. yield StreamItem.message_delta("final after event limit")
  304. class RoundStatsChatClient:
  305. async def stream_chat(
  306. self,
  307. messages: list[ChatMessage],
  308. tools: list[dict],
  309. params: AgentParams,
  310. tool_choice: dict[str, Any] | None = None,
  311. ) -> AsyncIterator[StreamItem]:
  312. yield StreamItem.message_delta("hello")
  313. yield StreamItem.usage_item(
  314. TokenUsage(
  315. prompt_tokens=10,
  316. completion_tokens=20,
  317. total_tokens=30,
  318. cached_tokens=5,
  319. )
  320. )
  321. class EventRoundStatsChatClient:
  322. def __init__(self) -> None:
  323. self.calls = 0
  324. async def stream_chat(
  325. self,
  326. messages: list[ChatMessage],
  327. tools: list[dict],
  328. params: AgentParams,
  329. tool_choice: dict[str, Any] | None = None,
  330. ) -> AsyncIterator[StreamItem]:
  331. if tools:
  332. yield StreamItem.raw_response_chunk(
  333. {
  334. "choices": [
  335. {
  336. "delta": {
  337. "tool_calls": [
  338. {
  339. "index": 0,
  340. "function": {
  341. "name": tools[0]["function"]["name"],
  342. },
  343. }
  344. ]
  345. },
  346. "finish_reason": None,
  347. }
  348. ]
  349. }
  350. )
  351. yield _event_tool_call_from_tools(tools, messages)
  352. return
  353. self.calls += 1
  354. if self.calls == 1:
  355. yield StreamItem.raw_response_chunk(
  356. {
  357. "choices": [
  358. {
  359. "delta": {
  360. "content": "<agent_events>handoff_note</agent_events>",
  361. },
  362. "finish_reason": None,
  363. }
  364. ]
  365. }
  366. )
  367. yield StreamItem.text_event(
  368. ToolCallEvent(
  369. id="call_1",
  370. name="handoff_note",
  371. arguments={},
  372. raw_arguments="{}",
  373. )
  374. )
  375. yield StreamItem.usage_item(
  376. TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3)
  377. )
  378. return
  379. yield StreamItem.raw_response_chunk(
  380. {
  381. "choices": [
  382. {
  383. "delta": {"content": "final answer"},
  384. "finish_reason": None,
  385. }
  386. ]
  387. }
  388. )
  389. yield StreamItem.message_delta("final answer")
  390. yield StreamItem.usage_item(
  391. TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
  392. )
  393. class ContextBoundaryChatClient:
  394. def __init__(self) -> None:
  395. self.calls = 0
  396. self.event_agent_messages: list[list[ChatMessage]] = []
  397. async def stream_chat(
  398. self,
  399. messages: list[ChatMessage],
  400. tools: list[dict],
  401. params: AgentParams,
  402. tool_choice: dict[str, Any] | None = None,
  403. ) -> AsyncIterator[StreamItem]:
  404. if tools:
  405. self.event_agent_messages.append(list(messages))
  406. yield _event_tool_call_from_tools(tools, messages)
  407. return
  408. self.calls += 1
  409. if self.calls == 1:
  410. yield StreamItem.message_delta("I will check.")
  411. yield StreamItem.text_event(
  412. ToolCallEvent(
  413. id="call_1",
  414. name="handoff_note",
  415. arguments={},
  416. raw_arguments="{}",
  417. )
  418. )
  419. return
  420. yield StreamItem.message_delta("final answer")
  421. class TwoTurnSessionChatClient:
  422. def __init__(self) -> None:
  423. self.calls = 0
  424. self.messages_by_call: list[list[ChatMessage]] = []
  425. async def stream_chat(
  426. self,
  427. messages: list[ChatMessage],
  428. tools: list[dict],
  429. params: AgentParams,
  430. tool_choice: dict[str, Any] | None = None,
  431. ) -> AsyncIterator[StreamItem]:
  432. if tools:
  433. yield _event_tool_call_from_tools(tools, messages)
  434. return
  435. self.calls += 1
  436. self.messages_by_call.append(list(messages))
  437. if self.calls in {1, 3}:
  438. yield StreamItem.text_event(
  439. ToolCallEvent(
  440. id=f"call_{self.calls}",
  441. name="handoff_note",
  442. arguments={},
  443. raw_arguments="{}",
  444. )
  445. )
  446. return
  447. yield StreamItem.message_delta(f"final answer {self.calls}")
  448. class SlowAfterEventChatClient:
  449. def __init__(self) -> None:
  450. self.calls = 0
  451. self.event_seen = asyncio.Event()
  452. self.release_stream = asyncio.Event()
  453. async def stream_chat(
  454. self,
  455. messages: list[ChatMessage],
  456. tools: list[dict],
  457. params: AgentParams,
  458. tool_choice: dict[str, Any] | None = None,
  459. ) -> AsyncIterator[StreamItem]:
  460. if tools:
  461. yield _event_tool_call_from_tools(tools, messages)
  462. return
  463. self.calls += 1
  464. if self.calls == 1:
  465. yield StreamItem.message_delta("Need event.")
  466. yield StreamItem.text_event(
  467. ToolCallEvent(
  468. id="call_1",
  469. name="handoff_note",
  470. arguments={},
  471. raw_arguments="{}",
  472. )
  473. )
  474. self.event_seen.set()
  475. await self.release_stream.wait()
  476. return
  477. yield StreamItem.message_delta("final answer")
  478. def test_runtime_queues_exposes_input_output_and_events_queues():
  479. RuntimeQueues = _runtime_queues_class()
  480. queues = RuntimeQueues()
  481. assert isinstance(queues.input, asyncio.Queue)
  482. assert isinstance(queues.output, asyncio.Queue)
  483. assert isinstance(queues.events, asyncio.Queue)
  484. assert queues.input is not queues.output
  485. assert queues.input is not queues.events
  486. assert queues.output is not queues.events
  487. @pytest.mark.asyncio
  488. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  489. request = DebugRunRequest(
  490. user_message="debug this",
  491. system_prompts=["You are a debugger."],
  492. pre_messages=[],
  493. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  494. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  495. )
  496. client = FakeChatClient()
  497. runtime = DebugRuntime(client)
  498. outputs = [message async for message in runtime.run(request)]
  499. business_outputs = _without_audit(outputs)
  500. assert client.calls == 2
  501. assert _message_types(outputs) == [
  502. "session_started",
  503. "event",
  504. "tool_result",
  505. "round_stats",
  506. "message_delta",
  507. "round_stats",
  508. "done",
  509. ]
  510. assert business_outputs[1]["event"]["name"] == "handoff_note"
  511. assert business_outputs[4]["content"] == "final answer"
  512. @pytest.mark.asyncio
  513. @pytest.mark.parametrize(
  514. "item",
  515. [
  516. StreamItem.provider_tool_call(
  517. ToolCallEvent(
  518. id="provider_call_1",
  519. name="handoff_note",
  520. arguments={"message": "wrong source"},
  521. raw_arguments='{"message":"wrong source"}',
  522. )
  523. ),
  524. StreamItem.event(
  525. ToolCallEvent(
  526. id="legacy_event_1",
  527. name="handoff_note",
  528. arguments={},
  529. raw_arguments="{}",
  530. )
  531. ),
  532. ],
  533. ids=["provider_tool_call", "legacy_event"],
  534. )
  535. async def test_runtime_ignores_non_text_event_sources(item: StreamItem):
  536. request = DebugRunRequest(
  537. user_message="debug this",
  538. system_prompts=[],
  539. pre_messages=[],
  540. chat_agent=AgentParams(model="fake-model"),
  541. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  542. )
  543. client = WrongSourceEventChatClient(item)
  544. outputs = [message async for message in DebugRuntime(client).run(request)]
  545. assert client.calls == 1
  546. assert "event" not in _message_types(outputs)
  547. assert "tool_result" not in _message_types(outputs)
  548. @pytest.mark.asyncio
  549. async def test_runtime_emits_audit_events_and_backend_logs(caplog):
  550. request = DebugRunRequest(
  551. user_message="debug this",
  552. system_prompts=[],
  553. pre_messages=[],
  554. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  555. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  556. )
  557. runtime = DebugRuntime(FakeChatClient())
  558. with caplog.at_level(logging.INFO, logger="agent_lab.application.runtime"):
  559. outputs = [message async for message in runtime.run(request)]
  560. audit_events = [
  561. message["event"]
  562. for message in outputs
  563. if message["type"] == "audit"
  564. ]
  565. assert audit_events == [
  566. "session_started",
  567. "chat_round_started",
  568. "chat_agent_request",
  569. "chat_event_detected",
  570. "chat_agent_response",
  571. "event_agent_request",
  572. "event_agent_response",
  573. "event_agent_completed",
  574. "chat_round_finished",
  575. "chat_round_started",
  576. "chat_agent_request",
  577. "chat_message_stream_started",
  578. "chat_message_stream_finished",
  579. "chat_agent_response",
  580. "chat_round_finished",
  581. "session_finished",
  582. ]
  583. assert "chat_event_detected" in caplog.text
  584. assert "event_agent_completed" in caplog.text
  585. @pytest.mark.asyncio
  586. async def test_runtime_audits_chat_message_stream_boundaries_in_output_order():
  587. request = DebugRunRequest(
  588. user_message="debug this",
  589. system_prompts=[],
  590. pre_messages=[],
  591. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  592. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  593. )
  594. runtime = DebugRuntime(RoundStatsChatClient())
  595. outputs = [message async for message in runtime.run(request)]
  596. ordered_labels = [
  597. message["event"] if message["type"] == "audit" else message["type"]
  598. for message in outputs
  599. ]
  600. assert ordered_labels.index("chat_message_stream_started") < ordered_labels.index(
  601. "message_delta"
  602. )
  603. assert ordered_labels.index("message_delta") < ordered_labels.index(
  604. "chat_message_stream_finished"
  605. )
  606. assert ordered_labels.index("chat_message_stream_finished") < ordered_labels.index(
  607. "chat_agent_response"
  608. )
  609. stream_started = next(
  610. message
  611. for message in outputs
  612. if message.get("event") == "chat_message_stream_started"
  613. )
  614. stream_finished = next(
  615. message
  616. for message in outputs
  617. if message.get("event") == "chat_message_stream_finished"
  618. )
  619. assert stream_started["details"]["agent"] == "chat_agent"
  620. assert stream_started["details"]["round_index"] == 1
  621. assert stream_finished["details"]["delta_count"] == 1
  622. assert stream_finished["details"]["content_length"] == len("hello")
  623. @pytest.mark.asyncio
  624. async def test_runtime_audit_events_include_turn_relative_elapsed_time():
  625. request = DebugRunRequest(
  626. user_message="debug this",
  627. system_prompts=[],
  628. pre_messages=[],
  629. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  630. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  631. )
  632. runtime = DebugRuntime(FakeChatClient(), clock=IncrementingClock())
  633. queues = runtime.start_session(request)
  634. outputs: list[dict[str, Any]] = []
  635. while not any(message["type"] == "turn_completed" for message in outputs):
  636. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  637. await runtime.aclose()
  638. turn_audits = [
  639. message
  640. for message in outputs
  641. if message["type"] == "audit"
  642. ]
  643. elapsed_values = [
  644. message["details"].get("turn_elapsed_ms") for message in turn_audits
  645. ]
  646. assert turn_audits
  647. assert all(isinstance(value, int) for value in elapsed_values)
  648. assert all(value >= 0 for value in elapsed_values)
  649. assert elapsed_values == sorted(elapsed_values)
  650. assert next(
  651. message
  652. for message in turn_audits
  653. if message["event"] == "event_agent_request"
  654. )["details"]["turn_elapsed_ms"] >= 0
  655. @pytest.mark.asyncio
  656. async def test_runtime_audit_includes_model_params_prompts_results_and_usage():
  657. request = DebugRunRequest(
  658. user_message="debug this",
  659. system_prompts=["You are a debugger."],
  660. pre_messages=[],
  661. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  662. event_agent=EventAgentParams(
  663. model="event-model",
  664. temperature=0.4,
  665. max_tokens=80,
  666. enabled_tools=["handoff_note"],
  667. max_event_loops=1,
  668. ),
  669. )
  670. runtime = DebugRuntime(EventRoundStatsChatClient())
  671. outputs = [message async for message in runtime.run(request)]
  672. audits = [message for message in outputs if message["type"] == "audit"]
  673. chat_request = next(
  674. message for message in audits if message["event"] == "chat_agent_request"
  675. )
  676. assert chat_request["details"]["agent"] == "chat_agent"
  677. assert chat_request["details"]["params"]["model"] == "chat-model"
  678. assert chat_request["details"]["params"]["temperature"] == 0.1
  679. assert chat_request["details"]["params"]["max_tokens"] == 200
  680. assert chat_request["details"]["tools"] == []
  681. assert chat_request["details"]["messages"][0] == {
  682. "role": "system",
  683. "content": "You are a debugger.",
  684. "name": None,
  685. "tool_call_id": None,
  686. "tool_calls": [],
  687. }
  688. chat_response = next(
  689. message for message in audits if message["event"] == "chat_agent_response"
  690. )
  691. assert chat_response["details"]["event_names"] == ["handoff_note"]
  692. assert chat_response["details"]["raw_chunks"] == [
  693. {
  694. "choices": [
  695. {
  696. "delta": {"content": "<agent_events>handoff_note</agent_events>"},
  697. "finish_reason": None,
  698. }
  699. ]
  700. }
  701. ]
  702. assert chat_response["details"]["usage"] == {
  703. "prompt_tokens": 3,
  704. "completion_tokens": 0,
  705. "total_tokens": 3,
  706. "cached_tokens": 0,
  707. }
  708. event_request = next(
  709. message for message in audits if message["event"] == "event_agent_request"
  710. )
  711. assert event_request["details"]["agent"] == "event_agent"
  712. assert event_request["details"]["params"]["model"] == "event-model"
  713. assert event_request["details"]["params"]["temperature"] == 0.4
  714. assert event_request["details"]["events"][0]["name"] == "handoff_note"
  715. assert event_request["details"]["tools"][0]["function"]["name"] == "handoff_note"
  716. assert "Authorization" not in str(event_request["details"])
  717. event_response = next(
  718. message for message in audits if message["event"] == "event_agent_response"
  719. )
  720. assert event_response["details"]["replies"][0]["role"] == "tool"
  721. assert '"tool": "handoff_note"' in event_response["details"]["replies"][0]["content"]
  722. assert event_response["details"]["raw_model_chunks"][0]["event_name"] == "handoff_note"
  723. assert event_response["details"]["raw_model_chunks"][0]["chunks"] == []
  724. @pytest.mark.asyncio
  725. async def test_runtime_round_started_separates_available_events_from_round_budget():
  726. request = DebugRunRequest(
  727. user_message="debug this",
  728. system_prompts=[],
  729. pre_messages=[],
  730. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  731. event_agent=EventAgentParams(
  732. enabled_tools=["handoff_note"],
  733. max_event_loops=1,
  734. ),
  735. )
  736. runtime = DebugRuntime(EventRoundStatsChatClient())
  737. outputs = [message async for message in runtime.run(request)]
  738. round_starts = [
  739. message
  740. for message in outputs
  741. if message.get("event") == "chat_round_started"
  742. ]
  743. assert round_starts[0]["details"]["events_enabled"] == ["handoff_note"]
  744. assert round_starts[0]["details"]["configured_events"] == ["handoff_note"]
  745. assert round_starts[0]["details"]["event_generation_enabled"] is True
  746. assert round_starts[0]["details"]["event_prompt_events"] == ["handoff_note"]
  747. assert round_starts[1]["details"]["events_enabled"] == []
  748. assert round_starts[1]["details"]["configured_events"] == ["handoff_note"]
  749. assert round_starts[1]["details"]["event_generation_enabled"] is False
  750. assert round_starts[1]["details"]["event_prompt_events"] == []
  751. @pytest.mark.asyncio
  752. async def test_runtime_event_agent_history_excludes_chat_agent_system_context():
  753. request = DebugRunRequest(
  754. user_message="debug this",
  755. system_prompts=["ChatAgent root prompt."],
  756. pre_messages=[
  757. ChatMessage(role="system", content="ChatAgent pre system."),
  758. ChatMessage(role="user", content="earlier user"),
  759. ChatMessage(role="assistant", content="earlier assistant"),
  760. ],
  761. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  762. event_agent=EventAgentParams(
  763. model="event-model",
  764. enabled_tools=["handoff_note"],
  765. max_event_loops=1,
  766. ),
  767. )
  768. client = ContextBoundaryChatClient()
  769. runtime = DebugRuntime(client)
  770. outputs = [message async for message in runtime.run(request)]
  771. event_request = next(
  772. message for message in outputs if message.get("event") == "event_agent_request"
  773. )
  774. assert [
  775. (message["role"], message["content"])
  776. for message in event_request["details"]["history"]
  777. ] == [
  778. ("user", "earlier user"),
  779. ("assistant", "earlier assistant"),
  780. ("user", "debug this"),
  781. ("assistant", "I will check."),
  782. ]
  783. assert not any(
  784. message["role"] == "system"
  785. for message in event_request["details"]["history"]
  786. )
  787. assert client.event_agent_messages == []
  788. @pytest.mark.asyncio
  789. async def test_runtime_session_event_agent_history_excludes_internal_replies():
  790. request = DebugRunRequest(
  791. user_message="debug this",
  792. system_prompts=["ChatAgent root prompt."],
  793. pre_messages=[ChatMessage(role="assistant", content="prior answer")],
  794. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  795. event_agent=EventAgentParams(
  796. model="event-model",
  797. enabled_tools=["handoff_note"],
  798. max_event_loops=1,
  799. ),
  800. )
  801. runtime = DebugRuntime(ContextBoundaryChatClient())
  802. queues = runtime.start_session(request)
  803. outputs: list[dict[str, Any]] = []
  804. while not any(message["type"] == "turn_completed" for message in outputs):
  805. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  806. await runtime.aclose()
  807. event_request = next(
  808. message for message in outputs if message.get("event") == "event_agent_request"
  809. )
  810. assert [
  811. (message["role"], message["content"], message["name"])
  812. for message in event_request["details"]["history"]
  813. ] == [
  814. ("assistant", "prior answer", None),
  815. ("user", "debug this", None),
  816. ("assistant", "I will check.", None),
  817. ]
  818. assert not any(
  819. message["name"] == "event_agent"
  820. for message in event_request["details"]["history"]
  821. )
  822. @pytest.mark.asyncio
  823. async def test_runtime_outputs_event_as_soon_as_chat_stream_detects_it():
  824. RuntimeQueues = _runtime_queues_class()
  825. queues = RuntimeQueues()
  826. client = SlowAfterEventChatClient()
  827. runtime = DebugRuntime(client, queues=queues)
  828. request = DebugRunRequest(
  829. user_message="debug this",
  830. system_prompts=[],
  831. pre_messages=[],
  832. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  833. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  834. )
  835. runtime.start(request)
  836. try:
  837. assert await _next_non_audit_from_queue(queues) == {"type": "session_started"}
  838. assert await _next_non_audit_from_queue(queues) == {
  839. "type": "message_delta",
  840. "content": "Need event.",
  841. }
  842. await asyncio.wait_for(client.event_seen.wait(), timeout=1)
  843. event_message = await asyncio.wait_for(
  844. _next_non_audit_from_queue(queues),
  845. timeout=0.2,
  846. )
  847. assert event_message["type"] == "event"
  848. assert event_message["event"]["name"] == "handoff_note"
  849. finally:
  850. client.release_stream.set()
  851. await runtime.aclose()
  852. @pytest.mark.asyncio
  853. async def test_runtime_batches_round_events_before_continuing_chat_agent():
  854. def resolve_from_history(
  855. event: ToolCallEvent,
  856. context: ToolExecutionContext,
  857. ) -> dict[str, Any]:
  858. return {"message": context.history[-1].content, "event": event.name}
  859. registry = ToolRegistry(
  860. [
  861. ToolDefinition(
  862. name="handoff_note",
  863. description="Send a handoff note.",
  864. parameters={
  865. "type": "object",
  866. "properties": {"message": {"type": "string"}},
  867. "required": ["message"],
  868. },
  869. handler=lambda event: {
  870. "tool": event.name,
  871. "message": event.arguments["message"],
  872. },
  873. argument_resolver=resolve_from_history,
  874. ),
  875. ToolDefinition(
  876. name="audit_note",
  877. description="Send an audit note.",
  878. parameters={
  879. "type": "object",
  880. "properties": {"message": {"type": "string"}},
  881. "required": ["message"],
  882. },
  883. handler=lambda event: {
  884. "tool": event.name,
  885. "message": event.arguments["message"],
  886. },
  887. argument_resolver=resolve_from_history,
  888. ),
  889. ]
  890. )
  891. request = DebugRunRequest(
  892. user_message="debug this",
  893. system_prompts=[],
  894. pre_messages=[],
  895. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  896. event_agent=EventAgentParams(
  897. enabled_tools=["handoff_note", "audit_note"],
  898. max_event_loops=2,
  899. ),
  900. )
  901. client = MultiEventChatClient()
  902. runtime = DebugRuntime(client, registry=registry)
  903. outputs = [message async for message in runtime.run(request)]
  904. assert _message_types(outputs) == [
  905. "session_started",
  906. "message_delta",
  907. "event",
  908. "event",
  909. "tool_result",
  910. "tool_result",
  911. "round_stats",
  912. "message_delta",
  913. "round_stats",
  914. "done",
  915. ]
  916. assert client.calls == 2
  917. assert [message.role for message in client.second_call_messages] == [
  918. "system",
  919. "user",
  920. "assistant",
  921. "user",
  922. ]
  923. assistant_message = client.second_call_messages[2]
  924. assert assistant_message.content == "Checking events."
  925. assert not any(message.role == "tool" for message in client.second_call_messages)
  926. assert client.second_call_messages[-1].content == (
  927. "EventAgent results:\n"
  928. '{"tool": "handoff_note", "message": "Checking events."}\n'
  929. '{"tool": "audit_note", "message": "Checking events."}'
  930. )
  931. assert client.second_call_messages[-1].name == "event_agent"
  932. @pytest.mark.asyncio
  933. async def test_runtime_start_returns_queues_for_downstream_output_consumer():
  934. request = DebugRunRequest(
  935. user_message="debug this",
  936. system_prompts=[],
  937. pre_messages=[],
  938. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  939. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  940. )
  941. runtime = DebugRuntime(RoundStatsChatClient())
  942. queues = runtime.start(request)
  943. outputs: list[dict[str, Any]] = []
  944. while True:
  945. message = await asyncio.wait_for(queues.output.get(), timeout=1)
  946. outputs.append(message)
  947. if message["type"] == "done":
  948. break
  949. assert _message_types(outputs) == [
  950. "session_started",
  951. "message_delta",
  952. "usage",
  953. "round_stats",
  954. "done",
  955. ]
  956. @pytest.mark.asyncio
  957. async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
  958. request = DebugRunRequest(
  959. user_message="debug this",
  960. system_prompts=[],
  961. pre_messages=[],
  962. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  963. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  964. )
  965. client = EventLoopLimitChatClient()
  966. runtime = DebugRuntime(client)
  967. outputs = [message async for message in runtime.run(request)]
  968. business_outputs = _without_audit(outputs)
  969. assert client.calls == 2
  970. assert client.tools_by_call[0] == []
  971. assert client.tools_by_call[1] == []
  972. assert "Available events:" in client.messages_by_call[0][0].content
  973. assert not any(
  974. "Available events:" in message.content
  975. for message in client.messages_by_call[1]
  976. if message.role == "system"
  977. )
  978. assert _message_types(outputs) == [
  979. "session_started",
  980. "event",
  981. "tool_result",
  982. "round_stats",
  983. "message_delta",
  984. "round_stats",
  985. "done",
  986. ]
  987. assert business_outputs[4]["content"] == "final after event limit"
  988. @pytest.mark.asyncio
  989. async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
  990. RuntimeQueues = _runtime_queues_class()
  991. queues = RuntimeQueues()
  992. request = DebugRunRequest(
  993. user_message="debug this",
  994. system_prompts=[],
  995. pre_messages=[],
  996. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  997. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  998. )
  999. client = StrictHistoryChatClient()
  1000. runtime = DebugRuntime(client, queues=queues)
  1001. stream = runtime.run(request)
  1002. assert await _next_non_audit(stream) == {"type": "session_started"}
  1003. event_message = await _next_non_audit(stream)
  1004. assert event_message["type"] == "event"
  1005. await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
  1006. remaining = [message async for message in stream]
  1007. assert remaining[-1] == {"type": "done"}
  1008. assert [message.role for message in client.second_call_messages] == [
  1009. "system",
  1010. "user",
  1011. "assistant",
  1012. "user",
  1013. "user",
  1014. ]
  1015. assert not any(message.role == "tool" for message in client.second_call_messages)
  1016. assert client.second_call_messages[3].content.startswith("EventAgent results:\n")
  1017. assert client.second_call_messages[3].name == "event_agent"
  1018. assert client.second_call_messages[4].content == "follow-up while tool runs"
  1019. @pytest.mark.asyncio
  1020. async def test_runtime_continues_when_event_agent_tool_handler_raises():
  1021. def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
  1022. raise RuntimeError("boom")
  1023. registry = ToolRegistry(
  1024. [
  1025. ToolDefinition(
  1026. name="handoff_note",
  1027. description="Broken handoff tool.",
  1028. parameters={"type": "object"},
  1029. handler=fail_tool,
  1030. )
  1031. ]
  1032. )
  1033. request = DebugRunRequest(
  1034. user_message="debug this",
  1035. system_prompts=[],
  1036. pre_messages=[],
  1037. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1038. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1039. )
  1040. runtime = DebugRuntime(FakeChatClient(), registry=registry)
  1041. outputs = await asyncio.wait_for(
  1042. _collect_outputs(runtime.run(request)),
  1043. timeout=1,
  1044. )
  1045. business_outputs = _without_audit(outputs)
  1046. assert _message_types(outputs) == [
  1047. "session_started",
  1048. "event",
  1049. "tool_result",
  1050. "round_stats",
  1051. "message_delta",
  1052. "round_stats",
  1053. "done",
  1054. ]
  1055. assert json.loads(business_outputs[2]["message"]["content"]) == {
  1056. "tool": "handoff_note",
  1057. "error": "tool handler failed: boom",
  1058. }
  1059. @pytest.mark.asyncio
  1060. async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
  1061. RuntimeQueues = _runtime_queues_class()
  1062. queue_log: list[tuple[str, str, str]] = []
  1063. queues = RuntimeQueues(
  1064. input=RecordingQueue("input", queue_log),
  1065. output=RecordingQueue("output", queue_log),
  1066. events=RecordingQueue("events", queue_log),
  1067. )
  1068. request = DebugRunRequest(
  1069. user_message="debug this",
  1070. system_prompts=["You are a debugger."],
  1071. pre_messages=[],
  1072. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1073. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  1074. )
  1075. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  1076. outputs = [message async for message in runtime.run(request)]
  1077. assert _message_types(outputs) == [
  1078. "session_started",
  1079. "event",
  1080. "tool_result",
  1081. "round_stats",
  1082. "message_delta",
  1083. "round_stats",
  1084. "done",
  1085. ]
  1086. assert queue_log.index(("input", "put", "user")) < queue_log.index(
  1087. ("input", "get", "user")
  1088. )
  1089. assert queue_log.index(
  1090. ("events", "put", "event_request:handoff_note:call_1")
  1091. ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
  1092. assert queue_log.index(
  1093. ("events", "get", "event_request:handoff_note:call_1")
  1094. ) < queue_log.index(("input", "put", "tool:call_1"))
  1095. assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
  1096. ("input", "get", "tool:call_1")
  1097. )
  1098. @pytest.mark.asyncio
  1099. async def test_runtime_run_consumes_output_queue_in_stream_order():
  1100. RuntimeQueues = _runtime_queues_class()
  1101. queue_log: list[tuple[str, str, str]] = []
  1102. queues = RuntimeQueues(
  1103. input=RecordingQueue("input", queue_log),
  1104. output=RecordingQueue("output", queue_log),
  1105. events=RecordingQueue("events", queue_log),
  1106. )
  1107. request = DebugRunRequest(
  1108. user_message="debug this",
  1109. system_prompts=["You are a debugger."],
  1110. pre_messages=[],
  1111. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1112. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  1113. )
  1114. runtime = DebugRuntime(FakeChatClient(), queues=queues)
  1115. outputs = [message async for message in runtime.run(request)]
  1116. assert _message_types(outputs) == [
  1117. "session_started",
  1118. "event",
  1119. "tool_result",
  1120. "round_stats",
  1121. "message_delta",
  1122. "round_stats",
  1123. "done",
  1124. ]
  1125. output_puts = [
  1126. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
  1127. ]
  1128. output_gets = [
  1129. entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
  1130. ]
  1131. assert [message for message in output_puts if message != "output:audit"] == [
  1132. "output:session_started",
  1133. "output:event",
  1134. "output:tool_result",
  1135. "output:round_stats",
  1136. "output:message_delta",
  1137. "output:round_stats",
  1138. "output:done",
  1139. ]
  1140. assert output_gets == output_puts
  1141. @pytest.mark.asyncio
  1142. async def test_runtime_continues_with_event_summary_without_tool_call_history():
  1143. request = DebugRunRequest(
  1144. user_message="debug this",
  1145. system_prompts=["You are a debugger."],
  1146. pre_messages=[],
  1147. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1148. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1149. )
  1150. client = StrictHistoryChatClient()
  1151. runtime = DebugRuntime(client)
  1152. outputs = [message async for message in runtime.run(request)]
  1153. assert client.calls == 2
  1154. assert [message.role for message in client.second_call_messages] == [
  1155. "system",
  1156. "system",
  1157. "user",
  1158. "assistant",
  1159. "user",
  1160. ]
  1161. assistant_message = client.second_call_messages[3]
  1162. assert assistant_message.content == ""
  1163. assert "handoff_note" in client.second_call_messages[1].content
  1164. assert not any(message.role == "tool" for message in client.second_call_messages)
  1165. assert client.second_call_messages[4].content.startswith("EventAgent results:\n")
  1166. assert client.second_call_messages[4].name == "event_agent"
  1167. assert outputs[-1] == {"type": "done"}
  1168. @pytest.mark.asyncio
  1169. async def test_runtime_passes_event_catalog_system_message_without_chat_tools():
  1170. registry = ToolRegistry(
  1171. [
  1172. ToolDefinition(
  1173. name="handoff_note",
  1174. description="Registry-owned handoff tool.",
  1175. parameters={
  1176. "type": "object",
  1177. "properties": {
  1178. "message": {"type": "string"},
  1179. "priority": {"type": "number"},
  1180. },
  1181. "required": ["message"],
  1182. },
  1183. handler=lambda event: {"tool": event.name, "message": "handled"},
  1184. )
  1185. ]
  1186. )
  1187. request = DebugRunRequest(
  1188. user_message="debug this",
  1189. system_prompts=[],
  1190. pre_messages=[],
  1191. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1192. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1193. )
  1194. client = ToolCapturingChatClient()
  1195. runtime = DebugRuntime(client, registry=registry)
  1196. outputs = [message async for message in runtime.run(request)]
  1197. assert outputs[-1] == {"type": "done"}
  1198. assert client.tools == []
  1199. assert client.messages[0].role == "system"
  1200. assert "Available events:" in client.messages[0].content
  1201. assert "- handoff_note: Registry-owned handoff tool." in client.messages[0].content
  1202. assert "message" not in client.messages[0].content
  1203. assert "priority" not in client.messages[0].content
  1204. @pytest.mark.asyncio
  1205. async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
  1206. request = DebugRunRequest(
  1207. user_message="debug this",
  1208. system_prompts=[],
  1209. pre_messages=[],
  1210. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1211. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1212. )
  1213. ticks = iter(
  1214. [
  1215. 0.9,
  1216. 0.91,
  1217. 1.0,
  1218. 1.01,
  1219. 1.02,
  1220. 1.123,
  1221. 1.2,
  1222. 1.25,
  1223. 1.3,
  1224. 1.456,
  1225. 1.7,
  1226. 1.8,
  1227. ]
  1228. )
  1229. runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
  1230. outputs = [message async for message in runtime.run(request)]
  1231. business_outputs = _without_audit(outputs)
  1232. assert _message_types(outputs) == [
  1233. "session_started",
  1234. "message_delta",
  1235. "usage",
  1236. "round_stats",
  1237. "done",
  1238. ]
  1239. assert business_outputs[3] == {
  1240. "type": "round_stats",
  1241. "round_index": 1,
  1242. "ttft_ms": 123,
  1243. "elapsed_ms": 456,
  1244. "prompt_tokens": 10,
  1245. "completion_tokens": 20,
  1246. "total_tokens": 30,
  1247. "cached_tokens": 5,
  1248. "had_event": False,
  1249. }
  1250. @pytest.mark.asyncio
  1251. async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
  1252. request = DebugRunRequest(
  1253. user_message="debug this",
  1254. system_prompts=[],
  1255. pre_messages=[],
  1256. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1257. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  1258. )
  1259. ticks = iter(
  1260. [
  1261. 1.9,
  1262. 1.91,
  1263. 2.0,
  1264. 2.01,
  1265. 2.02,
  1266. 2.03,
  1267. 2.04,
  1268. 2.05,
  1269. 2.06,
  1270. 2.07,
  1271. 2.25,
  1272. 2.26,
  1273. 3.0,
  1274. 3.01,
  1275. 3.02,
  1276. 3.05,
  1277. 3.1,
  1278. 3.15,
  1279. 3.18,
  1280. 3.2,
  1281. 3.23,
  1282. 3.24,
  1283. 3.25,
  1284. 3.26,
  1285. ]
  1286. )
  1287. client = EventRoundStatsChatClient()
  1288. runtime = DebugRuntime(client, clock=lambda: next(ticks))
  1289. outputs = [message async for message in runtime.run(request)]
  1290. stats = [message for message in outputs if message["type"] == "round_stats"]
  1291. assert client.calls == 2
  1292. assert stats == [
  1293. {
  1294. "type": "round_stats",
  1295. "round_index": 1,
  1296. "ttft_ms": None,
  1297. "elapsed_ms": 250,
  1298. "prompt_tokens": 3,
  1299. "completion_tokens": 0,
  1300. "total_tokens": 3,
  1301. "cached_tokens": 0,
  1302. "had_event": True,
  1303. },
  1304. {
  1305. "type": "round_stats",
  1306. "round_index": 2,
  1307. "ttft_ms": 50,
  1308. "elapsed_ms": 200,
  1309. "prompt_tokens": 4,
  1310. "completion_tokens": 6,
  1311. "total_tokens": 10,
  1312. "cached_tokens": 0,
  1313. "had_event": False,
  1314. },
  1315. ]
  1316. @pytest.mark.asyncio
  1317. async def test_runtime_session_resets_event_budget_for_each_user_turn():
  1318. request = DebugRunRequest(
  1319. user_message="first turn",
  1320. system_prompts=[],
  1321. pre_messages=[],
  1322. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  1323. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
  1324. )
  1325. client = TwoTurnSessionChatClient()
  1326. runtime = DebugRuntime(client)
  1327. queues = runtime.start_session(request)
  1328. outputs: list[dict[str, Any]] = []
  1329. while len([message for message in outputs if message["type"] == "turn_completed"]) < 1:
  1330. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1331. await queues.input.put(ChatMessage(role="user", content="second turn"))
  1332. while len([message for message in outputs if message["type"] == "turn_completed"]) < 2:
  1333. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1334. await runtime.aclose()
  1335. business_types = [message["type"] for message in _without_audit(outputs)]
  1336. assert business_types.count("turn_started") == 2
  1337. assert business_types.count("turn_completed") == 2
  1338. assert client.calls == 4
  1339. assert "Available events:" in client.messages_by_call[0][0].content
  1340. assert "Available events:" in client.messages_by_call[2][0].content
  1341. @pytest.mark.asyncio
  1342. async def test_runtime_persists_session_turn_messages_audit_and_usage(tmp_path):
  1343. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1344. request = DebugRunRequest(
  1345. session_id="session-1",
  1346. user_message="persist this",
  1347. system_prompts=["You are a debugger."],
  1348. pre_messages=[],
  1349. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  1350. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1351. )
  1352. runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
  1353. queues = runtime.start_session(request)
  1354. outputs: list[dict[str, Any]] = []
  1355. while not any(message["type"] == "turn_completed" for message in outputs):
  1356. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1357. await runtime.aclose()
  1358. assert _without_audit(outputs)[0] == {
  1359. "type": "session_started",
  1360. "session_id": "session-1",
  1361. }
  1362. assert store.get_session("session-1")["config"]["chat_agent"]["model"] == "chat-model"
  1363. assert [
  1364. (message["turn_index"], message["role"], message["content"])
  1365. for message in store.list_messages("session-1")
  1366. ] == [
  1367. (1, "user", "persist this"),
  1368. (1, "assistant", "hello"),
  1369. ]
  1370. audit_events = [
  1371. audit["event"]
  1372. for audit in store.list_audit_logs("session-1")
  1373. ]
  1374. assert "session_started" in audit_events
  1375. assert "chat_agent_request" in audit_events
  1376. assert "turn_completed" in audit_events
  1377. usage = store.usage_summary("session-1")
  1378. assert usage["calls"][0]["total_tokens"] == 30
  1379. assert usage["turns"][0]["turn_index"] == 1
  1380. assert usage["session"]["total_tokens"] == 30
  1381. @pytest.mark.asyncio
  1382. async def test_runtime_continues_persisted_turn_indexes_for_existing_session(tmp_path):
  1383. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1384. session_id = store.create_session(title="existing", config={})
  1385. store.start_turn(session_id, turn_index=1, user_message="old")
  1386. request = DebugRunRequest(
  1387. session_id=session_id,
  1388. user_message="new",
  1389. system_prompts=[],
  1390. pre_messages=[],
  1391. chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
  1392. event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
  1393. )
  1394. runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
  1395. queues = runtime.start_session(request)
  1396. outputs: list[dict[str, Any]] = []
  1397. while not any(message["type"] == "turn_completed" for message in outputs):
  1398. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1399. await runtime.aclose()
  1400. assert [
  1401. (message["turn_index"], message["content"])
  1402. for message in store.list_messages(session_id)
  1403. ] == [
  1404. (2, "new"),
  1405. (2, "hello"),
  1406. ]
  1407. class ScriptedChatClient:
  1408. def __init__(self, rounds: list[list[StreamItem]]) -> None:
  1409. self.rounds = rounds
  1410. self.calls = 0
  1411. self.messages_by_call: list[list[ChatMessage]] = []
  1412. self.tools_by_call: list[list[dict[str, Any]]] = []
  1413. async def stream_chat(
  1414. self,
  1415. messages: list[ChatMessage],
  1416. tools: list[dict],
  1417. params: AgentParams,
  1418. tool_choice: dict[str, Any] | None = None,
  1419. ) -> AsyncIterator[StreamItem]:
  1420. self.messages_by_call.append(list(messages))
  1421. self.tools_by_call.append(list(tools))
  1422. round_items = self.rounds[self.calls]
  1423. self.calls += 1
  1424. for item in round_items:
  1425. yield item
  1426. def _direct_request(
  1427. *,
  1428. enabled_tools: list[str],
  1429. max_event_loops: int = 1,
  1430. session_id: str | None = None,
  1431. ) -> DebugRunRequest:
  1432. return DebugRunRequest(
  1433. session_id=session_id,
  1434. user_message="debug direct tools",
  1435. system_prompts=["You are a debugger."],
  1436. pre_messages=[],
  1437. chat_agent=AgentParams(model="chat-model"),
  1438. event_agent=EventAgentParams(
  1439. enabled_tools=enabled_tools,
  1440. max_event_loops=max_event_loops,
  1441. ),
  1442. tool_invocation_mode="chat_agent_tools",
  1443. )
  1444. @pytest.mark.asyncio
  1445. async def test_direct_mode_uses_enabled_schemas_and_valid_ordered_provider_transcript():
  1446. executed: list[tuple[str, dict[str, Any], str]] = []
  1447. def handle(event: ToolCallEvent) -> dict[str, Any]:
  1448. executed.append((event.name, event.arguments, event.raw_arguments))
  1449. return {"tool": event.name, "value": event.arguments["value"]}
  1450. registry = ToolRegistry(
  1451. [
  1452. ToolDefinition(
  1453. name="first_tool",
  1454. description="First direct tool.",
  1455. parameters={
  1456. "type": "object",
  1457. "properties": {"value": {"type": "string"}},
  1458. "required": ["value"],
  1459. },
  1460. handler=handle,
  1461. ),
  1462. ToolDefinition(
  1463. name="second_tool",
  1464. description="Second direct tool.",
  1465. parameters={
  1466. "type": "object",
  1467. "properties": {"value": {"type": "string"}},
  1468. "required": ["value"],
  1469. },
  1470. handler=handle,
  1471. ),
  1472. ]
  1473. )
  1474. calls = [
  1475. ToolCallEvent(
  1476. id="provider-1",
  1477. name="first_tool",
  1478. arguments={"value": "one"},
  1479. raw_arguments='{"value":"one"}',
  1480. ),
  1481. ToolCallEvent(
  1482. id="provider-2",
  1483. name="second_tool",
  1484. arguments={"value": "two"},
  1485. raw_arguments='{"value":"two"}',
  1486. ),
  1487. ]
  1488. ignored_text_event = ToolCallEvent(
  1489. id="text-ignored",
  1490. name="first_tool",
  1491. arguments={},
  1492. raw_arguments="{}",
  1493. )
  1494. client = ScriptedChatClient(
  1495. [
  1496. [
  1497. StreamItem.message_delta("Visible before tools."),
  1498. StreamItem.provider_tool_call(calls[0]),
  1499. StreamItem.text_event(ignored_text_event),
  1500. StreamItem.provider_tool_call(calls[1]),
  1501. ],
  1502. [StreamItem.message_delta("Final answer.")],
  1503. ]
  1504. )
  1505. outputs = await _collect_outputs(
  1506. DebugRuntime(client, registry=registry).run(
  1507. _direct_request(enabled_tools=["first_tool", "second_tool"])
  1508. )
  1509. )
  1510. assert [tool["function"]["name"] for tool in client.tools_by_call[0]] == [
  1511. "first_tool",
  1512. "second_tool",
  1513. ]
  1514. assert client.tools_by_call[1] == []
  1515. assert not any(
  1516. "Available events:" in message.content
  1517. for message in client.messages_by_call[0]
  1518. if message.role == "system"
  1519. )
  1520. transcript = client.messages_by_call[1]
  1521. assert [message.role for message in transcript[-4:]] == [
  1522. "user",
  1523. "assistant",
  1524. "tool",
  1525. "tool",
  1526. ]
  1527. assert transcript[-3].content == "Visible before tools."
  1528. assert transcript[-3].tool_calls == calls
  1529. assert [message.tool_call_id for message in transcript[-2:]] == [
  1530. "provider-1",
  1531. "provider-2",
  1532. ]
  1533. assert [message.name for message in transcript[-2:]] == [
  1534. "first_tool",
  1535. "second_tool",
  1536. ]
  1537. assert executed == [
  1538. ("first_tool", {"value": "one"}, '{"value":"one"}'),
  1539. ("second_tool", {"value": "two"}, '{"value":"two"}'),
  1540. ]
  1541. emitted_events = [
  1542. message["event"] for message in outputs if message["type"] == "event"
  1543. ]
  1544. assert emitted_events == [call.model_dump() for call in calls]
  1545. assert "event_agent_request" not in [
  1546. message.get("event") for message in outputs if message["type"] == "audit"
  1547. ]
  1548. assert not any(
  1549. message.get("event", {}).get("id") == "text-ignored"
  1550. for message in outputs
  1551. if message["type"] == "event"
  1552. )
  1553. assert all(
  1554. message["details"]["tool_invocation_mode"] == "chat_agent_tools"
  1555. for message in outputs
  1556. if message["type"] == "audit"
  1557. and message["event"] in {"chat_round_started", "chat_agent_request"}
  1558. )
  1559. @pytest.mark.asyncio
  1560. async def test_direct_mode_does_not_execute_unknown_disabled_or_failed_handlers():
  1561. executed: list[str] = []
  1562. def disabled_handler(event: ToolCallEvent) -> dict[str, Any]:
  1563. executed.append(event.name)
  1564. return {"tool": event.name}
  1565. def failing_handler(event: ToolCallEvent) -> dict[str, Any]:
  1566. executed.append(event.name)
  1567. raise RuntimeError("direct boom")
  1568. registry = ToolRegistry(
  1569. [
  1570. ToolDefinition(
  1571. name="enabled_tool",
  1572. description="Enabled.",
  1573. parameters={"type": "object"},
  1574. handler=lambda event: {"tool": event.name},
  1575. ),
  1576. ToolDefinition(
  1577. name="disabled_tool",
  1578. description="Disabled.",
  1579. parameters={"type": "object"},
  1580. handler=disabled_handler,
  1581. ),
  1582. ToolDefinition(
  1583. name="failing_tool",
  1584. description="Fails.",
  1585. parameters={"type": "object"},
  1586. handler=failing_handler,
  1587. ),
  1588. ]
  1589. )
  1590. client = ScriptedChatClient(
  1591. [
  1592. [
  1593. StreamItem.provider_tool_call(
  1594. ToolCallEvent(
  1595. id="unknown-1",
  1596. name="unknown_tool",
  1597. arguments={"kept": True},
  1598. raw_arguments='{"kept":true}',
  1599. )
  1600. ),
  1601. StreamItem.provider_tool_call(
  1602. ToolCallEvent(
  1603. id="disabled-1",
  1604. name="disabled_tool",
  1605. arguments={},
  1606. raw_arguments="{}",
  1607. )
  1608. ),
  1609. StreamItem.provider_tool_call(
  1610. ToolCallEvent(
  1611. id="failed-1",
  1612. name="failing_tool",
  1613. arguments={},
  1614. raw_arguments="{}",
  1615. )
  1616. ),
  1617. ],
  1618. [StreamItem.message_delta("continued")],
  1619. ]
  1620. )
  1621. outputs = await _collect_outputs(
  1622. DebugRuntime(client, registry=registry).run(
  1623. _direct_request(enabled_tools=["enabled_tool", "failing_tool"])
  1624. )
  1625. )
  1626. assert [tool["function"]["name"] for tool in client.tools_by_call[0]] == [
  1627. "enabled_tool",
  1628. "failing_tool",
  1629. ]
  1630. assert executed == ["failing_tool"]
  1631. results = [
  1632. json.loads(message["message"]["content"])
  1633. for message in outputs
  1634. if message["type"] == "tool_result"
  1635. ]
  1636. assert results == [
  1637. {"tool": "unknown_tool", "error": "unknown tool"},
  1638. {"tool": "disabled_tool", "error": "tool disabled"},
  1639. {"tool": "failing_tool", "error": "tool handler failed: direct boom"},
  1640. ]
  1641. @pytest.mark.asyncio
  1642. async def test_direct_mode_ignores_provider_calls_after_event_budget_exhaustion():
  1643. executed: list[str] = []
  1644. registry = ToolRegistry(
  1645. [
  1646. ToolDefinition(
  1647. name="once_tool",
  1648. description="Run once.",
  1649. parameters={"type": "object"},
  1650. handler=lambda event: executed.append(event.id) or {"tool": event.name},
  1651. )
  1652. ]
  1653. )
  1654. client = ScriptedChatClient(
  1655. [
  1656. [
  1657. StreamItem.provider_tool_call(
  1658. ToolCallEvent(
  1659. id="accepted",
  1660. name="once_tool",
  1661. arguments={},
  1662. raw_arguments="{}",
  1663. )
  1664. )
  1665. ],
  1666. [
  1667. StreamItem.provider_tool_call(
  1668. ToolCallEvent(
  1669. id="ignored",
  1670. name="once_tool",
  1671. arguments={},
  1672. raw_arguments="{}",
  1673. )
  1674. ),
  1675. StreamItem.message_delta("budget exhausted"),
  1676. ],
  1677. ]
  1678. )
  1679. outputs = await _collect_outputs(
  1680. DebugRuntime(client, registry=registry).run(
  1681. _direct_request(enabled_tools=["once_tool"], max_event_loops=1)
  1682. )
  1683. )
  1684. assert client.tools_by_call == [
  1685. [registry.tool_schema("once_tool")],
  1686. [],
  1687. ]
  1688. assert executed == ["accepted"]
  1689. assert [
  1690. message["event"]["id"]
  1691. for message in outputs
  1692. if message["type"] == "event"
  1693. ] == ["accepted"]
  1694. assert outputs[-1] == {"type": "done"}
  1695. @pytest.mark.asyncio
  1696. async def test_direct_mode_session_turns_reset_budget_and_snapshot_mode(tmp_path):
  1697. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  1698. registry = ToolRegistry(
  1699. [
  1700. ToolDefinition(
  1701. name="turn_tool",
  1702. description="Per-turn tool.",
  1703. parameters={"type": "object"},
  1704. handler=lambda event: {"tool": event.name, "call_id": event.id},
  1705. )
  1706. ]
  1707. )
  1708. client = ScriptedChatClient(
  1709. [
  1710. [
  1711. StreamItem.provider_tool_call(
  1712. ToolCallEvent(
  1713. id="turn-1-call",
  1714. name="turn_tool",
  1715. arguments={},
  1716. raw_arguments="{}",
  1717. )
  1718. )
  1719. ],
  1720. [StreamItem.message_delta("turn one done")],
  1721. [
  1722. StreamItem.provider_tool_call(
  1723. ToolCallEvent(
  1724. id="turn-2-call",
  1725. name="turn_tool",
  1726. arguments={},
  1727. raw_arguments="{}",
  1728. )
  1729. )
  1730. ],
  1731. [StreamItem.message_delta("turn two done")],
  1732. ]
  1733. )
  1734. runtime = DebugRuntime(client, registry=registry, session_store=store)
  1735. request = _direct_request(
  1736. enabled_tools=["turn_tool"],
  1737. max_event_loops=1,
  1738. session_id="direct-session",
  1739. )
  1740. queues = runtime.start_session(request)
  1741. outputs: list[dict[str, Any]] = []
  1742. while sum(message["type"] == "turn_completed" for message in outputs) < 1:
  1743. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1744. await queues.input.put(ChatMessage(role="user", content="second direct turn"))
  1745. while sum(message["type"] == "turn_completed" for message in outputs) < 2:
  1746. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1747. await runtime.aclose()
  1748. assert [bool(tools) for tools in client.tools_by_call] == [True, False, True, False]
  1749. assert [message.role for message in client.messages_by_call[1][-3:]] == [
  1750. "user",
  1751. "assistant",
  1752. "tool",
  1753. ]
  1754. second_turn_transcript = client.messages_by_call[3]
  1755. assert not any(message.name == "event_agent" for message in second_turn_transcript)
  1756. assert [
  1757. message.tool_call_id
  1758. for message in second_turn_transcript
  1759. if message.role == "tool"
  1760. ] == ["turn-1-call", "turn-2-call"]
  1761. session = store.get_session("direct-session")
  1762. assert session is not None
  1763. assert session["config"]["tool_invocation_mode"] == "chat_agent_tools"
  1764. request_audits = [
  1765. audit
  1766. for audit in store.list_audit_logs("direct-session")
  1767. if audit["event"] == "chat_agent_request"
  1768. ]
  1769. assert request_audits
  1770. assert all(
  1771. audit["details"]["tool_invocation_mode"] == "chat_agent_tools"
  1772. for audit in request_audits
  1773. )
  1774. def _single_tool_registry(
  1775. handler: Any,
  1776. *,
  1777. name: str = "safe_tool",
  1778. ) -> ToolRegistry:
  1779. return ToolRegistry(
  1780. [
  1781. ToolDefinition(
  1782. name=name,
  1783. description="Tool used by execution-safety tests.",
  1784. parameters={"type": "object"},
  1785. handler=handler,
  1786. )
  1787. ]
  1788. )
  1789. def _tool_call(call_id: str, *, name: str = "safe_tool") -> ToolCallEvent:
  1790. return ToolCallEvent(
  1791. id=call_id,
  1792. name=name,
  1793. arguments={},
  1794. raw_arguments="{}",
  1795. )
  1796. @pytest.mark.asyncio
  1797. async def test_direct_mode_rejects_provider_call_id_reused_across_one_shot_rounds():
  1798. executed: list[str] = []
  1799. registry = _single_tool_registry(
  1800. lambda event: executed.append(event.id) or {"tool": event.name}
  1801. )
  1802. client = ScriptedChatClient(
  1803. [
  1804. [StreamItem.provider_tool_call(_tool_call("duplicate-call"))],
  1805. [StreamItem.provider_tool_call(_tool_call("duplicate-call"))],
  1806. ]
  1807. )
  1808. outputs = await _collect_outputs(
  1809. DebugRuntime(client, registry=registry).run(
  1810. _direct_request(enabled_tools=["safe_tool"], max_event_loops=2)
  1811. )
  1812. )
  1813. assert executed == ["duplicate-call"]
  1814. assert [
  1815. message["message"]["tool_call_id"]
  1816. for message in outputs
  1817. if message["type"] == "tool_result"
  1818. ] == ["duplicate-call"]
  1819. assert outputs[-1]["type"] == "error"
  1820. assert "duplicate provider tool-call ID: duplicate-call" in outputs[-1]["message"]
  1821. @pytest.mark.asyncio
  1822. async def test_direct_mode_rejects_provider_call_id_reused_in_session_rounds():
  1823. executed: list[str] = []
  1824. registry = _single_tool_registry(
  1825. lambda event: executed.append(event.id) or {"tool": event.name}
  1826. )
  1827. client = ScriptedChatClient(
  1828. [
  1829. [StreamItem.provider_tool_call(_tool_call("session-duplicate"))],
  1830. [StreamItem.provider_tool_call(_tool_call("session-duplicate"))],
  1831. ]
  1832. )
  1833. runtime = DebugRuntime(client, registry=registry)
  1834. queues = runtime.start_session(
  1835. _direct_request(enabled_tools=["safe_tool"], max_event_loops=2)
  1836. )
  1837. outputs: list[dict[str, Any]] = []
  1838. while not any(message["type"] == "error" for message in outputs):
  1839. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1840. await runtime.aclose()
  1841. assert executed == ["session-duplicate"]
  1842. assert [
  1843. message["message"]["tool_call_id"]
  1844. for message in outputs
  1845. if message["type"] == "tool_result"
  1846. ] == ["session-duplicate"]
  1847. assert "duplicate provider tool-call ID: session-duplicate" in outputs[-1][
  1848. "message"
  1849. ]
  1850. @pytest.mark.asyncio
  1851. async def test_direct_mode_rejects_duplicate_provider_call_ids_within_batch():
  1852. executed: list[str] = []
  1853. registry = _single_tool_registry(
  1854. lambda event: executed.append(event.id) or {"tool": event.name}
  1855. )
  1856. client = ScriptedChatClient(
  1857. [
  1858. [
  1859. StreamItem.provider_tool_call(_tool_call("same-batch")),
  1860. StreamItem.provider_tool_call(_tool_call("same-batch")),
  1861. ]
  1862. ]
  1863. )
  1864. outputs = await _collect_outputs(
  1865. DebugRuntime(client, registry=registry).run(
  1866. _direct_request(enabled_tools=["safe_tool"])
  1867. )
  1868. )
  1869. assert executed == []
  1870. assert outputs[-1]["type"] == "error"
  1871. assert "duplicate provider tool-call ID: same-batch" in outputs[-1]["message"]
  1872. @pytest.mark.asyncio
  1873. async def test_direct_mode_executes_batch_concurrently_but_replies_in_provider_order():
  1874. both_started = asyncio.Event()
  1875. first_release = asyncio.Event()
  1876. second_release = asyncio.Event()
  1877. second_finished = asyncio.Event()
  1878. started: list[str] = []
  1879. async def handler(event: ToolCallEvent) -> dict[str, Any]:
  1880. started.append(event.id)
  1881. if len(started) == 2:
  1882. both_started.set()
  1883. if event.id == "first-call":
  1884. await first_release.wait()
  1885. else:
  1886. await second_release.wait()
  1887. second_finished.set()
  1888. return {"tool": event.name, "call_id": event.id}
  1889. registry = _single_tool_registry(handler)
  1890. client = ScriptedChatClient(
  1891. [
  1892. [
  1893. StreamItem.provider_tool_call(_tool_call("first-call")),
  1894. StreamItem.provider_tool_call(_tool_call("second-call")),
  1895. ],
  1896. [StreamItem.message_delta("finished")],
  1897. ]
  1898. )
  1899. runtime = DebugRuntime(client, registry=registry)
  1900. run_task = asyncio.create_task(
  1901. _collect_outputs(
  1902. runtime.run(_direct_request(enabled_tools=["safe_tool"]))
  1903. )
  1904. )
  1905. overlapped = False
  1906. try:
  1907. try:
  1908. await asyncio.wait_for(both_started.wait(), timeout=1)
  1909. except TimeoutError:
  1910. pass
  1911. else:
  1912. overlapped = True
  1913. second_release.set()
  1914. await asyncio.wait_for(second_finished.wait(), timeout=1)
  1915. assert not run_task.done()
  1916. finally:
  1917. first_release.set()
  1918. second_release.set()
  1919. outputs = await asyncio.wait_for(run_task, timeout=1)
  1920. assert overlapped is True
  1921. assert started == ["first-call", "second-call"]
  1922. assert [
  1923. message["message"]["tool_call_id"]
  1924. for message in outputs
  1925. if message["type"] == "tool_result"
  1926. ] == ["first-call", "second-call"]
  1927. def _dual_budget_request() -> DebugRunRequest:
  1928. return DebugRunRequest(
  1929. user_message="dual budget",
  1930. system_prompts=[],
  1931. pre_messages=[],
  1932. chat_agent=AgentParams(model="chat-model"),
  1933. event_agent=EventAgentParams(
  1934. enabled_tools=["safe_tool"],
  1935. max_event_loops=1,
  1936. ),
  1937. )
  1938. @pytest.mark.asyncio
  1939. async def test_dual_mode_ignores_text_event_after_one_shot_budget_exhaustion():
  1940. executed: list[str] = []
  1941. registry = _single_tool_registry(
  1942. lambda event: executed.append(event.id) or {"tool": event.name}
  1943. )
  1944. client = ScriptedChatClient(
  1945. [
  1946. [StreamItem.text_event(_tool_call("accepted-text-event"))],
  1947. [
  1948. StreamItem.text_event(_tool_call("ignored-text-event")),
  1949. StreamItem.message_delta("final dual answer"),
  1950. ],
  1951. ]
  1952. )
  1953. outputs = await _collect_outputs(
  1954. DebugRuntime(client, registry=registry).run(_dual_budget_request())
  1955. )
  1956. assert executed == ["accepted-text-event"]
  1957. assert [
  1958. message["event"]["id"]
  1959. for message in outputs
  1960. if message["type"] == "event"
  1961. ] == ["accepted-text-event"]
  1962. assert outputs[-1] == {"type": "done"}
  1963. @pytest.mark.asyncio
  1964. async def test_dual_mode_session_main_path_ignores_text_event_after_budget_exhaustion():
  1965. executed: list[str] = []
  1966. registry = _single_tool_registry(
  1967. lambda event: executed.append(event.id) or {"tool": event.name}
  1968. )
  1969. client = ScriptedChatClient(
  1970. [
  1971. [StreamItem.text_event(_tool_call("session-accepted"))],
  1972. [
  1973. StreamItem.text_event(_tool_call("session-ignored")),
  1974. StreamItem.message_delta("session final"),
  1975. ],
  1976. ]
  1977. )
  1978. runtime = DebugRuntime(client, registry=registry)
  1979. queues = runtime.start_session(_dual_budget_request())
  1980. outputs: list[dict[str, Any]] = []
  1981. while not any(
  1982. message["type"] in {"turn_completed", "error"} for message in outputs
  1983. ):
  1984. outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
  1985. await runtime.aclose()
  1986. assert executed == ["session-accepted"]
  1987. assert [
  1988. message["event"]["id"]
  1989. for message in outputs
  1990. if message["type"] == "event"
  1991. ] == ["session-accepted"]
  1992. assert outputs[-1]["type"] == "turn_completed"