runtime.py 42 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117
  1. import asyncio
  2. import json
  3. import logging
  4. import time
  5. from collections.abc import AsyncIterator, Callable
  6. from typing import Any, Protocol
  7. from agent_lab.application.contracts import AgentParams, DebugRunRequest
  8. from agent_lab.application.event_agent import EventAgent, EventAgentRequest
  9. from agent_lab.application.queues import RuntimeQueues
  10. from agent_lab.application.session_store import SessionStore
  11. from agent_lab.application.tools import ToolRegistry, build_default_tool_registry
  12. from agent_lab.domain.events import EVENT_BLOCK_START, ToolCallEvent
  13. from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
  14. logger = logging.getLogger(__name__)
  15. class ChatClient(Protocol):
  16. async def stream_chat(
  17. self,
  18. messages: list[ChatMessage],
  19. tools: list[dict[str, Any]],
  20. params: AgentParams,
  21. tool_choice: dict[str, Any] | None = None,
  22. ) -> AsyncIterator[StreamItem]:
  23. ...
  24. class DebugRuntime:
  25. def __init__(
  26. self,
  27. chat_client: ChatClient,
  28. queues: RuntimeQueues | None = None,
  29. registry: ToolRegistry | None = None,
  30. session_store: SessionStore | None = None,
  31. clock: Callable[[], float] = time.perf_counter,
  32. ) -> None:
  33. self.chat_client = chat_client
  34. self.queues = queues
  35. self.registry = registry or build_default_tool_registry()
  36. self.session_store = session_store
  37. self.clock = clock
  38. self._tasks: list[asyncio.Task[None]] = []
  39. def start(self, request: DebugRunRequest) -> RuntimeQueues:
  40. queues = self.queues or RuntimeQueues()
  41. task = asyncio.create_task(self._produce(request, queues))
  42. self._tasks.append(task)
  43. return queues
  44. def start_session(self, request: DebugRunRequest) -> RuntimeQueues:
  45. queues = self.queues or RuntimeQueues()
  46. task = asyncio.create_task(self._produce_session(request, queues))
  47. self._tasks.append(task)
  48. return queues
  49. async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]:
  50. queues = self.start(request)
  51. async for message in self.output_messages(queues):
  52. yield message
  53. await self._wait_for_tasks()
  54. async def output_messages(
  55. self,
  56. queues: RuntimeQueues,
  57. ) -> AsyncIterator[dict[str, Any]]:
  58. while True:
  59. message = await queues.output.get()
  60. yield message
  61. if message.get("type") in {"done", "error"}:
  62. break
  63. async def aclose(self) -> None:
  64. for task in self._tasks:
  65. if not task.done():
  66. task.cancel()
  67. await self._wait_for_tasks()
  68. async def _wait_for_tasks(self) -> None:
  69. if not self._tasks:
  70. return
  71. await asyncio.gather(*self._tasks, return_exceptions=True)
  72. self._tasks = [task for task in self._tasks if not task.done()]
  73. async def _produce(self, request: DebugRunRequest, queues: RuntimeQueues) -> None:
  74. event_worker = self._start_event_worker(request, queues)
  75. try:
  76. await self._run_chat_agent(request, queues)
  77. except Exception as exc:
  78. logger.exception("runtime session failed")
  79. await queues.output.put({"type": "error", "message": str(exc)})
  80. finally:
  81. if event_worker is not None:
  82. await queues.events.put(None)
  83. await asyncio.gather(event_worker, return_exceptions=True)
  84. async def _produce_session(
  85. self,
  86. request: DebugRunRequest,
  87. queues: RuntimeQueues,
  88. ) -> None:
  89. event_worker = self._start_event_worker(request, queues)
  90. try:
  91. await self._run_chat_session(request, queues)
  92. except asyncio.CancelledError:
  93. raise
  94. except Exception as exc:
  95. logger.exception("runtime session failed")
  96. await queues.output.put({"type": "error", "message": str(exc)})
  97. finally:
  98. if event_worker is not None:
  99. await queues.events.put(None)
  100. await asyncio.gather(event_worker, return_exceptions=True)
  101. def _start_event_worker(
  102. self,
  103. request: DebugRunRequest,
  104. queues: RuntimeQueues,
  105. ) -> asyncio.Task[None] | None:
  106. if request.tool_invocation_mode != "dual_agent":
  107. return None
  108. event_agent = EventAgent(
  109. request.event_agent.enabled_tools,
  110. registry=self.registry,
  111. chat_client=self.chat_client,
  112. params=request.event_agent,
  113. )
  114. return asyncio.create_task(self._consume_events(queues, event_agent))
  115. async def _run_chat_agent(
  116. self,
  117. request: DebugRunRequest,
  118. queues: RuntimeQueues,
  119. ) -> None:
  120. messages = self._build_initial_messages(request)
  121. turn_started_at = self.clock()
  122. await queues.input.put(ChatMessage(role="user", content=request.user_message))
  123. await queues.output.put({"type": "session_started"})
  124. await self._audit(
  125. queues,
  126. "session_started",
  127. turn_started_at=turn_started_at,
  128. tool_invocation_mode=request.tool_invocation_mode,
  129. )
  130. event_loops = 0
  131. round_index = 0
  132. while True:
  133. await self._drain_input(queues, messages)
  134. round_index += 1
  135. started_at = self.clock()
  136. ttft_ms: int | None = None
  137. prompt_tokens = 0
  138. completion_tokens = 0
  139. total_tokens = 0
  140. cached_tokens = 0
  141. saw_event = False
  142. assistant_content: list[str] = []
  143. events: list[ToolCallEvent] = []
  144. tool_replies: list[ChatMessage] = []
  145. message_stream_started = False
  146. message_delta_count = 0
  147. configured_events = request.event_agent.enabled_tools
  148. event_prompt_events = (
  149. request.event_agent.enabled_tools
  150. if event_loops < request.event_agent.max_event_loops
  151. else []
  152. )
  153. chat_tools = self._chat_tools_for_round(request, event_prompt_events)
  154. raw_chunks: list[dict[str, Any]] = []
  155. chat_messages = self._chat_messages_for_mode(
  156. request,
  157. messages,
  158. event_prompt_events,
  159. )
  160. await self._audit(
  161. queues,
  162. "chat_round_started",
  163. turn_started_at=turn_started_at,
  164. round_index=round_index,
  165. tool_invocation_mode=request.tool_invocation_mode,
  166. events_enabled=event_prompt_events,
  167. configured_events=configured_events,
  168. event_generation_enabled=bool(event_prompt_events),
  169. event_prompt_events=event_prompt_events,
  170. )
  171. await self._audit(
  172. queues,
  173. "chat_agent_request",
  174. turn_started_at=turn_started_at,
  175. agent="chat_agent",
  176. round_index=round_index,
  177. tool_invocation_mode=request.tool_invocation_mode,
  178. params=self._params_snapshot(request.chat_agent),
  179. messages=self._message_snapshots(chat_messages),
  180. tools=chat_tools,
  181. )
  182. async for item in self.chat_client.stream_chat(
  183. messages=chat_messages,
  184. tools=chat_tools,
  185. params=request.chat_agent,
  186. ):
  187. if item.kind == "raw_chunk" and item.raw_chunk is not None:
  188. raw_chunks.append(item.raw_chunk)
  189. continue
  190. if item.kind == "message_delta":
  191. if ttft_ms is None:
  192. ttft_ms = self._elapsed_ms(started_at)
  193. if not message_stream_started:
  194. message_stream_started = True
  195. await self._audit(
  196. queues,
  197. "chat_message_stream_started",
  198. turn_started_at=turn_started_at,
  199. agent="chat_agent",
  200. round_index=round_index,
  201. ttft_ms=ttft_ms,
  202. )
  203. message_delta_count += 1
  204. assistant_content.append(item.content or "")
  205. await queues.output.put(
  206. {"type": "message_delta", "content": item.content}
  207. )
  208. continue
  209. if item.kind == "usage" and item.usage is not None:
  210. prompt_tokens = item.usage.prompt_tokens
  211. completion_tokens = item.usage.completion_tokens
  212. total_tokens = item.usage.total_tokens
  213. cached_tokens = item.usage.cached_tokens
  214. await queues.output.put(
  215. {"type": "usage", "usage": item.usage.model_dump()}
  216. )
  217. continue
  218. event = self._accepted_chat_event(
  219. request,
  220. item,
  221. budget_available=event_loops < request.event_agent.max_event_loops,
  222. )
  223. if event is not None:
  224. saw_event = True
  225. events.append(event)
  226. await queues.output.put(
  227. {"type": "event", "event": event.model_dump()}
  228. )
  229. await self._audit(
  230. queues,
  231. "chat_event_detected",
  232. turn_started_at=turn_started_at,
  233. round_index=round_index,
  234. event_id=event.id,
  235. event_name=event.name,
  236. event_source=(
  237. "provider_resolved"
  238. if request.tool_invocation_mode == "chat_agent_tools"
  239. else "text_event"
  240. ),
  241. tool_invocation_mode=request.tool_invocation_mode,
  242. )
  243. if message_stream_started:
  244. await self._audit(
  245. queues,
  246. "chat_message_stream_finished",
  247. turn_started_at=turn_started_at,
  248. agent="chat_agent",
  249. round_index=round_index,
  250. delta_count=message_delta_count,
  251. content_length=len("".join(assistant_content)),
  252. )
  253. await self._audit(
  254. queues,
  255. "chat_agent_response",
  256. turn_started_at=turn_started_at,
  257. agent="chat_agent",
  258. round_index=round_index,
  259. tool_invocation_mode=request.tool_invocation_mode,
  260. content="".join(assistant_content),
  261. event_names=[event.name for event in events],
  262. events=self._event_snapshots(events),
  263. usage={
  264. "prompt_tokens": prompt_tokens,
  265. "completion_tokens": completion_tokens,
  266. "total_tokens": total_tokens,
  267. "cached_tokens": cached_tokens,
  268. },
  269. raw_chunks=raw_chunks,
  270. )
  271. if events and request.tool_invocation_mode == "chat_agent_tools":
  272. self._validate_provider_tool_call_ids(messages, events)
  273. if assistant_content or events:
  274. assistant_message = ChatMessage(
  275. role="assistant",
  276. content="".join(assistant_content),
  277. tool_calls=(
  278. events
  279. if request.tool_invocation_mode == "chat_agent_tools"
  280. else []
  281. ),
  282. )
  283. messages.append(assistant_message)
  284. if events:
  285. if request.tool_invocation_mode == "chat_agent_tools":
  286. tool_replies = await self._execute_provider_tools(
  287. events,
  288. enabled_names=request.event_agent.enabled_tools,
  289. )
  290. messages.extend(tool_replies)
  291. else:
  292. await queues.events.put(
  293. EventAgentRequest(
  294. events=events,
  295. history=self._event_agent_history(messages),
  296. system_prompt=request.event_agent.system_prompt,
  297. extra_body=request.event_agent.extra_body,
  298. turn_started_at=turn_started_at,
  299. )
  300. )
  301. tool_replies = await self._wait_for_tool_replies(
  302. queues,
  303. [event.id for event in events],
  304. )
  305. for reply in tool_replies:
  306. await queues.output.put(
  307. {"type": "tool_result", "message": reply.model_dump()}
  308. )
  309. if request.tool_invocation_mode == "chat_agent_tools":
  310. await self._audit(
  311. queues,
  312. "provider_tools_completed",
  313. turn_started_at=turn_started_at,
  314. round_index=round_index,
  315. tool_invocation_mode=request.tool_invocation_mode,
  316. event_source="provider_resolved",
  317. events=self._event_snapshots(events),
  318. result_count=len(tool_replies),
  319. )
  320. else:
  321. await self._audit(
  322. queues,
  323. "event_agent_completed",
  324. turn_started_at=turn_started_at,
  325. round_index=round_index,
  326. event_names=[event.name for event in events],
  327. result_count=len(tool_replies),
  328. result_summary="\n".join(
  329. reply.content for reply in tool_replies
  330. ),
  331. )
  332. elapsed_ms = self._elapsed_ms(started_at)
  333. await queues.output.put(
  334. {
  335. "type": "round_stats",
  336. "round_index": round_index,
  337. "ttft_ms": ttft_ms,
  338. "elapsed_ms": elapsed_ms,
  339. "prompt_tokens": prompt_tokens,
  340. "completion_tokens": completion_tokens,
  341. "total_tokens": total_tokens,
  342. "cached_tokens": cached_tokens,
  343. "had_event": saw_event,
  344. },
  345. )
  346. await self._audit(
  347. queues,
  348. "chat_round_finished",
  349. turn_started_at=turn_started_at,
  350. round_index=round_index,
  351. tool_invocation_mode=request.tool_invocation_mode,
  352. had_event=saw_event,
  353. elapsed_ms=elapsed_ms,
  354. )
  355. if not saw_event:
  356. break
  357. event_loops += 1
  358. await self._audit(
  359. queues,
  360. "session_finished",
  361. turn_started_at=turn_started_at,
  362. round_count=round_index,
  363. )
  364. await queues.output.put({"type": "done"})
  365. async def _run_chat_session(
  366. self,
  367. request: DebugRunRequest,
  368. queues: RuntimeQueues,
  369. ) -> None:
  370. session_id = self._ensure_persisted_session(request)
  371. messages = self._build_initial_messages(request)
  372. initial_turn_started_at = self.clock()
  373. session_message = {"type": "session_started"}
  374. if session_id is not None:
  375. session_message["session_id"] = session_id
  376. await queues.output.put(session_message)
  377. await self._audit(
  378. queues,
  379. "session_started",
  380. turn_started_at=initial_turn_started_at,
  381. session_id=session_id,
  382. tool_invocation_mode=request.tool_invocation_mode,
  383. )
  384. await queues.input.put(ChatMessage(role="user", content=request.user_message))
  385. turn_index = self._initial_turn_index(session_id) - 1
  386. next_turn_started_at: float | None = initial_turn_started_at
  387. while True:
  388. user_message = await queues.input.get()
  389. if user_message.role != "user" or user_message.name == "event_agent":
  390. messages.append(user_message)
  391. continue
  392. turn_index += 1
  393. turn_started_at = next_turn_started_at or self.clock()
  394. next_turn_started_at = None
  395. await self._run_chat_turn(
  396. request,
  397. queues,
  398. messages,
  399. user_message,
  400. turn_index,
  401. session_id,
  402. turn_started_at,
  403. )
  404. async def _run_chat_turn(
  405. self,
  406. request: DebugRunRequest,
  407. queues: RuntimeQueues,
  408. messages: list[ChatMessage],
  409. user_message: ChatMessage,
  410. turn_index: int,
  411. session_id: str | None = None,
  412. turn_started_at: float | None = None,
  413. ) -> None:
  414. turn_started_at = turn_started_at or self.clock()
  415. messages.append(user_message)
  416. self._start_persisted_turn(session_id, turn_index, user_message)
  417. await queues.output.put({"type": "turn_started", "turn_index": turn_index})
  418. await self._audit(
  419. queues,
  420. "turn_started",
  421. turn_started_at=turn_started_at,
  422. session_id=session_id,
  423. turn_index=turn_index,
  424. )
  425. event_loops = 0
  426. round_index = 0
  427. deferred_user_messages: list[ChatMessage] = []
  428. while True:
  429. if round_index:
  430. await self._drain_input(
  431. queues,
  432. messages,
  433. deferred_user_messages=deferred_user_messages,
  434. )
  435. round_index += 1
  436. started_at = self.clock()
  437. ttft_ms: int | None = None
  438. prompt_tokens = 0
  439. completion_tokens = 0
  440. total_tokens = 0
  441. cached_tokens = 0
  442. saw_event = False
  443. assistant_content: list[str] = []
  444. events: list[ToolCallEvent] = []
  445. message_stream_started = False
  446. message_delta_count = 0
  447. configured_events = request.event_agent.enabled_tools
  448. event_prompt_events = (
  449. request.event_agent.enabled_tools
  450. if event_loops < request.event_agent.max_event_loops
  451. else []
  452. )
  453. chat_tools = self._chat_tools_for_round(request, event_prompt_events)
  454. raw_chunks: list[dict[str, Any]] = []
  455. chat_messages = self._chat_messages_for_mode(
  456. request,
  457. messages,
  458. event_prompt_events,
  459. )
  460. await self._audit(
  461. queues,
  462. "chat_round_started",
  463. turn_started_at=turn_started_at,
  464. session_id=session_id,
  465. turn_index=turn_index,
  466. round_index=round_index,
  467. tool_invocation_mode=request.tool_invocation_mode,
  468. events_enabled=event_prompt_events,
  469. configured_events=configured_events,
  470. event_generation_enabled=bool(event_prompt_events),
  471. event_prompt_events=event_prompt_events,
  472. )
  473. await self._audit(
  474. queues,
  475. "chat_agent_request",
  476. turn_started_at=turn_started_at,
  477. session_id=session_id,
  478. agent="chat_agent",
  479. turn_index=turn_index,
  480. round_index=round_index,
  481. tool_invocation_mode=request.tool_invocation_mode,
  482. params=self._params_snapshot(request.chat_agent),
  483. messages=self._message_snapshots(chat_messages),
  484. tools=chat_tools,
  485. )
  486. async for item in self.chat_client.stream_chat(
  487. messages=chat_messages,
  488. tools=chat_tools,
  489. params=request.chat_agent,
  490. ):
  491. if item.kind == "raw_chunk" and item.raw_chunk is not None:
  492. raw_chunks.append(item.raw_chunk)
  493. continue
  494. if item.kind == "message_delta":
  495. if ttft_ms is None:
  496. ttft_ms = self._elapsed_ms(started_at)
  497. if not message_stream_started:
  498. message_stream_started = True
  499. await self._audit(
  500. queues,
  501. "chat_message_stream_started",
  502. turn_started_at=turn_started_at,
  503. session_id=session_id,
  504. agent="chat_agent",
  505. turn_index=turn_index,
  506. round_index=round_index,
  507. ttft_ms=ttft_ms,
  508. )
  509. message_delta_count += 1
  510. assistant_content.append(item.content or "")
  511. await queues.output.put(
  512. {"type": "message_delta", "content": item.content}
  513. )
  514. continue
  515. if item.kind == "usage" and item.usage is not None:
  516. prompt_tokens = item.usage.prompt_tokens
  517. completion_tokens = item.usage.completion_tokens
  518. total_tokens = item.usage.total_tokens
  519. cached_tokens = item.usage.cached_tokens
  520. await queues.output.put(
  521. {"type": "usage", "usage": item.usage.model_dump()}
  522. )
  523. continue
  524. event = self._accepted_chat_event(
  525. request,
  526. item,
  527. budget_available=event_loops < request.event_agent.max_event_loops,
  528. )
  529. if event is not None:
  530. saw_event = True
  531. events.append(event)
  532. await queues.output.put(
  533. {"type": "event", "event": event.model_dump()}
  534. )
  535. await self._audit(
  536. queues,
  537. "chat_event_detected",
  538. turn_started_at=turn_started_at,
  539. session_id=session_id,
  540. turn_index=turn_index,
  541. round_index=round_index,
  542. event_id=event.id,
  543. event_name=event.name,
  544. event_source=(
  545. "provider_resolved"
  546. if request.tool_invocation_mode == "chat_agent_tools"
  547. else "text_event"
  548. ),
  549. tool_invocation_mode=request.tool_invocation_mode,
  550. )
  551. if message_stream_started:
  552. await self._audit(
  553. queues,
  554. "chat_message_stream_finished",
  555. turn_started_at=turn_started_at,
  556. session_id=session_id,
  557. agent="chat_agent",
  558. turn_index=turn_index,
  559. round_index=round_index,
  560. delta_count=message_delta_count,
  561. content_length=len("".join(assistant_content)),
  562. )
  563. await self._audit(
  564. queues,
  565. "chat_agent_response",
  566. turn_started_at=turn_started_at,
  567. session_id=session_id,
  568. agent="chat_agent",
  569. turn_index=turn_index,
  570. round_index=round_index,
  571. tool_invocation_mode=request.tool_invocation_mode,
  572. content="".join(assistant_content),
  573. event_names=[event.name for event in events],
  574. events=self._event_snapshots(events),
  575. usage={
  576. "prompt_tokens": prompt_tokens,
  577. "completion_tokens": completion_tokens,
  578. "total_tokens": total_tokens,
  579. "cached_tokens": cached_tokens,
  580. },
  581. raw_chunks=raw_chunks,
  582. )
  583. if events and request.tool_invocation_mode == "chat_agent_tools":
  584. self._validate_provider_tool_call_ids(messages, events)
  585. if assistant_content or events:
  586. assistant_message = ChatMessage(
  587. role="assistant",
  588. content="".join(assistant_content),
  589. tool_calls=(
  590. events
  591. if request.tool_invocation_mode == "chat_agent_tools"
  592. else []
  593. ),
  594. )
  595. messages.append(assistant_message)
  596. self._append_persisted_message(
  597. session_id,
  598. turn_index,
  599. assistant_message,
  600. )
  601. if events:
  602. if request.tool_invocation_mode == "chat_agent_tools":
  603. tool_replies = await self._execute_provider_tools(
  604. events,
  605. enabled_names=request.event_agent.enabled_tools,
  606. )
  607. messages.extend(tool_replies)
  608. else:
  609. await queues.events.put(
  610. EventAgentRequest(
  611. events=events,
  612. history=self._event_agent_history(messages),
  613. system_prompt=request.event_agent.system_prompt,
  614. extra_body=request.event_agent.extra_body,
  615. session_id=session_id,
  616. turn_index=turn_index,
  617. round_index=round_index,
  618. turn_started_at=turn_started_at,
  619. )
  620. )
  621. tool_replies = await self._wait_for_tool_replies(
  622. queues,
  623. [event.id for event in events],
  624. )
  625. for reply in tool_replies:
  626. await queues.output.put(
  627. {"type": "tool_result", "message": reply.model_dump()}
  628. )
  629. if request.tool_invocation_mode == "chat_agent_tools":
  630. await self._audit(
  631. queues,
  632. "provider_tools_completed",
  633. turn_started_at=turn_started_at,
  634. session_id=session_id,
  635. turn_index=turn_index,
  636. round_index=round_index,
  637. tool_invocation_mode=request.tool_invocation_mode,
  638. event_source="provider_resolved",
  639. events=self._event_snapshots(events),
  640. result_count=len(tool_replies),
  641. )
  642. else:
  643. await self._audit(
  644. queues,
  645. "event_agent_completed",
  646. turn_started_at=turn_started_at,
  647. session_id=session_id,
  648. turn_index=turn_index,
  649. round_index=round_index,
  650. event_names=[event.name for event in events],
  651. result_count=len(tool_replies),
  652. result_summary="\n".join(
  653. reply.content for reply in tool_replies
  654. ),
  655. )
  656. elapsed_ms = self._elapsed_ms(started_at)
  657. self._append_persisted_usage(
  658. session_id,
  659. turn_index,
  660. round_index,
  661. usage=TokenUsage(
  662. prompt_tokens=prompt_tokens,
  663. completion_tokens=completion_tokens,
  664. total_tokens=total_tokens,
  665. cached_tokens=cached_tokens,
  666. ),
  667. ttft_ms=ttft_ms,
  668. elapsed_ms=elapsed_ms,
  669. )
  670. await queues.output.put(
  671. {
  672. "type": "round_stats",
  673. "round_index": round_index,
  674. "ttft_ms": ttft_ms,
  675. "elapsed_ms": elapsed_ms,
  676. "prompt_tokens": prompt_tokens,
  677. "completion_tokens": completion_tokens,
  678. "total_tokens": total_tokens,
  679. "cached_tokens": cached_tokens,
  680. "had_event": saw_event,
  681. },
  682. )
  683. await self._audit(
  684. queues,
  685. "chat_round_finished",
  686. turn_started_at=turn_started_at,
  687. session_id=session_id,
  688. turn_index=turn_index,
  689. round_index=round_index,
  690. tool_invocation_mode=request.tool_invocation_mode,
  691. had_event=saw_event,
  692. elapsed_ms=elapsed_ms,
  693. )
  694. if not saw_event:
  695. break
  696. event_loops += 1
  697. for deferred_message in deferred_user_messages:
  698. await queues.input.put(deferred_message)
  699. await self._audit(
  700. queues,
  701. "turn_completed",
  702. turn_started_at=turn_started_at,
  703. session_id=session_id,
  704. turn_index=turn_index,
  705. round_count=round_index,
  706. )
  707. self._complete_persisted_turn(session_id, turn_index)
  708. await queues.output.put(
  709. {
  710. "type": "turn_completed",
  711. "turn_index": turn_index,
  712. "round_count": round_index,
  713. }
  714. )
  715. async def _consume_events(
  716. self,
  717. queues: RuntimeQueues,
  718. event_agent: EventAgent,
  719. ) -> None:
  720. while True:
  721. request = await queues.events.get()
  722. if request is None:
  723. return
  724. await self._audit(
  725. queues,
  726. "event_agent_request",
  727. turn_started_at=request.turn_started_at,
  728. session_id=request.session_id,
  729. agent="event_agent",
  730. turn_index=request.turn_index,
  731. round_index=request.round_index,
  732. params=self._params_snapshot(event_agent.params),
  733. events=self._event_snapshots(request.events),
  734. history=self._message_snapshots(request.history),
  735. system_prompt=request.system_prompt,
  736. extra_body=request.extra_body or {},
  737. tools=[
  738. tool
  739. for tool in (
  740. event_agent.registry.tool_schema(event.name)
  741. for event in request.events
  742. )
  743. if tool is not None
  744. ],
  745. )
  746. replies = await event_agent.handle_many(
  747. request.events,
  748. history=request.history,
  749. system_prompt=request.system_prompt,
  750. extra_body=request.extra_body,
  751. )
  752. await self._audit(
  753. queues,
  754. "event_agent_response",
  755. turn_started_at=request.turn_started_at,
  756. session_id=request.session_id,
  757. agent="event_agent",
  758. turn_index=request.turn_index,
  759. round_index=request.round_index,
  760. replies=[reply.model_dump() for reply in replies],
  761. raw_model_chunks=event_agent.raw_model_chunks(request.events),
  762. )
  763. for reply in replies:
  764. await queues.input.put(reply)
  765. summary = event_agent.summarize_replies(replies)
  766. if summary is not None:
  767. await queues.input.put(summary)
  768. async def _wait_for_tool_replies(
  769. self,
  770. queues: RuntimeQueues,
  771. tool_call_ids: list[str],
  772. ) -> list[ChatMessage]:
  773. pending = set(tool_call_ids)
  774. replies: dict[str, ChatMessage] = {}
  775. deferred: list[ChatMessage] = []
  776. while pending:
  777. message = await queues.input.get()
  778. if message.role == "tool" and message.tool_call_id in pending:
  779. replies[message.tool_call_id] = message
  780. pending.remove(message.tool_call_id)
  781. continue
  782. deferred.append(message)
  783. for deferred_message in deferred:
  784. await queues.input.put(deferred_message)
  785. return [replies[tool_call_id] for tool_call_id in tool_call_ids]
  786. def _event_name_only(self, event: ToolCallEvent) -> ToolCallEvent:
  787. return event.model_copy(update={"arguments": {}, "raw_arguments": "{}"})
  788. def _build_initial_messages(self, request: DebugRunRequest) -> list[ChatMessage]:
  789. messages = [
  790. ChatMessage(role="system", content=prompt)
  791. for prompt in request.system_prompts
  792. ]
  793. if request.chat_agent.system_prompt.strip():
  794. messages.append(
  795. ChatMessage(role="system", content=request.chat_agent.system_prompt)
  796. )
  797. messages.extend(request.pre_messages)
  798. return messages
  799. def _chat_messages_for_mode(
  800. self,
  801. request: DebugRunRequest,
  802. messages: list[ChatMessage],
  803. enabled_events: list[str],
  804. ) -> list[ChatMessage]:
  805. if request.tool_invocation_mode == "chat_agent_tools":
  806. return list(messages)
  807. return self._chat_messages_for_round(messages, enabled_events)
  808. def _chat_tools_for_round(
  809. self,
  810. request: DebugRunRequest,
  811. enabled_events: list[str],
  812. ) -> list[dict[str, Any]]:
  813. if request.tool_invocation_mode != "chat_agent_tools":
  814. return []
  815. return self.registry.provider_tool_schemas(enabled_events)
  816. def _accepted_chat_event(
  817. self,
  818. request: DebugRunRequest,
  819. item: StreamItem,
  820. *,
  821. budget_available: bool,
  822. ) -> ToolCallEvent | None:
  823. if not budget_available:
  824. return None
  825. if request.tool_invocation_mode == "chat_agent_tools":
  826. if (
  827. item.kind == "provider_tool_call"
  828. and item.event is not None
  829. ):
  830. return item.event
  831. return None
  832. if item.kind == "text_event" and item.event is not None:
  833. return self._event_name_only(item.event)
  834. return None
  835. def _validate_provider_tool_call_ids(
  836. self,
  837. messages: list[ChatMessage],
  838. events: list[ToolCallEvent],
  839. ) -> None:
  840. declared_ids = {
  841. tool_call.id
  842. for message in messages
  843. for tool_call in message.tool_calls
  844. }
  845. batch_ids: set[str] = set()
  846. for event in events:
  847. if event.id in declared_ids or event.id in batch_ids:
  848. raise ValueError(f"duplicate provider tool-call ID: {event.id}")
  849. batch_ids.add(event.id)
  850. async def _execute_provider_tools(
  851. self,
  852. events: list[ToolCallEvent],
  853. *,
  854. enabled_names: list[str],
  855. ) -> list[ChatMessage]:
  856. payloads = await asyncio.gather(
  857. *[
  858. self.registry.execute_async(
  859. event,
  860. enabled_names=enabled_names,
  861. )
  862. for event in events
  863. ]
  864. )
  865. return [
  866. ChatMessage(
  867. role="tool",
  868. content=json.dumps(payload, ensure_ascii=False),
  869. name=event.name,
  870. tool_call_id=event.id,
  871. )
  872. for event, payload in zip(events, payloads, strict=True)
  873. ]
  874. def _chat_messages_for_round(
  875. self,
  876. messages: list[ChatMessage],
  877. enabled_events: list[str],
  878. ) -> list[ChatMessage]:
  879. if self._has_event_instructions(messages):
  880. return list(messages)
  881. event_prompt = self.registry.chat_event_system_message(enabled_events)
  882. if not event_prompt:
  883. return list(messages)
  884. insert_at = 0
  885. while insert_at < len(messages) and messages[insert_at].role == "system":
  886. insert_at += 1
  887. return [
  888. *messages[:insert_at],
  889. ChatMessage(role="system", content=event_prompt),
  890. *messages[insert_at:],
  891. ]
  892. def _has_event_instructions(self, messages: list[ChatMessage]) -> bool:
  893. return any(
  894. message.role == "system"
  895. and "Available events:" in message.content
  896. and EVENT_BLOCK_START in message.content
  897. for message in messages
  898. )
  899. def _event_agent_history(self, messages: list[ChatMessage]) -> list[ChatMessage]:
  900. return [
  901. message
  902. for message in messages
  903. if message.role in {"user", "assistant"} and message.name != "event_agent"
  904. ]
  905. async def _drain_input(
  906. self,
  907. queues: RuntimeQueues,
  908. messages: list[ChatMessage],
  909. deferred_user_messages: list[ChatMessage] | None = None,
  910. ) -> None:
  911. while not queues.input.empty():
  912. message = await queues.input.get()
  913. if (
  914. deferred_user_messages is not None
  915. and message.role == "user"
  916. and message.name != "event_agent"
  917. ):
  918. deferred_user_messages.append(message)
  919. continue
  920. messages.append(message)
  921. def _elapsed_ms(self, started_at: float) -> int:
  922. return round((self.clock() - started_at) * 1000)
  923. def _ensure_persisted_session(self, request: DebugRunRequest) -> str | None:
  924. if self.session_store is None:
  925. return None
  926. return self.session_store.ensure_session(
  927. request.session_id,
  928. title=self._session_title(request.user_message),
  929. config=self._session_config_snapshot(request),
  930. )
  931. def _initial_turn_index(self, session_id: str | None) -> int:
  932. if self.session_store is None or session_id is None:
  933. return 1
  934. return self.session_store.next_turn_index(session_id)
  935. def _start_persisted_turn(
  936. self,
  937. session_id: str | None,
  938. turn_index: int,
  939. user_message: ChatMessage,
  940. ) -> None:
  941. if self.session_store is None or session_id is None:
  942. return
  943. self.session_store.start_turn(
  944. session_id,
  945. turn_index=turn_index,
  946. user_message=user_message.content,
  947. )
  948. self.session_store.append_message(
  949. session_id,
  950. turn_index=turn_index,
  951. message=user_message,
  952. )
  953. def _complete_persisted_turn(
  954. self,
  955. session_id: str | None,
  956. turn_index: int,
  957. ) -> None:
  958. if self.session_store is None or session_id is None:
  959. return
  960. self.session_store.complete_turn(session_id, turn_index=turn_index)
  961. def _append_persisted_message(
  962. self,
  963. session_id: str | None,
  964. turn_index: int,
  965. message: ChatMessage,
  966. ) -> None:
  967. if self.session_store is None or session_id is None:
  968. return
  969. self.session_store.append_message(
  970. session_id,
  971. turn_index=turn_index,
  972. message=message,
  973. )
  974. def _append_persisted_usage(
  975. self,
  976. session_id: str | None,
  977. turn_index: int,
  978. round_index: int,
  979. *,
  980. usage: TokenUsage,
  981. ttft_ms: int | None,
  982. elapsed_ms: int,
  983. ) -> None:
  984. if self.session_store is None or session_id is None:
  985. return
  986. self.session_store.append_usage(
  987. session_id,
  988. turn_index=turn_index,
  989. round_index=round_index,
  990. usage=usage,
  991. ttft_ms=ttft_ms,
  992. elapsed_ms=elapsed_ms,
  993. )
  994. def _session_title(self, user_message: str) -> str:
  995. title = " ".join(user_message.split())
  996. if not title:
  997. return "Debug session"
  998. return title[:80]
  999. def _session_config_snapshot(self, request: DebugRunRequest) -> dict[str, Any]:
  1000. return {
  1001. "tool_invocation_mode": request.tool_invocation_mode,
  1002. "system_prompts": list(request.system_prompts),
  1003. "pre_messages": self._message_snapshots(request.pre_messages),
  1004. "chat_agent": request.chat_agent.model_dump(),
  1005. "event_agent": request.event_agent.model_dump(),
  1006. }
  1007. def _params_snapshot(self, params: AgentParams) -> dict[str, Any]:
  1008. return params.model_dump()
  1009. def _message_snapshots(self, messages: list[ChatMessage] | tuple[ChatMessage, ...]) -> list[dict[str, Any]]:
  1010. return [message.model_dump() for message in messages]
  1011. def _event_snapshots(self, events: list[ToolCallEvent]) -> list[dict[str, Any]]:
  1012. return [event.model_dump() for event in events]
  1013. async def _audit(
  1014. self,
  1015. queues: RuntimeQueues,
  1016. event: str,
  1017. *,
  1018. session_id: str | None = None,
  1019. turn_started_at: float | None = None,
  1020. **details: Any,
  1021. ) -> None:
  1022. if turn_started_at is not None:
  1023. details["turn_elapsed_ms"] = max(0, self._elapsed_ms(turn_started_at))
  1024. logger.info("audit event=%s details=%s", event, details)
  1025. if self.session_store is not None and session_id is not None:
  1026. turn_index = details.get("turn_index")
  1027. round_index = details.get("round_index")
  1028. self.session_store.append_audit(
  1029. session_id,
  1030. event=event,
  1031. details=details,
  1032. turn_index=turn_index if isinstance(turn_index, int) else None,
  1033. round_index=round_index if isinstance(round_index, int) else None,
  1034. )
  1035. await queues.output.put(
  1036. {
  1037. "type": "audit",
  1038. "event": event,
  1039. "details": details,
  1040. }
  1041. )