refactor(agent): unify internal turn lifecycle (#4993)

This commit is contained in:
chengyongru
2026-07-21 00:14:27 +08:00
committed by GitHub
parent b67f4b1371
commit 9db0d9f3c9
3 changed files with 275 additions and 185 deletions
+132 -180
View File
@@ -110,6 +110,20 @@ class TurnState(Enum):
DONE = auto() 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 @dataclass
class StateTraceEntry: class StateTraceEntry:
state: TurnState state: TurnState
@@ -126,6 +140,8 @@ class TurnContext:
state: TurnState state: TurnState
turn_id: str turn_id: str
runtime: LLMRuntime runtime: LLMRuntime
kind: TurnKind
route: TurnRoute
original_user_text: str | None = None original_user_text: str | None = None
session: Session | None = None session: Session | None = None
@@ -576,9 +592,28 @@ class AgentLoop:
self._runtime_context_providers.append(provider) self._runtime_context_providers.append(provider)
@staticmethod @staticmethod
def _runtime_chat_id(msg: InboundMessage) -> str: def _turn_route(msg: InboundMessage, session_key: str) -> TurnRoute:
"""Return the chat id shown in runtime metadata for the model.""" """Resolve response routing without mixing it into execution metadata."""
return str(msg.metadata.get("context_chat_id") or msg.chat_id) 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( async def _build_bus_progress_callback(
self, msg: InboundMessage self, msg: InboundMessage
@@ -660,38 +695,38 @@ class AgentLoop:
return True return True
return False return False
def _build_initial_messages( def _build_initial_messages(self, ctx: TurnContext) -> list[dict[str, Any]]:
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]]:
"""Build the initial message list for the LLM turn.""" """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( return self.context.build_messages(
history=history, history=ctx.history,
current_message=msg.content, current_message="" if is_subagent else ctx.msg.content,
media=msg.media if msg.media else None, media=ctx.msg.media if ctx.kind is TurnKind.USER and ctx.msg.media else None,
channel=msg.channel, channel=ctx.route.channel,
chat_id=self._runtime_chat_id(msg), chat_id=str(ctx.msg.metadata.get("context_chat_id") or ctx.route.chat_id),
sender_id=msg.sender_id, current_role="assistant" if is_subagent else "user",
session_summary=pending_summary, sender_id=ctx.msg.sender_id,
session_metadata=session.metadata, session_summary=ctx.pending_summary,
session_metadata=ctx.session.metadata,
workspace=scope.project_path, workspace=scope.project_path,
runtime_context_blocks=runtime_context_blocks, runtime_context_blocks=ctx.runtime_context_blocks,
include_memory_recent_history=include_memory_recent_history, include_memory_recent_history=not ctx.ephemeral,
session_key=session.key, session_key=ctx.session.key,
unified_session=self._unified_session, unified_session=self._unified_session,
) )
def _request_context_for_turn(self, ctx: TurnContext) -> RequestContext: 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( return RequestContext(
channel=ctx.msg.channel, channel=ctx.route.channel,
chat_id=ctx.msg.chat_id, chat_id=ctx.route.chat_id,
message_id=ctx.msg.metadata.get("message_id"), message_id=ctx.msg.metadata.get("message_id"),
session_key=ctx.session_key, session_key=ctx.session_key,
original_user_text=ctx.original_user_text, original_user_text=ctx.original_user_text,
@@ -1258,110 +1293,6 @@ class AgentLoop:
self._running = False self._running = False
logger.info("Agent loop stopping") 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( async def _process_message(
self, self,
msg: InboundMessage, msg: InboundMessage,
@@ -1381,19 +1312,15 @@ class AgentLoop:
if runtime is None: if runtime is None:
runtime = self.llm_runtime() runtime = self.llm_runtime()
if msg.channel == "system": kind = TurnKind.SYSTEM if msg.channel == "system" else TurnKind.USER
return await self._process_system_message( if kind is TurnKind.SYSTEM:
msg, destination = (
runtime=runtime, msg.chat_id.split(":", 1) if ":" in msg.chat_id else ("cli", msg.chat_id)
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,
) )
key = session_key or msg.session_key_override or f"{destination[0]}:{destination[1]}"
key = session_key or msg.session_key else:
key = session_key or msg.session_key
route = self._turn_route(msg, key)
t0 = time.time() t0 = time.time()
ctx = TurnContext( ctx = TurnContext(
msg=msg, msg=msg,
@@ -1402,9 +1329,12 @@ class AgentLoop:
state=TurnState.RESTORE, state=TurnState.RESTORE,
turn_id=f"{key}:{time.time_ns()}", turn_id=f"{key}:{time.time_ns()}",
runtime=runtime, runtime=runtime,
kind=kind,
route=route,
original_user_text=( original_user_text=(
None 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 else msg.content
), ),
turn_wall_started_at=t0, turn_wall_started_at=t0,
@@ -1515,20 +1445,24 @@ class AgentLoop:
"""Restore checkpoint / pending user turn; extract documents.""" """Restore checkpoint / pending user turn; extract documents."""
msg = ctx.msg 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) new_content, image_only = self._prepare_message_media(msg.content, msg.media)
ctx.msg = dataclasses.replace(msg, content=new_content, media=image_only) ctx.msg = dataclasses.replace(msg, content=new_content, media=image_only)
msg = ctx.msg msg = ctx.msg
preview = msg.content[:80] + "..." if len(msg.content) > 80 else msg.content 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 # Session is already fetched by the caller (_process_message) but
# ensure it exists in case this handler is invoked independently. # ensure it exists in case this handler is invoked independently.
if ctx.session is None: if ctx.session is None:
ctx.session = self.sessions.get_or_create(ctx.session_key) ctx.session = self.sessions.get_or_create(ctx.session_key)
await self._runtime_events().session_turn_started(msg, ctx.session_key) if ctx.kind is TurnKind.USER:
self.workspace_scopes.persist_message_scope(ctx.session, msg) 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): if self._restore_runtime_checkpoint(ctx.session):
self.sessions.save(ctx.session) self.sessions.save(ctx.session)
@@ -1553,6 +1487,8 @@ class AgentLoop:
return "ok" return "ok"
async def _state_command(self, ctx: TurnContext) -> str: async def _state_command(self, ctx: TurnContext) -> str:
if ctx.kind is TurnKind.SYSTEM:
return "dispatch"
raw = ctx.msg.content.strip() raw = ctx.msg.content.strip()
_, automation_metadata = automation_history_overrides(ctx.msg.metadata) _, automation_metadata = automation_history_overrides(ctx.msg.metadata)
is_user_turn = ( is_user_turn = (
@@ -1601,14 +1537,19 @@ class AgentLoop:
runtime=ctx.runtime, runtime=ctx.runtime,
replay_max_messages=replay_max_messages, 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): if isinstance(message_tool, MessageTool):
message_tool.start_turn() message_tool.start_turn()
_hist_kwargs: dict[str, Any] = { _hist_kwargs: dict[str, Any] = {
"max_messages": replay_max_messages, "max_messages": replay_max_messages,
"max_tokens": self._replay_token_budget(ctx.runtime), "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) ctx.history = ctx.session.get_history(**_hist_kwargs)
self._runtime_events().record_turn_runtime( self._runtime_events().record_turn_runtime(
@@ -1617,37 +1558,33 @@ class AgentLoop:
) )
ctx.request_context = self._request_context_for_turn(ctx) ctx.request_context = self._request_context_for_turn(ctx)
ctx.runtime_context_blocks = await self._resolve_runtime_context_for_turn(ctx) if ctx.kind is TurnKind.USER:
ctx.initial_messages = self._build_initial_messages( ctx.runtime_context_blocks = await self._resolve_runtime_context_for_turn(ctx)
ctx.msg, ctx.initial_messages = self._build_initial_messages(ctx)
ctx.session, if ctx.kind is TurnKind.USER:
ctx.history, ctx.user_persisted_early = self._persist_user_message_early(
ctx.pending_summary, ctx.msg,
include_memory_recent_history=not ctx.ephemeral, ctx.session,
runtime_context_blocks=ctx.runtime_context_blocks, 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.on_progress is None: if ctx.on_progress is None:
ctx.on_progress = await self._build_bus_progress_callback(ctx.msg) ctx.on_progress = await self._build_bus_progress_callback(ctx.msg)
if ctx.on_retry_wait is None: if ctx.on_retry_wait is None:
ctx.on_retry_wait = await self._build_retry_wait_callback(ctx.msg) ctx.on_retry_wait = await self._build_retry_wait_callback(ctx.msg)
return "ok" return "ok"
async def _state_run(self, ctx: TurnContext) -> str: async def _state_run(self, ctx: TurnContext) -> str:
if ctx.visible_run_started_at is None: if ctx.visible_run_started_at is None:
ctx.visible_run_started_at = time.time() ctx.visible_run_started_at = time.time()
await self._runtime_events().run_status_changed( if ctx.kind is TurnKind.USER:
ctx.msg, await self._runtime_events().run_status_changed(
ctx.session_key, ctx.msg,
"running", ctx.session_key,
started_at=ctx.visible_run_started_at, "running",
) started_at=ctx.visible_run_started_at,
)
result = await self._run_agent_loop( result = await self._run_agent_loop(
ctx.initial_messages, ctx.initial_messages,
runtime=ctx.runtime, runtime=ctx.runtime,
@@ -1656,8 +1593,8 @@ class AgentLoop:
on_stream_end=ctx.on_stream_end, on_stream_end=ctx.on_stream_end,
on_retry_wait=ctx.on_retry_wait, on_retry_wait=ctx.on_retry_wait,
session=ctx.session, session=ctx.session,
channel=ctx.msg.channel, channel=ctx.route.channel,
chat_id=ctx.msg.chat_id, chat_id=ctx.route.chat_id,
message_id=ctx.msg.metadata.get("message_id"), message_id=ctx.msg.metadata.get("message_id"),
metadata=ctx.msg.metadata, metadata=ctx.msg.metadata,
session_key=ctx.session_key, session_key=ctx.session_key,
@@ -1677,21 +1614,26 @@ class AgentLoop:
ctx.all_messages = all_msgs ctx.all_messages = all_msgs
ctx.stop_reason = stop_reason ctx.stop_reason = stop_reason
ctx.had_injections = had_injections 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" return "ok"
async def _state_save(self, ctx: TurnContext) -> str: async def _state_save(self, ctx: TurnContext) -> str:
turn_continuation.prepare_save_boundary(ctx) turn_continuation.prepare_save_boundary(ctx)
if ( 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 and not ctx.suppress_response
): ):
ctx.final_content = EMPTY_FINAL_RESPONSE_MESSAGE ctx.final_content = EMPTY_FINAL_RESPONSE_MESSAGE
latency_started_at = ( latency_started_at = (
ctx.visible_run_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 and ctx.visible_run_started_at is not None
else ctx.turn_wall_started_at else ctx.turn_wall_started_at
) )
@@ -1726,6 +1668,14 @@ class AgentLoop:
if ctx.suppress_response: if ctx.suppress_response:
ctx.outbound = None ctx.outbound = None
return "ok" 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.outbound = self._assemble_outbound(
ctx.msg, ctx.msg,
ctx.final_content, ctx.final_content,
@@ -1985,7 +1935,9 @@ class AgentLoop:
persist_user_message: bool = True, persist_user_message: bool = True,
runtime: LLMRuntime | None = None, runtime: LLMRuntime | None = None,
) -> OutboundMessage | 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() await self._connect_mcp()
metadata: dict[str, Any] = {} metadata: dict[str, Any] = {}
if not persist_user_message: if not persist_user_message:
@@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from nanobot.agent.loop import AgentLoop, TurnContext, TurnState from nanobot.agent.loop import AgentLoop, TurnContext, TurnKind, TurnRoute, TurnState
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.schema import ChannelsConfig from nanobot.config.schema import ChannelsConfig
@@ -54,6 +54,8 @@ async def test_state_restore_extracts_documents_by_default(
state=TurnState.RESTORE, state=TurnState.RESTORE,
turn_id="turn-1", turn_id="turn-1",
runtime=loop.llm_runtime(), runtime=loop.llm_runtime(),
kind=TurnKind.USER,
route=TurnRoute(channel="cli", chat_id="c"),
) )
assert await loop._state_restore(ctx) == "ok" assert await loop._state_restore(ctx) == "ok"
@@ -89,6 +91,8 @@ async def test_state_restore_references_documents_when_extraction_disabled(
state=TurnState.RESTORE, state=TurnState.RESTORE,
turn_id="turn-1", turn_id="turn-1",
runtime=loop.llm_runtime(), runtime=loop.llm_runtime(),
kind=TurnKind.USER,
route=TurnRoute(channel="cli", chat_id="c"),
) )
assert await loop._state_restore(ctx) == "ok" assert await loop._state_restore(ctx) == "ok"
+138 -4
View File
@@ -4,9 +4,10 @@ from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from loguru import logger
from nanobot.agent.context import ContextBuilder from nanobot.agent.context import ContextBuilder
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop, TurnState
from nanobot.agent.tools.context import RequestContext, request_context from nanobot.agent.tools.context import RequestContext, request_context
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
from nanobot.bus.outbound_events import ( from nanobot.bus.outbound_events import (
@@ -1143,6 +1144,19 @@ async def test_run_agent_loop_goal_continue_message_reads_latest_metadata(
assert "Goal created during this runner call." in (seen["goal_continue"] or "") assert "Goal created during this runner call." in (seen["goal_continue"] or "")
@pytest.mark.asyncio
async def test_process_direct_rejects_reserved_system_channel(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path)
loop._connect_mcp = AsyncMock() # type: ignore[method-assign]
loop._process_message = AsyncMock(return_value=None) # type: ignore[method-assign]
with pytest.raises(ValueError, match="reserved for internal messages"):
await loop.process_direct("external input", channel="system")
loop._connect_mcp.assert_not_awaited()
loop._process_message.assert_not_awaited()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_process_direct_skip_user_persist_does_not_save_retry_user( async def test_process_direct_skip_user_persist_does_not_save_retry_user(
tmp_path: Path, tmp_path: Path,
@@ -1363,6 +1377,7 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
async def fake_run_agent_loop(initial_messages, **kwargs): async def fake_run_agent_loop(initial_messages, **kwargs):
seen["initial_messages"] = initial_messages seen["initial_messages"] = initial_messages
seen["runtime"] = kwargs["runtime"] seen["runtime"] = kwargs["runtime"]
seen["request_context"] = kwargs["request_context"]
return ( return (
"done", "done",
[], [],
@@ -1385,6 +1400,15 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
) )
assert seen["runtime"] is runtime assert seen["runtime"] is runtime
request = seen["request_context"]
assert isinstance(request, RequestContext)
assert request.channel == "cli"
assert request.chat_id == "test"
assert request.session_key == "cli:test"
assert request.original_user_text is None
assert request.sender_id == "subagent"
assert request.metadata == {"subagent_task_id": "sub-1"}
assert request.turn_id
record_runtime.assert_called_once_with("cli:test", runtime) record_runtime.assert_called_once_with("cli:test", runtime)
assert len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2 assert len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2
assert all( assert all(
@@ -1421,6 +1445,104 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
] ]
@pytest.mark.asyncio
async def test_system_subagent_followup_does_not_log_content(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock( # type: ignore[method-assign]
return_value=False
)
async def fake_run_agent_loop(initial_messages, **_kwargs):
return (
"done",
[],
[*initial_messages, {"role": "assistant", "content": "done"}],
"stop",
False,
)
loop._run_agent_loop = fake_run_agent_loop # type: ignore[method-assign]
secret = "LEAKME42"
content = f"[Subagent 'research' completed]\n\nTask: inspect logs\n\nResult:\n{secret}"
logs: list[str] = []
sink_id = logger.add(logs.append, level="INFO", format="{message}")
try:
await loop._process_message(
InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:logs",
content=content,
metadata={"subagent_task_id": "sub-logs"},
)
)
finally:
logger.remove(sink_id)
logged = "".join(logs)
assert "Processing system message from subagent" in logged
assert secret not in logged
@pytest.mark.asyncio
async def test_system_subagent_followup_uses_common_turn_state_machine(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock( # type: ignore[method-assign]
return_value=False
)
visited: list[TurnState] = []
for state in (
TurnState.RESTORE,
TurnState.COMPACT,
TurnState.COMMAND,
TurnState.BUILD,
TurnState.RUN,
TurnState.SAVE,
TurnState.RESPOND,
):
name = f"_state_{state.name.lower()}"
original = getattr(loop, name)
async def record(ctx, *, _original=original, _state=state):
visited.append(_state)
return await _original(ctx)
setattr(loop, name, record)
async def fake_run_agent_loop(initial_messages, **_kwargs):
return (
"done",
[],
[*initial_messages, {"role": "assistant", "content": "done"}],
"stop",
False,
)
loop._run_agent_loop = fake_run_agent_loop # type: ignore[method-assign]
await loop._process_message(
InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:test",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
)
assert visited == [
TurnState.RESTORE,
TurnState.COMPACT,
TurnState.COMMAND,
TurnState.BUILD,
TurnState.RUN,
TurnState.SAVE,
TurnState.RESPOND,
]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_multiple_subagent_followups_all_persist_as_standalone_history(tmp_path: Path) -> None: async def test_multiple_subagent_followups_all_persist_as_standalone_history(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path) loop = _make_full_loop(tmp_path)
@@ -1556,10 +1678,11 @@ async def test_system_subagent_followup_uses_thread_session_and_slack_metadata(t
thread_session.add_message("user", "thread question") thread_session.add_message("user", "thread question")
loop.sessions.save(thread_session) loop.sessions.save(thread_session)
seen: dict[str, list[dict]] = {} seen: dict[str, object] = {}
async def fake_run_agent_loop(initial_messages, **_kwargs): async def fake_run_agent_loop(initial_messages, **kwargs):
seen["initial_messages"] = initial_messages seen["initial_messages"] = initial_messages
seen["request_context"] = kwargs["request_context"]
return ( return (
"done", "done",
[], [],
@@ -1588,7 +1711,18 @@ async def test_system_subagent_followup_uses_thread_session_and_slack_metadata(t
"slack": {"thread_ts": "1700.42"}, "slack": {"thread_ts": "1700.42"},
"origin_message_id": "msg-123", "origin_message_id": "msg-123",
} }
assert "thread question" in seen["initial_messages"][1]["content"] request = seen["request_context"]
assert isinstance(request, RequestContext)
assert request.channel == "slack"
assert request.chat_id == "C123"
assert request.metadata == {
"subagent_task_id": "sub-1",
"origin_message_id": "msg-123",
}
assert "slack" not in request.metadata
initial_messages = seen["initial_messages"]
assert isinstance(initial_messages, list)
assert "thread question" in initial_messages[1]["content"]
loop.sessions.invalidate("slack:C123:1700.42") loop.sessions.invalidate("slack:C123:1700.42")
persisted = loop.sessions.get_or_create("slack:C123:1700.42") persisted = loop.sessions.get_or_create("slack:C123:1700.42")