feat(agent): make model presets session-scoped (#4866)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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 == []
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user