refactor(agent): make resolver sole runtime owner

This commit is contained in:
chengyongru
2026-07-10 17:54:34 +08:00
committed by Xubin Ren
parent c9d3e74342
commit 21f58cbabf
16 changed files with 395 additions and 443 deletions
+83 -125
View File
@@ -6,6 +6,7 @@ import asyncio
import dataclasses
import os
import time
from collections.abc import Mapping
from contextlib import AsyncExitStack, nullcontext, suppress
from dataclasses import dataclass, field
from enum import Enum, auto
@@ -174,37 +175,48 @@ class AgentLoop:
def tool_names(self) -> list[str]:
return self.tools.tool_names
@property
def provider(self) -> LLMProvider:
"""Provider selected for future turn admissions."""
return self.runtime_resolver.runtime.provider
@property
def model(self) -> str:
"""Model selected for future turn admissions."""
return self.runtime_resolver.runtime.model
@property
def context_window_tokens(self) -> int:
"""Context limit selected for future turn admissions."""
return self.runtime_resolver.runtime.context_window_tokens
@property
def model_presets(self) -> Mapping[str, ModelPresetConfig]:
"""Configured model presets exposed for selection and display."""
return self.runtime_resolver.model_presets
@property
def model_preset(self) -> str | None:
return self.runtime_resolver.model_preset
@model_preset.setter
def model_preset(self, name: str | None) -> None:
self.set_model_preset(name)
def llm_runtime(self) -> LLMRuntime:
"""Resolve the immutable default used to admit the next turn."""
self._refresh_provider_snapshot()
runtime = self.runtime_resolver.current()
captured = LLMRuntime.capture(
self.provider,
self.model,
context_window_tokens=self.context_window_tokens,
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.
previous = self.runtime_resolver.runtime
try:
runtime = self.runtime_resolver.current(refresh=True)
except Exception:
logger.exception("Failed to refresh model runtime")
return previous
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
runtime.model != previous.model
or runtime.model_preset != previous.model_preset
or runtime.snapshot_signature != previous.snapshot_signature
):
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,
)
self._publish_runtime_selection(runtime)
return runtime
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
@@ -271,32 +283,26 @@ class AgentLoop:
self.runtime_event_publisher = RuntimeEventPublisher(self.runtime_events)
self.channels_config = channels_config
self.restart_mode = restart_mode
self.provider = provider
self._provider_snapshot_loader = provider_snapshot_loader
self._preset_snapshot_loader = preset_snapshot_loader
self._runtime_model_publisher = runtime_model_publisher
self._provider_signature = provider_signature
self._default_selection_signature = preset_helpers.default_selection_signature(provider_signature)
self.workspace = workspace
self.model = model or provider.get_default_model()
initial_model = model or provider.get_default_model()
self.max_iterations = (
max_iterations if max_iterations is not None else defaults.max_tool_iterations
)
self.context_window_tokens = (
initial_context_window = (
context_window_tokens
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
configured_presets = model_presets or {}
self.runtime_resolver = ModelRuntimeResolver(
LLMRuntime.capture(
provider,
self.model,
context_window_tokens=self.context_window_tokens,
initial_model,
context_window_tokens=initial_context_window,
snapshot_signature=provider_signature,
),
model_presets=self.model_presets,
model_presets=configured_presets,
provider_snapshot_loader=provider_snapshot_loader,
preset_snapshot_loader=preset_snapshot_loader,
)
@@ -352,7 +358,6 @@ class AgentLoop:
llm_wall_timeout_for_session=lambda sk: runner_wall_llm_timeout_s(self.sessions, sk),
)
self._unified_session = unified_session
self._max_messages = replay_max_messages_for_context(self.context_window_tokens)
self._running = False
self._mcp_servers = mcp_servers or {}
self._mcp_stacks: dict[str, AsyncExitStack] = {}
@@ -401,7 +406,7 @@ class AgentLoop:
)
if model_preset:
self.set_model_preset(model_preset, publish_update=False)
self._register_default_tools()
self._register_default_tools(provider_snapshot_loader=provider_snapshot_loader)
self._runtime_vars: dict[str, Any] = {}
self._current_iteration: int = 0
self.commands = CommandRouter()
@@ -468,98 +473,51 @@ class AgentLoop:
"""Keep subagent runtime limits aligned with mutable loop settings."""
self.subagents.max_iterations = self.max_iterations
def _apply_provider_snapshot(
def _publish_runtime_selection(
self,
snapshot: ProviderSnapshot,
runtime: LLMRuntime,
*,
publish_update: bool = True,
model_preset: str | None = None,
) -> None:
"""Swap model/provider for future turns without disturbing an active one."""
runtime = self.runtime_resolver.adopt_snapshot(
snapshot,
model_preset=model_preset,
if not publish_update:
return
if self._runtime_model_publisher is not None:
self._runtime_model_publisher(runtime.model, runtime.model_preset)
self._runtime_events().runtime_model_changed(
runtime.model,
runtime.model_preset,
)
provider = runtime.provider
model = runtime.model
context_window_tokens = runtime.context_window_tokens
provider.generation = runtime.generation
def set_model_preset(
self,
name: str | None,
*,
publish_update: bool = True,
) -> LLMRuntime:
"""Select a named default runtime for future turns."""
old_model = self.model
self.provider = provider
self.model = model
self.context_window_tokens = context_window_tokens
self._sync_replay_max_messages()
self._provider_signature = snapshot.signature
if publish_update and self._runtime_model_publisher is not None:
self._runtime_model_publisher(
self.model,
model_preset if model_preset is not None else self.model_preset,
)
if publish_update:
self._runtime_events().runtime_model_changed(
self.model,
model_preset if model_preset is not None else self.model_preset,
)
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
def _sync_replay_max_messages(self) -> None:
self._max_messages = replay_max_messages_for_context(self.context_window_tokens)
def _refresh_provider_snapshot(self) -> None:
if self._provider_snapshot_loader is None:
return
try:
snapshot = self._provider_snapshot_loader()
except Exception:
logger.exception("Failed to refresh provider config")
return
default_selection = preset_helpers.default_selection_signature(snapshot.signature)
if self._active_preset and self._default_selection_signature in (None, default_selection):
self._default_selection_signature = default_selection
try:
snapshot = self._build_model_preset_snapshot(self._active_preset)
except Exception:
logger.exception("Failed to refresh active model preset")
return
else:
self._active_preset = None
self._default_selection_signature = default_selection
if snapshot.signature == self._provider_signature:
return
self._default_selection_signature = preset_helpers.default_selection_signature(snapshot.signature)
self._apply_provider_snapshot(snapshot)
@property
def model_preset(self) -> str | None:
return self._active_preset
@model_preset.setter
def model_preset(self, name: str | None) -> None:
self.set_model_preset(name)
def _build_model_preset_snapshot(self, name: str) -> ProviderSnapshot:
return preset_helpers.build_runtime_preset_snapshot(
name=name,
presets=self.model_presets,
provider=self.provider,
loader=self._preset_snapshot_loader,
)
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)
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._publish_runtime_selection(runtime, publish_update=publish_update)
logger.info(
"Runtime model switched for next turn: {} -> {}",
old_model,
runtime.model,
)
self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name)
self._active_preset = name
return runtime
def _register_default_tools(self) -> None:
def set_runtime_model(self, model: str) -> LLMRuntime:
"""Select a model on the current provider for future turns."""
return self.runtime_resolver.select_model(model)
def set_runtime_context_window(self, context_window_tokens: int) -> LLMRuntime:
"""Select a context limit for future turns."""
return self.runtime_resolver.select_context_window(context_window_tokens)
def _register_default_tools(
self,
*,
provider_snapshot_loader: Callable[..., ProviderSnapshot] | None,
) -> None:
"""Register the default set of tools via plugin loader."""
from nanobot.agent.tools.context import ToolContext
from nanobot.agent.tools.loader import ToolLoader
@@ -571,7 +529,7 @@ class AgentLoop:
subagent_manager=self.subagents,
cron_service=self.cron_service,
sessions=self.sessions,
provider_snapshot_loader=self._provider_snapshot_loader,
provider_snapshot_loader=provider_snapshot_loader,
image_generation_provider_configs=self._image_generation_provider_configs,
timezone=self.context.timezone or "UTC",
workspace_sandbox=self.workspace_scopes.sandbox_status,