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,
+4 -2
View File
@@ -56,7 +56,9 @@ class RuntimeState(Protocol):
def _sync_subagent_runtime_limits(self) -> None: ...
def set_runtime_model(self, model: str) -> Any: ...
def set_runtime_context_window(self, context_window_tokens: int) -> Any: ...
@property
def model_preset(self) -> str | None: ...
_active_preset: str | None
+33 -9
View File
@@ -57,7 +57,7 @@ class MyTool(Tool):
BLOCKED = frozenset({
# Core infrastructure
"bus", "provider", "_running", "tools",
"bus", "provider", "runtime_resolver", "_running", "tools",
# Config management
"_runtime_vars",
# Subsystems
@@ -107,6 +107,11 @@ class MyTool(Tool):
}
_MAX_RUNTIME_KEYS = 64
_MODEL_RUNTIME_FIELDS = frozenset({
"model",
"model_preset",
"context_window_tokens",
})
def __init__(self, runtime_state: RuntimeState, modify_allowed: bool = True) -> None:
self._runtime_state = runtime_state
@@ -325,9 +330,20 @@ class MyTool(Tool):
# -- inspect --
def _current_runtime_value(self, key: str) -> tuple[bool, Any]:
request_ctx = current_request_context()
runtime = request_ctx.runtime if request_ctx is not None else None
if runtime is None or key not in self._MODEL_RUNTIME_FIELDS:
return False, None
return True, getattr(runtime, key)
def _inspect(self, key: str | None) -> str:
if not key:
return self._inspect_all()
if "." not in key:
found, value = self._current_runtime_value(key)
if found:
return self._format_value(value, key)
top = key.split(".")[0]
if top in self._DENIED_ATTRS or top.startswith("__"):
return ToolResult.error(f"Error: '{top}' is not accessible")
@@ -353,8 +369,13 @@ class MyTool(Tool):
parts: list[str] = []
# RESTRICTED keys
for k in self.RESTRICTED:
parts.append(self._format_value(getattr(state, k, None), k))
parts.append(self._format_value(state.model_preset, "model_preset"))
found, value = self._current_runtime_value(k)
parts.append(self._format_value(value if found else getattr(state, k, None), k))
found, value = self._current_runtime_value("model_preset")
parts.append(self._format_value(
value if found else state.model_preset,
"model_preset",
))
# Other useful top-level keys shown in description
for k in ("workspace", "provider_retry_mode", "max_tool_result_chars", "_current_iteration", "web_config", "exec_config", "workspace_sandbox", "subagents"):
if _has_real_attr(state, k):
@@ -432,13 +453,16 @@ class MyTool(Tool):
return ToolResult.error(f"Error: '{key}' must be <= {spec['max']}")
if "min_len" in spec and len(str(value)) < spec["min_len"]:
return ToolResult.error(f"Error: '{key}' must be at least {spec['min_len']} characters")
setattr(self._runtime_state, key, value)
if key == "model":
self._runtime_state._active_preset = None
sync_replay = getattr(self._runtime_state, "_sync_replay_max_messages", None)
if key == "context_window_tokens" and callable(sync_replay):
sync_replay()
if key == "max_iterations" and hasattr(self._runtime_state, "_sync_subagent_runtime_limits"):
self._runtime_state.set_runtime_model(value)
elif key == "context_window_tokens":
self._runtime_state.set_runtime_context_window(value)
else:
setattr(self._runtime_state, key, value)
if key == "max_iterations" and hasattr(
self._runtime_state,
"_sync_subagent_runtime_limits",
):
self._runtime_state._sync_subagent_runtime_limits()
self._audit("modify", f"{key}: {old!r} -> {value!r}")
return f"Set {key} = {value!r} (was {old!r})"
+5 -5
View File
@@ -349,7 +349,7 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
name = parts[0]
try:
loop.set_model_preset(name)
runtime = loop.set_model_preset(name)
except (KeyError, ValueError) as exc:
names = _model_preset_names(loop)
return OutboundMessage(
@@ -362,11 +362,11 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
metadata=metadata,
)
max_tokens = getattr(getattr(loop.provider, "generation", None), "max_tokens", None)
max_tokens = runtime.generation.max_tokens
lines = [
f"Switched model preset to `{loop.model_preset}`.",
f"- Model: `{loop.model}`",
f"- Context window: {loop.context_window_tokens}",
f"Switched model preset to `{runtime.model_preset}`.",
f"- Model: `{runtime.model}`",
f"- Context window: {runtime.context_window_tokens}",
]
if max_tokens is not None:
lines.append(f"- Max output tokens: {max_tokens}")
+69 -63
View File
@@ -14,7 +14,6 @@ from nanobot.config.schema import Config
from nanobot.providers.image_generation import image_gen_provider_configs
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
from nanobot.sdk.runtime import (
SDKRuntimeController,
build_process_direct_kwargs,
ensure_single_model_selector,
)
@@ -74,7 +73,6 @@ class Nanobot:
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
self._loop = loop
self._config = config
self._runtime_overrides = SDKRuntimeController(loop, config=config)
self.sessions = SessionClient(loop)
self.memory = MemoryClient(loop)
self.runtime = RuntimeClient(loop)
@@ -156,20 +154,26 @@ class Nanobot:
"""
capture = SDKCaptureHook()
per_run_hooks = [capture, *(hooks or [])]
async with self._runtime_overrides.override(model=model, model_preset=model_preset):
kwargs = build_process_direct_kwargs(
session_key=session_key,
channel=channel,
chat_id=chat_id,
sender_id=sender_id,
media=media,
ephemeral=ephemeral,
)
response = await self._loop.process_direct(
message,
**kwargs,
hooks=per_run_hooks,
)
runtime = self._loop.runtime_resolver.resolve_override(
model=model,
model_preset=model_preset,
config=self._config,
)
kwargs = build_process_direct_kwargs(
session_key=session_key,
channel=channel,
chat_id=chat_id,
sender_id=sender_id,
media=media,
ephemeral=ephemeral,
)
if runtime is not None:
kwargs["runtime"] = runtime
response = await self._loop.process_direct(
message,
**kwargs,
hooks=per_run_hooks,
)
return result_from_response(response, capture)
@@ -188,7 +192,11 @@ class Nanobot:
model_preset: str | None = None,
) -> RunStream:
"""Start a streamed run and return a handle for events and final result."""
ensure_single_model_selector(model=model, model_preset=model_preset)
runtime = self._loop.runtime_resolver.resolve_override(
model=model,
model_preset=model_preset,
config=self._config,
) or self._loop.llm_runtime()
queue: asyncio.Queue[StreamEvent | object] = asyncio.Queue(maxsize=256)
emitter = SDKStreamEmitter(queue)
stream_hook = SDKStreamingHook(emitter)
@@ -202,55 +210,53 @@ class Nanobot:
await emitter.text_completed(resuming=resuming)
async def _run() -> RunResult:
async with self._runtime_overrides.override(model=model, model_preset=model_preset):
kwargs = build_process_direct_kwargs(
session_key=session_key,
channel=channel,
chat_id=chat_id,
sender_id=sender_id,
media=media,
ephemeral=ephemeral,
on_stream=_on_stream,
on_stream_end=_on_stream_end,
kwargs = build_process_direct_kwargs(
session_key=session_key,
channel=channel,
chat_id=chat_id,
sender_id=sender_id,
media=media,
ephemeral=ephemeral,
on_stream=_on_stream,
on_stream_end=_on_stream_end,
)
kwargs["runtime"] = runtime
await emitter.emit(StreamEvent(
type=STREAM_EVENT_RUN_STARTED,
metadata={
"session_key": session_key,
"channel": channel,
"chat_id": chat_id,
"sender_id": sender_id,
"model": runtime.model,
"model_preset": runtime.model_preset,
},
))
try:
response = await self._loop.process_direct(
message,
**kwargs,
hooks=per_run_hooks,
)
await emitter.text_completed(resuming=False, force=False)
result = result_from_response(response, capture)
await emitter.emit(StreamEvent(
type=STREAM_EVENT_RUN_STARTED,
metadata={
"session_key": session_key,
"channel": channel,
"chat_id": chat_id,
"sender_id": sender_id,
"model": self._loop.model,
"model_preset": (
model_preset if model_preset is not None else self._loop.model_preset
),
},
type=STREAM_EVENT_RUN_COMPLETED,
content=result.content,
result=result,
usage=dict(result.usage),
metadata=dict(result.metadata),
))
try:
response = await self._loop.process_direct(
message,
**kwargs,
hooks=per_run_hooks,
)
await emitter.text_completed(resuming=False, force=False)
result = result_from_response(response, capture)
await emitter.emit(StreamEvent(
type=STREAM_EVENT_RUN_COMPLETED,
content=result.content,
result=result,
usage=dict(result.usage),
metadata=dict(result.metadata),
))
return result
except Exception as exc:
await emitter.emit(StreamEvent(
type=STREAM_EVENT_RUN_FAILED,
error=str(exc),
metadata={"exception_type": type(exc).__name__},
))
raise
finally:
emitter.close()
return result
except Exception as exc:
await emitter.emit(StreamEvent(
type=STREAM_EVENT_RUN_FAILED,
error=str(exc),
metadata={"exception_type": type(exc).__name__},
))
raise
finally:
emitter.close()
task = asyncio.create_task(_run())
return RunStream(task, queue)
+1 -149
View File
@@ -2,16 +2,7 @@
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, Any
from nanobot.config.schema import Config, ModelPresetConfig
from nanobot.providers.factory import ProviderSnapshot, build_provider_snapshot
if TYPE_CHECKING:
from nanobot.agent.loop import AgentLoop
from typing import Any
def ensure_single_model_selector(
@@ -51,142 +42,3 @@ def build_process_direct_kwargs(
if on_stream_end is not None:
kwargs["on_stream_end"] = on_stream_end
return kwargs
class SDKRuntimeGate:
"""Allow normal SDK runs to overlap while model overrides stay exclusive."""
def __init__(self) -> None:
self._condition = asyncio.Condition()
self._readers = 0
self._writer_active = False
self._writers_waiting = 0
def slot(self, *, exclusive: bool) -> SDKRuntimeGateSlot:
return SDKRuntimeGateSlot(self, exclusive=exclusive)
async def _acquire(self, *, exclusive: bool) -> None:
async with self._condition:
if exclusive:
self._writers_waiting += 1
try:
await self._condition.wait_for(
lambda: not self._writer_active and self._readers == 0
)
self._writer_active = True
finally:
self._writers_waiting -= 1
self._condition.notify_all()
return
await self._condition.wait_for(
lambda: not self._writer_active and self._writers_waiting == 0
)
self._readers += 1
async def _release(self, *, exclusive: bool) -> None:
async with self._condition:
if exclusive:
self._writer_active = False
else:
self._readers = max(0, self._readers - 1)
self._condition.notify_all()
class SDKRuntimeGateSlot:
def __init__(self, gate: SDKRuntimeGate, *, exclusive: bool) -> None:
self._gate = gate
self._exclusive = exclusive
async def __aenter__(self) -> None:
await self._gate._acquire(exclusive=self._exclusive)
async def __aexit__(self, *exc: object) -> None:
await self._gate._release(exclusive=self._exclusive)
class SDKRuntimeController:
"""Apply per-run SDK model overrides without leaking global runtime state."""
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
self._loop = loop
self._config = config
self._gate = SDKRuntimeGate()
@asynccontextmanager
async def override(
self,
*,
model: str | None,
model_preset: str | None,
) -> AsyncIterator[None]:
ensure_single_model_selector(model=model, model_preset=model_preset)
exclusive = model is not None or model_preset is not None
async with self._gate.slot(exclusive=exclusive):
override = self.model_override_snapshot(model=model, model_preset=model_preset)
restore = self._current_snapshot() if override is not None else None
restore_signature = self._loop._provider_signature
if override is not None:
self._loop._apply_provider_snapshot(
override,
publish_update=False,
model_preset=model_preset,
)
try:
yield
finally:
if restore is not None:
self._restore_snapshot(
restore,
provider_signature=restore_signature,
)
def model_override_snapshot(
self,
*,
model: str | None,
model_preset: str | None,
) -> ProviderSnapshot | None:
ensure_single_model_selector(model=model, model_preset=model_preset)
if model_preset is not None:
return self._loop._build_model_preset_snapshot(model_preset)
if model is None:
return None
if self._config is not None:
base = self._config.resolve_preset(self._loop.model_preset)
preset = base.model_copy(update={"model": model, "provider": "auto"})
return build_provider_snapshot(self._config, preset=preset)
generation = getattr(self._loop.provider, "generation", None)
preset = ModelPresetConfig(
model=model,
provider="auto",
max_tokens=getattr(generation, "max_tokens", 8192),
context_window_tokens=self._loop.context_window_tokens,
temperature=getattr(generation, "temperature", 0.1),
reasoning_effort=getattr(generation, "reasoning_effort", None),
)
from nanobot.agent.model_presets import build_static_preset_snapshot
return build_static_preset_snapshot(self._loop.provider, "sdk:override", preset)
def _current_snapshot(self) -> ProviderSnapshot:
signature = self._loop._provider_signature
if signature is None:
signature = ("sdk:runtime", id(self._loop.provider), self._loop.model)
return ProviderSnapshot(
provider=self._loop.provider,
model=self._loop.model,
context_window_tokens=self._loop.context_window_tokens,
signature=signature,
)
def _restore_snapshot(
self,
snapshot: ProviderSnapshot,
*,
provider_signature: tuple[object, ...] | None,
) -> None:
self._loop._apply_provider_snapshot(snapshot, publish_update=False)
self._loop._provider_signature = provider_signature