refactor(agent): capture runtime at turn admission

This commit is contained in:
chengyongru
2026-07-10 17:54:34 +08:00
committed by Xubin Ren
parent 45ace23580
commit 198fd9f869
13 changed files with 209 additions and 54 deletions
+96 -27
View File
@@ -23,6 +23,7 @@ from nanobot.agent.context import ContextBuilder
from nanobot.agent.cron_turns import CronTurnCoordinator
from nanobot.agent.hook import AgentHook, AgentTurnHookFactory
from nanobot.agent.memory import Consolidator
from nanobot.agent.model_runtime import ModelRuntimeResolver
from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRunSpec
from nanobot.agent.subagent import SubagentManager
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
@@ -113,6 +114,7 @@ class TurnContext:
session_key: str
state: TurnState
turn_id: str
runtime: LLMRuntime
original_user_text: str | None = None
session: Session | None = None
@@ -173,15 +175,37 @@ class AgentLoop:
return self.tools.tool_names
def llm_runtime(self) -> LLMRuntime:
"""Capture the current provider/model settings owned by this loop."""
"""Resolve the immutable default used to admit the next turn."""
self._refresh_provider_snapshot()
return LLMRuntime.capture(
runtime = self.runtime_resolver.current()
captured = LLMRuntime.capture(
self.provider,
self.model,
context_window_tokens=self.context_window_tokens,
model_preset=self.model_preset,
model_preset=self._active_preset,
snapshot_signature=self._provider_signature,
)
# Temporary compatibility for MyTool's legacy direct mutations. Round 9
# moves those writes behind the resolver and deletes these projections.
if (
runtime.provider is not self.provider
or runtime.model != self.model
or runtime.generation != captured.generation
or runtime.context_window_tokens != self.context_window_tokens
or runtime.model_preset != self._active_preset
):
snapshot = ProviderSnapshot(
provider=self.provider,
model=self.model,
context_window_tokens=self.context_window_tokens,
signature=self._provider_signature or ("legacy_loop_runtime", self.model),
generation=captured.generation,
)
runtime = self.runtime_resolver.adopt_snapshot(
snapshot,
model_preset=self._active_preset,
)
return runtime
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
_PENDING_USER_TURN_KEY = "pending_user_turn"
@@ -263,6 +287,19 @@ class AgentLoop:
if context_window_tokens is not None
else defaults.context_window_tokens
)
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {}
self._active_preset: str | None = None
self.runtime_resolver = ModelRuntimeResolver(
LLMRuntime.capture(
provider,
self.model,
context_window_tokens=self.context_window_tokens,
snapshot_signature=provider_signature,
),
model_presets=self.model_presets,
provider_snapshot_loader=provider_snapshot_loader,
preset_snapshot_loader=preset_snapshot_loader,
)
self.context_block_limit = context_block_limit
self.max_tool_result_chars = (
max_tool_result_chars
@@ -369,8 +406,6 @@ class AgentLoop:
consolidator=self.consolidator,
session_ttl_minutes=session_ttl_minutes,
)
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {}
self._active_preset: str | None = None
if model_preset:
self.set_model_preset(model_preset, publish_update=False)
self._register_default_tools()
@@ -448,11 +483,14 @@ class AgentLoop:
model_preset: str | None = None,
) -> None:
"""Swap model/provider for future turns without disturbing an active one."""
provider = snapshot.provider
model = snapshot.model
context_window_tokens = snapshot.context_window_tokens
if snapshot.generation is not None:
provider.generation = snapshot.generation
runtime = self.runtime_resolver.adopt_snapshot(
snapshot,
model_preset=model_preset,
)
provider = runtime.provider
model = runtime.model
context_window_tokens = runtime.context_window_tokens
provider.generation = runtime.generation
old_model = self.model
self.provider = provider
self.model = model
@@ -519,7 +557,14 @@ class AgentLoop:
def set_model_preset(self, name: str | None, *, publish_update: bool = True) -> None:
"""Resolve a preset by name and apply all runtime model dependents."""
name = preset_helpers.normalize_preset_name(name, self.model_presets)
snapshot = self._build_model_preset_snapshot(name)
runtime = self.runtime_resolver.select_preset(name)
snapshot = ProviderSnapshot(
provider=runtime.provider,
model=runtime.model,
context_window_tokens=runtime.context_window_tokens,
signature=runtime.snapshot_signature or ("model_preset", name),
generation=runtime.generation,
)
self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name)
self._active_preset = name
@@ -696,17 +741,18 @@ class AgentLoop:
return UNIFIED_SESSION_KEY
return msg.session_key
def _replay_token_budget(self) -> int:
@staticmethod
def _replay_token_budget(runtime: LLMRuntime) -> int:
"""Derive a token budget for session history replay from the context window."""
if self.context_window_tokens <= 0:
if runtime.context_window_tokens <= 0:
return 0
max_output = getattr(getattr(self.provider, "generation", None), "max_tokens", 4096)
max_output = runtime.generation.max_tokens
try:
reserved_output = int(max_output)
except (TypeError, ValueError):
reserved_output = 4096
budget = self.context_window_tokens - max(1, reserved_output) - 1024
return budget if budget > 0 else max(128, self.context_window_tokens // 2)
budget = runtime.context_window_tokens - max(1, reserved_output) - 1024
return budget if budget > 0 else max(128, runtime.context_window_tokens // 2)
async def _run_agent_loop(
self,
@@ -716,6 +762,7 @@ class AgentLoop:
on_stream_end: Callable[..., Awaitable[None]] | None = None,
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
*,
runtime: LLMRuntime,
session: Session | None = None,
channel: str = "cli",
chat_id: str = "direct",
@@ -823,6 +870,7 @@ class AgentLoop:
message_id=message_id,
session_key=active_session_key,
original_user_text=original_user_text,
runtime=runtime,
metadata=dict(metadata or {}),
)
file_state_token = bind_file_states(self._file_state_store.for_session(active_session_key))
@@ -864,7 +912,7 @@ class AgentLoop:
result = await self.runner.run(AgentRunSpec(
initial_messages=initial_messages,
tools=effective_tools,
runtime=self.llm_runtime(),
runtime=runtime,
max_iterations=self.max_iterations,
max_tool_result_chars=self.max_tool_result_chars,
hook=hook,
@@ -1200,6 +1248,8 @@ class AgentLoop:
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,
@@ -1214,6 +1264,7 @@ class AgentLoop:
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):
@@ -1225,7 +1276,9 @@ class AgentLoop:
await self.consolidator.maybe_consolidate_by_tokens(
session,
replay_max_messages=self._max_messages,
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):
@@ -1233,8 +1286,8 @@ class AgentLoop:
self.sessions.save(session)
current_role = "assistant" if is_subagent else "user"
_hist_kwargs: dict[str, Any] = {
"max_messages": self._max_messages,
"max_tokens": self._replay_token_budget(),
"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)
@@ -1259,6 +1312,7 @@ class AgentLoop:
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,
@@ -1278,7 +1332,9 @@ class AgentLoop:
self._schedule_background(
self.consolidator.maybe_consolidate_by_tokens(
session,
replay_max_messages=self._max_messages,
replay_max_messages=replay_max_messages_for_context(
runtime.context_window_tokens
),
)
)
content = final_content or "Background task completed."
@@ -1307,13 +1363,16 @@ class AgentLoop:
hooks: list[AgentHook] | None = None,
hook_factories: list[AgentTurnHookFactory] | None = None,
tools: ToolRegistry | None = None,
runtime: LLMRuntime | None = None,
) -> OutboundMessage | None:
"""Process a single inbound message and return the response."""
self._refresh_provider_snapshot()
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,
@@ -1330,6 +1389,7 @@ class AgentLoop:
session_key=key,
state=TurnState.RESTORE,
turn_id=f"{key}:{time.time_ns()}",
runtime=runtime,
original_user_text=(
None
if turn_continuation.internal_continuation_inbound(msg.metadata)
@@ -1506,24 +1566,27 @@ class AgentLoop:
return "dispatch"
async def _state_build(self, ctx: TurnContext) -> str:
replay_max_messages = replay_max_messages_for_context(
ctx.runtime.context_window_tokens
)
if not ctx.ephemeral:
await self.consolidator.maybe_consolidate_by_tokens(
ctx.session,
replay_max_messages=self._max_messages,
replay_max_messages=replay_max_messages,
)
if message_tool := self.tools.get("message"):
if isinstance(message_tool, MessageTool):
message_tool.start_turn()
_hist_kwargs: dict[str, Any] = {
"max_messages": self._max_messages,
"max_tokens": self._replay_token_budget(),
"max_messages": replay_max_messages,
"max_tokens": self._replay_token_budget(ctx.runtime),
"extend_to_user": False,
}
ctx.history = ctx.session.get_history(**_hist_kwargs)
self._runtime_events().record_turn_runtime(
ctx.session_key,
self.llm_runtime(),
ctx.runtime,
)
ctx.initial_messages = self._build_initial_messages(
@@ -1555,6 +1618,7 @@ class AgentLoop:
)
result = await self._run_agent_loop(
ctx.initial_messages,
runtime=ctx.runtime,
on_progress=ctx.on_progress,
on_stream=ctx.on_stream,
on_stream_end=ctx.on_stream_end,
@@ -1613,7 +1677,9 @@ class AgentLoop:
self._schedule_background(
self.consolidator.maybe_consolidate_by_tokens(
ctx.session,
replay_max_messages=self._max_messages,
replay_max_messages=replay_max_messages_for_context(
ctx.runtime.context_window_tokens
),
)
)
self._clear_pending_user_turn(ctx.session)
@@ -1891,6 +1957,7 @@ class AgentLoop:
hook_factories: list[AgentTurnHookFactory] | None = None,
tools: ToolRegistry | None = None,
persist_user_message: bool = True,
runtime: LLMRuntime | None = None,
) -> OutboundMessage | None:
"""Process a message directly and return the outbound payload."""
await self._connect_mcp()
@@ -1920,6 +1987,8 @@ class AgentLoop:
kwargs["hook_factories"] = hook_factories
if tools is not None:
kwargs["tools"] = tools
if runtime is not None:
kwargs["runtime"] = runtime
return await self._process_message(
msg,
**kwargs,