| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233 |
- 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)
|