refactor(agent): make resolver sole runtime owner

This commit is contained in:
chengyongru
2026-07-10 17:54:34 +08:00
committed by Xubin Ren
parent c9d3e74342
commit 21f58cbabf
16 changed files with 395 additions and 443 deletions
+83 -125
View File
@@ -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,
+4 -2
View File
@@ -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
+33 -9
View File
@@ -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})"
+5 -5
View File
@@ -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}")
+69 -63
View File
@@ -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)
+1 -149
View File
@@ -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
+16 -8
View File
@@ -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()
@@ -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(
+9 -2
View File
@@ -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]
+12 -2
View File
@@ -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
+15 -11
View File
@@ -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")
+27 -3
View File
@@ -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
+11 -8
View File
@@ -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
+8 -2
View File
@@ -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):
+2 -2
View File
@@ -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,
+98 -51
View File
@@ -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)