fix(agent): preserve runtime compatibility contracts
This commit is contained in:
@@ -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))
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user