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):
|
||||
|
||||
Reference in New Issue
Block a user