feat(agent): make model presets session-scoped (#4866)

This commit is contained in:
chengyongru
2026-07-23 00:38:49 +08:00
committed by GitHub
parent 66690fdb0c
commit c22efb5f7a
41 changed files with 1140 additions and 155 deletions
+4 -4
View File
@@ -307,7 +307,7 @@ class TestAutoCompact:
loop.sessions.save(s2)
loop.consolidator.compact_idle_session = _make_fake_compact(loop)
loop.auto_compact.check_expired(loop._schedule_background, loop.llm_runtime)
loop.auto_compact.check_expired(loop._schedule_background, loop.runtime_for_session)
await _drain_background_tasks(loop)
active_after = loop.sessions.get_or_create("cli:active")
@@ -776,7 +776,7 @@ class TestProactiveAutoCompact:
"""Helper: run check_expired via callback and wait for background tasks."""
loop.auto_compact.check_expired(
loop._schedule_background,
loop.llm_runtime,
loop.runtime_for_session,
active_session_keys=active_session_keys,
)
await _drain_background_tasks(loop)
@@ -878,12 +878,12 @@ class TestProactiveAutoCompact:
loop.consolidator.compact_idle_session = _slow_compact
# First call starts archiving via callback
loop.auto_compact.check_expired(loop._schedule_background, loop.llm_runtime)
loop.auto_compact.check_expired(loop._schedule_background, loop.runtime_for_session)
await started.wait()
assert archive_count == 1
# Second call should skip (key is in _archiving)
loop.auto_compact.check_expired(loop._schedule_background, loop.llm_runtime)
loop.auto_compact.check_expired(loop._schedule_background, loop.runtime_for_session)
assert archive_count == 1
# Clean up
+51 -2
View File
@@ -9,7 +9,7 @@ from nanobot.agent.autocompact import AutoCompact
from nanobot.session.manager import Session, SessionManager
def _runtime():
def _runtime(_session: Session | None = None):
return MagicMock(name="runtime")
@@ -240,13 +240,62 @@ class TestCheckExpired:
resolve_runtime.return_value = replacement
await scheduled[0]
resolve_runtime.assert_called_once_with()
resolve_runtime.assert_called_once_with(session)
ac.consolidator.compact_idle_session.assert_awaited_once_with(
"cli:old",
runtime=admitted,
max_suffix=ac._RECENT_SUFFIX_MESSAGES,
)
@pytest.mark.parametrize("resolution_error", [KeyError, ValueError])
def test_invalid_preset_is_isolated_to_one_session(self, resolution_error):
ac = _make_autocompact(ttl=15)
old_dt = datetime.now() - timedelta(minutes=20)
sessions = {
key: _make_session(key, updated_at=old_dt)
for key in ("cli:removed", "cli:healthy")
}
for session in sessions.values():
_add_turns(session, 5)
ac.sessions.list_sessions.return_value = [
{"key": key, "updated_at": old_dt.isoformat()}
for key in sessions
]
ac.sessions.get_or_create.side_effect = sessions.__getitem__
healthy_runtime = _runtime()
def resolve_runtime(session: Session):
if session.key == "cli:removed":
raise resolution_error("model preset cannot be resolved")
return healthy_runtime
scheduled = []
def scheduler(coro):
scheduled.append(coro)
coro.close()
ac.check_expired(scheduler, resolve_runtime)
assert len(scheduled) == 1
assert ac._archiving == {"cli:healthy"}
def test_unexpected_runtime_resolution_failure_propagates(self):
ac = _make_autocompact(ttl=15)
old_dt = datetime.now() - timedelta(minutes=20)
session = _make_session("cli:old", updated_at=old_dt)
_add_turns(session, 5)
ac.sessions.list_sessions.return_value = [
{"key": session.key, "updated_at": old_dt.isoformat()}
]
ac.sessions.get_or_create.return_value = session
def fail(_session: Session):
raise RuntimeError("unexpected resolver failure")
with pytest.raises(RuntimeError, match="unexpected resolver failure"):
ac.check_expired(MagicMock(), fail)
def test_active_session_key_skips(self):
"""Session in active_session_keys should be skipped."""
ac = _make_autocompact(ttl=15)
+97 -1
View File
@@ -52,9 +52,10 @@ def test_provider_snapshot_has_one_canonical_runtime_conversion() -> None:
model="snapshot-model",
context_window_tokens=32_768,
signature=("snapshot-model", "openai"),
model_preset="fast",
)
runtime = runtime_from_provider_snapshot(snapshot, model_preset="fast")
runtime = runtime_from_provider_snapshot(snapshot)
assert runtime.provider is provider
assert runtime.model == "snapshot-model"
@@ -94,6 +95,101 @@ def test_resolver_resolves_preset_without_mutating_selected_runtime() -> None:
assert resolved.generation == GenerationSettings(0.5, 512, None)
def test_resolver_reuses_preset_until_runtime_config_is_invalidated() -> None:
initial = _runtime()
preset = ModelPresetConfig(model="fast-model")
load_count = 0
preset_signature = ("fast-model", "auto", "initial")
def load_preset(_name: str) -> ProviderSnapshot:
nonlocal load_count
load_count += 1
return ProviderSnapshot(
provider=_provider(),
model="fast-model",
context_window_tokens=20_000,
signature=preset_signature,
)
resolver = ModelRuntimeResolver(
initial,
model_presets={"fast": preset},
preset_snapshot_loader=load_preset,
)
first = resolver.resolve_preset("fast")
second = resolver.resolve_preset("fast")
assert first is second
assert load_count == 1
preset_signature = ("fast-model", "auto", "new-credential")
resolver.invalidate()
refreshed = resolver.resolve_preset("fast")
assert refreshed is not first
assert load_count == 2
def test_resolver_refreshes_preset_catalog_after_invalidation() -> None:
provider = _provider()
catalog = {
"old": ModelPresetConfig(model="old-model", provider="openai"),
}
default_name = "old"
def load_preset(name: str) -> ProviderSnapshot:
preset = catalog[name]
return ProviderSnapshot(
provider=provider,
model=preset.model,
context_window_tokens=preset.context_window_tokens,
signature=(preset.model, preset.provider),
model_preset=name,
)
resolver = ModelRuntimeResolver(
runtime_from_provider_snapshot(load_preset("old")),
model_presets=catalog,
preset_catalog_loader=lambda: catalog,
configured_default_preset="old",
provider_snapshot_loader=lambda: load_preset(default_name),
preset_snapshot_loader=load_preset,
)
catalog["new"] = ModelPresetConfig(model="new-model", provider="openai")
default_name = "new"
resolver.invalidate()
assert resolver.admit().model_preset == "new"
assert set(resolver.model_presets) == {"old", "new"}
del catalog["old"]
resolver.invalidate()
assert resolver.admit().model_preset == "new"
assert set(resolver.model_presets) == {"new"}
def test_resolver_model_presets_are_read_only() -> None:
resolver = ModelRuntimeResolver(
_runtime(),
model_presets={"fast": ModelPresetConfig(model="fast-model")},
)
exposed = resolver.model_presets
with pytest.raises(TypeError):
exposed["other"] = ModelPresetConfig( # type: ignore[index]
model="other-model"
)
exposed["fast"].model = "mutated-model"
assert set(resolver.model_presets) == {"fast"}
assert resolver.model_presets["fast"].model == "fast-model"
assert resolver.resolve_preset("fast").model == "fast-model"
def test_resolver_model_override_is_derived_without_default_mutation() -> None:
initial = _runtime()
resolver = ModelRuntimeResolver(initial)
+99
View File
@@ -1,3 +1,4 @@
import asyncio
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -6,10 +7,15 @@ import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.bus.queue import MessageBus
from nanobot.bus.runtime_events import RuntimeModelChanged
from nanobot.config.loader import save_config
from nanobot.config.schema import Config, ModelPresetConfig
from nanobot.providers.base import GenerationSettings
from nanobot.providers.factory import ProviderSnapshot, load_provider_snapshot
from nanobot.session.model_selection import (
SESSION_MODEL_PRESET_METADATA_KEY,
model_preset_from_metadata,
)
from nanobot.webui.settings_api import update_agent_settings
@@ -153,6 +159,99 @@ def test_same_snapshot_default_clears_preset_and_publishes_update(tmp_path: Path
assert published == [("fast-model", None)]
def test_named_default_refresh_is_used_by_sessions_without_override(tmp_path: Path) -> None:
provider = _provider("shared-model")
shared_signature = ("shared-model", "auto", "same-settings")
snapshots = {
"fast": ProviderSnapshot(
provider=provider,
model="shared-model",
context_window_tokens=16_000,
signature=shared_signature,
model_preset="fast",
),
"deep": ProviderSnapshot(
provider=provider,
model="shared-model",
context_window_tokens=16_000,
signature=shared_signature,
model_preset="deep",
),
}
default_snapshot = snapshots["deep"]
loop = AgentLoop(
bus=MessageBus(),
provider=provider,
workspace=tmp_path,
model="shared-model",
context_window_tokens=16_000,
provider_signature=shared_signature,
provider_snapshot_loader=lambda: default_snapshot,
model_presets={
name: ModelPresetConfig(model=snapshot.model)
for name, snapshot in snapshots.items()
},
model_preset="fast",
preset_snapshot_loader=snapshots.__getitem__,
)
loop.runtime_resolver.invalidate()
runtime = loop.llm_runtime()
session = loop.sessions.get_or_create("sdk:new-after-refresh")
assert runtime.model_preset == "deep"
assert loop.runtime_for_session(session).model_preset == "deep"
assert model_preset_from_metadata(session.metadata) is None
@pytest.mark.asyncio
async def test_config_invalidation_notifies_clients_before_session_runtime_refresh(
tmp_path: Path,
) -> None:
provider = _provider("model-a")
catalog = {"fast": ModelPresetConfig(model="model-a")}
current_model = "model-a"
published: list[RuntimeModelChanged] = []
def load_preset(_name: str) -> ProviderSnapshot:
return ProviderSnapshot(
provider=provider,
model=current_model,
context_window_tokens=16_000,
signature=(current_model, "auto"),
model_preset="fast",
)
loop = AgentLoop(
bus=MessageBus(),
provider=provider,
workspace=tmp_path,
model="model-a",
context_window_tokens=16_000,
provider_signature=("model-a", "auto"),
provider_snapshot_loader=lambda: load_preset("fast"),
model_presets=catalog,
preset_catalog_loader=lambda: catalog,
model_preset="fast",
preset_snapshot_loader=load_preset,
)
loop.runtime_events.subscribe(published.append, RuntimeModelChanged)
session = loop.sessions.get_or_create("websocket:chat")
session.metadata[SESSION_MODEL_PRESET_METADATA_KEY] = "fast"
current_model = "model-b"
catalog["fast"] = ModelPresetConfig(model="model-b")
loop.invalidate_runtime_config()
runtime = loop.runtime_for_session(session)
await asyncio.sleep(0)
assert [(event.model, event.model_preset) for event in published] == [
("model-a", "fast"),
]
assert runtime.model == "model-b"
assert loop.model_presets["fast"].model == "model-b"
def test_next_turn_captures_generation_changed_after_previous_admission(
tmp_path: Path,
) -> None:
+74
View File
@@ -4,10 +4,12 @@ from unittest.mock import MagicMock
import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.context import RequestContext, request_context
from nanobot.agent.tools.self import MyTool
from nanobot.bus.queue import MessageBus
from nanobot.config.schema import ModelPresetConfig
from nanobot.providers.factory import ProviderSnapshot
from nanobot.session.model_selection import model_preset_from_metadata
def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
@@ -287,6 +289,78 @@ def test_self_tool_set_model_preset_unknown_lists_available(tmp_path) -> None:
assert loop.model == "base-model"
def test_self_tool_sets_model_preset_for_current_session(tmp_path) -> None:
presets = {
"default": ModelPresetConfig(model="base-model"),
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
}
loop = _make_loop(tmp_path, presets=presets)
tool = MyTool(runtime_state=loop, modify_allowed=True)
with request_context(RequestContext(
channel="cli",
chat_id="one",
session_key="cli:one",
metadata={"source": "self-tool"},
)):
result = tool._modify("model_preset", "fast")
assert "for the next turn" in result
assert model_preset_from_metadata(
loop.sessions.get_or_create("cli:one").metadata
) == "fast"
assert loop.model_preset is None
assert loop.model == "base-model"
def test_self_tool_reports_session_preset_provider_configuration_error(tmp_path) -> None:
loop = _make_loop(tmp_path)
loop.set_session_model_preset = MagicMock(
side_effect=ValueError("No API key configured for provider 'openai'.")
)
tool = MyTool(runtime_state=loop, modify_allowed=True)
with request_context(RequestContext(
channel="cli",
chat_id="one",
session_key="cli:one",
)):
result = tool._modify("model_preset", "broken")
assert result == "Error: No API key configured for provider 'openai'."
@pytest.mark.parametrize(
("key", "value"),
[
("model", "other-model"),
("context_window_tokens", 8_192),
],
)
def test_self_tool_rejects_instance_runtime_changes_in_session(
tmp_path,
key: str,
value: object,
) -> None:
loop = _make_loop(tmp_path)
tool = MyTool(runtime_state=loop, modify_allowed=True)
session = loop.sessions.get_or_create("cli:one")
with request_context(RequestContext(
channel="cli",
chat_id="one",
session_key=session.key,
runtime=loop.runtime_for_session(session),
)):
result = tool._modify(key, value)
other_runtime = loop.runtime_for_session(loop.sessions.get_or_create("cli:two"))
assert "instance-wide and disabled" in result
assert "model_preset" in result
assert other_runtime.model == "base-model"
assert other_runtime.context_window_tokens == 1000
def test_self_tool_set_model_clears_active_preset(tmp_path) -> None:
presets = {
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
+228
View File
@@ -0,0 +1,228 @@
import asyncio
import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.bus.queue import MessageBus
from nanobot.config.schema import ModelPresetConfig
from nanobot.nanobot import Nanobot
from nanobot.providers.base import GenerationSettings, LLMProvider, LLMResponse
from nanobot.providers.factory import ProviderSnapshot
from nanobot.sdk.types import SessionSnapshot
from nanobot.session.model_selection import (
SESSION_MODEL_PRESET_METADATA_KEY,
model_preset_from_metadata,
)
from nanobot.utils.llm_runtime import LLMRuntime
class RecordingProvider(LLMProvider):
def __init__(self, name: str) -> None:
super().__init__()
self.name = name
self.generation = GenerationSettings(max_tokens=256, temperature=0.1)
self.calls: list[str | None] = []
async def chat(self, messages, tools=None, model=None, **kwargs):
await asyncio.sleep(0)
self.calls.append(model)
return LLMResponse(content=f"reply from {self.name}", finish_reason="stop")
def get_default_model(self) -> str:
return self.name
@pytest.mark.asyncio
async def test_sessions_run_concurrently_with_isolated_model_presets(tmp_path) -> None:
base = RecordingProvider("base-model")
fast = RecordingProvider("fast-model")
deep = RecordingProvider("deep-model")
providers = {"fast": fast, "deep": deep}
load_counts = {"fast": 0, "deep": 0}
presets = {
"default": ModelPresetConfig(model="base-model", context_window_tokens=8_000),
"fast": ModelPresetConfig(model="fast-model", context_window_tokens=16_000),
"deep": ModelPresetConfig(model="deep-model", context_window_tokens=32_000),
}
def load_preset(name: str) -> ProviderSnapshot:
load_counts[name] += 1
preset = presets[name]
provider = base if name == "default" else providers[name]
return ProviderSnapshot(
provider=provider,
model=preset.model,
context_window_tokens=preset.context_window_tokens,
signature=(name, preset.model),
)
loop = AgentLoop(
bus=MessageBus(),
provider=base,
workspace=tmp_path,
model="base-model",
context_window_tokens=8_000,
model_presets=presets,
preset_snapshot_loader=load_preset,
)
loop._schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
loop.set_session_model_preset("sdk:fast", "fast")
loop.set_session_model_preset("sdk:deep", "deep")
fast_reply, deep_reply = await asyncio.gather(
loop.process_direct("hello", session_key="sdk:fast"),
loop.process_direct("hello", session_key="sdk:deep"),
)
assert fast_reply is not None and fast_reply.content == "reply from fast-model"
assert deep_reply is not None and deep_reply.content == "reply from deep-model"
assert fast.calls == ["fast-model"]
assert deep.calls == ["deep-model"]
assert base.calls == []
assert loop.provider is base
assert loop.model == "base-model"
assert load_counts == {"fast": 1, "deep": 1}
loop.sessions.invalidate("sdk:fast")
restored = loop.sessions.get_or_create("sdk:fast")
assert model_preset_from_metadata(restored.metadata) == "fast"
override = RecordingProvider("override-model")
override_runtime = LLMRuntime.capture(
override,
"override-model",
context_window_tokens=24_000,
)
override_reply = await loop.process_direct(
"hello",
session_key="sdk:fast",
runtime=override_runtime,
)
assert override_reply is not None
assert override_reply.content == "reply from override-model"
assert override.calls == ["override-model"]
assert fast.calls == ["fast-model"]
assert load_counts == {"fast": 1, "deep": 1}
@pytest.mark.asyncio
async def test_streamed_sdk_resolves_session_runtime_after_lock_admission(tmp_path) -> None:
base = RecordingProvider("base-model")
fast = RecordingProvider("fast-model")
deep = RecordingProvider("deep-model")
providers = {"fast": fast, "deep": deep}
presets = {
"fast": ModelPresetConfig(model="fast-model", context_window_tokens=16_000),
"deep": ModelPresetConfig(model="deep-model", context_window_tokens=32_000),
}
def load_preset(name: str) -> ProviderSnapshot:
preset = presets[name]
return ProviderSnapshot(
provider=providers[name],
model=preset.model,
context_window_tokens=preset.context_window_tokens,
signature=(name, preset.model),
)
loop = AgentLoop(
bus=MessageBus(),
provider=base,
workspace=tmp_path,
model="base-model",
context_window_tokens=8_000,
model_presets=presets,
preset_snapshot_loader=load_preset,
)
loop._schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
session_key = "sdk:queued"
loop.set_session_model_preset(session_key, "fast")
lock = loop._session_locks.setdefault(session_key, asyncio.Lock())
await lock.acquire()
try:
run = await Nanobot(loop).run_streamed("hello", session_key=session_key)
loop.set_session_model_preset(session_key, "deep")
finally:
lock.release()
events = [event async for event in run.stream_events()]
result = await run.wait()
assert result.content == "reply from deep-model"
assert fast.calls == []
assert deep.calls == ["deep-model"]
assert events[0].type == "run.started"
assert events[0].metadata["model"] == "deep-model"
assert events[0].metadata["model_preset"] == "deep"
@pytest.mark.parametrize("custom_value", ["legacy-tag", 7])
@pytest.mark.asyncio
async def test_sdk_custom_model_preset_metadata_does_not_select_runtime(
tmp_path,
custom_value,
) -> None:
base = RecordingProvider("base-model")
loop = AgentLoop(
bus=MessageBus(),
provider=base,
workspace=tmp_path,
model="base-model",
context_window_tokens=8_000,
)
loop._schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
bot = Nanobot(loop)
await bot.sessions.ingest(
"sdk:custom-metadata",
[],
metadata={"model_preset": custom_value},
)
ingested_result = await bot.run("hello", session_key="sdk:custom-metadata")
exported = bot.sessions.export("sdk:custom-metadata")
restored = await bot.sessions.restore(
SessionSnapshot(
key="sdk:restored-metadata",
messages=[],
metadata={"model_preset": custom_value},
)
)
restored_result = await bot.run("hello", session_key=restored.key)
assert ingested_result.content == "reply from base-model"
assert restored_result.content == "reply from base-model"
assert base.calls == ["base-model", "base-model"]
assert exported is not None
assert exported.metadata["model_preset"] == custom_value
assert restored.metadata["model_preset"] == custom_value
@pytest.mark.parametrize("invalid_value", [{"invalid": True}, " "])
@pytest.mark.asyncio
async def test_sdk_invalid_internal_model_preset_metadata_fails_explicitly(
tmp_path,
invalid_value,
) -> None:
base = RecordingProvider("base-model")
loop = AgentLoop(
bus=MessageBus(),
provider=base,
workspace=tmp_path,
model="base-model",
context_window_tokens=8_000,
)
loop._schedule_background = lambda coro: coro.close() # type: ignore[method-assign]
bot = Nanobot(loop)
await bot.sessions.ingest(
"sdk:invalid-internal-metadata",
[],
metadata={SESSION_MODEL_PRESET_METADATA_KEY: invalid_value},
)
with pytest.raises(ValueError, match="session model preset must be a non-empty string"):
await bot.run("hello", session_key="sdk:invalid-internal-metadata")
assert base.calls == []
+1 -1
View File
@@ -302,7 +302,7 @@ class TestCmdNewUnifiedSession:
sessions=sessions,
consolidator=SimpleNamespace(archive=AsyncMock(return_value=True)),
_cancel_active_tasks=AsyncMock(return_value=0),
llm_runtime=MagicMock(return_value=MagicMock()),
runtime_for_session=MagicMock(return_value=MagicMock()),
)
loop._schedule_background = lambda coro: asyncio.ensure_future(coro)
+28
View File
@@ -4,6 +4,7 @@ from __future__ import annotations
import time
from pathlib import Path
from types import MappingProxyType
from unittest.mock import MagicMock, patch
import pytest
@@ -11,6 +12,7 @@ from pydantic import BaseModel
from nanobot.agent.tools.context import RequestContext, request_context
from nanobot.agent.tools.self import MyTool
from nanobot.config.schema import ModelPresetConfig
# ---------------------------------------------------------------------------
# Helpers
@@ -1078,6 +1080,32 @@ class TestSecurityAttributeProtection:
result = await tool.execute(action="set", key="web_config.enable", value=False)
assert "read-only" in result
@pytest.mark.asyncio
async def test_modify_model_presets_dotpath_blocked(self):
"""The config-derived model preset catalog is inspectable but not mutable."""
presets = {"fast": {"model": "fast-model"}}
tool = _make_tool(runtime_state=_make_mock_loop(model_presets=presets))
result = await tool.execute(
action="set",
key="model_presets.other",
value={"model": "other-model"},
)
assert "read-only" in result
assert presets == {"fast": {"model": "fast-model"}}
@pytest.mark.asyncio
async def test_inspect_read_only_model_preset_dotpath(self):
presets = MappingProxyType({
"fast": ModelPresetConfig(model="fast-model"),
})
tool = _make_tool(runtime_state=_make_mock_loop(model_presets=presets))
result = await tool.execute(action="check", key="model_presets.fast.model")
assert result == "model_presets.fast.model: 'fast-model'"
# ---------------------------------------------------------------------------
# current iteration count (Fix #2)
+77 -10
View File
@@ -16,6 +16,10 @@ from nanobot.command.builtin import (
)
from nanobot.command.router import CommandContext, CommandRouter
from nanobot.config.schema import ModelPresetConfig
from nanobot.session.model_selection import (
SESSION_MODEL_PRESET_METADATA_KEY,
model_preset_from_metadata,
)
def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
@@ -29,7 +33,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
return provider
def _make_loop(tmp_path) -> AgentLoop:
def _make_loop(tmp_path, *, preset_snapshot_loader=None) -> AgentLoop:
return AgentLoop(
bus=MessageBus(),
provider=_provider("base-model", max_tokens=123),
@@ -48,6 +52,7 @@ def _make_loop(tmp_path) -> AgentLoop:
context_window_tokens=32_768,
),
},
preset_snapshot_loader=preset_snapshot_loader,
)
@@ -64,6 +69,11 @@ def _ctx_session(loop: AgentLoop, raw: str, args: str = "") -> CommandContext:
)
def _saved_model_preset(loop: AgentLoop, session_key: str = "cli:direct") -> str | None:
session = loop.sessions.get_or_create(session_key)
return model_preset_from_metadata(session.metadata)
@pytest.mark.asyncio
async def test_model_command_lists_current_and_available_presets(tmp_path) -> None:
loop = _make_loop(tmp_path)
@@ -84,23 +94,28 @@ async def test_model_command_switches_preset(tmp_path) -> None:
out = await cmd_model(_ctx(loop, "/model fast", args="fast"))
assert "Switched model preset to `fast`." in out.content
assert "Scope: current session" in out.content
assert "Model: `openai/gpt-4.1`" in out.content
assert loop.model_preset == "fast"
assert loop.model == "openai/gpt-4.1"
assert not hasattr(loop.subagents, "model")
assert not hasattr(loop.consolidator, "model")
assert loop.llm_runtime().model == "openai/gpt-4.1"
assert _saved_model_preset(loop) == "fast"
assert loop.model_preset is None
assert loop.model == "base-model"
await loop.process_direct("/new", session_key="cli:direct")
assert _saved_model_preset(loop) == "fast"
status = await loop.process_direct("/status", session_key="cli:direct")
assert status is not None and "openai/gpt-4.1" in status.content
@pytest.mark.asyncio
async def test_model_command_switches_back_to_default(tmp_path) -> None:
loop = _make_loop(tmp_path)
loop.set_model_preset("fast")
await cmd_model(_ctx(loop, "/model fast", args="fast"))
out = await cmd_model(_ctx(loop, "/model default", args="default"))
assert "Switched model preset to `default`." in out.content
assert loop.model_preset == "default"
assert _saved_model_preset(loop) == "default"
assert loop.model_preset is None
assert loop.model == "base-model"
assert loop.context_window_tokens == 1000
@@ -118,6 +133,24 @@ async def test_model_command_unknown_preset_keeps_old_state(tmp_path) -> None:
assert loop.model == "base-model"
@pytest.mark.asyncio
async def test_model_command_reports_provider_configuration_errors(tmp_path) -> None:
def fail_preset(_name: str):
raise ValueError("No API key configured for provider 'openai'.")
loop = _make_loop(tmp_path, preset_snapshot_loader=fail_preset)
switched = await cmd_model(_ctx(loop, "/model fast", args="fast"))
session = loop.sessions.get_or_create("cli:direct")
session.metadata[SESSION_MODEL_PRESET_METADATA_KEY] = "fast"
status = await cmd_model(_ctx(loop, "/model"))
assert "Could not switch model preset" in switched.content
assert "No API key configured for provider 'openai'." in switched.content
assert "Current selection error" in status.content
assert "No API key configured for provider 'openai'." in status.content
@pytest.mark.asyncio
async def test_model_command_does_not_depend_on_my_allow_set(tmp_path) -> None:
loop = _make_loop(tmp_path)
@@ -125,7 +158,7 @@ async def test_model_command_does_not_depend_on_my_allow_set(tmp_path) -> None:
await cmd_model(_ctx(loop, "/model fast", args="fast"))
assert loop.model_preset == "fast"
assert _saved_model_preset(loop) == "fast"
@pytest.mark.asyncio
@@ -142,11 +175,45 @@ async def test_model_command_registered_as_exact_and_prefix(tmp_path) -> None:
assert out.metadata == {"render_as": "text"}
assert out.content == "\n".join([
"Switched model preset to `fast`.",
"- Scope: current session",
"- Model: `openai/gpt-4.1`",
"- Context window: 32768",
"- Max output tokens: 4096",
])
assert loop.model_preset == "fast"
assert _saved_model_preset(loop) == "fast"
@pytest.mark.asyncio
async def test_model_command_does_not_change_another_session(tmp_path) -> None:
loop = _make_loop(tmp_path)
await cmd_model(_ctx(loop, "/model fast", args="fast"))
other = InboundMessage(channel="cli", sender_id="user", chat_id="other", content="/model")
out = await cmd_model(
CommandContext(msg=other, session=None, key=other.session_key, raw="/model", loop=loop)
)
assert "Current preset: `default`" in out.content
assert _saved_model_preset(loop) == "fast"
@pytest.mark.asyncio
async def test_model_command_reports_and_recovers_removed_session_preset(tmp_path) -> None:
loop = _make_loop(tmp_path)
session = loop.sessions.get_or_create("cli:direct")
session.metadata[SESSION_MODEL_PRESET_METADATA_KEY] = "removed"
loop.sessions.save(session)
status = await loop.process_direct("/model", session_key="cli:direct")
switched = await loop.process_direct("/model default", session_key="cli:direct")
assert status is not None
assert "model_preset 'removed' not found" in status.content
assert "Available presets: `default`, `fast`" in status.content
assert "Switch with `/model <preset>`" in status.content
assert switched is not None
assert "Switched model preset to `default`." in switched.content
assert _saved_model_preset(loop) == "default"
def test_model_command_in_help_and_palette() -> None:
+32 -9
View File
@@ -62,6 +62,13 @@ def _fake_provider(name: str, *, max_tokens: int = 8192) -> MagicMock:
return provider
async def _admit_fake_runtime(bot, session_key, callback, runtime=None) -> None:
admitted = runtime or bot._loop.runtime_for_session(
bot._loop.sessions.get_or_create(session_key)
)
await callback(admitted)
def test_from_config_missing_file():
with pytest.raises(FileNotFoundError):
Nanobot.from_config("/nonexistent/config.json")
@@ -633,8 +640,9 @@ async def test_run_model_preset_override_is_per_run(tmp_path):
model="openai/gpt-4.1-mini",
context_window_tokens=2048,
signature=("preset", "fast"),
model_preset="fast",
)
override_runtime = runtime_from_provider_snapshot(override, model_preset="fast")
override_runtime = runtime_from_provider_snapshot(override)
bot._loop.runtime_resolver.resolve_override = MagicMock(
return_value=override_runtime
)
@@ -775,10 +783,12 @@ async def test_stream_yields_text_events_in_order(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, runtime
message, *, session_key, on_stream, on_stream_end, hooks, on_runtime_admitted,
runtime=None
):
assert message == "hi"
assert session_key == "sdk:default"
await _admit_fake_runtime(bot, session_key, on_runtime_admitted, runtime)
await on_stream("Hel")
await on_stream("lo")
await on_stream_end(resuming=False)
@@ -812,8 +822,10 @@ async def test_run_streamed_wait_returns_full_result_without_consuming_events(tm
bot = Nanobot.from_config(config_path, workspace=tmp_path)
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
message, *, session_key, on_stream, on_stream_end, hooks, on_runtime_admitted,
runtime=None
):
await _admit_fake_runtime(bot, session_key, on_runtime_admitted, runtime)
await on_stream("done")
await on_stream_end(resuming=False)
ctx = AgentRunHookContext(
@@ -856,8 +868,10 @@ async def test_run_streamed_cancel_releases_full_queue_without_consuming(tmp_pat
bot = Nanobot.from_config(config_path, workspace=tmp_path)
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
message, *, session_key, on_stream, on_stream_end, hooks, on_runtime_admitted,
runtime=None
):
await _admit_fake_runtime(bot, session_key, on_runtime_admitted, runtime)
for i in range(400):
await on_stream(str(i))
await on_stream_end(resuming=False)
@@ -921,7 +935,8 @@ 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
assert "runtime" not in kwargs
assert callable(kwargs["on_runtime_admitted"])
@pytest.mark.asyncio
@@ -951,10 +966,12 @@ async def test_run_streamed_model_override_reports_admitted_runtime(tmp_path):
on_stream,
on_stream_end,
hooks,
on_runtime_admitted,
runtime,
):
assert runtime is override_runtime
assert bot._loop.runtime_resolver.runtime is original_runtime
await on_runtime_admitted(runtime)
await on_stream("ok")
await on_stream_end(resuming=False)
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
@@ -995,8 +1012,10 @@ async def test_run_streamed_emits_tool_events(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, runtime
message, *, session_key, on_stream, on_stream_end, hooks, on_runtime_admitted,
runtime=None
):
await _admit_fake_runtime(bot, session_key, on_runtime_admitted, runtime)
calls = [
ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}),
ToolCallRequest(id="call_bad", name="exec", arguments={"cmd": "false"}),
@@ -1043,8 +1062,10 @@ async def test_run_streamed_emits_reasoning_events(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, runtime
message, *, session_key, on_stream, on_stream_end, hooks, on_runtime_admitted,
runtime=None
):
await _admit_fake_runtime(bot, session_key, on_runtime_admitted, runtime)
for hook in hooks:
await hook.emit_reasoning("thinking")
await hook.emit_reasoning_end()
@@ -1072,8 +1093,10 @@ async def test_stream_generator_break_cancels_underlying_run(tmp_path):
cancelled = asyncio.Event()
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
message, *, session_key, on_stream, on_stream_end, hooks, on_runtime_admitted,
runtime=None
):
await _admit_fake_runtime(bot, session_key, on_runtime_admitted, runtime)
try:
await on_stream("first")
await asyncio.sleep(10)
@@ -1387,7 +1410,7 @@ async def test_runtime_helpers_expose_model_workspace_and_compact(tmp_path):
bot = Nanobot.from_config(config_path, workspace=tmp_path)
await bot.sessions.ingest("sdk:history", [{"role": "user", "content": "hello"}])
runtime = bot._loop.llm_runtime()
bot._loop.llm_runtime = MagicMock(return_value=runtime) # type: ignore[method-assign]
bot._loop.runtime_for_session = MagicMock(return_value=runtime) # type: ignore[method-assign]
bot._loop.consolidator.maybe_consolidate_by_tokens = AsyncMock()
snapshot = await bot.runtime.compact_session("sdk:history")
+22
View File
@@ -4,11 +4,14 @@ import os
from datetime import datetime
from pathlib import Path
import pytest
import nanobot.webui.session_list_index as session_list_index
from nanobot.cron.session_turns import CRON_HISTORY_META
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
from nanobot.session.manager import SessionManager
from nanobot.session.model_selection import SESSION_MODEL_PRESET_METADATA_KEY
def test_webui_session_list_reuses_valid_index_without_scanning_files(
@@ -17,10 +20,12 @@ def test_webui_session_list_reuses_valid_index_without_scanning_files(
) -> None:
manager = SessionManager(tmp_path)
session = manager.get_or_create("websocket:indexed")
session.metadata[SESSION_MODEL_PRESET_METADATA_KEY] = "fast"
session.add_message("user", "indexed preview")
manager.save(session)
assert list_webui_sessions(manager)[0]["preview"] == "indexed preview"
assert list_webui_sessions(manager)[0]["model_preset"] == "fast"
def fail_scan(session_manager: SessionManager, path: Path) -> None:
raise AssertionError(f"unexpected session file scan: {path}")
@@ -31,6 +36,23 @@ def test_webui_session_list_reuses_valid_index_without_scanning_files(
assert rows[0]["key"] == "websocket:indexed"
assert rows[0]["preview"] == "indexed preview"
assert rows[0]["model_preset"] == "fast"
def test_webui_session_list_rejects_invalid_internal_model_preset_metadata(
tmp_path: Path,
) -> None:
manager = SessionManager(tmp_path)
session = manager.get_or_create("websocket:custom-metadata")
session.metadata["model_preset"] = 7
session.metadata[SESSION_MODEL_PRESET_METADATA_KEY] = {"invalid": True}
session.add_message("user", "custom metadata")
manager.save(session)
with pytest.raises(ValueError, match="session model preset must be a non-empty string"):
list_webui_sessions(manager)
assert manager.get_or_create(session.key).metadata["model_preset"] == 7
def test_webui_session_list_rescans_only_changed_file(tmp_path: Path, monkeypatch) -> None:
+20
View File
@@ -464,6 +464,26 @@ def test_settings_payload_includes_dynamic_custom_provider(
assert providers[DYNAMIC_PROVIDER_NAME]["api_base"] == DYNAMIC_PROVIDER_API_BASE
def test_settings_payload_resolves_provider_for_each_auto_preset(
tmp_path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
config_path = tmp_path / "config.json"
config = _dynamic_provider_config()
config.model_presets["fast"] = ModelPresetConfig(
provider="auto",
model=f"{DYNAMIC_PROVIDER_NAME}/gpt-4",
)
save_config(config, config_path)
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
payload = settings_payload()
presets = {row["name"]: row for row in payload["model_presets"]}
assert presets["fast"]["provider"] == "auto"
assert presets["fast"]["resolved_provider"] == DYNAMIC_PROVIDER_NAME
def test_settings_payload_groups_opencode_compatibility_alias(tmp_path, monkeypatch) -> None:
config_path = tmp_path / "config.json"
save_config(Config(), config_path)