runtime.py 8.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233
  1. import asyncio
  2. import time
  3. from collections.abc import AsyncIterator, Callable
  4. from typing import Any, Protocol
  5. from agent_lab.application.contracts import AgentParams, DebugRunRequest
  6. from agent_lab.application.event_agent import EventAgent
  7. from agent_lab.application.queues import RuntimeQueues
  8. from agent_lab.application.tools import ToolRegistry, build_default_tool_registry
  9. from agent_lab.domain.messages import ChatMessage, StreamItem
  10. class ChatClient(Protocol):
  11. async def stream_chat(
  12. self,
  13. messages: list[ChatMessage],
  14. tools: list[dict[str, Any]],
  15. params: AgentParams,
  16. ) -> AsyncIterator[StreamItem]:
  17. ...
  18. class DebugRuntime:
  19. def __init__(
  20. self,
  21. chat_client: ChatClient,
  22. queues: RuntimeQueues | None = None,
  23. registry: ToolRegistry | None = None,
  24. clock: Callable[[], float] = time.perf_counter,
  25. ) -> None:
  26. self.chat_client = chat_client
  27. self.queues = queues
  28. self.registry = registry or build_default_tool_registry()
  29. self.clock = clock
  30. self._tasks: list[asyncio.Task[None]] = []
  31. def start(self, request: DebugRunRequest) -> RuntimeQueues:
  32. queues = self.queues or RuntimeQueues()
  33. task = asyncio.create_task(self._produce(request, queues))
  34. self._tasks.append(task)
  35. return queues
  36. async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]:
  37. queues = self.start(request)
  38. async for message in self.output_messages(queues):
  39. yield message
  40. await self._wait_for_tasks()
  41. async def output_messages(
  42. self,
  43. queues: RuntimeQueues,
  44. ) -> AsyncIterator[dict[str, Any]]:
  45. while True:
  46. message = await queues.output.get()
  47. yield message
  48. if message.get("type") in {"done", "error"}:
  49. break
  50. async def aclose(self) -> None:
  51. for task in self._tasks:
  52. if not task.done():
  53. task.cancel()
  54. await self._wait_for_tasks()
  55. async def _wait_for_tasks(self) -> None:
  56. if not self._tasks:
  57. return
  58. await asyncio.gather(*self._tasks, return_exceptions=True)
  59. self._tasks = [task for task in self._tasks if not task.done()]
  60. async def _produce(self, request: DebugRunRequest, queues: RuntimeQueues) -> None:
  61. event_agent = EventAgent(
  62. request.event_agent.enabled_tools,
  63. registry=self.registry,
  64. )
  65. event_worker = asyncio.create_task(self._consume_events(queues, event_agent))
  66. try:
  67. await self._run_chat_agent(request, queues)
  68. except Exception as exc:
  69. await queues.output.put({"type": "error", "message": str(exc)})
  70. finally:
  71. await queues.events.put(None)
  72. await asyncio.gather(event_worker, return_exceptions=True)
  73. async def _run_chat_agent(
  74. self,
  75. request: DebugRunRequest,
  76. queues: RuntimeQueues,
  77. ) -> None:
  78. messages = self._build_initial_messages(request)
  79. await queues.input.put(ChatMessage(role="user", content=request.user_message))
  80. tools = self.registry.chat_tools(request.event_agent.enabled_tools)
  81. await queues.output.put({"type": "session_started"})
  82. event_loops = 0
  83. round_index = 0
  84. while True:
  85. await self._drain_input(queues, messages)
  86. round_index += 1
  87. started_at = self.clock()
  88. ttft_ms: int | None = None
  89. prompt_tokens = 0
  90. completion_tokens = 0
  91. total_tokens = 0
  92. cached_tokens = 0
  93. saw_event = False
  94. assistant_content: list[str] = []
  95. assistant_tool_calls: list[dict[str, Any]] = []
  96. tool_replies: list[ChatMessage] = []
  97. async for item in self.chat_client.stream_chat(
  98. messages=messages,
  99. tools=tools,
  100. params=request.chat_agent,
  101. ):
  102. if item.kind == "message_delta":
  103. if ttft_ms is None:
  104. ttft_ms = self._elapsed_ms(started_at)
  105. assistant_content.append(item.content or "")
  106. await queues.output.put(
  107. {"type": "message_delta", "content": item.content}
  108. )
  109. continue
  110. if item.kind == "usage" and item.usage is not None:
  111. prompt_tokens = item.usage.prompt_tokens
  112. completion_tokens = item.usage.completion_tokens
  113. total_tokens = item.usage.total_tokens
  114. cached_tokens = item.usage.cached_tokens
  115. await queues.output.put(
  116. {"type": "usage", "usage": item.usage.model_dump()}
  117. )
  118. continue
  119. if item.kind == "event" and item.event is not None:
  120. saw_event = True
  121. assistant_tool_calls.append(
  122. {
  123. "id": item.event.id,
  124. "type": "function",
  125. "function": {
  126. "name": item.event.name,
  127. "arguments": item.event.raw_arguments,
  128. },
  129. }
  130. )
  131. await queues.events.put(item.event)
  132. await queues.output.put(
  133. {"type": "event", "event": item.event.model_dump()}
  134. )
  135. reply = await self._wait_for_tool_reply(queues, item.event.id)
  136. tool_replies.append(reply)
  137. await queues.output.put(
  138. {"type": "tool_result", "message": reply.model_dump()}
  139. )
  140. if assistant_content or assistant_tool_calls:
  141. messages.append(
  142. ChatMessage(
  143. role="assistant",
  144. content="".join(assistant_content),
  145. tool_calls=assistant_tool_calls or None,
  146. )
  147. )
  148. messages.extend(tool_replies)
  149. await queues.output.put(
  150. {
  151. "type": "round_stats",
  152. "round_index": round_index,
  153. "ttft_ms": ttft_ms,
  154. "elapsed_ms": self._elapsed_ms(started_at),
  155. "prompt_tokens": prompt_tokens,
  156. "completion_tokens": completion_tokens,
  157. "total_tokens": total_tokens,
  158. "cached_tokens": cached_tokens,
  159. "had_event": saw_event,
  160. },
  161. )
  162. if not saw_event:
  163. break
  164. event_loops += 1
  165. if event_loops >= request.event_agent.max_event_loops:
  166. break
  167. await queues.output.put({"type": "done"})
  168. async def _consume_events(
  169. self,
  170. queues: RuntimeQueues,
  171. event_agent: EventAgent,
  172. ) -> None:
  173. while True:
  174. event = await queues.events.get()
  175. if event is None:
  176. return
  177. await queues.input.put(await event_agent.handle(event))
  178. async def _wait_for_tool_reply(
  179. self,
  180. queues: RuntimeQueues,
  181. tool_call_id: str,
  182. ) -> ChatMessage:
  183. deferred: list[ChatMessage] = []
  184. while True:
  185. message = await queues.input.get()
  186. if message.role == "tool" and message.tool_call_id == tool_call_id:
  187. for deferred_message in deferred:
  188. await queues.input.put(deferred_message)
  189. return message
  190. deferred.append(message)
  191. def _build_initial_messages(self, request: DebugRunRequest) -> list[ChatMessage]:
  192. messages = [
  193. ChatMessage(role="system", content=prompt)
  194. for prompt in request.system_prompts
  195. ]
  196. messages.extend(request.pre_messages)
  197. return messages
  198. async def _drain_input(
  199. self,
  200. queues: RuntimeQueues,
  201. messages: list[ChatMessage],
  202. ) -> None:
  203. while not queues.input.empty():
  204. messages.append(await queues.input.get())
  205. def _elapsed_ms(self, started_at: float) -> int:
  206. return round((self.clock() - started_at) * 1000)