refactor(agent): make resolver sole runtime owner
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user