fix(agent): preserve runtime compatibility contracts

This commit is contained in:
chengyongru
2026-07-10 17:54:34 +08:00
committed by Xubin Ren
parent 10930f3902
commit fe7d94359b
5 changed files with 223 additions and 22 deletions
+17 -16
View File
@@ -31,7 +31,6 @@ class ModelRuntimeResolver:
self._model_presets = dict(model_presets or {}) self._model_presets = dict(model_presets or {})
self._provider_snapshot_loader = provider_snapshot_loader self._provider_snapshot_loader = provider_snapshot_loader
self._preset_snapshot_loader = preset_snapshot_loader self._preset_snapshot_loader = preset_snapshot_loader
self._active_preset = initial_runtime.model_preset
self._tracks_provider_generation = initial_runtime.model_preset is None self._tracks_provider_generation = initial_runtime.model_preset is None
self._default_selection_signature = preset_helpers.default_selection_signature( self._default_selection_signature = preset_helpers.default_selection_signature(
initial_runtime.snapshot_signature initial_runtime.snapshot_signature
@@ -48,7 +47,7 @@ class ModelRuntimeResolver:
@property @property
def model_preset(self) -> str | None: def model_preset(self) -> str | None:
return self._active_preset return self._runtime.model_preset
@property @property
def provider_signature(self) -> tuple[object, ...] | None: def provider_signature(self) -> tuple[object, ...] | None:
@@ -79,7 +78,6 @@ class ModelRuntimeResolver:
"""Select a snapshot as the default for future turns.""" """Select a snapshot as the default for future turns."""
runtime = self.resolve_snapshot(snapshot, model_preset=model_preset) runtime = self.resolve_snapshot(snapshot, model_preset=model_preset)
self._runtime = runtime self._runtime = runtime
self._active_preset = model_preset
self._tracks_provider_generation = model_preset is None self._tracks_provider_generation = model_preset is None
self._default_selection_signature = preset_helpers.default_selection_signature( self._default_selection_signature = preset_helpers.default_selection_signature(
runtime.snapshot_signature runtime.snapshot_signature
@@ -101,7 +99,6 @@ class ModelRuntimeResolver:
"""Select a named preset as the default for future turns.""" """Select a named preset as the default for future turns."""
runtime = self.resolve_preset(name) runtime = self.resolve_preset(name)
self._runtime = runtime self._runtime = runtime
self._active_preset = runtime.model_preset
self._tracks_provider_generation = False self._tracks_provider_generation = False
return runtime return runtime
@@ -114,7 +111,6 @@ class ModelRuntimeResolver:
model=model.strip(), model=model.strip(),
model_preset=None, model_preset=None,
) )
self._active_preset = None
return self._runtime return self._runtime
def select_context_window(self, context_window_tokens: int) -> LLMRuntime: def select_context_window(self, context_window_tokens: int) -> LLMRuntime:
@@ -154,23 +150,28 @@ class ModelRuntimeResolver:
snapshot = self._provider_snapshot_loader() snapshot = self._provider_snapshot_loader()
default_selection = preset_helpers.default_selection_signature(snapshot.signature) default_selection = preset_helpers.default_selection_signature(snapshot.signature)
active_preset = self._active_preset active_preset = self._runtime.model_preset
if active_preset and self._default_selection_signature in (None, default_selection): if active_preset and self._default_selection_signature in (None, default_selection):
self._default_selection_signature = default_selection
runtime = self.resolve_preset(active_preset) runtime = self.resolve_preset(active_preset)
else: else:
active_preset = None active_preset = None
self._active_preset = None
self._default_selection_signature = default_selection
runtime = self.resolve_snapshot(snapshot) runtime = self.resolve_snapshot(snapshot)
if runtime.snapshot_signature == self._runtime.snapshot_signature: unchanged = (
runtime.snapshot_signature == self._runtime.snapshot_signature
and runtime.model_preset == self._runtime.model_preset
)
if unchanged:
self._default_selection_signature = default_selection
return None return None
self._runtime = runtime (
self._active_preset = active_preset self._runtime,
self._tracks_provider_generation = active_preset is None self._tracks_provider_generation,
self._default_selection_signature = preset_helpers.default_selection_signature( self._default_selection_signature,
runtime.snapshot_signature ) = (
runtime,
active_preset is None,
default_selection,
) )
return runtime return runtime
@@ -197,6 +198,6 @@ class ModelRuntimeResolver:
snapshot_signature=("model_override", model), snapshot_signature=("model_override", model),
) )
base = config.resolve_preset(self._active_preset) base = config.resolve_preset(self.model_preset)
preset = base.model_copy(update={"model": model, "provider": "auto"}) preset = base.model_copy(update={"model": model, "provider": "auto"})
return self.resolve_snapshot(build_provider_snapshot(config, preset=preset)) return self.resolve_snapshot(build_provider_snapshot(config, preset=preset))
+69 -4
View File
@@ -4,6 +4,7 @@ import asyncio
import json import json
import time import time
import uuid import uuid
import warnings
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Callable from typing import Any, Callable
@@ -24,6 +25,7 @@ from nanobot.agent.tools.registry import ToolRegistry
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 AgentDefaults, ToolsConfig from nanobot.config.schema import AgentDefaults, ToolsConfig
from nanobot.providers.base import LLMProvider
from nanobot.security.workspace_access import ( from nanobot.security.workspace_access import (
WorkspaceScope, WorkspaceScope,
bind_workspace_scope, bind_workspace_scope,
@@ -81,9 +83,11 @@ class SubagentManager:
def __init__( def __init__(
self, self,
workspace: Path, provider: LLMProvider | None = None,
bus: MessageBus, workspace: Path | None = None,
max_tool_result_chars: int, bus: MessageBus | None = None,
max_tool_result_chars: int | None = None,
model: str | None = None,
tools_config: ToolsConfig | None = None, tools_config: ToolsConfig | None = None,
restrict_to_workspace: bool = False, restrict_to_workspace: bool = False,
disabled_skills: list[str] | None = None, disabled_skills: list[str] | None = None,
@@ -92,7 +96,31 @@ class SubagentManager:
fail_on_tool_error: bool | None = None, fail_on_tool_error: bool | None = None,
llm_wall_timeout_for_session: Callable[[str | None], float | None] | None = None, llm_wall_timeout_for_session: Callable[[str | None], float | None] | None = None,
): ):
if workspace is None:
raise TypeError("SubagentManager.__init__() missing required argument: 'workspace'")
if bus is None:
raise TypeError("SubagentManager.__init__() missing required argument: 'bus'")
if max_tool_result_chars is None:
raise TypeError(
"SubagentManager.__init__() missing required argument: 'max_tool_result_chars'"
)
if model is not None and provider is None:
raise TypeError("SubagentManager model compatibility argument requires provider")
defaults = AgentDefaults() defaults = AgentDefaults()
self._compat_runtime: LLMRuntime | None = None
if provider is not None:
warnings.warn(
"SubagentManager provider/model constructor arguments are deprecated; "
"pass runtime=... to spawn() instead",
DeprecationWarning,
stacklevel=2,
)
self._compat_runtime = LLMRuntime.capture(
provider,
model or provider.get_default_model(),
context_window_tokens=defaults.context_window_tokens,
)
self.workspace = workspace self.workspace = workspace
self.bus = bus self.bus = bus
self.tools_config = tools_config or ToolsConfig() self.tools_config = tools_config or ToolsConfig()
@@ -120,6 +148,41 @@ class SubagentManager:
self._task_statuses: dict[str, SubagentStatus] = {} self._task_statuses: dict[str, SubagentStatus] = {}
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...} self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
def set_provider(self, provider: LLMProvider, model: str) -> None:
"""Update the deprecated runtime source used by legacy ``spawn`` calls."""
warnings.warn(
"SubagentManager.set_provider() is deprecated; pass runtime=... to spawn() instead",
DeprecationWarning,
stacklevel=2,
)
context_window_tokens = (
self._compat_runtime.context_window_tokens
if self._compat_runtime is not None
else AgentDefaults().context_window_tokens
)
self._compat_runtime = LLMRuntime.capture(
provider,
model,
context_window_tokens=context_window_tokens,
)
def _compat_spawn_runtime(self) -> LLMRuntime:
runtime = self._compat_runtime
if runtime is None:
raise TypeError(
"SubagentManager.spawn() missing required keyword-only argument: 'runtime'"
)
warnings.warn(
"SubagentManager.spawn() without runtime is deprecated; pass runtime=... explicitly",
DeprecationWarning,
stacklevel=3,
)
return LLMRuntime.capture(
runtime.provider,
runtime.model,
context_window_tokens=runtime.context_window_tokens,
)
def _subagent_tools_config(self) -> ToolsConfig: def _subagent_tools_config(self) -> ToolsConfig:
"""Build a ToolsConfig scoped for subagent use.""" """Build a ToolsConfig scoped for subagent use."""
return ToolsConfig( return ToolsConfig(
@@ -161,9 +224,11 @@ class SubagentManager:
temperature: float | None = None, temperature: float | None = None,
workspace_scope: WorkspaceScope | None = None, workspace_scope: WorkspaceScope | None = None,
*, *,
runtime: LLMRuntime, runtime: LLMRuntime | None = None,
) -> str: ) -> str:
"""Spawn a subagent to execute a task in the background.""" """Spawn a subagent to execute a task in the background."""
if runtime is None:
runtime = self._compat_spawn_runtime()
if temperature is not None: if temperature is not None:
runtime = runtime.with_generation_overrides(temperature=temperature) runtime = runtime.with_generation_overrides(temperature=temperature)
task_id = str(uuid.uuid4())[:8] task_id = str(uuid.uuid4())[:8]
@@ -139,6 +139,38 @@ def test_resolver_refresh_preserves_unchanged_active_preset() -> None:
assert resolver.model_preset == "fast" assert resolver.model_preset == "fast"
def test_refresh_clears_preset_when_new_default_has_same_snapshot_signature() -> None:
initial = _runtime()
preset_provider = _provider()
preset_snapshot = ProviderSnapshot(
provider=preset_provider,
model="fast-model",
context_window_tokens=20_000,
signature=("fast-model", "auto", "same-runtime"),
)
default_snapshot = ProviderSnapshot(
provider=preset_provider,
model="fast-model",
context_window_tokens=20_000,
signature=preset_snapshot.signature,
)
resolver = ModelRuntimeResolver(
initial,
model_presets={"fast": ModelPresetConfig(model="fast-model")},
provider_snapshot_loader=lambda: default_snapshot,
preset_snapshot_loader=lambda _name: preset_snapshot,
)
resolver.select_preset("fast")
refreshed = resolver.refresh()
assert refreshed is resolver.runtime
assert refreshed is not None
assert resolver.model_preset is None
assert resolver.runtime.model_preset is None
assert "_active_preset" not in resolver.__dict__
def test_resolver_refreshes_provider_generation_for_next_default_turn() -> None: def test_resolver_refreshes_provider_generation_for_next_default_turn() -> None:
provider = _provider(temperature=0.2, max_tokens=2048) provider = _provider(temperature=0.2, max_tokens=2048)
resolver = ModelRuntimeResolver(_runtime(provider)) resolver = ModelRuntimeResolver(_runtime(provider))
+32 -1
View File
@@ -5,7 +5,7 @@ from unittest.mock import MagicMock
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.loader import save_config from nanobot.config.loader import save_config
from nanobot.config.schema import Config from nanobot.config.schema import Config, ModelPresetConfig
from nanobot.providers.base import GenerationSettings from nanobot.providers.base import GenerationSettings
from nanobot.providers.factory import ProviderSnapshot, load_provider_snapshot from nanobot.providers.factory import ProviderSnapshot, load_provider_snapshot
from nanobot.webui.settings_api import update_agent_settings from nanobot.webui.settings_api import update_agent_settings
@@ -99,6 +99,37 @@ def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
assert not hasattr(loop.runner, "provider") assert not hasattr(loop.runner, "provider")
def test_same_snapshot_default_clears_preset_and_publishes_update(tmp_path: Path) -> None:
base_provider = _provider("base-model")
fast_provider = _provider("fast-model")
fast_snapshot = ProviderSnapshot(
provider=fast_provider,
model="fast-model",
context_window_tokens=2000,
signature=("fast-model", "auto", "same-runtime"),
)
published: list[tuple[str, str | None]] = []
loop = AgentLoop(
bus=MessageBus(),
provider=base_provider,
workspace=tmp_path,
model="base-model",
context_window_tokens=1000,
provider_signature=("base-model", "auto", "initial"),
provider_snapshot_loader=lambda: fast_snapshot,
model_presets={"fast": ModelPresetConfig(model="fast-model")},
model_preset="fast",
preset_snapshot_loader=lambda _name: fast_snapshot,
runtime_model_publisher=lambda model, preset: published.append((model, preset)),
)
runtime = loop.llm_runtime()
assert runtime.model_preset is None
assert loop.model_preset is None
assert published == [("fast-model", None)]
def test_next_turn_captures_generation_changed_after_previous_admission( def test_next_turn_captures_generation_changed_after_previous_admission(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
+73 -1
View File
@@ -7,10 +7,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from nanobot.agent import SubagentManager
from nanobot.agent.hook import AgentHookContext from nanobot.agent.hook import AgentHookContext
from nanobot.agent.runner import AgentRunResult from nanobot.agent.runner import AgentRunResult
from nanobot.agent.subagent import ( from nanobot.agent.subagent import (
SubagentManager,
SubagentStatus, SubagentStatus,
_SubagentHook, _SubagentHook,
) )
@@ -98,6 +98,78 @@ class TestRuntimeOwnership:
assert not hasattr(sm.runner, "provider") assert not hasattr(sm.runner, "provider")
class TestLegacyCompatibility:
def test_accepts_exported_legacy_constructor_positionally(self, tmp_path):
provider = MagicMock(spec=LLMProvider)
provider.generation = GenerationSettings(temperature=0.2, max_tokens=2048)
with pytest.warns(DeprecationWarning, match="provider"):
sm = SubagentManager(
provider,
tmp_path,
MessageBus(),
16_000,
"legacy-model",
)
assert sm.workspace == tmp_path
assert sm.max_tool_result_chars == 16_000
assert not hasattr(sm, "provider")
assert not hasattr(sm, "model")
@pytest.mark.asyncio
async def test_legacy_spawn_captures_runtime_at_admission(self, tmp_path):
provider = MagicMock(spec=LLMProvider)
provider.generation = GenerationSettings(temperature=0.2, max_tokens=2048)
with pytest.warns(DeprecationWarning, match="provider"):
sm = SubagentManager(
provider=provider,
workspace=tmp_path,
bus=MessageBus(),
max_tool_result_chars=16_000,
model="legacy-model",
)
sm.runner.run = AsyncMock(return_value=AgentRunResult(
final_content="done", messages=[], stop_reason="completed",
))
provider.generation = GenerationSettings(temperature=0.8, max_tokens=512)
with pytest.warns(DeprecationWarning, match="runtime"):
await sm.spawn("legacy task")
await _drain_subagent_tasks(sm)
runtime = sm.runner.run.await_args.args[0].runtime
assert runtime.provider is provider
assert runtime.model == "legacy-model"
assert runtime.generation == GenerationSettings(0.8, 512, None)
@pytest.mark.asyncio
async def test_set_provider_supports_future_legacy_spawns(self, tmp_path):
sm = _manager(tmp_path)
provider = MagicMock(spec=LLMProvider)
provider.generation = GenerationSettings(temperature=0.3, max_tokens=1024)
with pytest.warns(DeprecationWarning, match="set_provider"):
sm.set_provider(provider, "replacement-model")
sm.runner.run = AsyncMock(return_value=AgentRunResult(
final_content="done", messages=[], stop_reason="completed",
))
with pytest.warns(DeprecationWarning, match="runtime"):
await sm.spawn("legacy task")
await _drain_subagent_tasks(sm)
runtime = sm.runner.run.await_args.args[0].runtime
assert runtime.provider is provider
assert runtime.model == "replacement-model"
@pytest.mark.asyncio
async def test_new_constructor_still_requires_explicit_spawn_runtime(self, tmp_path):
sm = _manager(tmp_path)
with pytest.raises(TypeError, match="runtime"):
await sm.spawn("task")
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# spawn # spawn
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------