refactor(agent): make resolver sole runtime owner

This commit is contained in:
chengyongru 2026-07-10 15:17:20 +08:00 committed by Xubin Ren
parent c9d3e74342
commit 21f58cbabf
16 changed files with 395 additions and 443 deletions

View File

@ -6,6 +6,7 @@ import asyncio
import dataclasses import dataclasses
import os import os
import time import time
from collections.abc import Mapping
from contextlib import AsyncExitStack, nullcontext, suppress from contextlib import AsyncExitStack, nullcontext, suppress
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum, auto from enum import Enum, auto
@ -174,37 +175,48 @@ class AgentLoop:
def tool_names(self) -> list[str]: def tool_names(self) -> list[str]:
return self.tools.tool_names 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: def llm_runtime(self) -> LLMRuntime:
"""Resolve the immutable default used to admit the next turn.""" """Resolve the immutable default used to admit the next turn."""
self._refresh_provider_snapshot() previous = self.runtime_resolver.runtime
runtime = self.runtime_resolver.current() try:
captured = LLMRuntime.capture( runtime = self.runtime_resolver.current(refresh=True)
self.provider, except Exception:
self.model, logger.exception("Failed to refresh model runtime")
context_window_tokens=self.context_window_tokens, return previous
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.
if ( if (
runtime.provider is not self.provider runtime.model != previous.model
or runtime.model != self.model or runtime.model_preset != previous.model_preset
or runtime.generation != captured.generation or runtime.snapshot_signature != previous.snapshot_signature
or runtime.context_window_tokens != self.context_window_tokens
or runtime.model_preset != self._active_preset
): ):
snapshot = ProviderSnapshot( self._publish_runtime_selection(runtime)
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,
)
return runtime return runtime
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint" _RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
@ -271,32 +283,26 @@ class AgentLoop:
self.runtime_event_publisher = RuntimeEventPublisher(self.runtime_events) self.runtime_event_publisher = RuntimeEventPublisher(self.runtime_events)
self.channels_config = channels_config self.channels_config = channels_config
self.restart_mode = restart_mode 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._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.workspace = workspace
self.model = model or provider.get_default_model() initial_model = model or provider.get_default_model()
self.max_iterations = ( self.max_iterations = (
max_iterations if max_iterations is not None else defaults.max_tool_iterations max_iterations if max_iterations is not None else defaults.max_tool_iterations
) )
self.context_window_tokens = ( initial_context_window = (
context_window_tokens context_window_tokens
if context_window_tokens is not None if context_window_tokens is not None
else defaults.context_window_tokens else defaults.context_window_tokens
) )
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {} configured_presets = model_presets or {}
self._active_preset: str | None = None
self.runtime_resolver = ModelRuntimeResolver( self.runtime_resolver = ModelRuntimeResolver(
LLMRuntime.capture( LLMRuntime.capture(
provider, provider,
self.model, initial_model,
context_window_tokens=self.context_window_tokens, context_window_tokens=initial_context_window,
snapshot_signature=provider_signature, snapshot_signature=provider_signature,
), ),
model_presets=self.model_presets, model_presets=configured_presets,
provider_snapshot_loader=provider_snapshot_loader, provider_snapshot_loader=provider_snapshot_loader,
preset_snapshot_loader=preset_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), llm_wall_timeout_for_session=lambda sk: runner_wall_llm_timeout_s(self.sessions, sk),
) )
self._unified_session = unified_session self._unified_session = unified_session
self._max_messages = replay_max_messages_for_context(self.context_window_tokens)
self._running = False self._running = False
self._mcp_servers = mcp_servers or {} self._mcp_servers = mcp_servers or {}
self._mcp_stacks: dict[str, AsyncExitStack] = {} self._mcp_stacks: dict[str, AsyncExitStack] = {}
@ -401,7 +406,7 @@ class AgentLoop:
) )
if model_preset: if model_preset:
self.set_model_preset(model_preset, publish_update=False) 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._runtime_vars: dict[str, Any] = {}
self._current_iteration: int = 0 self._current_iteration: int = 0
self.commands = CommandRouter() self.commands = CommandRouter()
@ -468,98 +473,51 @@ class AgentLoop:
"""Keep subagent runtime limits aligned with mutable loop settings.""" """Keep subagent runtime limits aligned with mutable loop settings."""
self.subagents.max_iterations = self.max_iterations self.subagents.max_iterations = self.max_iterations
def _apply_provider_snapshot( def _publish_runtime_selection(
self, self,
snapshot: ProviderSnapshot, runtime: LLMRuntime,
*, *,
publish_update: bool = True, publish_update: bool = True,
model_preset: str | None = None,
) -> None: ) -> None:
"""Swap model/provider for future turns without disturbing an active one.""" if not publish_update:
runtime = self.runtime_resolver.adopt_snapshot( return
snapshot, if self._runtime_model_publisher is not None:
model_preset=model_preset, 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 def set_model_preset(
context_window_tokens = runtime.context_window_tokens self,
provider.generation = runtime.generation name: str | None,
*,
publish_update: bool = True,
) -> LLMRuntime:
"""Select a named default runtime for future turns."""
old_model = self.model 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) runtime = self.runtime_resolver.select_preset(name)
snapshot = ProviderSnapshot( self._publish_runtime_selection(runtime, publish_update=publish_update)
provider=runtime.provider, logger.info(
model=runtime.model, "Runtime model switched for next turn: {} -> {}",
context_window_tokens=runtime.context_window_tokens, old_model,
signature=runtime.snapshot_signature or ("model_preset", name), runtime.model,
generation=runtime.generation,
) )
self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name) return runtime
self._active_preset = name
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.""" """Register the default set of tools via plugin loader."""
from nanobot.agent.tools.context import ToolContext from nanobot.agent.tools.context import ToolContext
from nanobot.agent.tools.loader import ToolLoader from nanobot.agent.tools.loader import ToolLoader
@ -571,7 +529,7 @@ class AgentLoop:
subagent_manager=self.subagents, subagent_manager=self.subagents,
cron_service=self.cron_service, cron_service=self.cron_service,
sessions=self.sessions, 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, image_generation_provider_configs=self._image_generation_provider_configs,
timezone=self.context.timezone or "UTC", timezone=self.context.timezone or "UTC",
workspace_sandbox=self.workspace_scopes.sandbox_status, workspace_sandbox=self.workspace_scopes.sandbox_status,

View File

@ -56,7 +56,9 @@ class RuntimeState(Protocol):
def _sync_subagent_runtime_limits(self) -> None: ... 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 @property
def model_preset(self) -> str | None: ... def model_preset(self) -> str | None: ...
_active_preset: str | None

View File

@ -57,7 +57,7 @@ class MyTool(Tool):
BLOCKED = frozenset({ BLOCKED = frozenset({
# Core infrastructure # Core infrastructure
"bus", "provider", "_running", "tools", "bus", "provider", "runtime_resolver", "_running", "tools",
# Config management # Config management
"_runtime_vars", "_runtime_vars",
# Subsystems # Subsystems
@ -107,6 +107,11 @@ class MyTool(Tool):
} }
_MAX_RUNTIME_KEYS = 64 _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: def __init__(self, runtime_state: RuntimeState, modify_allowed: bool = True) -> None:
self._runtime_state = runtime_state self._runtime_state = runtime_state
@ -325,9 +330,20 @@ class MyTool(Tool):
# -- inspect -- # -- 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: def _inspect(self, key: str | None) -> str:
if not key: if not key:
return self._inspect_all() 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] top = key.split(".")[0]
if top in self._DENIED_ATTRS or top.startswith("__"): if top in self._DENIED_ATTRS or top.startswith("__"):
return ToolResult.error(f"Error: '{top}' is not accessible") return ToolResult.error(f"Error: '{top}' is not accessible")
@ -353,8 +369,13 @@ class MyTool(Tool):
parts: list[str] = [] parts: list[str] = []
# RESTRICTED keys # RESTRICTED keys
for k in self.RESTRICTED: for k in self.RESTRICTED:
parts.append(self._format_value(getattr(state, k, None), k)) found, value = self._current_runtime_value(k)
parts.append(self._format_value(state.model_preset, "model_preset")) 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 # 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"): 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): if _has_real_attr(state, k):
@ -432,13 +453,16 @@ class MyTool(Tool):
return ToolResult.error(f"Error: '{key}' must be <= {spec['max']}") return ToolResult.error(f"Error: '{key}' must be <= {spec['max']}")
if "min_len" in spec and len(str(value)) < spec["min_len"]: 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") return ToolResult.error(f"Error: '{key}' must be at least {spec['min_len']} characters")
setattr(self._runtime_state, key, value)
if key == "model": if key == "model":
self._runtime_state._active_preset = None self._runtime_state.set_runtime_model(value)
sync_replay = getattr(self._runtime_state, "_sync_replay_max_messages", None) elif key == "context_window_tokens":
if key == "context_window_tokens" and callable(sync_replay): self._runtime_state.set_runtime_context_window(value)
sync_replay() else:
if key == "max_iterations" and hasattr(self._runtime_state, "_sync_subagent_runtime_limits"): 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._runtime_state._sync_subagent_runtime_limits()
self._audit("modify", f"{key}: {old!r} -> {value!r}") self._audit("modify", f"{key}: {old!r} -> {value!r}")
return f"Set {key} = {value!r} (was {old!r})" return f"Set {key} = {value!r} (was {old!r})"

View File

@ -349,7 +349,7 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
name = parts[0] name = parts[0]
try: try:
loop.set_model_preset(name) runtime = loop.set_model_preset(name)
except (KeyError, ValueError) as exc: except (KeyError, ValueError) as exc:
names = _model_preset_names(loop) names = _model_preset_names(loop)
return OutboundMessage( return OutboundMessage(
@ -362,11 +362,11 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
metadata=metadata, metadata=metadata,
) )
max_tokens = getattr(getattr(loop.provider, "generation", None), "max_tokens", None) max_tokens = runtime.generation.max_tokens
lines = [ lines = [
f"Switched model preset to `{loop.model_preset}`.", f"Switched model preset to `{runtime.model_preset}`.",
f"- Model: `{loop.model}`", f"- Model: `{runtime.model}`",
f"- Context window: {loop.context_window_tokens}", f"- Context window: {runtime.context_window_tokens}",
] ]
if max_tokens is not None: if max_tokens is not None:
lines.append(f"- Max output tokens: {max_tokens}") lines.append(f"- Max output tokens: {max_tokens}")

View File

@ -14,7 +14,6 @@ from nanobot.config.schema import Config
from nanobot.providers.image_generation import image_gen_provider_configs from nanobot.providers.image_generation import image_gen_provider_configs
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
from nanobot.sdk.runtime import ( from nanobot.sdk.runtime import (
SDKRuntimeController,
build_process_direct_kwargs, build_process_direct_kwargs,
ensure_single_model_selector, ensure_single_model_selector,
) )
@ -74,7 +73,6 @@ class Nanobot:
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None: def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
self._loop = loop self._loop = loop
self._config = config self._config = config
self._runtime_overrides = SDKRuntimeController(loop, config=config)
self.sessions = SessionClient(loop) self.sessions = SessionClient(loop)
self.memory = MemoryClient(loop) self.memory = MemoryClient(loop)
self.runtime = RuntimeClient(loop) self.runtime = RuntimeClient(loop)
@ -156,20 +154,26 @@ class Nanobot:
""" """
capture = SDKCaptureHook() capture = SDKCaptureHook()
per_run_hooks = [capture, *(hooks or [])] per_run_hooks = [capture, *(hooks or [])]
async with self._runtime_overrides.override(model=model, model_preset=model_preset): runtime = self._loop.runtime_resolver.resolve_override(
kwargs = build_process_direct_kwargs( model=model,
session_key=session_key, model_preset=model_preset,
channel=channel, config=self._config,
chat_id=chat_id, )
sender_id=sender_id, kwargs = build_process_direct_kwargs(
media=media, session_key=session_key,
ephemeral=ephemeral, channel=channel,
) chat_id=chat_id,
response = await self._loop.process_direct( sender_id=sender_id,
message, media=media,
**kwargs, ephemeral=ephemeral,
hooks=per_run_hooks, )
) 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) return result_from_response(response, capture)
@ -188,7 +192,11 @@ class Nanobot:
model_preset: str | None = None, model_preset: str | None = None,
) -> RunStream: ) -> RunStream:
"""Start a streamed run and return a handle for events and final result.""" """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) queue: asyncio.Queue[StreamEvent | object] = asyncio.Queue(maxsize=256)
emitter = SDKStreamEmitter(queue) emitter = SDKStreamEmitter(queue)
stream_hook = SDKStreamingHook(emitter) stream_hook = SDKStreamingHook(emitter)
@ -202,55 +210,53 @@ class Nanobot:
await emitter.text_completed(resuming=resuming) await emitter.text_completed(resuming=resuming)
async def _run() -> RunResult: async def _run() -> RunResult:
async with self._runtime_overrides.override(model=model, model_preset=model_preset): kwargs = build_process_direct_kwargs(
kwargs = build_process_direct_kwargs( session_key=session_key,
session_key=session_key, channel=channel,
channel=channel, chat_id=chat_id,
chat_id=chat_id, sender_id=sender_id,
sender_id=sender_id, media=media,
media=media, ephemeral=ephemeral,
ephemeral=ephemeral, on_stream=_on_stream,
on_stream=_on_stream, on_stream_end=_on_stream_end,
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( await emitter.emit(StreamEvent(
type=STREAM_EVENT_RUN_STARTED, type=STREAM_EVENT_RUN_COMPLETED,
metadata={ content=result.content,
"session_key": session_key, result=result,
"channel": channel, usage=dict(result.usage),
"chat_id": chat_id, metadata=dict(result.metadata),
"sender_id": sender_id,
"model": self._loop.model,
"model_preset": (
model_preset if model_preset is not None else self._loop.model_preset
),
},
)) ))
try: return result
response = await self._loop.process_direct( except Exception as exc:
message, await emitter.emit(StreamEvent(
**kwargs, type=STREAM_EVENT_RUN_FAILED,
hooks=per_run_hooks, error=str(exc),
) metadata={"exception_type": type(exc).__name__},
await emitter.text_completed(resuming=False, force=False) ))
result = result_from_response(response, capture) raise
await emitter.emit(StreamEvent( finally:
type=STREAM_EVENT_RUN_COMPLETED, emitter.close()
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()
task = asyncio.create_task(_run()) task = asyncio.create_task(_run())
return RunStream(task, queue) return RunStream(task, queue)

View File

@ -2,16 +2,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio from typing import Any
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
def ensure_single_model_selector( def ensure_single_model_selector(
@ -51,142 +42,3 @@ def build_process_direct_kwargs(
if on_stream_end is not None: if on_stream_end is not None:
kwargs["on_stream_end"] = on_stream_end kwargs["on_stream_end"] = on_stream_end
return kwargs 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

View File

@ -518,9 +518,11 @@ class TestNewCommandArchival:
loop.sessions.save(session) loop.sessions.save(session)
call_count = 0 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 nonlocal call_count
assert runtime is expected_runtime
assert session_key == "cli:test" assert session_key == "cli:test"
call_count += 1 call_count += 1
return False return False
@ -528,7 +530,7 @@ class TestNewCommandArchival:
loop.consolidator.archive = _failing_summarize # type: ignore[method-assign] loop.consolidator.archive = _failing_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") 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 response is not None
assert "new session started" in response.content.lower() assert "new session started" in response.content.lower()
@ -553,9 +555,11 @@ class TestNewCommandArchival:
archived_count = -1 archived_count = -1
archived_session_key = None 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 nonlocal archived_count, archived_session_key
assert runtime is expected_runtime
archived_count = len(messages) archived_count = len(messages)
archived_session_key = session_key archived_session_key = session_key
return True return True
@ -563,7 +567,7 @@ class TestNewCommandArchival:
loop.consolidator.archive = _fake_summarize # type: ignore[method-assign] loop.consolidator.archive = _fake_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") 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 response is not None
assert "new session started" in response.content.lower() assert "new session started" in response.content.lower()
@ -582,15 +586,17 @@ class TestNewCommandArchival:
session.add_message("user", f"msg{i}") session.add_message("user", f"msg{i}")
session.add_message("assistant", f"resp{i}") session.add_message("assistant", f"resp{i}")
loop.sessions.save(session) 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" assert session_key == "cli:test"
return True return True
loop.consolidator.archive = _ok_summarize # type: ignore[method-assign] loop.consolidator.archive = _ok_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") 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 response is not None
assert "new session started" in response.content.lower() assert "new session started" in response.content.lower()
@ -610,8 +616,10 @@ class TestNewCommandArchival:
archived = asyncio.Event() archived = asyncio.Event()
release_archive = 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" assert session_key == "cli:test"
await release_archive.wait() await release_archive.wait()
archived.set() archived.set()
@ -620,7 +628,7 @@ class TestNewCommandArchival:
loop.consolidator.archive = _slow_summarize # type: ignore[method-assign] loop.consolidator.archive = _slow_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new") 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() assert not archived.is_set()
release_archive.set() release_archive.set()

View File

@ -6,6 +6,7 @@ import nanobot.agent.memory as memory_module
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.providers.base import LLMResponse 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: 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( loop.consolidator.maybe_consolidate_by_tokens.assert_any_await(
session, session,
runtime=runtime, 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 len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2
assert all( assert all(

View File

@ -21,6 +21,7 @@ from nanobot.bus.outbound_events import (
) )
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMResponse, ToolCallRequest from nanobot.providers.base import LLMResponse, ToolCallRequest
from nanobot.providers.factory import ProviderSnapshot
from nanobot.session.webui_turns import WebuiTurnCoordinator from nanobot.session.webui_turns import WebuiTurnCoordinator
from nanobot.utils.progress_events import ( from nanobot.utils.progress_events import (
invoke_file_edit_progress, invoke_file_edit_progress,
@ -815,8 +816,14 @@ class TestToolEventProgress:
)) ))
assert len(scheduled_title) == 1 assert len(scheduled_title) == 1
loop.provider = MagicMock() next_provider = MagicMock()
loop.model = "switched-after-turn" 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] await scheduled_title[0] # type: ignore[misc]

View File

@ -19,6 +19,7 @@ from nanobot.bus.outbound_events import (
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
from nanobot.providers.base import LLMResponse from nanobot.providers.base import LLMResponse
from nanobot.providers.factory import ProviderSnapshot
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
from nanobot.session.goal_state import GOAL_STATE_KEY from nanobot.session.goal_state import GOAL_STATE_KEY
from nanobot.session.manager import Session, SessionManager 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" assert runtime.model == "test-model"
next_provider = MagicMock() next_provider = MagicMock()
loop.provider = next_provider next_provider.generation = SimpleNamespace(
loop.model = "next-model" 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() runtime = loop.llm_runtime()
assert runtime.provider is next_provider assert runtime.provider is next_provider

View File

@ -2,6 +2,7 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import replace
from pathlib import Path from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch 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: def test_default_for_200k_context_reaches_file_cap(self, tmp_path: Path) -> None:
loop = _make_loop(tmp_path) 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: def test_default_scales_with_context_window(self, tmp_path: Path) -> None:
loop = _make_loop(tmp_path, context_window_tokens=32_768) 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: def test_provider_refresh_resyncs_context_derived_limit(self, tmp_path: Path) -> None:
old_provider = MagicMock() old_provider = MagicMock()
@ -93,9 +96,10 @@ class TestMaxMessagesInit:
), ),
) )
assert loop._max_messages == 327 initial = loop.runtime_resolver.runtime
loop._refresh_provider_snapshot() assert replay_max_messages_for_context(initial.context_window_tokens) == 327
assert loop._max_messages == FILE_MAX_MESSAGES refreshed = loop.llm_runtime()
assert replay_max_messages_for_context(refreshed.context_window_tokens) == FILE_MAX_MESSAGES
class TestGetHistoryWithMaxMessages: class TestGetHistoryWithMaxMessages:
@ -136,7 +140,7 @@ class TestMaxMessagesIntegration:
async def test_process_message_passes_limit_to_history_call(self, tmp_path: Path) -> None: 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.""" """The real message path should pass max_messages into session history replay."""
loop = _make_loop(tmp_path) 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( loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="ok", tool_calls=[], usage={}) return_value=LLMResponse(content="ok", tool_calls=[], usage={})
) )
@ -146,12 +150,13 @@ class TestMaxMessagesIntegration:
session = loop.sessions.get_or_create("cli:test") session = loop.sessions.get_or_create("cli:test")
with patch.object(session, "get_history", wraps=session.get_history) as mock_hist: with patch.object(session, "get_history", wraps=session.get_history) as mock_hist:
result = await loop._process_message( 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 result is not None
assert mock_hist.call_count == 1 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 assert mock_hist.call_args.kwargs["extend_to_user"] is False
@pytest.mark.asyncio @pytest.mark.asyncio
@ -182,8 +187,7 @@ class TestMaxMessagesIntegration:
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
"""A live user turn should not extend history to an older long tool turn.""" """A live user turn should not extend history to an older long tool turn."""
loop = _make_loop(tmp_path) loop = _make_loop(tmp_path, context_window_tokens=8_000)
loop._max_messages = 6
loop.provider.chat_with_retry = AsyncMock( loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="ok", tool_calls=[], usage={}) return_value=LLMResponse(content="ok", tool_calls=[], usage={})
) )
@ -194,7 +198,7 @@ class TestMaxMessagesIntegration:
session.add_message("user", "old") session.add_message("user", "old")
session.add_message("assistant", "old answer") session.add_message("assistant", "old answer")
session.add_message("user", "long older turn") 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.messages.extend(_tool_round(f"older-{i}"))
session.add_message("assistant", "older final") session.add_message("assistant", "older final")

View File

@ -18,7 +18,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
return provider 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") old_provider = _provider("old-model")
new_provider = _provider("new-model", max_tokens=456) new_provider = _provider("new-model", max_tokens=456)
loop = AgentLoop( 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.provider is new_provider
assert loop.model == "new-model" assert loop.model == "new-model"
assert loop.context_window_tokens == 2000 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") 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: def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
old_provider = _provider("old-model") old_provider = _provider("old-model")
new_provider = _provider("new-model", max_tokens=456) 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) loop = AgentLoop.from_config(config, provider_snapshot_loader=loader)
payload = update_agent_settings({"context_window_tokens": ["262144"]}) payload = update_agent_settings({"context_window_tokens": ["262144"]})
loop._refresh_provider_snapshot() loop.llm_runtime()
assert payload["requires_restart"] is False assert payload["requires_restart"] is False
assert loop.context_window_tokens == 262_144 assert loop.context_window_tokens == 262_144

View File

@ -54,9 +54,10 @@ def test_model_preset_setter_updates_state(tmp_path) -> None:
assert loop.model_preset == "fast" assert loop.model_preset == "fast"
assert loop.model == "openai/gpt-4.1" assert loop.model == "openai/gpt-4.1"
assert loop.context_window_tokens == 32_768 assert loop.context_window_tokens == 32_768
assert loop.provider.generation.temperature == 0.5 runtime = loop.llm_runtime()
assert loop.provider.generation.max_tokens == 4096 assert runtime.generation.temperature == 0.5
assert loop.provider.generation.reasoning_effort == "low" assert runtime.generation.max_tokens == 4096
assert runtime.generation.reasoning_effort == "low"
assert not hasattr(loop.subagents, "model") assert not hasattr(loop.subagents, "model")
assert not hasattr(loop.consolidator, "model") assert not hasattr(loop.consolidator, "model")
assert not hasattr(loop.consolidator, "context_window_tokens") 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.set_model_preset("fast")
loop._refresh_provider_snapshot() loop.llm_runtime()
assert loop.model_preset == "fast" assert loop.model_preset == "fast"
assert loop.provider is fast_provider 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.set_model_preset("fast")
loop._refresh_provider_snapshot() loop.llm_runtime()
assert loop.model_preset is None assert loop.model_preset is None
assert loop.provider is webui_provider 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) tool = MyTool(runtime_state=loop, modify_allowed=True)
result = tool._modify("model", "anthropic/claude-opus-4-5") result = tool._modify("model", "anthropic/claude-opus-4-5")
assert "Error" not in result 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" 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") fake_provider = _provider("openai/gpt-4.1")
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider): with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
loop = AgentLoop.from_config(config) loop = AgentLoop.from_config(config)
assert loop._provider_snapshot_loader is None default_runtime = loop.runtime_resolver.runtime
assert loop._preset_snapshot_loader is not None resolved = loop.runtime_resolver.resolve_preset("fast")
assert resolved.model == "openai/gpt-4.1-mini"
assert loop.runtime_resolver.runtime is default_runtime

View File

@ -35,6 +35,12 @@ def _make_mock_loop(**overrides):
loop._concurrency_gate = None loop._concurrency_gate = None
loop._unified_session = False loop._unified_session = False
loop._extra_hooks = [] 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 # web_config mock — needed for check tests
loop.web_config = MagicMock() loop.web_config = MagicMock()
@ -237,12 +243,12 @@ class TestModifyRestricted:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_modify_context_window_valid(self): 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) tool = _make_tool(runtime_state=loop)
result = await tool.execute(action="set", key="context_window_tokens", value=131072) result = await tool.execute(action="set", key="context_window_tokens", value=131072)
assert "Set context_window_tokens" in result assert "Set context_window_tokens" in result
assert loop.context_window_tokens == 131072 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 @pytest.mark.asyncio
async def test_modify_none_value_for_restricted_int(self): async def test_modify_none_value_for_restricted_int(self):

View File

@ -245,8 +245,8 @@ class TestRestartCommand:
msg = InboundMessage(channel="telegram", sender_id="u1", chat_id="c1", content="/status") msg = InboundMessage(channel="telegram", sender_id="u1", chat_id="c1", content="/status")
runtime = loop.llm_runtime() runtime = loop.llm_runtime()
loop.model = "replacement-model" loop.set_runtime_model("replacement-model")
loop.context_window_tokens = 10 loop.set_runtime_context_window(10)
loop.provider.generation = SimpleNamespace( loop.provider.generation = SimpleNamespace(
temperature=1.0, temperature=1.0,
max_tokens=1, max_tokens=1,

View File

@ -30,6 +30,7 @@ from nanobot.nanobot import (
StreamEvent, StreamEvent,
StreamEventType, StreamEventType,
) )
from nanobot.utils.llm_runtime import runtime_from_provider_snapshot
def _write_config(tmp_path: Path, overrides: dict | None = None) -> Path: 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 @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.bus.events import OutboundMessage
from nanobot.providers.factory import ProviderSnapshot from nanobot.providers.factory import ProviderSnapshot
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path)
original_model = bot._loop.model assert not hasattr(bot, "_runtime_overrides")
active_models: list[str] = [] original_runtime = bot._loop.runtime_resolver.runtime
snapshot_base_models: list[str] = [] active_runs: list[tuple[str, str]] = []
first_entered = asyncio.Event() first_entered = asyncio.Event()
both_entered = asyncio.Event()
release_first = 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 is not None
assert model_preset is None assert model_preset is None
snapshot_base_models.append(bot._loop.model) assert config is bot._config
return ProviderSnapshot( return runtime_from_provider_snapshot(ProviderSnapshot(
provider=_fake_provider(model, max_tokens=2048), provider=_fake_provider(model, max_tokens=2048),
model=model, model=model,
context_window_tokens=4096, context_window_tokens=4096,
signature=("sdk", model), 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): async def fake_process_direct(message, *, session_key, hooks, runtime):
active_models.append(bot._loop.model) active_runs.append((session_key, runtime.model))
assert bot._loop.runtime_resolver.runtime is original_runtime
if message == "first": if message == "first":
first_entered.set() 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) return OutboundMessage(channel="cli", chat_id="direct", content=message)
bot._loop.process_direct = fake_process_direct 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) await asyncio.wait_for(first_entered.wait(), timeout=1)
second = asyncio.create_task(bot.run("second", model="model:second")) second = asyncio.create_task(bot.run(
await asyncio.sleep(0) "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() assert not second.done()
release_first.set() 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 first_result.content == "first"
assert second_result.content == "second" assert second_result.content == "second"
assert active_models == ["model:first", "model:second"] assert set(active_runs) == {
assert snapshot_base_models == [original_model, original_model] ("sdk:first", "model:first"),
assert bot._loop.model == original_model ("sdk:second", "model:second"),
}
assert bot._loop.runtime_resolver.runtime is original_runtime
@pytest.mark.asyncio @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.bus.events import OutboundMessage
from nanobot.providers.factory import ProviderSnapshot from nanobot.providers.factory import ProviderSnapshot
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path)
original_provider = bot._loop.provider original_runtime = bot._loop.runtime_resolver.runtime
original_model = bot._loop.model
original_signature = bot._loop._provider_signature
override_provider = _fake_provider("override-provider", max_tokens=2048) override_provider = _fake_provider("override-provider", max_tokens=2048)
override = ProviderSnapshot( override = ProviderSnapshot(
provider=override_provider, 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, context_window_tokens=4096,
signature=("sdk", "override"), 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): async def fake_process_direct(message, *, session_key, hooks, runtime):
assert bot._loop.provider is override_provider assert runtime is override_runtime
assert not hasattr(bot._loop.runner, "provider") assert not hasattr(bot._loop.runner, "provider")
assert bot._loop.model == "openai/gpt-4.1-mini" assert runtime.model == "openai/gpt-4.1-mini"
assert bot._loop.context_window_tokens == 4096 assert runtime.context_window_tokens == 4096
assert bot._loop.runtime_resolver.runtime is original_runtime
return OutboundMessage(channel="cli", chat_id="direct", content="ok") return OutboundMessage(channel="cli", chat_id="direct", content="ok")
bot._loop.process_direct = fake_process_direct 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") result = await bot.run("hi", model="openai/gpt-4.1-mini")
assert result.content == "ok" 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="openai/gpt-4.1-mini",
model_preset=None, model_preset=None,
config=bot._config,
) )
assert bot._loop.provider is original_provider
assert not hasattr(bot._loop.runner, "provider") assert not hasattr(bot._loop.runner, "provider")
assert bot._loop.model == original_model assert bot._loop.runtime_resolver.runtime is original_runtime
assert bot._loop._provider_signature == original_signature
@pytest.mark.asyncio @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) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=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_provider = _fake_provider("preset-provider", max_tokens=1024)
override = ProviderSnapshot( override = ProviderSnapshot(
provider=override_provider, provider=override_provider,
@ -611,19 +628,25 @@ async def test_run_model_preset_override_is_per_run(tmp_path):
context_window_tokens=2048, context_window_tokens=2048,
signature=("preset", "fast"), 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): async def fake_process_direct(message, *, session_key, hooks, runtime):
assert bot._loop.provider is override_provider assert runtime is override_runtime
assert bot._loop.model == "openai/gpt-4.1-mini"
return OutboundMessage(channel="cli", chat_id="direct", content="ok") return OutboundMessage(channel="cli", chat_id="direct", content="ok")
bot._loop.process_direct = fake_process_direct bot._loop.process_direct = fake_process_direct
await bot.run("hi", model_preset="fast") await bot.run("hi", model_preset="fast")
bot._loop._build_model_preset_snapshot.assert_called_once_with("fast") bot._loop.runtime_resolver.resolve_override.assert_called_once_with(
assert bot._loop.model == original_model model=None,
model_preset="fast",
config=bot._config,
)
assert bot._loop.runtime_resolver.runtime is original_runtime
assert bot._loop.model_preset is None 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) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=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 message == "hi"
assert session_key == "sdk:default" assert session_key == "sdk:default"
await on_stream("Hel") 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) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=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("done")
await on_stream_end(resuming=False) await on_stream_end(resuming=False)
ctx = AgentRunHookContext( 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) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=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): for i in range(400):
await on_stream(str(i)) await on_stream(str(i))
await on_stream_end(resuming=False) 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"])
assert callable(kwargs["on_stream_end"]) assert callable(kwargs["on_stream_end"])
assert kwargs["hooks"] assert kwargs["hooks"]
assert kwargs["runtime"] is bot._loop.runtime_resolver.runtime
@pytest.mark.asyncio @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.bus.events import OutboundMessage
from nanobot.providers.factory import ProviderSnapshot from nanobot.providers.factory import ProviderSnapshot
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=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_provider = _fake_provider("stream-provider", max_tokens=2048)
override = ProviderSnapshot( override = ProviderSnapshot(
provider=override_provider, provider=override_provider,
@ -903,11 +933,22 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
context_window_tokens=4096, context_window_tokens=4096,
signature=("sdk", "stream"), 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): async def fake_process_direct(
assert bot._loop.provider is override_provider message,
assert bot._loop.model == "openai/gpt-4.1-mini" *,
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("ok")
await on_stream_end(resuming=False) await on_stream_end(resuming=False)
return OutboundMessage(channel="cli", chat_id="direct", content="ok") 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].type == "run.started"
assert events[0].metadata["model"] == "openai/gpt-4.1-mini" assert events[0].metadata["model"] == "openai/gpt-4.1-mini"
assert events[0].metadata["model_preset"] is None 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 @pytest.mark.asyncio
@ -947,7 +988,9 @@ async def test_run_streamed_emits_tool_events(tmp_path):
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=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 = [ calls = [
ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}), ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}),
ToolCallRequest(id="call_bad", name="exec", arguments={"cmd": "false"}), 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) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=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: for hook in hooks:
await hook.emit_reasoning("thinking") await hook.emit_reasoning("thinking")
await hook.emit_reasoning_end() 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) bot = Nanobot.from_config(config_path, workspace=tmp_path)
cancelled = asyncio.Event() 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: try:
await on_stream("first") await on_stream("first")
await asyncio.sleep(10) await asyncio.sleep(10)