|
|
@@ -1,3 +1,4 @@
|
|
|
+import asyncio
|
|
|
import time
|
|
|
from collections.abc import AsyncIterator, Callable
|
|
|
from typing import Any, Protocol
|
|
|
@@ -31,18 +32,66 @@ class DebugRuntime:
|
|
|
self.queues = queues
|
|
|
self.registry = registry or build_default_tool_registry()
|
|
|
self.clock = clock
|
|
|
+ self._tasks: list[asyncio.Task[None]] = []
|
|
|
|
|
|
- async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]:
|
|
|
+ def start(self, request: DebugRunRequest) -> RuntimeQueues:
|
|
|
queues = self.queues or RuntimeQueues()
|
|
|
- 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)
|
|
|
+ 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)
|
|
|
|
|
|
- yield await self._emit(queues, {"type": "session_started"})
|
|
|
+ await queues.output.put({"type": "session_started"})
|
|
|
|
|
|
event_loops = 0
|
|
|
round_index = 0
|
|
|
@@ -58,6 +107,7 @@ class DebugRuntime:
|
|
|
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,
|
|
|
@@ -68,9 +118,8 @@ class DebugRuntime:
|
|
|
if ttft_ms is None:
|
|
|
ttft_ms = self._elapsed_ms(started_at)
|
|
|
assistant_content.append(item.content or "")
|
|
|
- yield await self._emit(
|
|
|
- queues,
|
|
|
- {"type": "message_delta", "content": item.content},
|
|
|
+ await queues.output.put(
|
|
|
+ {"type": "message_delta", "content": item.content}
|
|
|
)
|
|
|
continue
|
|
|
|
|
|
@@ -79,9 +128,8 @@ class DebugRuntime:
|
|
|
completion_tokens = item.usage.completion_tokens
|
|
|
total_tokens = item.usage.total_tokens
|
|
|
cached_tokens = item.usage.cached_tokens
|
|
|
- yield await self._emit(
|
|
|
- queues,
|
|
|
- {"type": "usage", "usage": item.usage.model_dump()},
|
|
|
+ await queues.output.put(
|
|
|
+ {"type": "usage", "usage": item.usage.model_dump()}
|
|
|
)
|
|
|
continue
|
|
|
|
|
|
@@ -98,16 +146,13 @@ class DebugRuntime:
|
|
|
}
|
|
|
)
|
|
|
await queues.events.put(item.event)
|
|
|
- yield await self._emit(
|
|
|
- queues,
|
|
|
- {"type": "event", "event": item.event.model_dump()},
|
|
|
+ await queues.output.put(
|
|
|
+ {"type": "event", "event": item.event.model_dump()}
|
|
|
)
|
|
|
- event = await queues.events.get()
|
|
|
- reply = await event_agent.handle(event)
|
|
|
- await queues.input.put(reply)
|
|
|
- yield await self._emit(
|
|
|
- queues,
|
|
|
- {"type": "tool_result", "message": reply.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:
|
|
|
@@ -118,9 +163,9 @@ class DebugRuntime:
|
|
|
tool_calls=assistant_tool_calls or None,
|
|
|
)
|
|
|
)
|
|
|
+ messages.extend(tool_replies)
|
|
|
|
|
|
- yield await self._emit(
|
|
|
- queues,
|
|
|
+ await queues.output.put(
|
|
|
{
|
|
|
"type": "round_stats",
|
|
|
"round_index": round_index,
|
|
|
@@ -141,7 +186,32 @@ class DebugRuntime:
|
|
|
if event_loops >= request.event_agent.max_event_loops:
|
|
|
break
|
|
|
|
|
|
- yield await self._emit(queues, {"type": "done"})
|
|
|
+ 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 = [
|
|
|
@@ -159,13 +229,5 @@ class DebugRuntime:
|
|
|
while not queues.input.empty():
|
|
|
messages.append(await queues.input.get())
|
|
|
|
|
|
- async def _emit(
|
|
|
- self,
|
|
|
- queues: RuntimeQueues,
|
|
|
- payload: dict[str, Any],
|
|
|
- ) -> dict[str, Any]:
|
|
|
- await queues.output.put(payload)
|
|
|
- return await queues.output.get()
|
|
|
-
|
|
|
def _elapsed_ms(self, started_at: float) -> int:
|
|
|
return round((self.clock() - started_at) * 1000)
|