From 21f58cbabf7fdba08ee9e3e68d7d8c84e13629c2 Mon Sep 17 00:00:00 2001 From: chengyongru Date: Fri, 10 Jul 2026 15:17:20 +0800 Subject: [PATCH] refactor(agent): make resolver sole runtime owner --- nanobot/agent/loop.py | 208 +++++++----------- nanobot/agent/tools/runtime_state.py | 6 +- nanobot/agent/tools/self.py | 42 +++- nanobot/command/builtin.py | 10 +- nanobot/nanobot.py | 132 +++++------ nanobot/sdk/runtime.py | 150 +------------ tests/agent/test_consolidate_offset.py | 24 +- tests/agent/test_loop_consolidation_tokens.py | 3 +- tests/agent/test_loop_progress.py | 11 +- tests/agent/test_loop_save_turn.py | 14 +- tests/agent/test_max_messages_config.py | 26 ++- tests/agent/test_runtime_refresh.py | 30 ++- tests/agent/test_self_model_preset.py | 19 +- tests/agent/tools/test_self_tool.py | 10 +- tests/cli/test_restart_command.py | 4 +- tests/test_nanobot_facade.py | 149 ++++++++----- 16 files changed, 395 insertions(+), 443 deletions(-) diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index 2919fd80..392ec900 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -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, diff --git a/nanobot/agent/tools/runtime_state.py b/nanobot/agent/tools/runtime_state.py index 449c7fe0..48245a8c 100644 --- a/nanobot/agent/tools/runtime_state.py +++ b/nanobot/agent/tools/runtime_state.py @@ -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 diff --git a/nanobot/agent/tools/self.py b/nanobot/agent/tools/self.py index d0310a9b..bfd1824c 100644 --- a/nanobot/agent/tools/self.py +++ b/nanobot/agent/tools/self.py @@ -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})" diff --git a/nanobot/command/builtin.py b/nanobot/command/builtin.py index 05011e4d..6598a1ca 100644 --- a/nanobot/command/builtin.py +++ b/nanobot/command/builtin.py @@ -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}") diff --git a/nanobot/nanobot.py b/nanobot/nanobot.py index 92a68084..aef50bac 100644 --- a/nanobot/nanobot.py +++ b/nanobot/nanobot.py @@ -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) diff --git a/nanobot/sdk/runtime.py b/nanobot/sdk/runtime.py index 38aae479..b0c0da15 100644 --- a/nanobot/sdk/runtime.py +++ b/nanobot/sdk/runtime.py @@ -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 diff --git a/tests/agent/test_consolidate_offset.py b/tests/agent/test_consolidate_offset.py index 74e79614..315e3eb6 100644 --- a/tests/agent/test_consolidate_offset.py +++ b/tests/agent/test_consolidate_offset.py @@ -518,9 +518,11 @@ class TestNewCommandArchival: loop.sessions.save(session) call_count = 0 + expected_runtime = loop.llm_runtime() - async def _failing_summarize(_messages, *, session_key=None) -> bool: + async def _failing_summarize(_messages, *, runtime, session_key=None) -> bool: nonlocal call_count + assert runtime is expected_runtime assert session_key == "cli:test" call_count += 1 return False @@ -528,7 +530,7 @@ class TestNewCommandArchival: loop.consolidator.archive = _failing_summarize # type: ignore[method-assign] new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") - response = await loop._process_message(new_msg) + response = await loop._process_message(new_msg, runtime=expected_runtime) assert response is not None assert "new session started" in response.content.lower() @@ -553,9 +555,11 @@ class TestNewCommandArchival: archived_count = -1 archived_session_key = None + expected_runtime = loop.llm_runtime() - async def _fake_summarize(messages, *, session_key=None) -> bool: + async def _fake_summarize(messages, *, runtime, session_key=None) -> bool: nonlocal archived_count, archived_session_key + assert runtime is expected_runtime archived_count = len(messages) archived_session_key = session_key return True @@ -563,7 +567,7 @@ class TestNewCommandArchival: loop.consolidator.archive = _fake_summarize # type: ignore[method-assign] new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") - response = await loop._process_message(new_msg) + response = await loop._process_message(new_msg, runtime=expected_runtime) assert response is not None assert "new session started" in response.content.lower() @@ -582,15 +586,17 @@ class TestNewCommandArchival: session.add_message("user", f"msg{i}") session.add_message("assistant", f"resp{i}") loop.sessions.save(session) + expected_runtime = loop.llm_runtime() - async def _ok_summarize(_messages, *, session_key=None) -> bool: + async def _ok_summarize(_messages, *, runtime, session_key=None) -> bool: + assert runtime is expected_runtime assert session_key == "cli:test" return True loop.consolidator.archive = _ok_summarize # type: ignore[method-assign] new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") - response = await loop._process_message(new_msg) + response = await loop._process_message(new_msg, runtime=expected_runtime) assert response is not None assert "new session started" in response.content.lower() @@ -610,8 +616,10 @@ class TestNewCommandArchival: archived = asyncio.Event() release_archive = asyncio.Event() + expected_runtime = loop.llm_runtime() - async def _slow_summarize(_messages, *, session_key=None) -> bool: + async def _slow_summarize(_messages, *, runtime, session_key=None) -> bool: + assert runtime is expected_runtime assert session_key == "cli:test" await release_archive.wait() archived.set() @@ -620,7 +628,7 @@ class TestNewCommandArchival: loop.consolidator.archive = _slow_summarize # type: ignore[method-assign] new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") - await loop._process_message(new_msg) + await loop._process_message(new_msg, runtime=expected_runtime) assert not archived.is_set() release_archive.set() diff --git a/tests/agent/test_loop_consolidation_tokens.py b/tests/agent/test_loop_consolidation_tokens.py index f113b14a..d60404fe 100644 --- a/tests/agent/test_loop_consolidation_tokens.py +++ b/tests/agent/test_loop_consolidation_tokens.py @@ -6,6 +6,7 @@ import nanobot.agent.memory as memory_module from nanobot.agent.loop import AgentLoop from nanobot.bus.queue import MessageBus from nanobot.providers.base import LLMResponse +from nanobot.session.manager import replay_max_messages_for_context def _make_loop(tmp_path, *, estimated_tokens: int, context_window_tokens: int) -> AgentLoop: @@ -222,7 +223,7 @@ async def test_preflight_consolidation_receives_pending_summary(tmp_path) -> Non loop.consolidator.maybe_consolidate_by_tokens.assert_any_await( session, runtime=runtime, - replay_max_messages=loop._max_messages, + replay_max_messages=replay_max_messages_for_context(runtime.context_window_tokens), ) assert len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2 assert all( diff --git a/tests/agent/test_loop_progress.py b/tests/agent/test_loop_progress.py index 33590405..d1ddbc98 100644 --- a/tests/agent/test_loop_progress.py +++ b/tests/agent/test_loop_progress.py @@ -21,6 +21,7 @@ from nanobot.bus.outbound_events import ( ) from nanobot.bus.queue import MessageBus from nanobot.providers.base import LLMResponse, ToolCallRequest +from nanobot.providers.factory import ProviderSnapshot from nanobot.session.webui_turns import WebuiTurnCoordinator from nanobot.utils.progress_events import ( invoke_file_edit_progress, @@ -815,8 +816,14 @@ class TestToolEventProgress: )) assert len(scheduled_title) == 1 - loop.provider = MagicMock() - loop.model = "switched-after-turn" + next_provider = MagicMock() + next_provider.generation = loop.llm_runtime().generation + loop.runtime_resolver.adopt_snapshot(ProviderSnapshot( + provider=next_provider, + model="switched-after-turn", + context_window_tokens=loop.context_window_tokens, + signature=("switched-after-turn",), + )) await scheduled_title[0] # type: ignore[misc] diff --git a/tests/agent/test_loop_save_turn.py b/tests/agent/test_loop_save_turn.py index 0726671d..fc4f2fc2 100644 --- a/tests/agent/test_loop_save_turn.py +++ b/tests/agent/test_loop_save_turn.py @@ -19,6 +19,7 @@ from nanobot.bus.outbound_events import ( from nanobot.bus.queue import MessageBus from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META from nanobot.providers.base import LLMResponse +from nanobot.providers.factory import ProviderSnapshot from nanobot.session.automation_turns import AUTOMATION_HISTORY_META from nanobot.session.goal_state import GOAL_STATE_KEY from nanobot.session.manager import Session, SessionManager @@ -69,8 +70,17 @@ def test_agent_loop_llm_runtime_reflects_current_provider_and_model(tmp_path: Pa assert runtime.model == "test-model" next_provider = MagicMock() - loop.provider = next_provider - loop.model = "next-model" + next_provider.generation = SimpleNamespace( + temperature=0.1, + max_tokens=4096, + reasoning_effort=None, + ) + loop.runtime_resolver.adopt_snapshot(ProviderSnapshot( + provider=next_provider, + model="next-model", + context_window_tokens=runtime.context_window_tokens, + signature=("next-model",), + )) runtime = loop.llm_runtime() assert runtime.provider is next_provider diff --git a/tests/agent/test_max_messages_config.py b/tests/agent/test_max_messages_config.py index 58101138..1f21200a 100644 --- a/tests/agent/test_max_messages_config.py +++ b/tests/agent/test_max_messages_config.py @@ -2,6 +2,7 @@ from __future__ import annotations +from dataclasses import replace from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch @@ -67,11 +68,13 @@ class TestMaxMessagesInit: def test_default_for_200k_context_reaches_file_cap(self, tmp_path: Path) -> None: loop = _make_loop(tmp_path) - assert loop._max_messages == FILE_MAX_MESSAGES + runtime = loop.runtime_resolver.runtime + assert replay_max_messages_for_context(runtime.context_window_tokens) == FILE_MAX_MESSAGES def test_default_scales_with_context_window(self, tmp_path: Path) -> None: loop = _make_loop(tmp_path, context_window_tokens=32_768) - assert loop._max_messages == 327 + runtime = loop.runtime_resolver.runtime + assert replay_max_messages_for_context(runtime.context_window_tokens) == 327 def test_provider_refresh_resyncs_context_derived_limit(self, tmp_path: Path) -> None: old_provider = MagicMock() @@ -93,9 +96,10 @@ class TestMaxMessagesInit: ), ) - assert loop._max_messages == 327 - loop._refresh_provider_snapshot() - assert loop._max_messages == FILE_MAX_MESSAGES + initial = loop.runtime_resolver.runtime + assert replay_max_messages_for_context(initial.context_window_tokens) == 327 + refreshed = loop.llm_runtime() + assert replay_max_messages_for_context(refreshed.context_window_tokens) == FILE_MAX_MESSAGES class TestGetHistoryWithMaxMessages: @@ -136,7 +140,7 @@ class TestMaxMessagesIntegration: async def test_process_message_passes_limit_to_history_call(self, tmp_path: Path) -> None: """The real message path should pass max_messages into session history replay.""" loop = _make_loop(tmp_path) - loop._max_messages = 25 + runtime = replace(loop.llm_runtime(), context_window_tokens=32_768) loop.provider.chat_with_retry = AsyncMock( return_value=LLMResponse(content="ok", tool_calls=[], usage={}) ) @@ -146,12 +150,13 @@ class TestMaxMessagesIntegration: session = loop.sessions.get_or_create("cli:test") with patch.object(session, "get_history", wraps=session.get_history) as mock_hist: result = await loop._process_message( - InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello") + InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello"), + runtime=runtime, ) assert result is not None assert mock_hist.call_count == 1 - assert mock_hist.call_args.kwargs["max_messages"] == 25 + assert mock_hist.call_args.kwargs["max_messages"] == 327 assert mock_hist.call_args.kwargs["extend_to_user"] is False @pytest.mark.asyncio @@ -182,8 +187,7 @@ class TestMaxMessagesIntegration: tmp_path: Path, ) -> None: """A live user turn should not extend history to an older long tool turn.""" - loop = _make_loop(tmp_path) - loop._max_messages = 6 + loop = _make_loop(tmp_path, context_window_tokens=8_000) loop.provider.chat_with_retry = AsyncMock( return_value=LLMResponse(content="ok", tool_calls=[], usage={}) ) @@ -194,7 +198,7 @@ class TestMaxMessagesIntegration: session.add_message("user", "old") session.add_message("assistant", "old answer") session.add_message("user", "long older turn") - for i in range(8): + for i in range(70): session.messages.extend(_tool_round(f"older-{i}")) session.add_message("assistant", "older final") diff --git a/tests/agent/test_runtime_refresh.py b/tests/agent/test_runtime_refresh.py index e387530d..d21aa3af 100644 --- a/tests/agent/test_runtime_refresh.py +++ b/tests/agent/test_runtime_refresh.py @@ -18,7 +18,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock: return provider -def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None: +def test_provider_refresh_updates_only_runtime_resolver(tmp_path: Path) -> None: old_provider = _provider("old-model") new_provider = _provider("new-model", max_tokens=456) loop = AgentLoop( @@ -35,8 +35,9 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None: ), ) - loop._refresh_provider_snapshot() + runtime = loop.llm_runtime() + assert runtime is loop.runtime_resolver.runtime assert loop.provider is new_provider assert loop.model == "new-model" assert loop.context_window_tokens == 2000 @@ -50,6 +51,29 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None: assert not hasattr(loop.consolidator, "max_completion_tokens") +def test_loop_has_no_mutable_runtime_mirrors_or_legacy_snapshot_api(tmp_path: Path) -> None: + loop = AgentLoop( + bus=MessageBus(), + provider=_provider("test-model"), + workspace=tmp_path, + model="test-model", + context_window_tokens=1000, + ) + + assert { + "provider", + "model", + "context_window_tokens", + "model_presets", + "_active_preset", + "_provider_signature", + "_max_messages", + }.isdisjoint(loop.__dict__) + assert not hasattr(loop, "_apply_provider_snapshot") + assert not hasattr(loop, "_build_model_preset_snapshot") + assert not hasattr(loop, "_sync_replay_max_messages") + + def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None: old_provider = _provider("old-model") new_provider = _provider("new-model", max_tokens=456) @@ -118,7 +142,7 @@ def test_settings_context_window_refreshes_runtime_state( loop = AgentLoop.from_config(config, provider_snapshot_loader=loader) payload = update_agent_settings({"context_window_tokens": ["262144"]}) - loop._refresh_provider_snapshot() + loop.llm_runtime() assert payload["requires_restart"] is False assert loop.context_window_tokens == 262_144 diff --git a/tests/agent/test_self_model_preset.py b/tests/agent/test_self_model_preset.py index b2e71625..ced0b819 100644 --- a/tests/agent/test_self_model_preset.py +++ b/tests/agent/test_self_model_preset.py @@ -54,9 +54,10 @@ def test_model_preset_setter_updates_state(tmp_path) -> None: assert loop.model_preset == "fast" assert loop.model == "openai/gpt-4.1" assert loop.context_window_tokens == 32_768 - assert loop.provider.generation.temperature == 0.5 - assert loop.provider.generation.max_tokens == 4096 - assert loop.provider.generation.reasoning_effort == "low" + runtime = loop.llm_runtime() + assert runtime.generation.temperature == 0.5 + assert runtime.generation.max_tokens == 4096 + assert runtime.generation.reasoning_effort == "low" assert not hasattr(loop.subagents, "model") assert not hasattr(loop.consolidator, "model") assert not hasattr(loop.consolidator, "context_window_tokens") @@ -174,7 +175,7 @@ def test_active_model_preset_survives_unchanged_config_refresh(tmp_path) -> None ) loop.set_model_preset("fast") - loop._refresh_provider_snapshot() + loop.llm_runtime() assert loop.model_preset == "fast" assert loop.provider is fast_provider @@ -210,7 +211,7 @@ def test_config_model_refresh_clears_active_model_preset(tmp_path) -> None: ) loop.set_model_preset("fast") - loop._refresh_provider_snapshot() + loop.llm_runtime() assert loop.model_preset is None assert loop.provider is webui_provider @@ -292,7 +293,7 @@ def test_self_tool_set_model_clears_active_preset(tmp_path) -> None: tool = MyTool(runtime_state=loop, modify_allowed=True) result = tool._modify("model", "anthropic/claude-opus-4-5") assert "Error" not in result - assert loop._active_preset is None + assert loop.model_preset is None assert loop.model == "anthropic/claude-opus-4-5" @@ -323,5 +324,7 @@ def test_from_config_static_preset_loader_does_not_enable_hot_reload(tmp_path) - fake_provider = _provider("openai/gpt-4.1") with patch("nanobot.providers.factory.make_provider", return_value=fake_provider): loop = AgentLoop.from_config(config) - assert loop._provider_snapshot_loader is None - assert loop._preset_snapshot_loader is not None + default_runtime = loop.runtime_resolver.runtime + resolved = loop.runtime_resolver.resolve_preset("fast") + assert resolved.model == "openai/gpt-4.1-mini" + assert loop.runtime_resolver.runtime is default_runtime diff --git a/tests/agent/tools/test_self_tool.py b/tests/agent/tools/test_self_tool.py index 8e777a09..f2f2541f 100644 --- a/tests/agent/tools/test_self_tool.py +++ b/tests/agent/tools/test_self_tool.py @@ -35,6 +35,12 @@ def _make_mock_loop(**overrides): loop._concurrency_gate = None loop._unified_session = False loop._extra_hooks = [] + loop.set_runtime_model.side_effect = lambda value: setattr(loop, "model", value) + loop.set_runtime_context_window.side_effect = lambda value: setattr( + loop, + "context_window_tokens", + value, + ) # web_config mock — needed for check tests loop.web_config = MagicMock() @@ -237,12 +243,12 @@ class TestModifyRestricted: @pytest.mark.asyncio async def test_modify_context_window_valid(self): - loop = _make_mock_loop(_sync_replay_max_messages=MagicMock()) + loop = _make_mock_loop() tool = _make_tool(runtime_state=loop) result = await tool.execute(action="set", key="context_window_tokens", value=131072) assert "Set context_window_tokens" in result assert loop.context_window_tokens == 131072 - loop._sync_replay_max_messages.assert_called_once_with() + loop.set_runtime_context_window.assert_called_once_with(131072) @pytest.mark.asyncio async def test_modify_none_value_for_restricted_int(self): diff --git a/tests/cli/test_restart_command.py b/tests/cli/test_restart_command.py index 8a227a86..2bd499bf 100644 --- a/tests/cli/test_restart_command.py +++ b/tests/cli/test_restart_command.py @@ -245,8 +245,8 @@ class TestRestartCommand: msg = InboundMessage(channel="telegram", sender_id="u1", chat_id="c1", content="/status") runtime = loop.llm_runtime() - loop.model = "replacement-model" - loop.context_window_tokens = 10 + loop.set_runtime_model("replacement-model") + loop.set_runtime_context_window(10) loop.provider.generation = SimpleNamespace( temperature=1.0, max_tokens=1, diff --git a/tests/test_nanobot_facade.py b/tests/test_nanobot_facade.py index a1a93b46..e2314999 100644 --- a/tests/test_nanobot_facade.py +++ b/tests/test_nanobot_facade.py @@ -30,6 +30,7 @@ from nanobot.nanobot import ( StreamEvent, StreamEventType, ) +from nanobot.utils.llm_runtime import runtime_from_provider_snapshot def _write_config(tmp_path: Path, overrides: dict | None = None) -> Path: @@ -504,45 +505,58 @@ async def test_run_allows_parallel_sessions_without_model_override(tmp_path): @pytest.mark.asyncio -async def test_run_model_overrides_are_serialized_before_snapshot_build(tmp_path): +async def test_run_model_overrides_can_overlap_without_default_mutation(tmp_path): from nanobot.bus.events import OutboundMessage from nanobot.providers.factory import ProviderSnapshot config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - original_model = bot._loop.model - active_models: list[str] = [] - snapshot_base_models: list[str] = [] + assert not hasattr(bot, "_runtime_overrides") + original_runtime = bot._loop.runtime_resolver.runtime + active_runs: list[tuple[str, str]] = [] first_entered = asyncio.Event() + both_entered = asyncio.Event() release_first = asyncio.Event() - def fake_snapshot(*, model, model_preset): + def fake_resolve(*, model, model_preset, config): assert model is not None assert model_preset is None - snapshot_base_models.append(bot._loop.model) - return ProviderSnapshot( + assert config is bot._config + return runtime_from_provider_snapshot(ProviderSnapshot( provider=_fake_provider(model, max_tokens=2048), model=model, context_window_tokens=4096, signature=("sdk", model), - ) + )) - bot._runtime_overrides.model_override_snapshot = MagicMock(side_effect=fake_snapshot) + bot._loop.runtime_resolver.resolve_override = MagicMock(side_effect=fake_resolve) - async def fake_process_direct(message, *, session_key, hooks): - active_models.append(bot._loop.model) + async def fake_process_direct(message, *, session_key, hooks, runtime): + active_runs.append((session_key, runtime.model)) + assert bot._loop.runtime_resolver.runtime is original_runtime if message == "first": first_entered.set() - await asyncio.wait_for(release_first.wait(), timeout=1) + if len(active_runs) == 2: + both_entered.set() + await asyncio.wait_for(release_first.wait(), timeout=1) return OutboundMessage(channel="cli", chat_id="direct", content=message) bot._loop.process_direct = fake_process_direct - first = asyncio.create_task(bot.run("first", model="model:first")) + first = asyncio.create_task(bot.run( + "first", + session_key="sdk:first", + model="model:first", + )) await asyncio.wait_for(first_entered.wait(), timeout=1) - second = asyncio.create_task(bot.run("second", model="model:second")) - await asyncio.sleep(0) + second = asyncio.create_task(bot.run( + "second", + session_key="sdk:second", + model="model:second", + )) + await asyncio.wait_for(both_entered.wait(), timeout=1) + assert not first.done() assert not second.done() release_first.set() @@ -550,21 +564,21 @@ async def test_run_model_overrides_are_serialized_before_snapshot_build(tmp_path assert first_result.content == "first" assert second_result.content == "second" - assert active_models == ["model:first", "model:second"] - assert snapshot_base_models == [original_model, original_model] - assert bot._loop.model == original_model + assert set(active_runs) == { + ("sdk:first", "model:first"), + ("sdk:second", "model:second"), + } + assert bot._loop.runtime_resolver.runtime is original_runtime @pytest.mark.asyncio -async def test_run_model_override_is_per_run_and_restores_default(tmp_path): +async def test_run_model_override_is_per_run_without_default_mutation(tmp_path): from nanobot.bus.events import OutboundMessage from nanobot.providers.factory import ProviderSnapshot config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - original_provider = bot._loop.provider - original_model = bot._loop.model - original_signature = bot._loop._provider_signature + original_runtime = bot._loop.runtime_resolver.runtime override_provider = _fake_provider("override-provider", max_tokens=2048) override = ProviderSnapshot( provider=override_provider, @@ -572,13 +586,17 @@ async def test_run_model_override_is_per_run_and_restores_default(tmp_path): context_window_tokens=4096, signature=("sdk", "override"), ) - bot._runtime_overrides.model_override_snapshot = MagicMock(return_value=override) + override_runtime = runtime_from_provider_snapshot(override) + bot._loop.runtime_resolver.resolve_override = MagicMock( + return_value=override_runtime + ) - async def fake_process_direct(message, *, session_key, hooks): - assert bot._loop.provider is override_provider + async def fake_process_direct(message, *, session_key, hooks, runtime): + assert runtime is override_runtime assert not hasattr(bot._loop.runner, "provider") - assert bot._loop.model == "openai/gpt-4.1-mini" - assert bot._loop.context_window_tokens == 4096 + assert runtime.model == "openai/gpt-4.1-mini" + assert runtime.context_window_tokens == 4096 + assert bot._loop.runtime_resolver.runtime is original_runtime return OutboundMessage(channel="cli", chat_id="direct", content="ok") bot._loop.process_direct = fake_process_direct @@ -586,14 +604,13 @@ async def test_run_model_override_is_per_run_and_restores_default(tmp_path): result = await bot.run("hi", model="openai/gpt-4.1-mini") assert result.content == "ok" - bot._runtime_overrides.model_override_snapshot.assert_called_once_with( + bot._loop.runtime_resolver.resolve_override.assert_called_once_with( model="openai/gpt-4.1-mini", model_preset=None, + config=bot._config, ) - assert bot._loop.provider is original_provider assert not hasattr(bot._loop.runner, "provider") - assert bot._loop.model == original_model - assert bot._loop._provider_signature == original_signature + assert bot._loop.runtime_resolver.runtime is original_runtime @pytest.mark.asyncio @@ -603,7 +620,7 @@ async def test_run_model_preset_override_is_per_run(tmp_path): config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - original_model = bot._loop.model + original_runtime = bot._loop.runtime_resolver.runtime override_provider = _fake_provider("preset-provider", max_tokens=1024) override = ProviderSnapshot( provider=override_provider, @@ -611,19 +628,25 @@ async def test_run_model_preset_override_is_per_run(tmp_path): context_window_tokens=2048, signature=("preset", "fast"), ) - bot._loop._build_model_preset_snapshot = MagicMock(return_value=override) + override_runtime = runtime_from_provider_snapshot(override, model_preset="fast") + bot._loop.runtime_resolver.resolve_override = MagicMock( + return_value=override_runtime + ) - async def fake_process_direct(message, *, session_key, hooks): - assert bot._loop.provider is override_provider - assert bot._loop.model == "openai/gpt-4.1-mini" + async def fake_process_direct(message, *, session_key, hooks, runtime): + assert runtime is override_runtime return OutboundMessage(channel="cli", chat_id="direct", content="ok") bot._loop.process_direct = fake_process_direct await bot.run("hi", model_preset="fast") - bot._loop._build_model_preset_snapshot.assert_called_once_with("fast") - assert bot._loop.model == original_model + bot._loop.runtime_resolver.resolve_override.assert_called_once_with( + model=None, + model_preset="fast", + config=bot._config, + ) + assert bot._loop.runtime_resolver.runtime is original_runtime assert bot._loop.model_preset is None @@ -745,7 +768,9 @@ async def test_stream_yields_text_events_in_order(tmp_path): config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks): + async def fake_process_direct( + message, *, session_key, on_stream, on_stream_end, hooks, runtime + ): assert message == "hi" assert session_key == "sdk:default" await on_stream("Hel") @@ -780,7 +805,9 @@ async def test_run_streamed_wait_returns_full_result_without_consuming_events(tm config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks): + async def fake_process_direct( + message, *, session_key, on_stream, on_stream_end, hooks, runtime + ): await on_stream("done") await on_stream_end(resuming=False) ctx = AgentRunHookContext( @@ -822,7 +849,9 @@ async def test_run_streamed_cancel_releases_full_queue_without_consuming(tmp_pat config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks): + async def fake_process_direct( + message, *, session_key, on_stream, on_stream_end, hooks, runtime + ): for i in range(400): await on_stream(str(i)) await on_stream_end(resuming=False) @@ -886,16 +915,17 @@ async def test_run_streamed_forwards_runtime_options(tmp_path): assert callable(kwargs["on_stream"]) assert callable(kwargs["on_stream_end"]) assert kwargs["hooks"] + assert kwargs["runtime"] is bot._loop.runtime_resolver.runtime @pytest.mark.asyncio -async def test_run_streamed_model_override_reports_model_and_restores(tmp_path): +async def test_run_streamed_model_override_reports_admitted_runtime(tmp_path): from nanobot.bus.events import OutboundMessage from nanobot.providers.factory import ProviderSnapshot config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - original_model = bot._loop.model + original_runtime = bot._loop.runtime_resolver.runtime override_provider = _fake_provider("stream-provider", max_tokens=2048) override = ProviderSnapshot( provider=override_provider, @@ -903,11 +933,22 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path): context_window_tokens=4096, signature=("sdk", "stream"), ) - bot._runtime_overrides.model_override_snapshot = MagicMock(return_value=override) + override_runtime = runtime_from_provider_snapshot(override) + bot._loop.runtime_resolver.resolve_override = MagicMock( + return_value=override_runtime + ) - async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks): - assert bot._loop.provider is override_provider - assert bot._loop.model == "openai/gpt-4.1-mini" + async def fake_process_direct( + message, + *, + session_key, + on_stream, + on_stream_end, + hooks, + runtime, + ): + assert runtime is override_runtime + assert bot._loop.runtime_resolver.runtime is original_runtime await on_stream("ok") await on_stream_end(resuming=False) return OutboundMessage(channel="cli", chat_id="direct", content="ok") @@ -922,7 +963,7 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path): assert events[0].type == "run.started" assert events[0].metadata["model"] == "openai/gpt-4.1-mini" assert events[0].metadata["model_preset"] is None - assert bot._loop.model == original_model + assert bot._loop.runtime_resolver.runtime is original_runtime @pytest.mark.asyncio @@ -947,7 +988,9 @@ async def test_run_streamed_emits_tool_events(tmp_path): config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks): + async def fake_process_direct( + message, *, session_key, on_stream, on_stream_end, hooks, runtime + ): calls = [ ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}), ToolCallRequest(id="call_bad", name="exec", arguments={"cmd": "false"}), @@ -993,7 +1036,9 @@ async def test_run_streamed_emits_reasoning_events(tmp_path): config_path = _write_config(tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path) - async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks): + async def fake_process_direct( + message, *, session_key, on_stream, on_stream_end, hooks, runtime + ): for hook in hooks: await hook.emit_reasoning("thinking") await hook.emit_reasoning_end() @@ -1020,7 +1065,9 @@ async def test_stream_generator_break_cancels_underlying_run(tmp_path): bot = Nanobot.from_config(config_path, workspace=tmp_path) cancelled = asyncio.Event() - async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks): + async def fake_process_direct( + message, *, session_key, on_stream, on_stream_end, hooks, runtime + ): try: await on_stream("first") await asyncio.sleep(10)