From fe7d94359bc77d35482979055091c41617b60268 Mon Sep 17 00:00:00 2001 From: chengyongru Date: Fri, 10 Jul 2026 17:28:33 +0800 Subject: [PATCH] fix(agent): preserve runtime compatibility contracts --- nanobot/agent/model_runtime.py | 33 +++++----- nanobot/agent/subagent.py | 73 +++++++++++++++++++-- tests/agent/test_model_runtime_resolver.py | 32 ++++++++++ tests/agent/test_runtime_refresh.py | 33 +++++++++- tests/agent/test_subagent_lifecycle.py | 74 +++++++++++++++++++++- 5 files changed, 223 insertions(+), 22 deletions(-) diff --git a/nanobot/agent/model_runtime.py b/nanobot/agent/model_runtime.py index 238a6181..d48d6c1f 100644 --- a/nanobot/agent/model_runtime.py +++ b/nanobot/agent/model_runtime.py @@ -31,7 +31,6 @@ class ModelRuntimeResolver: self._model_presets = dict(model_presets or {}) self._provider_snapshot_loader = provider_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._default_selection_signature = preset_helpers.default_selection_signature( initial_runtime.snapshot_signature @@ -48,7 +47,7 @@ class ModelRuntimeResolver: @property def model_preset(self) -> str | None: - return self._active_preset + return self._runtime.model_preset @property def provider_signature(self) -> tuple[object, ...] | None: @@ -79,7 +78,6 @@ class ModelRuntimeResolver: """Select a snapshot as the default for future turns.""" runtime = self.resolve_snapshot(snapshot, model_preset=model_preset) self._runtime = runtime - self._active_preset = model_preset self._tracks_provider_generation = model_preset is None self._default_selection_signature = preset_helpers.default_selection_signature( runtime.snapshot_signature @@ -101,7 +99,6 @@ class ModelRuntimeResolver: """Select a named preset as the default for future turns.""" runtime = self.resolve_preset(name) self._runtime = runtime - self._active_preset = runtime.model_preset self._tracks_provider_generation = False return runtime @@ -114,7 +111,6 @@ class ModelRuntimeResolver: model=model.strip(), model_preset=None, ) - self._active_preset = None return self._runtime def select_context_window(self, context_window_tokens: int) -> LLMRuntime: @@ -154,23 +150,28 @@ class ModelRuntimeResolver: snapshot = self._provider_snapshot_loader() 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): - self._default_selection_signature = default_selection runtime = self.resolve_preset(active_preset) else: active_preset = None - self._active_preset = None - self._default_selection_signature = default_selection 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 - self._runtime = runtime - self._active_preset = active_preset - self._tracks_provider_generation = active_preset is None - self._default_selection_signature = preset_helpers.default_selection_signature( - runtime.snapshot_signature + ( + self._runtime, + self._tracks_provider_generation, + self._default_selection_signature, + ) = ( + runtime, + active_preset is None, + default_selection, ) return runtime @@ -197,6 +198,6 @@ class ModelRuntimeResolver: 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"}) return self.resolve_snapshot(build_provider_snapshot(config, preset=preset)) diff --git a/nanobot/agent/subagent.py b/nanobot/agent/subagent.py index 355ad169..29f8588c 100644 --- a/nanobot/agent/subagent.py +++ b/nanobot/agent/subagent.py @@ -4,6 +4,7 @@ import asyncio import json import time import uuid +import warnings from dataclasses import dataclass, field from pathlib import Path 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.queue import MessageBus from nanobot.config.schema import AgentDefaults, ToolsConfig +from nanobot.providers.base import LLMProvider from nanobot.security.workspace_access import ( WorkspaceScope, bind_workspace_scope, @@ -81,9 +83,11 @@ class SubagentManager: def __init__( self, - workspace: Path, - bus: MessageBus, - max_tool_result_chars: int, + provider: LLMProvider | None = None, + workspace: Path | None = None, + bus: MessageBus | None = None, + max_tool_result_chars: int | None = None, + model: str | None = None, tools_config: ToolsConfig | None = None, restrict_to_workspace: bool = False, disabled_skills: list[str] | None = None, @@ -92,7 +96,31 @@ class SubagentManager: fail_on_tool_error: bool | 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() + 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.bus = bus self.tools_config = tools_config or ToolsConfig() @@ -120,6 +148,41 @@ class SubagentManager: self._task_statuses: dict[str, SubagentStatus] = {} 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: """Build a ToolsConfig scoped for subagent use.""" return ToolsConfig( @@ -161,9 +224,11 @@ class SubagentManager: temperature: float | None = None, workspace_scope: WorkspaceScope | None = None, *, - runtime: LLMRuntime, + runtime: LLMRuntime | None = None, ) -> str: """Spawn a subagent to execute a task in the background.""" + if runtime is None: + runtime = self._compat_spawn_runtime() if temperature is not None: runtime = runtime.with_generation_overrides(temperature=temperature) task_id = str(uuid.uuid4())[:8] diff --git a/tests/agent/test_model_runtime_resolver.py b/tests/agent/test_model_runtime_resolver.py index 14e45607..e6e8f401 100644 --- a/tests/agent/test_model_runtime_resolver.py +++ b/tests/agent/test_model_runtime_resolver.py @@ -139,6 +139,38 @@ def test_resolver_refresh_preserves_unchanged_active_preset() -> None: 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: provider = _provider(temperature=0.2, max_tokens=2048) resolver = ModelRuntimeResolver(_runtime(provider)) diff --git a/tests/agent/test_runtime_refresh.py b/tests/agent/test_runtime_refresh.py index d21aa3af..4154b110 100644 --- a/tests/agent/test_runtime_refresh.py +++ b/tests/agent/test_runtime_refresh.py @@ -5,7 +5,7 @@ from unittest.mock import MagicMock from nanobot.agent.loop import AgentLoop from nanobot.bus.queue import MessageBus 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.factory import ProviderSnapshot, load_provider_snapshot 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") +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( tmp_path: Path, ) -> None: diff --git a/tests/agent/test_subagent_lifecycle.py b/tests/agent/test_subagent_lifecycle.py index 907065d5..daa4df0a 100644 --- a/tests/agent/test_subagent_lifecycle.py +++ b/tests/agent/test_subagent_lifecycle.py @@ -7,10 +7,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from nanobot.agent import SubagentManager from nanobot.agent.hook import AgentHookContext from nanobot.agent.runner import AgentRunResult from nanobot.agent.subagent import ( - SubagentManager, SubagentStatus, _SubagentHook, ) @@ -98,6 +98,78 @@ class TestRuntimeOwnership: 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 # ---------------------------------------------------------------------------