refactor(agent): unify internal turn lifecycle (#4993)
This commit is contained in:
+132
-180
@@ -110,6 +110,20 @@ class TurnState(Enum):
|
||||
DONE = auto()
|
||||
|
||||
|
||||
class TurnKind(Enum):
|
||||
USER = auto()
|
||||
SYSTEM = auto()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TurnRoute:
|
||||
"""Where a turn response is delivered, separate from its execution input."""
|
||||
|
||||
channel: str
|
||||
chat_id: str
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class StateTraceEntry:
|
||||
state: TurnState
|
||||
@@ -126,6 +140,8 @@ class TurnContext:
|
||||
state: TurnState
|
||||
turn_id: str
|
||||
runtime: LLMRuntime
|
||||
kind: TurnKind
|
||||
route: TurnRoute
|
||||
original_user_text: str | None = None
|
||||
session: Session | None = None
|
||||
|
||||
@@ -576,9 +592,28 @@ class AgentLoop:
|
||||
self._runtime_context_providers.append(provider)
|
||||
|
||||
@staticmethod
|
||||
def _runtime_chat_id(msg: InboundMessage) -> str:
|
||||
"""Return the chat id shown in runtime metadata for the model."""
|
||||
return str(msg.metadata.get("context_chat_id") or msg.chat_id)
|
||||
def _turn_route(msg: InboundMessage, session_key: str) -> TurnRoute:
|
||||
"""Resolve response routing without mixing it into execution metadata."""
|
||||
if msg.channel != "system":
|
||||
return TurnRoute(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
metadata=dict(msg.metadata or {}),
|
||||
)
|
||||
|
||||
channel, chat_id = (
|
||||
msg.chat_id.split(":", 1) if ":" in msg.chat_id else ("cli", msg.chat_id)
|
||||
)
|
||||
metadata: dict[str, Any] = {}
|
||||
if (
|
||||
channel == "slack"
|
||||
and session_key.startswith("slack:")
|
||||
and session_key.count(":") >= 2
|
||||
):
|
||||
metadata["slack"] = {"thread_ts": session_key.split(":", 2)[2]}
|
||||
if origin_message_id := msg.metadata.get("origin_message_id"):
|
||||
metadata["origin_message_id"] = origin_message_id
|
||||
return TurnRoute(channel=channel, chat_id=chat_id, metadata=metadata)
|
||||
|
||||
async def _build_bus_progress_callback(
|
||||
self, msg: InboundMessage
|
||||
@@ -660,38 +695,38 @@ class AgentLoop:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _build_initial_messages(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
session: Session,
|
||||
history: list[dict[str, Any]],
|
||||
pending_summary: str | None,
|
||||
include_memory_recent_history: bool = True,
|
||||
runtime_context_blocks: list[RuntimeContextBlock] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
def _build_initial_messages(self, ctx: TurnContext) -> list[dict[str, Any]]:
|
||||
"""Build the initial message list for the LLM turn."""
|
||||
scope = self.workspace_scopes.for_message(msg, session.metadata)
|
||||
assert ctx.session is not None
|
||||
is_subagent = ctx.kind is TurnKind.SYSTEM and ctx.msg.sender_id == "subagent"
|
||||
scope = self.workspace_scopes.for_message(ctx.msg, ctx.session.metadata)
|
||||
return self.context.build_messages(
|
||||
history=history,
|
||||
current_message=msg.content,
|
||||
media=msg.media if msg.media else None,
|
||||
channel=msg.channel,
|
||||
chat_id=self._runtime_chat_id(msg),
|
||||
sender_id=msg.sender_id,
|
||||
session_summary=pending_summary,
|
||||
session_metadata=session.metadata,
|
||||
history=ctx.history,
|
||||
current_message="" if is_subagent else ctx.msg.content,
|
||||
media=ctx.msg.media if ctx.kind is TurnKind.USER and ctx.msg.media else None,
|
||||
channel=ctx.route.channel,
|
||||
chat_id=str(ctx.msg.metadata.get("context_chat_id") or ctx.route.chat_id),
|
||||
current_role="assistant" if is_subagent else "user",
|
||||
sender_id=ctx.msg.sender_id,
|
||||
session_summary=ctx.pending_summary,
|
||||
session_metadata=ctx.session.metadata,
|
||||
workspace=scope.project_path,
|
||||
runtime_context_blocks=runtime_context_blocks,
|
||||
include_memory_recent_history=include_memory_recent_history,
|
||||
session_key=session.key,
|
||||
runtime_context_blocks=ctx.runtime_context_blocks,
|
||||
include_memory_recent_history=not ctx.ephemeral,
|
||||
session_key=ctx.session.key,
|
||||
unified_session=self._unified_session,
|
||||
)
|
||||
|
||||
def _request_context_for_turn(self, ctx: TurnContext) -> RequestContext:
|
||||
scope = self.workspace_scopes.for_message(ctx.msg, ctx.session.metadata)
|
||||
assert ctx.session is not None
|
||||
scope = self.workspace_scopes.for_turn(
|
||||
channel=ctx.route.channel,
|
||||
message_metadata=ctx.msg.metadata,
|
||||
session_metadata=ctx.session.metadata,
|
||||
)
|
||||
return RequestContext(
|
||||
channel=ctx.msg.channel,
|
||||
chat_id=ctx.msg.chat_id,
|
||||
channel=ctx.route.channel,
|
||||
chat_id=ctx.route.chat_id,
|
||||
message_id=ctx.msg.metadata.get("message_id"),
|
||||
session_key=ctx.session_key,
|
||||
original_user_text=ctx.original_user_text,
|
||||
@@ -1258,110 +1293,6 @@ class AgentLoop:
|
||||
self._running = False
|
||||
logger.info("Agent loop stopping")
|
||||
|
||||
async def _process_system_message(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
runtime: LLMRuntime,
|
||||
session_key: str | None = None,
|
||||
on_progress: Callable[..., Awaitable[None]] | None = None,
|
||||
on_stream: Callable[[str], Awaitable[None]] | None = None,
|
||||
on_stream_end: Callable[..., Awaitable[None]] | None = None,
|
||||
pending_queue: asyncio.Queue | None = None,
|
||||
hook_factories: list[AgentTurnHookFactory] | None = None,
|
||||
) -> OutboundMessage | None:
|
||||
"""Process a system inbound message (e.g. subagent announce)."""
|
||||
channel, chat_id = (
|
||||
msg.chat_id.split(":", 1) if ":" in msg.chat_id else ("cli", msg.chat_id)
|
||||
)
|
||||
logger.info("Processing system message from {}", msg.sender_id)
|
||||
key = msg.session_key_override or f"{channel}:{chat_id}"
|
||||
session = self.sessions.get_or_create(key)
|
||||
self._runtime_events().record_turn_runtime(key, runtime)
|
||||
if self._restore_runtime_checkpoint(session):
|
||||
self.sessions.save(session)
|
||||
if self._restore_pending_user_turn(session):
|
||||
self.sessions.save(session)
|
||||
|
||||
session, pending = self.auto_compact.prepare_session(session, key)
|
||||
if pending:
|
||||
logger.info("Memory compact triggered for session {}", key)
|
||||
|
||||
await self.consolidator.maybe_consolidate_by_tokens(
|
||||
session,
|
||||
runtime=runtime,
|
||||
replay_max_messages=replay_max_messages_for_context(
|
||||
runtime.context_window_tokens
|
||||
),
|
||||
)
|
||||
is_subagent = msg.sender_id == "subagent"
|
||||
if is_subagent and self._persist_subagent_followup(session, msg):
|
||||
logger.debug("Subagent result persisted for session {}", key)
|
||||
self.sessions.save(session)
|
||||
current_role = "assistant" if is_subagent else "user"
|
||||
_hist_kwargs: dict[str, Any] = {
|
||||
"max_messages": replay_max_messages_for_context(runtime.context_window_tokens),
|
||||
"max_tokens": self._replay_token_budget(runtime),
|
||||
"extend_to_user": is_subagent,
|
||||
}
|
||||
history = session.get_history(**_hist_kwargs)
|
||||
workspace_scope = self.workspace_scopes.for_message(msg, session.metadata)
|
||||
|
||||
messages = self.context.build_messages(
|
||||
history=history,
|
||||
current_message="" if is_subagent else msg.content,
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
current_role=current_role,
|
||||
sender_id=msg.sender_id,
|
||||
session_summary=pending,
|
||||
session_metadata=session.metadata,
|
||||
workspace=workspace_scope.project_path,
|
||||
session_key=key,
|
||||
unified_session=self._unified_session,
|
||||
)
|
||||
t_wall = time.time()
|
||||
final_content, _, all_msgs, stop_reason, _ = await self._run_agent_loop(
|
||||
messages, session=session, channel=channel, chat_id=chat_id,
|
||||
runtime=runtime,
|
||||
message_id=msg.metadata.get("message_id"),
|
||||
metadata=msg.metadata,
|
||||
session_key=key,
|
||||
original_user_text=None,
|
||||
pending_queue=pending_queue,
|
||||
hook_factories=hook_factories,
|
||||
)
|
||||
wall_done = time.time()
|
||||
latency_ms = max(0, int((wall_done - t_wall) * 1000))
|
||||
self._save_turn(session, all_msgs, 1 + len(history), turn_latency_ms=latency_ms)
|
||||
self._runtime_events().record_turn_latency(key, latency_ms)
|
||||
session.enforce_file_cap(
|
||||
on_archive=partial(self.context.memory.raw_archive, session_key=key)
|
||||
)
|
||||
self._clear_runtime_checkpoint(session)
|
||||
self.sessions.save(session)
|
||||
self._schedule_background(
|
||||
self.consolidator.maybe_consolidate_by_tokens(
|
||||
session,
|
||||
runtime=runtime,
|
||||
replay_max_messages=replay_max_messages_for_context(
|
||||
runtime.context_window_tokens
|
||||
),
|
||||
)
|
||||
)
|
||||
content = final_content or "Background task completed."
|
||||
outbound_metadata: dict[str, Any] = {}
|
||||
if channel == "slack" and key.startswith("slack:") and key.count(":") >= 2:
|
||||
outbound_metadata["slack"] = {"thread_ts": key.split(":", 2)[2]}
|
||||
if origin_message_id := msg.metadata.get("origin_message_id"):
|
||||
outbound_metadata["origin_message_id"] = origin_message_id
|
||||
return OutboundMessage(
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
content=content,
|
||||
metadata=outbound_metadata,
|
||||
)
|
||||
|
||||
async def _process_message(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
@@ -1381,19 +1312,15 @@ class AgentLoop:
|
||||
if runtime is None:
|
||||
runtime = self.llm_runtime()
|
||||
|
||||
if msg.channel == "system":
|
||||
return await self._process_system_message(
|
||||
msg,
|
||||
runtime=runtime,
|
||||
session_key=session_key,
|
||||
on_progress=on_progress,
|
||||
on_stream=on_stream,
|
||||
on_stream_end=on_stream_end,
|
||||
pending_queue=pending_queue,
|
||||
hook_factories=hook_factories,
|
||||
kind = TurnKind.SYSTEM if msg.channel == "system" else TurnKind.USER
|
||||
if kind is TurnKind.SYSTEM:
|
||||
destination = (
|
||||
msg.chat_id.split(":", 1) if ":" in msg.chat_id else ("cli", msg.chat_id)
|
||||
)
|
||||
|
||||
key = session_key or msg.session_key
|
||||
key = session_key or msg.session_key_override or f"{destination[0]}:{destination[1]}"
|
||||
else:
|
||||
key = session_key or msg.session_key
|
||||
route = self._turn_route(msg, key)
|
||||
t0 = time.time()
|
||||
ctx = TurnContext(
|
||||
msg=msg,
|
||||
@@ -1402,9 +1329,12 @@ class AgentLoop:
|
||||
state=TurnState.RESTORE,
|
||||
turn_id=f"{key}:{time.time_ns()}",
|
||||
runtime=runtime,
|
||||
kind=kind,
|
||||
route=route,
|
||||
original_user_text=(
|
||||
None
|
||||
if turn_continuation.internal_continuation_inbound(msg.metadata)
|
||||
if kind is TurnKind.SYSTEM
|
||||
or turn_continuation.internal_continuation_inbound(msg.metadata)
|
||||
else msg.content
|
||||
),
|
||||
turn_wall_started_at=t0,
|
||||
@@ -1515,20 +1445,24 @@ class AgentLoop:
|
||||
"""Restore checkpoint / pending user turn; extract documents."""
|
||||
msg = ctx.msg
|
||||
|
||||
if msg.media:
|
||||
if ctx.kind is TurnKind.USER and msg.media:
|
||||
new_content, image_only = self._prepare_message_media(msg.content, msg.media)
|
||||
ctx.msg = dataclasses.replace(msg, content=new_content, media=image_only)
|
||||
msg = ctx.msg
|
||||
|
||||
preview = msg.content[:80] + "..." if len(msg.content) > 80 else msg.content
|
||||
logger.info("Processing message from {}:{}: {}", msg.channel, msg.sender_id, preview)
|
||||
if ctx.kind is TurnKind.SYSTEM:
|
||||
logger.info("Processing system message from {}", msg.sender_id)
|
||||
else:
|
||||
logger.info("Processing message from {}:{}: {}", msg.channel, msg.sender_id, preview)
|
||||
|
||||
# Session is already fetched by the caller (_process_message) but
|
||||
# ensure it exists in case this handler is invoked independently.
|
||||
if ctx.session is None:
|
||||
ctx.session = self.sessions.get_or_create(ctx.session_key)
|
||||
await self._runtime_events().session_turn_started(msg, ctx.session_key)
|
||||
self.workspace_scopes.persist_message_scope(ctx.session, msg)
|
||||
if ctx.kind is TurnKind.USER:
|
||||
await self._runtime_events().session_turn_started(msg, ctx.session_key)
|
||||
self.workspace_scopes.persist_message_scope(ctx.session, msg)
|
||||
|
||||
if self._restore_runtime_checkpoint(ctx.session):
|
||||
self.sessions.save(ctx.session)
|
||||
@@ -1553,6 +1487,8 @@ class AgentLoop:
|
||||
return "ok"
|
||||
|
||||
async def _state_command(self, ctx: TurnContext) -> str:
|
||||
if ctx.kind is TurnKind.SYSTEM:
|
||||
return "dispatch"
|
||||
raw = ctx.msg.content.strip()
|
||||
_, automation_metadata = automation_history_overrides(ctx.msg.metadata)
|
||||
is_user_turn = (
|
||||
@@ -1601,14 +1537,19 @@ class AgentLoop:
|
||||
runtime=ctx.runtime,
|
||||
replay_max_messages=replay_max_messages,
|
||||
)
|
||||
if message_tool := self.tools.get("message"):
|
||||
is_subagent = ctx.kind is TurnKind.SYSTEM and ctx.msg.sender_id == "subagent"
|
||||
if is_subagent and self._persist_subagent_followup(ctx.session, ctx.msg):
|
||||
logger.debug("Subagent result persisted for session {}", ctx.session_key)
|
||||
self.sessions.save(ctx.session)
|
||||
|
||||
if ctx.kind is TurnKind.USER and (message_tool := self.tools.get("message")):
|
||||
if isinstance(message_tool, MessageTool):
|
||||
message_tool.start_turn()
|
||||
|
||||
_hist_kwargs: dict[str, Any] = {
|
||||
"max_messages": replay_max_messages,
|
||||
"max_tokens": self._replay_token_budget(ctx.runtime),
|
||||
"extend_to_user": False,
|
||||
"extend_to_user": is_subagent,
|
||||
}
|
||||
ctx.history = ctx.session.get_history(**_hist_kwargs)
|
||||
self._runtime_events().record_turn_runtime(
|
||||
@@ -1617,37 +1558,33 @@ class AgentLoop:
|
||||
)
|
||||
|
||||
ctx.request_context = self._request_context_for_turn(ctx)
|
||||
ctx.runtime_context_blocks = await self._resolve_runtime_context_for_turn(ctx)
|
||||
ctx.initial_messages = self._build_initial_messages(
|
||||
ctx.msg,
|
||||
ctx.session,
|
||||
ctx.history,
|
||||
ctx.pending_summary,
|
||||
include_memory_recent_history=not ctx.ephemeral,
|
||||
runtime_context_blocks=ctx.runtime_context_blocks,
|
||||
)
|
||||
ctx.user_persisted_early = self._persist_user_message_early(
|
||||
ctx.msg,
|
||||
ctx.session,
|
||||
runtime_context_blocks=ctx.runtime_context_blocks,
|
||||
)
|
||||
if ctx.kind is TurnKind.USER:
|
||||
ctx.runtime_context_blocks = await self._resolve_runtime_context_for_turn(ctx)
|
||||
ctx.initial_messages = self._build_initial_messages(ctx)
|
||||
if ctx.kind is TurnKind.USER:
|
||||
ctx.user_persisted_early = self._persist_user_message_early(
|
||||
ctx.msg,
|
||||
ctx.session,
|
||||
runtime_context_blocks=ctx.runtime_context_blocks,
|
||||
)
|
||||
|
||||
if ctx.on_progress is None:
|
||||
ctx.on_progress = await self._build_bus_progress_callback(ctx.msg)
|
||||
if ctx.on_retry_wait is None:
|
||||
ctx.on_retry_wait = await self._build_retry_wait_callback(ctx.msg)
|
||||
if ctx.on_progress is None:
|
||||
ctx.on_progress = await self._build_bus_progress_callback(ctx.msg)
|
||||
if ctx.on_retry_wait is None:
|
||||
ctx.on_retry_wait = await self._build_retry_wait_callback(ctx.msg)
|
||||
|
||||
return "ok"
|
||||
|
||||
async def _state_run(self, ctx: TurnContext) -> str:
|
||||
if ctx.visible_run_started_at is None:
|
||||
ctx.visible_run_started_at = time.time()
|
||||
await self._runtime_events().run_status_changed(
|
||||
ctx.msg,
|
||||
ctx.session_key,
|
||||
"running",
|
||||
started_at=ctx.visible_run_started_at,
|
||||
)
|
||||
if ctx.kind is TurnKind.USER:
|
||||
await self._runtime_events().run_status_changed(
|
||||
ctx.msg,
|
||||
ctx.session_key,
|
||||
"running",
|
||||
started_at=ctx.visible_run_started_at,
|
||||
)
|
||||
result = await self._run_agent_loop(
|
||||
ctx.initial_messages,
|
||||
runtime=ctx.runtime,
|
||||
@@ -1656,8 +1593,8 @@ class AgentLoop:
|
||||
on_stream_end=ctx.on_stream_end,
|
||||
on_retry_wait=ctx.on_retry_wait,
|
||||
session=ctx.session,
|
||||
channel=ctx.msg.channel,
|
||||
chat_id=ctx.msg.chat_id,
|
||||
channel=ctx.route.channel,
|
||||
chat_id=ctx.route.chat_id,
|
||||
message_id=ctx.msg.metadata.get("message_id"),
|
||||
metadata=ctx.msg.metadata,
|
||||
session_key=ctx.session_key,
|
||||
@@ -1677,21 +1614,26 @@ class AgentLoop:
|
||||
ctx.all_messages = all_msgs
|
||||
ctx.stop_reason = stop_reason
|
||||
ctx.had_injections = had_injections
|
||||
await turn_continuation.maybe_continue_turn(ctx)
|
||||
if ctx.kind is TurnKind.USER:
|
||||
await turn_continuation.maybe_continue_turn(ctx)
|
||||
return "ok"
|
||||
|
||||
async def _state_save(self, ctx: TurnContext) -> str:
|
||||
turn_continuation.prepare_save_boundary(ctx)
|
||||
|
||||
if (
|
||||
(ctx.final_content is None or not ctx.final_content.strip())
|
||||
ctx.kind is TurnKind.USER
|
||||
and (ctx.final_content is None or not ctx.final_content.strip())
|
||||
and not ctx.suppress_response
|
||||
):
|
||||
ctx.final_content = EMPTY_FINAL_RESPONSE_MESSAGE
|
||||
|
||||
latency_started_at = (
|
||||
ctx.visible_run_started_at
|
||||
if turn_continuation.internal_continuation_inbound(ctx.msg.metadata)
|
||||
if (
|
||||
ctx.kind is TurnKind.SYSTEM
|
||||
or turn_continuation.internal_continuation_inbound(ctx.msg.metadata)
|
||||
)
|
||||
and ctx.visible_run_started_at is not None
|
||||
else ctx.turn_wall_started_at
|
||||
)
|
||||
@@ -1726,6 +1668,14 @@ class AgentLoop:
|
||||
if ctx.suppress_response:
|
||||
ctx.outbound = None
|
||||
return "ok"
|
||||
if ctx.kind is TurnKind.SYSTEM:
|
||||
ctx.outbound = OutboundMessage(
|
||||
channel=ctx.route.channel,
|
||||
chat_id=ctx.route.chat_id,
|
||||
content=ctx.final_content or "Background task completed.",
|
||||
metadata=dict(ctx.route.metadata),
|
||||
)
|
||||
return "ok"
|
||||
ctx.outbound = self._assemble_outbound(
|
||||
ctx.msg,
|
||||
ctx.final_content,
|
||||
@@ -1985,7 +1935,9 @@ class AgentLoop:
|
||||
persist_user_message: bool = True,
|
||||
runtime: LLMRuntime | None = None,
|
||||
) -> OutboundMessage | None:
|
||||
"""Process a message directly and return the outbound payload."""
|
||||
"""Process an external message directly and return the outbound payload."""
|
||||
if channel == "system":
|
||||
raise ValueError("channel 'system' is reserved for internal messages")
|
||||
await self._connect_mcp()
|
||||
metadata: dict[str, Any] = {}
|
||||
if not persist_user_message:
|
||||
|
||||
Reference in New Issue
Block a user