refactor(agent): make resolver sole runtime owner

This commit is contained in:
chengyongru
2026-07-10 17:54:34 +08:00
committed by Xubin Ren
parent c9d3e74342
commit 21f58cbabf
16 changed files with 395 additions and 443 deletions
+16 -8
View File
@@ -518,9 +518,11 @@ class TestNewCommandArchival:
loop.sessions.save(session)
call_count = 0
expected_runtime = loop.llm_runtime()
async def _failing_summarize(_messages, *, session_key=None) -> bool:
async def _failing_summarize(_messages, *, runtime, session_key=None) -> bool:
nonlocal call_count
assert runtime is expected_runtime
assert session_key == "cli:test"
call_count += 1
return False
@@ -528,7 +530,7 @@ class TestNewCommandArchival:
loop.consolidator.archive = _failing_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
response = await loop._process_message(new_msg)
response = await loop._process_message(new_msg, runtime=expected_runtime)
assert response is not None
assert "new session started" in response.content.lower()
@@ -553,9 +555,11 @@ class TestNewCommandArchival:
archived_count = -1
archived_session_key = None
expected_runtime = loop.llm_runtime()
async def _fake_summarize(messages, *, session_key=None) -> bool:
async def _fake_summarize(messages, *, runtime, session_key=None) -> bool:
nonlocal archived_count, archived_session_key
assert runtime is expected_runtime
archived_count = len(messages)
archived_session_key = session_key
return True
@@ -563,7 +567,7 @@ class TestNewCommandArchival:
loop.consolidator.archive = _fake_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
response = await loop._process_message(new_msg)
response = await loop._process_message(new_msg, runtime=expected_runtime)
assert response is not None
assert "new session started" in response.content.lower()
@@ -582,15 +586,17 @@ class TestNewCommandArchival:
session.add_message("user", f"msg{i}")
session.add_message("assistant", f"resp{i}")
loop.sessions.save(session)
expected_runtime = loop.llm_runtime()
async def _ok_summarize(_messages, *, session_key=None) -> bool:
async def _ok_summarize(_messages, *, runtime, session_key=None) -> bool:
assert runtime is expected_runtime
assert session_key == "cli:test"
return True
loop.consolidator.archive = _ok_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
response = await loop._process_message(new_msg)
response = await loop._process_message(new_msg, runtime=expected_runtime)
assert response is not None
assert "new session started" in response.content.lower()
@@ -610,8 +616,10 @@ class TestNewCommandArchival:
archived = asyncio.Event()
release_archive = asyncio.Event()
expected_runtime = loop.llm_runtime()
async def _slow_summarize(_messages, *, session_key=None) -> bool:
async def _slow_summarize(_messages, *, runtime, session_key=None) -> bool:
assert runtime is expected_runtime
assert session_key == "cli:test"
await release_archive.wait()
archived.set()
@@ -620,7 +628,7 @@ class TestNewCommandArchival:
loop.consolidator.archive = _slow_summarize # type: ignore[method-assign]
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
await loop._process_message(new_msg)
await loop._process_message(new_msg, runtime=expected_runtime)
assert not archived.is_set()
release_archive.set()
@@ -6,6 +6,7 @@ import nanobot.agent.memory as memory_module
from nanobot.agent.loop import AgentLoop
from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMResponse
from nanobot.session.manager import replay_max_messages_for_context
def _make_loop(tmp_path, *, estimated_tokens: int, context_window_tokens: int) -> AgentLoop:
@@ -222,7 +223,7 @@ async def test_preflight_consolidation_receives_pending_summary(tmp_path) -> Non
loop.consolidator.maybe_consolidate_by_tokens.assert_any_await(
session,
runtime=runtime,
replay_max_messages=loop._max_messages,
replay_max_messages=replay_max_messages_for_context(runtime.context_window_tokens),
)
assert len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2
assert all(
+9 -2
View File
@@ -21,6 +21,7 @@ from nanobot.bus.outbound_events import (
)
from nanobot.bus.queue import MessageBus
from nanobot.providers.base import LLMResponse, ToolCallRequest
from nanobot.providers.factory import ProviderSnapshot
from nanobot.session.webui_turns import WebuiTurnCoordinator
from nanobot.utils.progress_events import (
invoke_file_edit_progress,
@@ -815,8 +816,14 @@ class TestToolEventProgress:
))
assert len(scheduled_title) == 1
loop.provider = MagicMock()
loop.model = "switched-after-turn"
next_provider = MagicMock()
next_provider.generation = loop.llm_runtime().generation
loop.runtime_resolver.adopt_snapshot(ProviderSnapshot(
provider=next_provider,
model="switched-after-turn",
context_window_tokens=loop.context_window_tokens,
signature=("switched-after-turn",),
))
await scheduled_title[0] # type: ignore[misc]
+12 -2
View File
@@ -19,6 +19,7 @@ from nanobot.bus.outbound_events import (
from nanobot.bus.queue import MessageBus
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
from nanobot.providers.base import LLMResponse
from nanobot.providers.factory import ProviderSnapshot
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
from nanobot.session.goal_state import GOAL_STATE_KEY
from nanobot.session.manager import Session, SessionManager
@@ -69,8 +70,17 @@ def test_agent_loop_llm_runtime_reflects_current_provider_and_model(tmp_path: Pa
assert runtime.model == "test-model"
next_provider = MagicMock()
loop.provider = next_provider
loop.model = "next-model"
next_provider.generation = SimpleNamespace(
temperature=0.1,
max_tokens=4096,
reasoning_effort=None,
)
loop.runtime_resolver.adopt_snapshot(ProviderSnapshot(
provider=next_provider,
model="next-model",
context_window_tokens=runtime.context_window_tokens,
signature=("next-model",),
))
runtime = loop.llm_runtime()
assert runtime.provider is next_provider
+15 -11
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
from dataclasses import replace
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch
@@ -67,11 +68,13 @@ class TestMaxMessagesInit:
def test_default_for_200k_context_reaches_file_cap(self, tmp_path: Path) -> None:
loop = _make_loop(tmp_path)
assert loop._max_messages == FILE_MAX_MESSAGES
runtime = loop.runtime_resolver.runtime
assert replay_max_messages_for_context(runtime.context_window_tokens) == FILE_MAX_MESSAGES
def test_default_scales_with_context_window(self, tmp_path: Path) -> None:
loop = _make_loop(tmp_path, context_window_tokens=32_768)
assert loop._max_messages == 327
runtime = loop.runtime_resolver.runtime
assert replay_max_messages_for_context(runtime.context_window_tokens) == 327
def test_provider_refresh_resyncs_context_derived_limit(self, tmp_path: Path) -> None:
old_provider = MagicMock()
@@ -93,9 +96,10 @@ class TestMaxMessagesInit:
),
)
assert loop._max_messages == 327
loop._refresh_provider_snapshot()
assert loop._max_messages == FILE_MAX_MESSAGES
initial = loop.runtime_resolver.runtime
assert replay_max_messages_for_context(initial.context_window_tokens) == 327
refreshed = loop.llm_runtime()
assert replay_max_messages_for_context(refreshed.context_window_tokens) == FILE_MAX_MESSAGES
class TestGetHistoryWithMaxMessages:
@@ -136,7 +140,7 @@ class TestMaxMessagesIntegration:
async def test_process_message_passes_limit_to_history_call(self, tmp_path: Path) -> None:
"""The real message path should pass max_messages into session history replay."""
loop = _make_loop(tmp_path)
loop._max_messages = 25
runtime = replace(loop.llm_runtime(), context_window_tokens=32_768)
loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
)
@@ -146,12 +150,13 @@ class TestMaxMessagesIntegration:
session = loop.sessions.get_or_create("cli:test")
with patch.object(session, "get_history", wraps=session.get_history) as mock_hist:
result = await loop._process_message(
InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello")
InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello"),
runtime=runtime,
)
assert result is not None
assert mock_hist.call_count == 1
assert mock_hist.call_args.kwargs["max_messages"] == 25
assert mock_hist.call_args.kwargs["max_messages"] == 327
assert mock_hist.call_args.kwargs["extend_to_user"] is False
@pytest.mark.asyncio
@@ -182,8 +187,7 @@ class TestMaxMessagesIntegration:
tmp_path: Path,
) -> None:
"""A live user turn should not extend history to an older long tool turn."""
loop = _make_loop(tmp_path)
loop._max_messages = 6
loop = _make_loop(tmp_path, context_window_tokens=8_000)
loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
)
@@ -194,7 +198,7 @@ class TestMaxMessagesIntegration:
session.add_message("user", "old")
session.add_message("assistant", "old answer")
session.add_message("user", "long older turn")
for i in range(8):
for i in range(70):
session.messages.extend(_tool_round(f"older-{i}"))
session.add_message("assistant", "older final")
+27 -3
View File
@@ -18,7 +18,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
return provider
def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
def test_provider_refresh_updates_only_runtime_resolver(tmp_path: Path) -> None:
old_provider = _provider("old-model")
new_provider = _provider("new-model", max_tokens=456)
loop = AgentLoop(
@@ -35,8 +35,9 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
),
)
loop._refresh_provider_snapshot()
runtime = loop.llm_runtime()
assert runtime is loop.runtime_resolver.runtime
assert loop.provider is new_provider
assert loop.model == "new-model"
assert loop.context_window_tokens == 2000
@@ -50,6 +51,29 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
assert not hasattr(loop.consolidator, "max_completion_tokens")
def test_loop_has_no_mutable_runtime_mirrors_or_legacy_snapshot_api(tmp_path: Path) -> None:
loop = AgentLoop(
bus=MessageBus(),
provider=_provider("test-model"),
workspace=tmp_path,
model="test-model",
context_window_tokens=1000,
)
assert {
"provider",
"model",
"context_window_tokens",
"model_presets",
"_active_preset",
"_provider_signature",
"_max_messages",
}.isdisjoint(loop.__dict__)
assert not hasattr(loop, "_apply_provider_snapshot")
assert not hasattr(loop, "_build_model_preset_snapshot")
assert not hasattr(loop, "_sync_replay_max_messages")
def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
old_provider = _provider("old-model")
new_provider = _provider("new-model", max_tokens=456)
@@ -118,7 +142,7 @@ def test_settings_context_window_refreshes_runtime_state(
loop = AgentLoop.from_config(config, provider_snapshot_loader=loader)
payload = update_agent_settings({"context_window_tokens": ["262144"]})
loop._refresh_provider_snapshot()
loop.llm_runtime()
assert payload["requires_restart"] is False
assert loop.context_window_tokens == 262_144
+11 -8
View File
@@ -54,9 +54,10 @@ def test_model_preset_setter_updates_state(tmp_path) -> None:
assert loop.model_preset == "fast"
assert loop.model == "openai/gpt-4.1"
assert loop.context_window_tokens == 32_768
assert loop.provider.generation.temperature == 0.5
assert loop.provider.generation.max_tokens == 4096
assert loop.provider.generation.reasoning_effort == "low"
runtime = loop.llm_runtime()
assert runtime.generation.temperature == 0.5
assert runtime.generation.max_tokens == 4096
assert runtime.generation.reasoning_effort == "low"
assert not hasattr(loop.subagents, "model")
assert not hasattr(loop.consolidator, "model")
assert not hasattr(loop.consolidator, "context_window_tokens")
@@ -174,7 +175,7 @@ def test_active_model_preset_survives_unchanged_config_refresh(tmp_path) -> None
)
loop.set_model_preset("fast")
loop._refresh_provider_snapshot()
loop.llm_runtime()
assert loop.model_preset == "fast"
assert loop.provider is fast_provider
@@ -210,7 +211,7 @@ def test_config_model_refresh_clears_active_model_preset(tmp_path) -> None:
)
loop.set_model_preset("fast")
loop._refresh_provider_snapshot()
loop.llm_runtime()
assert loop.model_preset is None
assert loop.provider is webui_provider
@@ -292,7 +293,7 @@ def test_self_tool_set_model_clears_active_preset(tmp_path) -> None:
tool = MyTool(runtime_state=loop, modify_allowed=True)
result = tool._modify("model", "anthropic/claude-opus-4-5")
assert "Error" not in result
assert loop._active_preset is None
assert loop.model_preset is None
assert loop.model == "anthropic/claude-opus-4-5"
@@ -323,5 +324,7 @@ def test_from_config_static_preset_loader_does_not_enable_hot_reload(tmp_path) -
fake_provider = _provider("openai/gpt-4.1")
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
loop = AgentLoop.from_config(config)
assert loop._provider_snapshot_loader is None
assert loop._preset_snapshot_loader is not None
default_runtime = loop.runtime_resolver.runtime
resolved = loop.runtime_resolver.resolve_preset("fast")
assert resolved.model == "openai/gpt-4.1-mini"
assert loop.runtime_resolver.runtime is default_runtime
+8 -2
View File
@@ -35,6 +35,12 @@ def _make_mock_loop(**overrides):
loop._concurrency_gate = None
loop._unified_session = False
loop._extra_hooks = []
loop.set_runtime_model.side_effect = lambda value: setattr(loop, "model", value)
loop.set_runtime_context_window.side_effect = lambda value: setattr(
loop,
"context_window_tokens",
value,
)
# web_config mock — needed for check tests
loop.web_config = MagicMock()
@@ -237,12 +243,12 @@ class TestModifyRestricted:
@pytest.mark.asyncio
async def test_modify_context_window_valid(self):
loop = _make_mock_loop(_sync_replay_max_messages=MagicMock())
loop = _make_mock_loop()
tool = _make_tool(runtime_state=loop)
result = await tool.execute(action="set", key="context_window_tokens", value=131072)
assert "Set context_window_tokens" in result
assert loop.context_window_tokens == 131072
loop._sync_replay_max_messages.assert_called_once_with()
loop.set_runtime_context_window.assert_called_once_with(131072)
@pytest.mark.asyncio
async def test_modify_none_value_for_restricted_int(self):
+2 -2
View File
@@ -245,8 +245,8 @@ class TestRestartCommand:
msg = InboundMessage(channel="telegram", sender_id="u1", chat_id="c1", content="/status")
runtime = loop.llm_runtime()
loop.model = "replacement-model"
loop.context_window_tokens = 10
loop.set_runtime_model("replacement-model")
loop.set_runtime_context_window(10)
loop.provider.generation = SimpleNamespace(
temperature=1.0,
max_tokens=1,
+98 -51
View File
@@ -30,6 +30,7 @@ from nanobot.nanobot import (
StreamEvent,
StreamEventType,
)
from nanobot.utils.llm_runtime import runtime_from_provider_snapshot
def _write_config(tmp_path: Path, overrides: dict | None = None) -> Path:
@@ -504,45 +505,58 @@ async def test_run_allows_parallel_sessions_without_model_override(tmp_path):
@pytest.mark.asyncio
async def test_run_model_overrides_are_serialized_before_snapshot_build(tmp_path):
async def test_run_model_overrides_can_overlap_without_default_mutation(tmp_path):
from nanobot.bus.events import OutboundMessage
from nanobot.providers.factory import ProviderSnapshot
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
original_model = bot._loop.model
active_models: list[str] = []
snapshot_base_models: list[str] = []
assert not hasattr(bot, "_runtime_overrides")
original_runtime = bot._loop.runtime_resolver.runtime
active_runs: list[tuple[str, str]] = []
first_entered = asyncio.Event()
both_entered = asyncio.Event()
release_first = asyncio.Event()
def fake_snapshot(*, model, model_preset):
def fake_resolve(*, model, model_preset, config):
assert model is not None
assert model_preset is None
snapshot_base_models.append(bot._loop.model)
return ProviderSnapshot(
assert config is bot._config
return runtime_from_provider_snapshot(ProviderSnapshot(
provider=_fake_provider(model, max_tokens=2048),
model=model,
context_window_tokens=4096,
signature=("sdk", model),
)
))
bot._runtime_overrides.model_override_snapshot = MagicMock(side_effect=fake_snapshot)
bot._loop.runtime_resolver.resolve_override = MagicMock(side_effect=fake_resolve)
async def fake_process_direct(message, *, session_key, hooks):
active_models.append(bot._loop.model)
async def fake_process_direct(message, *, session_key, hooks, runtime):
active_runs.append((session_key, runtime.model))
assert bot._loop.runtime_resolver.runtime is original_runtime
if message == "first":
first_entered.set()
await asyncio.wait_for(release_first.wait(), timeout=1)
if len(active_runs) == 2:
both_entered.set()
await asyncio.wait_for(release_first.wait(), timeout=1)
return OutboundMessage(channel="cli", chat_id="direct", content=message)
bot._loop.process_direct = fake_process_direct
first = asyncio.create_task(bot.run("first", model="model:first"))
first = asyncio.create_task(bot.run(
"first",
session_key="sdk:first",
model="model:first",
))
await asyncio.wait_for(first_entered.wait(), timeout=1)
second = asyncio.create_task(bot.run("second", model="model:second"))
await asyncio.sleep(0)
second = asyncio.create_task(bot.run(
"second",
session_key="sdk:second",
model="model:second",
))
await asyncio.wait_for(both_entered.wait(), timeout=1)
assert not first.done()
assert not second.done()
release_first.set()
@@ -550,21 +564,21 @@ async def test_run_model_overrides_are_serialized_before_snapshot_build(tmp_path
assert first_result.content == "first"
assert second_result.content == "second"
assert active_models == ["model:first", "model:second"]
assert snapshot_base_models == [original_model, original_model]
assert bot._loop.model == original_model
assert set(active_runs) == {
("sdk:first", "model:first"),
("sdk:second", "model:second"),
}
assert bot._loop.runtime_resolver.runtime is original_runtime
@pytest.mark.asyncio
async def test_run_model_override_is_per_run_and_restores_default(tmp_path):
async def test_run_model_override_is_per_run_without_default_mutation(tmp_path):
from nanobot.bus.events import OutboundMessage
from nanobot.providers.factory import ProviderSnapshot
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
original_provider = bot._loop.provider
original_model = bot._loop.model
original_signature = bot._loop._provider_signature
original_runtime = bot._loop.runtime_resolver.runtime
override_provider = _fake_provider("override-provider", max_tokens=2048)
override = ProviderSnapshot(
provider=override_provider,
@@ -572,13 +586,17 @@ async def test_run_model_override_is_per_run_and_restores_default(tmp_path):
context_window_tokens=4096,
signature=("sdk", "override"),
)
bot._runtime_overrides.model_override_snapshot = MagicMock(return_value=override)
override_runtime = runtime_from_provider_snapshot(override)
bot._loop.runtime_resolver.resolve_override = MagicMock(
return_value=override_runtime
)
async def fake_process_direct(message, *, session_key, hooks):
assert bot._loop.provider is override_provider
async def fake_process_direct(message, *, session_key, hooks, runtime):
assert runtime is override_runtime
assert not hasattr(bot._loop.runner, "provider")
assert bot._loop.model == "openai/gpt-4.1-mini"
assert bot._loop.context_window_tokens == 4096
assert runtime.model == "openai/gpt-4.1-mini"
assert runtime.context_window_tokens == 4096
assert bot._loop.runtime_resolver.runtime is original_runtime
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
bot._loop.process_direct = fake_process_direct
@@ -586,14 +604,13 @@ async def test_run_model_override_is_per_run_and_restores_default(tmp_path):
result = await bot.run("hi", model="openai/gpt-4.1-mini")
assert result.content == "ok"
bot._runtime_overrides.model_override_snapshot.assert_called_once_with(
bot._loop.runtime_resolver.resolve_override.assert_called_once_with(
model="openai/gpt-4.1-mini",
model_preset=None,
config=bot._config,
)
assert bot._loop.provider is original_provider
assert not hasattr(bot._loop.runner, "provider")
assert bot._loop.model == original_model
assert bot._loop._provider_signature == original_signature
assert bot._loop.runtime_resolver.runtime is original_runtime
@pytest.mark.asyncio
@@ -603,7 +620,7 @@ async def test_run_model_preset_override_is_per_run(tmp_path):
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
original_model = bot._loop.model
original_runtime = bot._loop.runtime_resolver.runtime
override_provider = _fake_provider("preset-provider", max_tokens=1024)
override = ProviderSnapshot(
provider=override_provider,
@@ -611,19 +628,25 @@ async def test_run_model_preset_override_is_per_run(tmp_path):
context_window_tokens=2048,
signature=("preset", "fast"),
)
bot._loop._build_model_preset_snapshot = MagicMock(return_value=override)
override_runtime = runtime_from_provider_snapshot(override, model_preset="fast")
bot._loop.runtime_resolver.resolve_override = MagicMock(
return_value=override_runtime
)
async def fake_process_direct(message, *, session_key, hooks):
assert bot._loop.provider is override_provider
assert bot._loop.model == "openai/gpt-4.1-mini"
async def fake_process_direct(message, *, session_key, hooks, runtime):
assert runtime is override_runtime
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
bot._loop.process_direct = fake_process_direct
await bot.run("hi", model_preset="fast")
bot._loop._build_model_preset_snapshot.assert_called_once_with("fast")
assert bot._loop.model == original_model
bot._loop.runtime_resolver.resolve_override.assert_called_once_with(
model=None,
model_preset="fast",
config=bot._config,
)
assert bot._loop.runtime_resolver.runtime is original_runtime
assert bot._loop.model_preset is None
@@ -745,7 +768,9 @@ async def test_stream_yields_text_events_in_order(tmp_path):
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
):
assert message == "hi"
assert session_key == "sdk:default"
await on_stream("Hel")
@@ -780,7 +805,9 @@ async def test_run_streamed_wait_returns_full_result_without_consuming_events(tm
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
):
await on_stream("done")
await on_stream_end(resuming=False)
ctx = AgentRunHookContext(
@@ -822,7 +849,9 @@ async def test_run_streamed_cancel_releases_full_queue_without_consuming(tmp_pat
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
):
for i in range(400):
await on_stream(str(i))
await on_stream_end(resuming=False)
@@ -886,16 +915,17 @@ async def test_run_streamed_forwards_runtime_options(tmp_path):
assert callable(kwargs["on_stream"])
assert callable(kwargs["on_stream_end"])
assert kwargs["hooks"]
assert kwargs["runtime"] is bot._loop.runtime_resolver.runtime
@pytest.mark.asyncio
async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
async def test_run_streamed_model_override_reports_admitted_runtime(tmp_path):
from nanobot.bus.events import OutboundMessage
from nanobot.providers.factory import ProviderSnapshot
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
original_model = bot._loop.model
original_runtime = bot._loop.runtime_resolver.runtime
override_provider = _fake_provider("stream-provider", max_tokens=2048)
override = ProviderSnapshot(
provider=override_provider,
@@ -903,11 +933,22 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
context_window_tokens=4096,
signature=("sdk", "stream"),
)
bot._runtime_overrides.model_override_snapshot = MagicMock(return_value=override)
override_runtime = runtime_from_provider_snapshot(override)
bot._loop.runtime_resolver.resolve_override = MagicMock(
return_value=override_runtime
)
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
assert bot._loop.provider is override_provider
assert bot._loop.model == "openai/gpt-4.1-mini"
async def fake_process_direct(
message,
*,
session_key,
on_stream,
on_stream_end,
hooks,
runtime,
):
assert runtime is override_runtime
assert bot._loop.runtime_resolver.runtime is original_runtime
await on_stream("ok")
await on_stream_end(resuming=False)
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
@@ -922,7 +963,7 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
assert events[0].type == "run.started"
assert events[0].metadata["model"] == "openai/gpt-4.1-mini"
assert events[0].metadata["model_preset"] is None
assert bot._loop.model == original_model
assert bot._loop.runtime_resolver.runtime is original_runtime
@pytest.mark.asyncio
@@ -947,7 +988,9 @@ async def test_run_streamed_emits_tool_events(tmp_path):
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
):
calls = [
ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}),
ToolCallRequest(id="call_bad", name="exec", arguments={"cmd": "false"}),
@@ -993,7 +1036,9 @@ async def test_run_streamed_emits_reasoning_events(tmp_path):
config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
):
for hook in hooks:
await hook.emit_reasoning("thinking")
await hook.emit_reasoning_end()
@@ -1020,7 +1065,9 @@ async def test_stream_generator_break_cancels_underlying_run(tmp_path):
bot = Nanobot.from_config(config_path, workspace=tmp_path)
cancelled = asyncio.Event()
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
async def fake_process_direct(
message, *, session_key, on_stream, on_stream_end, hooks, runtime
):
try:
await on_stream("first")
await asyncio.sleep(10)