import asyncio import time from collections.abc import AsyncIterator, Callable from typing import Any, Protocol from agent_lab.application.contracts import AgentParams, DebugRunRequest from agent_lab.application.event_agent import EventAgent from agent_lab.application.queues import RuntimeQueues from agent_lab.application.tools import ToolRegistry, build_default_tool_registry from agent_lab.domain.messages import ChatMessage, StreamItem class ChatClient(Protocol): async def stream_chat( self, messages: list[ChatMessage], tools: list[dict[str, Any]], params: AgentParams, ) -> AsyncIterator[StreamItem]: ... class DebugRuntime: def __init__( self, chat_client: ChatClient, queues: RuntimeQueues | None = None, registry: ToolRegistry | None = None, clock: Callable[[], float] = time.perf_counter, ) -> None: self.chat_client = chat_client self.queues = queues self.registry = registry or build_default_tool_registry() self.clock = clock self._tasks: list[asyncio.Task[None]] = [] def start(self, request: DebugRunRequest) -> RuntimeQueues: queues = self.queues or RuntimeQueues() task = asyncio.create_task(self._produce(request, queues)) self._tasks.append(task) return queues async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]: queues = self.start(request) async for message in self.output_messages(queues): yield message await self._wait_for_tasks() async def output_messages( self, queues: RuntimeQueues, ) -> AsyncIterator[dict[str, Any]]: while True: message = await queues.output.get() yield message if message.get("type") in {"done", "error"}: break async def aclose(self) -> None: for task in self._tasks: if not task.done(): task.cancel() await self._wait_for_tasks() async def _wait_for_tasks(self) -> None: if not self._tasks: return await asyncio.gather(*self._tasks, return_exceptions=True) self._tasks = [task for task in self._tasks if not task.done()] async def _produce(self, request: DebugRunRequest, queues: RuntimeQueues) -> None: event_agent = EventAgent( request.event_agent.enabled_tools, registry=self.registry, ) event_worker = asyncio.create_task(self._consume_events(queues, event_agent)) try: await self._run_chat_agent(request, queues) except Exception as exc: await queues.output.put({"type": "error", "message": str(exc)}) finally: await queues.events.put(None) await asyncio.gather(event_worker, return_exceptions=True) async def _run_chat_agent( self, request: DebugRunRequest, queues: RuntimeQueues, ) -> None: messages = self._build_initial_messages(request) await queues.input.put(ChatMessage(role="user", content=request.user_message)) tools = self.registry.chat_tools(request.event_agent.enabled_tools) await queues.output.put({"type": "session_started"}) event_loops = 0 round_index = 0 while True: await self._drain_input(queues, messages) round_index += 1 started_at = self.clock() ttft_ms: int | None = None prompt_tokens = 0 completion_tokens = 0 total_tokens = 0 cached_tokens = 0 saw_event = False assistant_content: list[str] = [] assistant_tool_calls: list[dict[str, Any]] = [] tool_replies: list[ChatMessage] = [] async for item in self.chat_client.stream_chat( messages=messages, tools=tools, params=request.chat_agent, ): if item.kind == "message_delta": if ttft_ms is None: ttft_ms = self._elapsed_ms(started_at) assistant_content.append(item.content or "") await queues.output.put( {"type": "message_delta", "content": item.content} ) continue if item.kind == "usage" and item.usage is not None: prompt_tokens = item.usage.prompt_tokens completion_tokens = item.usage.completion_tokens total_tokens = item.usage.total_tokens cached_tokens = item.usage.cached_tokens await queues.output.put( {"type": "usage", "usage": item.usage.model_dump()} ) continue if item.kind == "event" and item.event is not None: saw_event = True assistant_tool_calls.append( { "id": item.event.id, "type": "function", "function": { "name": item.event.name, "arguments": item.event.raw_arguments, }, } ) await queues.events.put(item.event) await queues.output.put( {"type": "event", "event": item.event.model_dump()} ) reply = await self._wait_for_tool_reply(queues, item.event.id) tool_replies.append(reply) await queues.output.put( {"type": "tool_result", "message": reply.model_dump()} ) if assistant_content or assistant_tool_calls: messages.append( ChatMessage( role="assistant", content="".join(assistant_content), tool_calls=assistant_tool_calls or None, ) ) messages.extend(tool_replies) await queues.output.put( { "type": "round_stats", "round_index": round_index, "ttft_ms": ttft_ms, "elapsed_ms": self._elapsed_ms(started_at), "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, "total_tokens": total_tokens, "cached_tokens": cached_tokens, "had_event": saw_event, }, ) if not saw_event: break event_loops += 1 if event_loops >= request.event_agent.max_event_loops: break await queues.output.put({"type": "done"}) async def _consume_events( self, queues: RuntimeQueues, event_agent: EventAgent, ) -> None: while True: event = await queues.events.get() if event is None: return await queues.input.put(await event_agent.handle(event)) async def _wait_for_tool_reply( self, queues: RuntimeQueues, tool_call_id: str, ) -> ChatMessage: deferred: list[ChatMessage] = [] while True: message = await queues.input.get() if message.role == "tool" and message.tool_call_id == tool_call_id: for deferred_message in deferred: await queues.input.put(deferred_message) return message deferred.append(message) def _build_initial_messages(self, request: DebugRunRequest) -> list[ChatMessage]: messages = [ ChatMessage(role="system", content=prompt) for prompt in request.system_prompts ] messages.extend(request.pre_messages) return messages async def _drain_input( self, queues: RuntimeQueues, messages: list[ChatMessage], ) -> None: while not queues.input.empty(): messages.append(await queues.input.get()) def _elapsed_ms(self, started_at: float) -> int: return round((self.clock() - started_at) * 1000)