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)