| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117 |
- import asyncio
- import json
- import logging
- 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, EventAgentRequest
- from agent_lab.application.queues import RuntimeQueues
- from agent_lab.application.session_store import SessionStore
- from agent_lab.application.tools import ToolRegistry, build_default_tool_registry
- from agent_lab.domain.events import EVENT_BLOCK_START, ToolCallEvent
- from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
- logger = logging.getLogger(__name__)
- class ChatClient(Protocol):
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict[str, Any]],
- params: AgentParams,
- tool_choice: dict[str, Any] | None = None,
- ) -> AsyncIterator[StreamItem]:
- ...
- class DebugRuntime:
- def __init__(
- self,
- chat_client: ChatClient,
- queues: RuntimeQueues | None = None,
- registry: ToolRegistry | None = None,
- session_store: SessionStore | 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.session_store = session_store
- 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
- 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):
- 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_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:
- 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_worker = self._start_event_worker(request, queues)
- 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:
- 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,
- request: DebugRunRequest,
- queues: RuntimeQueues,
- ) -> None:
- messages = self._build_initial_messages(request)
- turn_started_at = self.clock()
- 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,
- tool_invocation_mode=request.tool_invocation_mode,
- )
- 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] = []
- events: list[ToolCallEvent] = []
- tool_replies: list[ChatMessage] = []
- message_stream_started = False
- message_delta_count = 0
- configured_events = request.event_agent.enabled_tools
- event_prompt_events = (
- request.event_agent.enabled_tools
- 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_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),
- event_prompt_events=event_prompt_events,
- )
- await self._audit(
- queues,
- "chat_agent_request",
- 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=chat_tools,
- )
- async for item in self.chat_client.stream_chat(
- messages=chat_messages,
- tools=chat_tools,
- params=request.chat_agent,
- ):
- if item.kind == "raw_chunk" and item.raw_chunk is not None:
- raw_chunks.append(item.raw_chunk)
- continue
- if item.kind == "message_delta":
- if ttft_ms is None:
- ttft_ms = self._elapsed_ms(started_at)
- if not message_stream_started:
- message_stream_started = True
- await self._audit(
- queues,
- "chat_message_stream_started",
- turn_started_at=turn_started_at,
- agent="chat_agent",
- round_index=round_index,
- ttft_ms=ttft_ms,
- )
- message_delta_count += 1
- 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
- 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
- events.append(event)
- await queues.output.put(
- {"type": "event", "event": event.model_dump()}
- )
- await self._audit(
- queues,
- "chat_event_detected",
- turn_started_at=turn_started_at,
- 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:
- await self._audit(
- queues,
- "chat_message_stream_finished",
- turn_started_at=turn_started_at,
- agent="chat_agent",
- round_index=round_index,
- delta_count=message_delta_count,
- content_length=len("".join(assistant_content)),
- )
- await self._audit(
- queues,
- "chat_agent_response",
- 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,
- "total_tokens": total_tokens,
- "cached_tokens": cached_tokens,
- },
- raw_chunks=raw_chunks,
- )
- if events and request.tool_invocation_mode == "chat_agent_tools":
- self._validate_provider_tool_call_ids(messages, events)
- if assistant_content or events:
- 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:
- 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],
- )
- for reply in tool_replies:
- await queues.output.put(
- {"type": "tool_result", "message": reply.model_dump()}
- )
- 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(
- {
- "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_started_at=turn_started_at,
- round_index=round_index,
- tool_invocation_mode=request.tool_invocation_mode,
- had_event=saw_event,
- elapsed_ms=elapsed_ms,
- )
- if not saw_event:
- break
- event_loops += 1
- await self._audit(
- queues,
- "session_finished",
- turn_started_at=turn_started_at,
- round_count=round_index,
- )
- await queues.output.put({"type": "done"})
- async def _run_chat_session(
- self,
- request: DebugRunRequest,
- queues: RuntimeQueues,
- ) -> None:
- session_id = self._ensure_persisted_session(request)
- messages = self._build_initial_messages(request)
- initial_turn_started_at = self.clock()
- session_message = {"type": "session_started"}
- if session_id is not None:
- session_message["session_id"] = session_id
- await queues.output.put(session_message)
- await self._audit(
- queues,
- "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))
- turn_index = self._initial_turn_index(session_id) - 1
- next_turn_started_at: float | None = initial_turn_started_at
- 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
- turn_started_at = next_turn_started_at or self.clock()
- next_turn_started_at = None
- await self._run_chat_turn(
- request,
- queues,
- messages,
- user_message,
- turn_index,
- session_id,
- turn_started_at,
- )
- async def _run_chat_turn(
- self,
- request: DebugRunRequest,
- queues: RuntimeQueues,
- messages: list[ChatMessage],
- user_message: ChatMessage,
- turn_index: int,
- session_id: str | None = None,
- turn_started_at: float | None = None,
- ) -> None:
- turn_started_at = turn_started_at or self.clock()
- messages.append(user_message)
- self._start_persisted_turn(session_id, turn_index, user_message)
- await queues.output.put({"type": "turn_started", "turn_index": turn_index})
- await self._audit(
- queues,
- "turn_started",
- turn_started_at=turn_started_at,
- session_id=session_id,
- 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] = []
- message_stream_started = False
- message_delta_count = 0
- configured_events = request.event_agent.enabled_tools
- event_prompt_events = (
- request.event_agent.enabled_tools
- 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_mode(
- request,
- messages,
- event_prompt_events,
- )
- await self._audit(
- queues,
- "chat_round_started",
- 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,
- events_enabled=event_prompt_events,
- configured_events=configured_events,
- event_generation_enabled=bool(event_prompt_events),
- event_prompt_events=event_prompt_events,
- )
- await self._audit(
- queues,
- "chat_agent_request",
- turn_started_at=turn_started_at,
- session_id=session_id,
- 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=chat_tools,
- )
- async for item in self.chat_client.stream_chat(
- messages=chat_messages,
- tools=chat_tools,
- params=request.chat_agent,
- ):
- if item.kind == "raw_chunk" and item.raw_chunk is not None:
- raw_chunks.append(item.raw_chunk)
- continue
- if item.kind == "message_delta":
- if ttft_ms is None:
- ttft_ms = self._elapsed_ms(started_at)
- if not message_stream_started:
- message_stream_started = True
- await self._audit(
- queues,
- "chat_message_stream_started",
- turn_started_at=turn_started_at,
- session_id=session_id,
- agent="chat_agent",
- turn_index=turn_index,
- round_index=round_index,
- ttft_ms=ttft_ms,
- )
- message_delta_count += 1
- 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
- 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
- events.append(event)
- await queues.output.put(
- {"type": "event", "event": event.model_dump()}
- )
- await self._audit(
- queues,
- "chat_event_detected",
- turn_started_at=turn_started_at,
- session_id=session_id,
- turn_index=turn_index,
- 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:
- await self._audit(
- queues,
- "chat_message_stream_finished",
- turn_started_at=turn_started_at,
- session_id=session_id,
- agent="chat_agent",
- turn_index=turn_index,
- round_index=round_index,
- delta_count=message_delta_count,
- content_length=len("".join(assistant_content)),
- )
- await self._audit(
- queues,
- "chat_agent_response",
- turn_started_at=turn_started_at,
- session_id=session_id,
- 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,
- "total_tokens": total_tokens,
- "cached_tokens": cached_tokens,
- },
- raw_chunks=raw_chunks,
- )
- if events and request.tool_invocation_mode == "chat_agent_tools":
- self._validate_provider_tool_call_ids(messages, events)
- if assistant_content or events:
- 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(
- session_id,
- turn_index,
- assistant_message,
- )
- if events:
- 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],
- )
- for reply in tool_replies:
- await queues.output.put(
- {"type": "tool_result", "message": reply.model_dump()}
- )
- 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(
- session_id,
- turn_index,
- round_index,
- usage=TokenUsage(
- prompt_tokens=prompt_tokens,
- completion_tokens=completion_tokens,
- total_tokens=total_tokens,
- cached_tokens=cached_tokens,
- ),
- ttft_ms=ttft_ms,
- elapsed_ms=elapsed_ms,
- )
- 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_started_at=turn_started_at,
- 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,
- )
- 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_started_at=turn_started_at,
- session_id=session_id,
- turn_index=turn_index,
- round_count=round_index,
- )
- self._complete_persisted_turn(session_id, turn_index)
- await queues.output.put(
- {
- "type": "turn_completed",
- "turn_index": turn_index,
- "round_count": round_index,
- }
- )
- async def _consume_events(
- self,
- queues: RuntimeQueues,
- event_agent: EventAgent,
- ) -> None:
- while True:
- request = await queues.events.get()
- if request is None:
- return
- await self._audit(
- queues,
- "event_agent_request",
- turn_started_at=request.turn_started_at,
- session_id=request.session_id,
- agent="event_agent",
- turn_index=request.turn_index,
- round_index=request.round_index,
- params=self._params_snapshot(event_agent.params),
- events=self._event_snapshots(request.events),
- history=self._message_snapshots(request.history),
- system_prompt=request.system_prompt,
- extra_body=request.extra_body or {},
- tools=[
- tool
- for tool in (
- event_agent.registry.tool_schema(event.name)
- for event in request.events
- )
- if tool is not None
- ],
- )
- replies = await event_agent.handle_many(
- request.events,
- history=request.history,
- system_prompt=request.system_prompt,
- extra_body=request.extra_body,
- )
- await self._audit(
- queues,
- "event_agent_response",
- turn_started_at=request.turn_started_at,
- session_id=request.session_id,
- agent="event_agent",
- turn_index=request.turn_index,
- round_index=request.round_index,
- replies=[reply.model_dump() for reply in replies],
- raw_model_chunks=event_agent.raw_model_chunks(request.events),
- )
- for reply in replies:
- await queues.input.put(reply)
- summary = event_agent.summarize_replies(replies)
- if summary is not None:
- await queues.input.put(summary)
- async def _wait_for_tool_replies(
- self,
- queues: RuntimeQueues,
- tool_call_ids: list[str],
- ) -> list[ChatMessage]:
- pending = set(tool_call_ids)
- replies: dict[str, ChatMessage] = {}
- deferred: list[ChatMessage] = []
- while pending:
- message = await queues.input.get()
- if message.role == "tool" and message.tool_call_id in pending:
- replies[message.tool_call_id] = message
- pending.remove(message.tool_call_id)
- continue
- deferred.append(message)
- for deferred_message in deferred:
- await queues.input.put(deferred_message)
- return [replies[tool_call_id] for tool_call_id in tool_call_ids]
- def _event_name_only(self, event: ToolCallEvent) -> ToolCallEvent:
- return event.model_copy(update={"arguments": {}, "raw_arguments": "{}"})
- def _build_initial_messages(self, request: DebugRunRequest) -> list[ChatMessage]:
- messages = [
- ChatMessage(role="system", content=prompt)
- for prompt in request.system_prompts
- ]
- if request.chat_agent.system_prompt.strip():
- messages.append(
- ChatMessage(role="system", content=request.chat_agent.system_prompt)
- )
- 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 not budget_available:
- return None
- if request.tool_invocation_mode == "chat_agent_tools":
- if (
- 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
- def _validate_provider_tool_call_ids(
- self,
- messages: list[ChatMessage],
- events: list[ToolCallEvent],
- ) -> None:
- declared_ids = {
- tool_call.id
- for message in messages
- for tool_call in message.tool_calls
- }
- batch_ids: set[str] = set()
- for event in events:
- if event.id in declared_ids or event.id in batch_ids:
- raise ValueError(f"duplicate provider tool-call ID: {event.id}")
- batch_ids.add(event.id)
- async def _execute_provider_tools(
- self,
- events: list[ToolCallEvent],
- *,
- enabled_names: list[str],
- ) -> list[ChatMessage]:
- payloads = await asyncio.gather(
- *[
- self.registry.execute_async(
- event,
- enabled_names=enabled_names,
- )
- for event in events
- ]
- )
- return [
- ChatMessage(
- role="tool",
- content=json.dumps(payload, ensure_ascii=False),
- name=event.name,
- tool_call_id=event.id,
- )
- for event, payload in zip(events, payloads, strict=True)
- ]
- def _chat_messages_for_round(
- self,
- messages: list[ChatMessage],
- enabled_events: list[str],
- ) -> list[ChatMessage]:
- if self._has_event_instructions(messages):
- return list(messages)
- event_prompt = self.registry.chat_event_system_message(enabled_events)
- if not event_prompt:
- return list(messages)
- insert_at = 0
- while insert_at < len(messages) and messages[insert_at].role == "system":
- insert_at += 1
- return [
- *messages[:insert_at],
- ChatMessage(role="system", content=event_prompt),
- *messages[insert_at:],
- ]
- def _has_event_instructions(self, messages: list[ChatMessage]) -> bool:
- return any(
- message.role == "system"
- and "Available events:" in message.content
- and EVENT_BLOCK_START in message.content
- for message in messages
- )
- def _event_agent_history(self, messages: list[ChatMessage]) -> list[ChatMessage]:
- return [
- message
- for message in messages
- if message.role in {"user", "assistant"} and message.name != "event_agent"
- ]
- async def _drain_input(
- self,
- queues: RuntimeQueues,
- messages: list[ChatMessage],
- deferred_user_messages: list[ChatMessage] | None = None,
- ) -> None:
- while not queues.input.empty():
- 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)
- def _ensure_persisted_session(self, request: DebugRunRequest) -> str | None:
- if self.session_store is None:
- return None
- return self.session_store.ensure_session(
- request.session_id,
- title=self._session_title(request.user_message),
- config=self._session_config_snapshot(request),
- )
- def _initial_turn_index(self, session_id: str | None) -> int:
- if self.session_store is None or session_id is None:
- return 1
- return self.session_store.next_turn_index(session_id)
- def _start_persisted_turn(
- self,
- session_id: str | None,
- turn_index: int,
- user_message: ChatMessage,
- ) -> None:
- if self.session_store is None or session_id is None:
- return
- self.session_store.start_turn(
- session_id,
- turn_index=turn_index,
- user_message=user_message.content,
- )
- self.session_store.append_message(
- session_id,
- turn_index=turn_index,
- message=user_message,
- )
- def _complete_persisted_turn(
- self,
- session_id: str | None,
- turn_index: int,
- ) -> None:
- if self.session_store is None or session_id is None:
- return
- self.session_store.complete_turn(session_id, turn_index=turn_index)
- def _append_persisted_message(
- self,
- session_id: str | None,
- turn_index: int,
- message: ChatMessage,
- ) -> None:
- if self.session_store is None or session_id is None:
- return
- self.session_store.append_message(
- session_id,
- turn_index=turn_index,
- message=message,
- )
- def _append_persisted_usage(
- self,
- session_id: str | None,
- turn_index: int,
- round_index: int,
- *,
- usage: TokenUsage,
- ttft_ms: int | None,
- elapsed_ms: int,
- ) -> None:
- if self.session_store is None or session_id is None:
- return
- self.session_store.append_usage(
- session_id,
- turn_index=turn_index,
- round_index=round_index,
- usage=usage,
- ttft_ms=ttft_ms,
- elapsed_ms=elapsed_ms,
- )
- def _session_title(self, user_message: str) -> str:
- title = " ".join(user_message.split())
- if not title:
- return "Debug session"
- return title[:80]
- 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(),
- "event_agent": request.event_agent.model_dump(),
- }
- def _params_snapshot(self, params: AgentParams) -> dict[str, Any]:
- return params.model_dump()
- def _message_snapshots(self, messages: list[ChatMessage] | tuple[ChatMessage, ...]) -> list[dict[str, Any]]:
- return [message.model_dump() for message in messages]
- def _event_snapshots(self, events: list[ToolCallEvent]) -> list[dict[str, Any]]:
- return [event.model_dump() for event in events]
- async def _audit(
- self,
- queues: RuntimeQueues,
- event: str,
- *,
- session_id: str | None = None,
- turn_started_at: float | None = None,
- **details: Any,
- ) -> None:
- if turn_started_at is not None:
- details["turn_elapsed_ms"] = max(0, self._elapsed_ms(turn_started_at))
- logger.info("audit event=%s details=%s", event, details)
- if self.session_store is not None and session_id is not None:
- turn_index = details.get("turn_index")
- round_index = details.get("round_index")
- self.session_store.append_audit(
- session_id,
- event=event,
- details=details,
- turn_index=turn_index if isinstance(turn_index, int) else None,
- round_index=round_index if isinstance(round_index, int) else None,
- )
- await queues.output.put(
- {
- "type": "audit",
- "event": event,
- "details": details,
- }
- )
|