|
|
@@ -1,4 +1,5 @@
|
|
|
import asyncio
|
|
|
+import json
|
|
|
import logging
|
|
|
import time
|
|
|
from collections.abc import AsyncIterator, Callable
|
|
|
@@ -84,34 +85,23 @@ class DebugRuntime:
|
|
|
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,
|
|
|
- chat_client=self.chat_client,
|
|
|
- params=request.event_agent,
|
|
|
- )
|
|
|
- event_worker = asyncio.create_task(self._consume_events(queues, event_agent))
|
|
|
+ event_worker = self._start_event_worker(request, queues)
|
|
|
try:
|
|
|
await self._run_chat_agent(request, queues)
|
|
|
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)
|
|
|
+ if event_worker is not None:
|
|
|
+ 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))
|
|
|
+ event_worker = self._start_event_worker(request, queues)
|
|
|
try:
|
|
|
await self._run_chat_session(request, queues)
|
|
|
except asyncio.CancelledError:
|
|
|
@@ -120,8 +110,24 @@ class DebugRuntime:
|
|
|
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)
|
|
|
+ if event_worker is not None:
|
|
|
+ await queues.events.put(None)
|
|
|
+ await asyncio.gather(event_worker, return_exceptions=True)
|
|
|
+
|
|
|
+ def _start_event_worker(
|
|
|
+ self,
|
|
|
+ request: DebugRunRequest,
|
|
|
+ queues: RuntimeQueues,
|
|
|
+ ) -> asyncio.Task[None] | None:
|
|
|
+ if request.tool_invocation_mode != "dual_agent":
|
|
|
+ return None
|
|
|
+ event_agent = EventAgent(
|
|
|
+ request.event_agent.enabled_tools,
|
|
|
+ registry=self.registry,
|
|
|
+ chat_client=self.chat_client,
|
|
|
+ params=request.event_agent,
|
|
|
+ )
|
|
|
+ return asyncio.create_task(self._consume_events(queues, event_agent))
|
|
|
|
|
|
async def _run_chat_agent(
|
|
|
self,
|
|
|
@@ -133,7 +139,12 @@ class DebugRuntime:
|
|
|
await queues.input.put(ChatMessage(role="user", content=request.user_message))
|
|
|
|
|
|
await queues.output.put({"type": "session_started"})
|
|
|
- await self._audit(queues, "session_started", turn_started_at=turn_started_at)
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "session_started",
|
|
|
+ turn_started_at=turn_started_at,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
+ )
|
|
|
|
|
|
event_loops = 0
|
|
|
round_index = 0
|
|
|
@@ -158,13 +169,19 @@ class DebugRuntime:
|
|
|
if event_loops < request.event_agent.max_event_loops
|
|
|
else []
|
|
|
)
|
|
|
+ chat_tools = self._chat_tools_for_round(request, event_prompt_events)
|
|
|
raw_chunks: list[dict[str, Any]] = []
|
|
|
- chat_messages = self._chat_messages_for_round(messages, event_prompt_events)
|
|
|
+ chat_messages = self._chat_messages_for_mode(
|
|
|
+ request,
|
|
|
+ messages,
|
|
|
+ event_prompt_events,
|
|
|
+ )
|
|
|
await self._audit(
|
|
|
queues,
|
|
|
"chat_round_started",
|
|
|
turn_started_at=turn_started_at,
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
events_enabled=event_prompt_events,
|
|
|
configured_events=configured_events,
|
|
|
event_generation_enabled=bool(event_prompt_events),
|
|
|
@@ -176,14 +193,15 @@ class DebugRuntime:
|
|
|
turn_started_at=turn_started_at,
|
|
|
agent="chat_agent",
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
params=self._params_snapshot(request.chat_agent),
|
|
|
messages=self._message_snapshots(chat_messages),
|
|
|
- tools=[],
|
|
|
+ tools=chat_tools,
|
|
|
)
|
|
|
|
|
|
async for item in self.chat_client.stream_chat(
|
|
|
messages=chat_messages,
|
|
|
- tools=[],
|
|
|
+ tools=chat_tools,
|
|
|
params=request.chat_agent,
|
|
|
):
|
|
|
if item.kind == "raw_chunk" and item.raw_chunk is not None:
|
|
|
@@ -220,9 +238,13 @@ class DebugRuntime:
|
|
|
)
|
|
|
continue
|
|
|
|
|
|
- if item.kind == "text_event" and item.event is not None:
|
|
|
+ event = self._accepted_chat_event(
|
|
|
+ request,
|
|
|
+ item,
|
|
|
+ budget_available=event_loops < request.event_agent.max_event_loops,
|
|
|
+ )
|
|
|
+ if 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()}
|
|
|
@@ -234,6 +256,12 @@ class DebugRuntime:
|
|
|
round_index=round_index,
|
|
|
event_id=event.id,
|
|
|
event_name=event.name,
|
|
|
+ event_source=(
|
|
|
+ "provider_resolved"
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools"
|
|
|
+ else "text_event"
|
|
|
+ ),
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
)
|
|
|
|
|
|
if message_stream_started:
|
|
|
@@ -252,8 +280,10 @@ class DebugRuntime:
|
|
|
turn_started_at=turn_started_at,
|
|
|
agent="chat_agent",
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
content="".join(assistant_content),
|
|
|
event_names=[event.name for event in events],
|
|
|
+ events=self._event_snapshots(events),
|
|
|
usage={
|
|
|
"prompt_tokens": prompt_tokens,
|
|
|
"completion_tokens": completion_tokens,
|
|
|
@@ -267,36 +297,62 @@ class DebugRuntime:
|
|
|
assistant_message = ChatMessage(
|
|
|
role="assistant",
|
|
|
content="".join(assistant_content),
|
|
|
+ tool_calls=(
|
|
|
+ events
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools"
|
|
|
+ else []
|
|
|
+ ),
|
|
|
)
|
|
|
messages.append(assistant_message)
|
|
|
|
|
|
if events:
|
|
|
- await queues.events.put(
|
|
|
- EventAgentRequest(
|
|
|
- events=events,
|
|
|
- history=self._event_agent_history(messages),
|
|
|
- system_prompt=request.event_agent.system_prompt,
|
|
|
- extra_body=request.event_agent.extra_body,
|
|
|
- turn_started_at=turn_started_at,
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools":
|
|
|
+ tool_replies = await self._execute_provider_tools(
|
|
|
+ events,
|
|
|
+ enabled_names=request.event_agent.enabled_tools,
|
|
|
+ )
|
|
|
+ messages.extend(tool_replies)
|
|
|
+ else:
|
|
|
+ await queues.events.put(
|
|
|
+ EventAgentRequest(
|
|
|
+ events=events,
|
|
|
+ history=self._event_agent_history(messages),
|
|
|
+ system_prompt=request.event_agent.system_prompt,
|
|
|
+ extra_body=request.event_agent.extra_body,
|
|
|
+ turn_started_at=turn_started_at,
|
|
|
+ )
|
|
|
+ )
|
|
|
+ tool_replies = await self._wait_for_tool_replies(
|
|
|
+ queues,
|
|
|
+ [event.id for event in events],
|
|
|
)
|
|
|
- )
|
|
|
- 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_started_at=turn_started_at,
|
|
|
- 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),
|
|
|
- )
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools":
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "provider_tools_completed",
|
|
|
+ turn_started_at=turn_started_at,
|
|
|
+ round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
+ event_source="provider_resolved",
|
|
|
+ events=self._event_snapshots(events),
|
|
|
+ result_count=len(tool_replies),
|
|
|
+ )
|
|
|
+ else:
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "event_agent_completed",
|
|
|
+ turn_started_at=turn_started_at,
|
|
|
+ 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(
|
|
|
@@ -317,6 +373,7 @@ class DebugRuntime:
|
|
|
"chat_round_finished",
|
|
|
turn_started_at=turn_started_at,
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
had_event=saw_event,
|
|
|
elapsed_ms=elapsed_ms,
|
|
|
)
|
|
|
@@ -351,6 +408,7 @@ class DebugRuntime:
|
|
|
"session_started",
|
|
|
turn_started_at=initial_turn_started_at,
|
|
|
session_id=session_id,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
)
|
|
|
await queues.input.put(ChatMessage(role="user", content=request.user_message))
|
|
|
|
|
|
@@ -424,8 +482,13 @@ class DebugRuntime:
|
|
|
if event_loops < request.event_agent.max_event_loops
|
|
|
else []
|
|
|
)
|
|
|
+ chat_tools = self._chat_tools_for_round(request, event_prompt_events)
|
|
|
raw_chunks: list[dict[str, Any]] = []
|
|
|
- chat_messages = self._chat_messages_for_round(messages, event_prompt_events)
|
|
|
+ chat_messages = self._chat_messages_for_mode(
|
|
|
+ request,
|
|
|
+ messages,
|
|
|
+ event_prompt_events,
|
|
|
+ )
|
|
|
await self._audit(
|
|
|
queues,
|
|
|
"chat_round_started",
|
|
|
@@ -433,6 +496,7 @@ class DebugRuntime:
|
|
|
session_id=session_id,
|
|
|
turn_index=turn_index,
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
events_enabled=event_prompt_events,
|
|
|
configured_events=configured_events,
|
|
|
event_generation_enabled=bool(event_prompt_events),
|
|
|
@@ -446,14 +510,15 @@ class DebugRuntime:
|
|
|
agent="chat_agent",
|
|
|
turn_index=turn_index,
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
params=self._params_snapshot(request.chat_agent),
|
|
|
messages=self._message_snapshots(chat_messages),
|
|
|
- tools=[],
|
|
|
+ tools=chat_tools,
|
|
|
)
|
|
|
|
|
|
async for item in self.chat_client.stream_chat(
|
|
|
messages=chat_messages,
|
|
|
- tools=[],
|
|
|
+ tools=chat_tools,
|
|
|
params=request.chat_agent,
|
|
|
):
|
|
|
if item.kind == "raw_chunk" and item.raw_chunk is not None:
|
|
|
@@ -492,9 +557,13 @@ class DebugRuntime:
|
|
|
)
|
|
|
continue
|
|
|
|
|
|
- if item.kind == "text_event" and item.event is not None:
|
|
|
+ event = self._accepted_chat_event(
|
|
|
+ request,
|
|
|
+ item,
|
|
|
+ budget_available=event_loops < request.event_agent.max_event_loops,
|
|
|
+ )
|
|
|
+ if 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()}
|
|
|
@@ -508,6 +577,12 @@ class DebugRuntime:
|
|
|
round_index=round_index,
|
|
|
event_id=event.id,
|
|
|
event_name=event.name,
|
|
|
+ event_source=(
|
|
|
+ "provider_resolved"
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools"
|
|
|
+ else "text_event"
|
|
|
+ ),
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
)
|
|
|
|
|
|
if message_stream_started:
|
|
|
@@ -530,8 +605,10 @@ class DebugRuntime:
|
|
|
agent="chat_agent",
|
|
|
turn_index=turn_index,
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
content="".join(assistant_content),
|
|
|
event_names=[event.name for event in events],
|
|
|
+ events=self._event_snapshots(events),
|
|
|
usage={
|
|
|
"prompt_tokens": prompt_tokens,
|
|
|
"completion_tokens": completion_tokens,
|
|
|
@@ -545,6 +622,11 @@ class DebugRuntime:
|
|
|
assistant_message = ChatMessage(
|
|
|
role="assistant",
|
|
|
content="".join(assistant_content),
|
|
|
+ tool_calls=(
|
|
|
+ events
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools"
|
|
|
+ else []
|
|
|
+ ),
|
|
|
)
|
|
|
messages.append(assistant_message)
|
|
|
self._append_persisted_message(
|
|
|
@@ -554,37 +636,60 @@ class DebugRuntime:
|
|
|
)
|
|
|
|
|
|
if events:
|
|
|
- await queues.events.put(
|
|
|
- EventAgentRequest(
|
|
|
- events=events,
|
|
|
- history=self._event_agent_history(messages),
|
|
|
- system_prompt=request.event_agent.system_prompt,
|
|
|
- extra_body=request.event_agent.extra_body,
|
|
|
- session_id=session_id,
|
|
|
- turn_index=turn_index,
|
|
|
- round_index=round_index,
|
|
|
- turn_started_at=turn_started_at,
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools":
|
|
|
+ tool_replies = await self._execute_provider_tools(
|
|
|
+ events,
|
|
|
+ enabled_names=request.event_agent.enabled_tools,
|
|
|
+ )
|
|
|
+ messages.extend(tool_replies)
|
|
|
+ else:
|
|
|
+ await queues.events.put(
|
|
|
+ EventAgentRequest(
|
|
|
+ events=events,
|
|
|
+ history=self._event_agent_history(messages),
|
|
|
+ system_prompt=request.event_agent.system_prompt,
|
|
|
+ extra_body=request.event_agent.extra_body,
|
|
|
+ session_id=session_id,
|
|
|
+ turn_index=turn_index,
|
|
|
+ round_index=round_index,
|
|
|
+ turn_started_at=turn_started_at,
|
|
|
+ )
|
|
|
+ )
|
|
|
+ tool_replies = await self._wait_for_tool_replies(
|
|
|
+ queues,
|
|
|
+ [event.id for event in events],
|
|
|
)
|
|
|
- )
|
|
|
- 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_started_at=turn_started_at,
|
|
|
- session_id=session_id,
|
|
|
- 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),
|
|
|
- )
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools":
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "provider_tools_completed",
|
|
|
+ turn_started_at=turn_started_at,
|
|
|
+ session_id=session_id,
|
|
|
+ turn_index=turn_index,
|
|
|
+ round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
+ event_source="provider_resolved",
|
|
|
+ events=self._event_snapshots(events),
|
|
|
+ result_count=len(tool_replies),
|
|
|
+ )
|
|
|
+ else:
|
|
|
+ await self._audit(
|
|
|
+ queues,
|
|
|
+ "event_agent_completed",
|
|
|
+ turn_started_at=turn_started_at,
|
|
|
+ session_id=session_id,
|
|
|
+ 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)
|
|
|
self._append_persisted_usage(
|
|
|
@@ -620,6 +725,7 @@ class DebugRuntime:
|
|
|
session_id=session_id,
|
|
|
turn_index=turn_index,
|
|
|
round_index=round_index,
|
|
|
+ tool_invocation_mode=request.tool_invocation_mode,
|
|
|
had_event=saw_event,
|
|
|
elapsed_ms=elapsed_ms,
|
|
|
)
|
|
|
@@ -736,6 +842,66 @@ class DebugRuntime:
|
|
|
messages.extend(request.pre_messages)
|
|
|
return messages
|
|
|
|
|
|
+ def _chat_messages_for_mode(
|
|
|
+ self,
|
|
|
+ request: DebugRunRequest,
|
|
|
+ messages: list[ChatMessage],
|
|
|
+ enabled_events: list[str],
|
|
|
+ ) -> list[ChatMessage]:
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools":
|
|
|
+ return list(messages)
|
|
|
+ return self._chat_messages_for_round(messages, enabled_events)
|
|
|
+
|
|
|
+ def _chat_tools_for_round(
|
|
|
+ self,
|
|
|
+ request: DebugRunRequest,
|
|
|
+ enabled_events: list[str],
|
|
|
+ ) -> list[dict[str, Any]]:
|
|
|
+ if request.tool_invocation_mode != "chat_agent_tools":
|
|
|
+ return []
|
|
|
+ return self.registry.provider_tool_schemas(enabled_events)
|
|
|
+
|
|
|
+ def _accepted_chat_event(
|
|
|
+ self,
|
|
|
+ request: DebugRunRequest,
|
|
|
+ item: StreamItem,
|
|
|
+ *,
|
|
|
+ budget_available: bool,
|
|
|
+ ) -> ToolCallEvent | None:
|
|
|
+ if request.tool_invocation_mode == "chat_agent_tools":
|
|
|
+ if (
|
|
|
+ budget_available
|
|
|
+ and item.kind == "provider_tool_call"
|
|
|
+ and item.event is not None
|
|
|
+ ):
|
|
|
+ return item.event
|
|
|
+ return None
|
|
|
+ if item.kind == "text_event" and item.event is not None:
|
|
|
+ return self._event_name_only(item.event)
|
|
|
+ return None
|
|
|
+
|
|
|
+ async def _execute_provider_tools(
|
|
|
+ self,
|
|
|
+ events: list[ToolCallEvent],
|
|
|
+ *,
|
|
|
+ enabled_names: list[str],
|
|
|
+ ) -> list[ChatMessage]:
|
|
|
+ replies: list[ChatMessage] = []
|
|
|
+ for event in events:
|
|
|
+ payload = await self.registry.execute_async(
|
|
|
+ event,
|
|
|
+ enabled_names=enabled_names,
|
|
|
+ )
|
|
|
+ replies.append(
|
|
|
+ ChatMessage(
|
|
|
+ role="tool",
|
|
|
+ content=json.dumps(payload, ensure_ascii=False),
|
|
|
+ name=event.name,
|
|
|
+ tool_call_id=event.id,
|
|
|
+ )
|
|
|
+ )
|
|
|
+ return replies
|
|
|
+
|
|
|
def _chat_messages_for_round(
|
|
|
self,
|
|
|
messages: list[ChatMessage],
|
|
|
@@ -878,6 +1044,7 @@ class DebugRuntime:
|
|
|
|
|
|
def _session_config_snapshot(self, request: DebugRunRequest) -> dict[str, Any]:
|
|
|
return {
|
|
|
+ "tool_invocation_mode": request.tool_invocation_mode,
|
|
|
"system_prompts": list(request.system_prompts),
|
|
|
"pre_messages": self._message_snapshots(request.pre_messages),
|
|
|
"chat_agent": request.chat_agent.model_dump(),
|