|
|
@@ -45,6 +45,12 @@ class DebugRuntime:
|
|
|
self._tasks.append(task)
|
|
|
return queues
|
|
|
|
|
|
+ def start_session(self, request: DebugRunRequest) -> RuntimeQueues:
|
|
|
+ queues = self.queues or RuntimeQueues()
|
|
|
+ task = asyncio.create_task(self._produce_session(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):
|
|
|
@@ -90,6 +96,29 @@ class DebugRuntime:
|
|
|
await queues.events.put(None)
|
|
|
await asyncio.gather(event_worker, return_exceptions=True)
|
|
|
|
|
|
+ async def _produce_session(
|
|
|
+ self,
|
|
|
+ request: DebugRunRequest,
|
|
|
+ queues: RuntimeQueues,
|
|
|
+ ) -> None:
|
|
|
+ event_agent = EventAgent(
|
|
|
+ request.event_agent.enabled_tools,
|
|
|
+ registry=self.registry,
|
|
|
+ chat_client=self.chat_client,
|
|
|
+ params=request.event_agent,
|
|
|
+ )
|
|
|
+ event_worker = asyncio.create_task(self._consume_events(queues, event_agent))
|
|
|
+ try:
|
|
|
+ await self._run_chat_session(request, queues)
|
|
|
+ except asyncio.CancelledError:
|
|
|
+ raise
|
|
|
+ except Exception as exc:
|
|
|
+ logger.exception("runtime session failed")
|
|
|
+ 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,
|
|
|
@@ -231,6 +260,196 @@ class DebugRuntime:
|
|
|
await self._audit(queues, "session_finished", round_count=round_index)
|
|
|
await queues.output.put({"type": "done"})
|
|
|
|
|
|
+ async def _run_chat_session(
|
|
|
+ self,
|
|
|
+ request: DebugRunRequest,
|
|
|
+ queues: RuntimeQueues,
|
|
|
+ ) -> None:
|
|
|
+ messages = self._build_initial_messages(request)
|
|
|
+ await queues.output.put({"type": "session_started"})
|
|
|
+ await self._audit(queues, "session_started")
|
|
|
+ await queues.input.put(ChatMessage(role="user", content=request.user_message))
|
|
|
+
|
|
|
+ turn_index = 0
|
|
|
+ while True:
|
|
|
+ user_message = await queues.input.get()
|
|
|
+ if user_message.role != "user" or user_message.name == "event_agent":
|
|
|
+ messages.append(user_message)
|
|
|
+ continue
|
|
|
+ turn_index += 1
|
|
|
+ await self._run_chat_turn(
|
|
|
+ request,
|
|
|
+ queues,
|
|
|
+ messages,
|
|
|
+ user_message,
|
|
|
+ turn_index,
|
|
|
+ )
|
|
|
+
|
|
|
+ async def _run_chat_turn(
|
|
|
+ self,
|
|
|
+ request: DebugRunRequest,
|
|
|
+ queues: RuntimeQueues,
|
|
|
+ messages: list[ChatMessage],
|
|
|
+ user_message: ChatMessage,
|
|
|
+ turn_index: int,
|
|
|
+ ) -> None:
|
|
|
+ messages.append(user_message)
|
|
|
+ await queues.output.put({"type": "turn_started", "turn_index": turn_index})
|
|
|
+ await self._audit(queues, "turn_started", turn_index=turn_index)
|
|
|
+
|
|
|
+ event_loops = 0
|
|
|
+ round_index = 0
|
|
|
+ deferred_user_messages: list[ChatMessage] = []
|
|
|
+ while True:
|
|
|
+ if round_index:
|
|
|
+ await self._drain_input(
|
|
|
+ queues,
|
|
|
+ messages,
|
|
|
+ deferred_user_messages=deferred_user_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] = []
|
|
|
+ events: list[ToolCallEvent] = []
|
|
|
+ enabled_events = (
|
|
|
+ request.event_agent.enabled_tools
|
|
|
+ if event_loops < request.event_agent.max_event_loops
|
|
|
+ else []
|
|
|
+ )
|
|
|
+ chat_messages = self._chat_messages_for_round(messages, enabled_events)
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "chat_round_started",
|
|
|
+ turn_index=turn_index,
|
|
|
+ round_index=round_index,
|
|
|
+ events_enabled=enabled_events,
|
|
|
+ )
|
|
|
+
|
|
|
+ async for item in self.chat_client.stream_chat(
|
|
|
+ messages=chat_messages,
|
|
|
+ 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
|
|
|
+ event = self._event_name_only(item.event)
|
|
|
+ events.append(event)
|
|
|
+ await queues.output.put(
|
|
|
+ {"type": "event", "event": event.model_dump()}
|
|
|
+ )
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "chat_event_detected",
|
|
|
+ turn_index=turn_index,
|
|
|
+ round_index=round_index,
|
|
|
+ event_id=event.id,
|
|
|
+ event_name=event.name,
|
|
|
+ )
|
|
|
+
|
|
|
+ if assistant_content or events:
|
|
|
+ messages.append(
|
|
|
+ ChatMessage(
|
|
|
+ role="assistant",
|
|
|
+ content="".join(assistant_content),
|
|
|
+ )
|
|
|
+ )
|
|
|
+
|
|
|
+ if events:
|
|
|
+ await queues.events.put(
|
|
|
+ EventAgentRequest(
|
|
|
+ events=events,
|
|
|
+ history=list(messages),
|
|
|
+ system_prompt=request.event_agent.system_prompt,
|
|
|
+ extra_body=request.event_agent.extra_body,
|
|
|
+ )
|
|
|
+ )
|
|
|
+ tool_replies = await self._wait_for_tool_replies(
|
|
|
+ queues,
|
|
|
+ [event.id for event in events],
|
|
|
+ )
|
|
|
+ for reply in tool_replies:
|
|
|
+ await queues.output.put(
|
|
|
+ {"type": "tool_result", "message": reply.model_dump()}
|
|
|
+ )
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "event_agent_completed",
|
|
|
+ turn_index=turn_index,
|
|
|
+ round_index=round_index,
|
|
|
+ event_names=[event.name for event in events],
|
|
|
+ result_count=len(tool_replies),
|
|
|
+ result_summary="\n".join(reply.content for reply in tool_replies),
|
|
|
+ )
|
|
|
+
|
|
|
+ elapsed_ms = self._elapsed_ms(started_at)
|
|
|
+ await queues.output.put(
|
|
|
+ {
|
|
|
+ "type": "round_stats",
|
|
|
+ "round_index": round_index,
|
|
|
+ "ttft_ms": ttft_ms,
|
|
|
+ "elapsed_ms": elapsed_ms,
|
|
|
+ "prompt_tokens": prompt_tokens,
|
|
|
+ "completion_tokens": completion_tokens,
|
|
|
+ "total_tokens": total_tokens,
|
|
|
+ "cached_tokens": cached_tokens,
|
|
|
+ "had_event": saw_event,
|
|
|
+ },
|
|
|
+ )
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "chat_round_finished",
|
|
|
+ turn_index=turn_index,
|
|
|
+ round_index=round_index,
|
|
|
+ had_event=saw_event,
|
|
|
+ elapsed_ms=elapsed_ms,
|
|
|
+ )
|
|
|
+
|
|
|
+ if not saw_event:
|
|
|
+ break
|
|
|
+
|
|
|
+ event_loops += 1
|
|
|
+
|
|
|
+ for deferred_message in deferred_user_messages:
|
|
|
+ await queues.input.put(deferred_message)
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "turn_completed",
|
|
|
+ turn_index=turn_index,
|
|
|
+ round_count=round_index,
|
|
|
+ )
|
|
|
+ await queues.output.put(
|
|
|
+ {
|
|
|
+ "type": "turn_completed",
|
|
|
+ "turn_index": turn_index,
|
|
|
+ "round_count": round_index,
|
|
|
+ }
|
|
|
+ )
|
|
|
+
|
|
|
async def _consume_events(
|
|
|
self,
|
|
|
queues: RuntimeQueues,
|
|
|
@@ -320,9 +539,18 @@ class DebugRuntime:
|
|
|
self,
|
|
|
queues: RuntimeQueues,
|
|
|
messages: list[ChatMessage],
|
|
|
+ deferred_user_messages: list[ChatMessage] | None = None,
|
|
|
) -> None:
|
|
|
while not queues.input.empty():
|
|
|
- messages.append(await queues.input.get())
|
|
|
+ message = await queues.input.get()
|
|
|
+ if (
|
|
|
+ deferred_user_messages is not None
|
|
|
+ and message.role == "user"
|
|
|
+ and message.name != "event_agent"
|
|
|
+ ):
|
|
|
+ deferred_user_messages.append(message)
|
|
|
+ continue
|
|
|
+ messages.append(message)
|
|
|
|
|
|
def _elapsed_ms(self, started_at: float) -> int:
|
|
|
return round((self.clock() - started_at) * 1000)
|