refactor(agent): capture runtime at turn admission
This commit is contained in:
+96
-27
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user