feat(config): watch runtime configuration changes (#5026)
This commit is contained in:
@@ -235,11 +235,7 @@ class AgentLoop:
|
|||||||
def llm_runtime(self) -> LLMRuntime:
|
def llm_runtime(self) -> LLMRuntime:
|
||||||
"""Resolve the immutable default used to admit the next turn."""
|
"""Resolve the immutable default used to admit the next turn."""
|
||||||
previous = self.runtime_resolver.runtime
|
previous = self.runtime_resolver.runtime
|
||||||
try:
|
runtime = self.runtime_resolver.admit()
|
||||||
runtime = self.runtime_resolver.current(refresh=True)
|
|
||||||
except Exception:
|
|
||||||
logger.exception("Failed to refresh model runtime")
|
|
||||||
return previous
|
|
||||||
if (
|
if (
|
||||||
runtime.model != previous.model
|
runtime.model != previous.model
|
||||||
or runtime.model_preset != previous.model_preset
|
or runtime.model_preset != previous.model_preset
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ class ModelRuntimeResolver:
|
|||||||
self._model_presets = dict(model_presets or {})
|
self._model_presets = dict(model_presets or {})
|
||||||
self._provider_snapshot_loader = provider_snapshot_loader
|
self._provider_snapshot_loader = provider_snapshot_loader
|
||||||
self._preset_snapshot_loader = preset_snapshot_loader
|
self._preset_snapshot_loader = preset_snapshot_loader
|
||||||
|
self._refresh_required = False
|
||||||
self._tracks_provider_generation = initial_runtime.model_preset is None
|
self._tracks_provider_generation = initial_runtime.model_preset is None
|
||||||
self._default_selection_signature = preset_helpers.default_selection_signature(
|
self._default_selection_signature = preset_helpers.default_selection_signature(
|
||||||
initial_runtime.snapshot_signature
|
initial_runtime.snapshot_signature
|
||||||
@@ -60,6 +61,17 @@ class ModelRuntimeResolver:
|
|||||||
self._refresh_provider_generation()
|
self._refresh_provider_generation()
|
||||||
return self._runtime
|
return self._runtime
|
||||||
|
|
||||||
|
def admit(self) -> LLMRuntime:
|
||||||
|
"""Resolve the immutable runtime for the next turn admission."""
|
||||||
|
if self._refresh_required:
|
||||||
|
self.refresh()
|
||||||
|
self._refresh_provider_generation()
|
||||||
|
return self._runtime
|
||||||
|
|
||||||
|
def invalidate(self) -> None:
|
||||||
|
"""Refresh configured runtime state on the next admission."""
|
||||||
|
self._refresh_required = True
|
||||||
|
|
||||||
def resolve_snapshot(
|
def resolve_snapshot(
|
||||||
self,
|
self,
|
||||||
snapshot: ProviderSnapshot,
|
snapshot: ProviderSnapshot,
|
||||||
@@ -146,6 +158,7 @@ class ModelRuntimeResolver:
|
|||||||
def refresh(self) -> LLMRuntime | None:
|
def refresh(self) -> LLMRuntime | None:
|
||||||
"""Refresh configured defaults and return the replacement when changed."""
|
"""Refresh configured defaults and return the replacement when changed."""
|
||||||
if self._provider_snapshot_loader is None:
|
if self._provider_snapshot_loader is None:
|
||||||
|
self._refresh_required = False
|
||||||
return None
|
return None
|
||||||
|
|
||||||
snapshot = self._provider_snapshot_loader()
|
snapshot = self._provider_snapshot_loader()
|
||||||
@@ -161,6 +174,7 @@ class ModelRuntimeResolver:
|
|||||||
runtime.snapshot_signature == self._runtime.snapshot_signature
|
runtime.snapshot_signature == self._runtime.snapshot_signature
|
||||||
and runtime.model_preset == self._runtime.model_preset
|
and runtime.model_preset == self._runtime.model_preset
|
||||||
)
|
)
|
||||||
|
self._refresh_required = False
|
||||||
if unchanged:
|
if unchanged:
|
||||||
self._default_selection_signature = default_selection
|
self._default_selection_signature = default_selection
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -1609,6 +1609,7 @@ def _run_gateway(
|
|||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.bus.runtime_events import RuntimeEventBus
|
from nanobot.bus.runtime_events import RuntimeEventBus
|
||||||
from nanobot.channels.manager import ChannelManager
|
from nanobot.channels.manager import ChannelManager
|
||||||
|
from nanobot.config.watcher import watch_config_file
|
||||||
from nanobot.cron.bound_runner import run_bound_cron_job
|
from nanobot.cron.bound_runner import run_bound_cron_job
|
||||||
from nanobot.cron.service import CronJobSkippedError, CronService
|
from nanobot.cron.service import CronJobSkippedError, CronService
|
||||||
from nanobot.cron.session_turns import is_bound_cron_job
|
from nanobot.cron.session_turns import is_bound_cron_job
|
||||||
@@ -2070,7 +2071,16 @@ def _run_gateway(
|
|||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
await cron.start()
|
await cron.start()
|
||||||
|
# Re-read once on first admission to close the watcher subscription window.
|
||||||
|
agent.runtime_resolver.invalidate()
|
||||||
tasks = [
|
tasks = [
|
||||||
|
asyncio.create_task(
|
||||||
|
watch_config_file(
|
||||||
|
Path(config_path),
|
||||||
|
lambda: agent.runtime_resolver.invalidate(),
|
||||||
|
),
|
||||||
|
name="nanobot-config-watcher",
|
||||||
|
),
|
||||||
asyncio.create_task(agent.run(), name="nanobot-agent-loop"),
|
asyncio.create_task(agent.run(), name="nanobot-agent-loop"),
|
||||||
asyncio.create_task(channels.start_all(), name="nanobot-channels"),
|
asyncio.create_task(channels.start_all(), name="nanobot-channels"),
|
||||||
asyncio.create_task(
|
asyncio.create_task(
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
"""System-level notification for config file changes."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from watchfiles import Change, awatch
|
||||||
|
|
||||||
|
|
||||||
|
async def watch_config_file(config_path: Path, on_change: Callable[[], None]) -> None:
|
||||||
|
"""Notify ``on_change`` after the configured file changes."""
|
||||||
|
target = config_path.resolve(strict=False)
|
||||||
|
|
||||||
|
def is_config_file(_change: Change, changed_path: str) -> bool:
|
||||||
|
return Path(changed_path).resolve(strict=False) == target
|
||||||
|
|
||||||
|
async for _changes in awatch(
|
||||||
|
target.parent,
|
||||||
|
watch_filter=is_config_file,
|
||||||
|
recursive=False,
|
||||||
|
):
|
||||||
|
on_change()
|
||||||
@@ -48,6 +48,7 @@ dependencies = [
|
|||||||
"dulwich>=0.22.0,<1.0.0",
|
"dulwich>=0.22.0,<1.0.0",
|
||||||
"pyyaml>=6.0,<7.0.0",
|
"pyyaml>=6.0,<7.0.0",
|
||||||
"filelock>=3.25.2",
|
"filelock>=3.25.2",
|
||||||
|
"watchfiles>=1.1.1,<2.0.0",
|
||||||
"packaging>=24.0",
|
"packaging>=24.0",
|
||||||
"tzdata>=2025.2; sys_platform == 'win32'",
|
"tzdata>=2025.2; sys_platform == 'win32'",
|
||||||
"defusedxml>=0.7.1,<1.0.0",
|
"defusedxml>=0.7.1,<1.0.0",
|
||||||
|
|||||||
@@ -98,6 +98,7 @@ class TestMaxMessagesInit:
|
|||||||
|
|
||||||
initial = loop.runtime_resolver.runtime
|
initial = loop.runtime_resolver.runtime
|
||||||
assert replay_max_messages_for_context(initial.context_window_tokens) == 327
|
assert replay_max_messages_for_context(initial.context_window_tokens) == 327
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
refreshed = loop.llm_runtime()
|
refreshed = loop.llm_runtime()
|
||||||
assert replay_max_messages_for_context(refreshed.context_window_tokens) == FILE_MAX_MESSAGES
|
assert replay_max_messages_for_context(refreshed.context_window_tokens) == FILE_MAX_MESSAGES
|
||||||
|
|
||||||
|
|||||||
@@ -177,12 +177,59 @@ def test_resolver_refreshes_provider_generation_for_next_default_turn() -> None:
|
|||||||
admitted = resolver.current()
|
admitted = resolver.current()
|
||||||
|
|
||||||
provider.generation = GenerationSettings(temperature=0.8, max_tokens=512)
|
provider.generation = GenerationSettings(temperature=0.8, max_tokens=512)
|
||||||
refreshed = resolver.current(refresh=True)
|
refreshed = resolver.admit()
|
||||||
|
|
||||||
assert admitted.generation == GenerationSettings(0.2, 2048, None)
|
assert admitted.generation == GenerationSettings(0.2, 2048, None)
|
||||||
assert refreshed.generation == GenerationSettings(0.8, 512, None)
|
assert refreshed.generation == GenerationSettings(0.8, 512, None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolver_admission_reloads_config_only_after_invalidation() -> None:
|
||||||
|
initial = _runtime()
|
||||||
|
refreshed_provider = _provider()
|
||||||
|
load_count = 0
|
||||||
|
|
||||||
|
def load_snapshot() -> ProviderSnapshot:
|
||||||
|
nonlocal load_count
|
||||||
|
load_count += 1
|
||||||
|
return ProviderSnapshot(
|
||||||
|
provider=refreshed_provider,
|
||||||
|
model="refreshed-model",
|
||||||
|
context_window_tokens=20_000,
|
||||||
|
signature=("refreshed-model", "auto"),
|
||||||
|
)
|
||||||
|
|
||||||
|
resolver = ModelRuntimeResolver(initial, provider_snapshot_loader=load_snapshot)
|
||||||
|
|
||||||
|
assert resolver.admit() is initial
|
||||||
|
assert load_count == 0
|
||||||
|
|
||||||
|
resolver.invalidate()
|
||||||
|
refreshed = resolver.admit()
|
||||||
|
|
||||||
|
assert refreshed.provider is refreshed_provider
|
||||||
|
assert refreshed.model == "refreshed-model"
|
||||||
|
assert resolver.admit() is refreshed
|
||||||
|
assert load_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_current_refresh_forces_config_reload() -> None:
|
||||||
|
initial = _runtime()
|
||||||
|
load_snapshot = MagicMock(
|
||||||
|
return_value=ProviderSnapshot(
|
||||||
|
provider=_provider(),
|
||||||
|
model="refreshed-model",
|
||||||
|
context_window_tokens=20_000,
|
||||||
|
signature=("refreshed-model", "auto"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
resolver = ModelRuntimeResolver(initial, provider_snapshot_loader=load_snapshot)
|
||||||
|
|
||||||
|
refreshed = resolver.current(refresh=True)
|
||||||
|
|
||||||
|
assert refreshed.model == "refreshed-model"
|
||||||
|
load_snapshot.assert_called_once_with()
|
||||||
|
|
||||||
|
|
||||||
def test_selected_preset_generation_does_not_fall_back_to_provider_defaults() -> None:
|
def test_selected_preset_generation_does_not_fall_back_to_provider_defaults() -> None:
|
||||||
provider = _provider(temperature=0.1, max_tokens=1024)
|
provider = _provider(temperature=0.1, max_tokens=1024)
|
||||||
resolver = ModelRuntimeResolver(
|
resolver = ModelRuntimeResolver(
|
||||||
@@ -198,7 +245,7 @@ def test_selected_preset_generation_does_not_fall_back_to_provider_defaults() ->
|
|||||||
selected = resolver.select_preset("creative")
|
selected = resolver.select_preset("creative")
|
||||||
|
|
||||||
provider.generation = GenerationSettings(temperature=0.9, max_tokens=64)
|
provider.generation = GenerationSettings(temperature=0.9, max_tokens=64)
|
||||||
refreshed = resolver.current(refresh=True)
|
refreshed = resolver.admit()
|
||||||
|
|
||||||
assert refreshed is selected
|
assert refreshed is selected
|
||||||
assert refreshed.generation == GenerationSettings(0.7, 4096, None)
|
assert refreshed.generation == GenerationSettings(0.7, 4096, None)
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ from pathlib import Path
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.config.loader import save_config
|
from nanobot.config.loader import save_config
|
||||||
@@ -34,6 +36,7 @@ def test_provider_refresh_updates_only_runtime_resolver(tmp_path: Path) -> None:
|
|||||||
signature=("new-model",),
|
signature=("new-model",),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
@@ -90,6 +93,7 @@ def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
|
|||||||
signature=("new-model",),
|
signature=("new-model",),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
@@ -99,6 +103,24 @@ def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
|
|||||||
assert not hasattr(loop.runner, "provider")
|
assert not hasattr(loop.runner, "provider")
|
||||||
|
|
||||||
|
|
||||||
|
def test_llm_runtime_surfaces_invalidated_config_errors(tmp_path: Path) -> None:
|
||||||
|
def fail_refresh() -> ProviderSnapshot:
|
||||||
|
raise ValueError("invalid config")
|
||||||
|
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=MessageBus(),
|
||||||
|
provider=_provider("old-model"),
|
||||||
|
workspace=tmp_path,
|
||||||
|
model="old-model",
|
||||||
|
context_window_tokens=1000,
|
||||||
|
provider_snapshot_loader=fail_refresh,
|
||||||
|
)
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="invalid config"):
|
||||||
|
loop.llm_runtime()
|
||||||
|
|
||||||
|
|
||||||
def test_same_snapshot_default_clears_preset_and_publishes_update(tmp_path: Path) -> None:
|
def test_same_snapshot_default_clears_preset_and_publishes_update(tmp_path: Path) -> None:
|
||||||
base_provider = _provider("base-model")
|
base_provider = _provider("base-model")
|
||||||
fast_provider = _provider("fast-model")
|
fast_provider = _provider("fast-model")
|
||||||
@@ -122,6 +144,7 @@ def test_same_snapshot_default_clears_preset_and_publishes_update(tmp_path: Path
|
|||||||
preset_snapshot_loader=lambda _name: fast_snapshot,
|
preset_snapshot_loader=lambda _name: fast_snapshot,
|
||||||
runtime_model_publisher=lambda model, preset: published.append((model, preset)),
|
runtime_model_publisher=lambda model, preset: published.append((model, preset)),
|
||||||
)
|
)
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
|
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
@@ -173,6 +196,7 @@ def test_settings_context_window_refreshes_runtime_state(
|
|||||||
loop = AgentLoop.from_config(config, provider_snapshot_loader=loader)
|
loop = AgentLoop.from_config(config, provider_snapshot_loader=loader)
|
||||||
|
|
||||||
payload = update_agent_settings({"context_window_tokens": ["262144"]})
|
payload = update_agent_settings({"context_window_tokens": ["262144"]})
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
loop.llm_runtime()
|
loop.llm_runtime()
|
||||||
|
|
||||||
assert payload["requires_restart"] is False
|
assert payload["requires_restart"] is False
|
||||||
|
|||||||
@@ -175,6 +175,7 @@ def test_active_model_preset_survives_unchanged_config_refresh(tmp_path) -> None
|
|||||||
)
|
)
|
||||||
|
|
||||||
loop.set_model_preset("fast")
|
loop.set_model_preset("fast")
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
loop.llm_runtime()
|
loop.llm_runtime()
|
||||||
|
|
||||||
assert loop.model_preset == "fast"
|
assert loop.model_preset == "fast"
|
||||||
@@ -211,6 +212,7 @@ def test_config_model_refresh_clears_active_model_preset(tmp_path) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
loop.set_model_preset("fast")
|
loop.set_model_preset("fast")
|
||||||
|
loop.runtime_resolver.invalidate()
|
||||||
loop.llm_runtime()
|
loop.llm_runtime()
|
||||||
|
|
||||||
assert loop.model_preset is None
|
assert loop.model_preset is None
|
||||||
|
|||||||
@@ -2737,12 +2737,14 @@ def test_gateway_local_trigger_queue_submits_agent_turns(
|
|||||||
self.context = _FakeContext()
|
self.context = _FakeContext()
|
||||||
self.sessions = kwargs["session_manager"]
|
self.sessions = kwargs["session_manager"]
|
||||||
self.submit_local_trigger_turn = AsyncMock()
|
self.submit_local_trigger_turn = AsyncMock()
|
||||||
|
self.runtime_resolver = MagicMock()
|
||||||
seen["agent"] = self
|
seen["agent"] = self
|
||||||
|
|
||||||
def _schedule_background(self, _coro) -> None:
|
def _schedule_background(self, _coro) -> None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def run(self) -> None:
|
async def run(self) -> None:
|
||||||
|
self.runtime_resolver.invalidate.assert_called_once_with()
|
||||||
await asyncio.Event().wait()
|
await asyncio.Event().wait()
|
||||||
|
|
||||||
async def close_mcp(self) -> None:
|
async def close_mcp(self) -> None:
|
||||||
@@ -2971,6 +2973,7 @@ def test_gateway_health_endpoint_binds_and_serves_expected_responses(
|
|||||||
self.model = "test-model"
|
self.model = "test-model"
|
||||||
self.provider = object()
|
self.provider = object()
|
||||||
self.sessions = _FakeSessionManager()
|
self.sessions = _FakeSessionManager()
|
||||||
|
self.runtime_resolver = MagicMock()
|
||||||
|
|
||||||
def llm_runtime(self) -> None:
|
def llm_runtime(self) -> None:
|
||||||
return None
|
return None
|
||||||
@@ -3164,6 +3167,7 @@ def test_gateway_shutdown_lets_agent_task_own_mcp_cleanup(
|
|||||||
self.model = "test-model"
|
self.model = "test-model"
|
||||||
self.provider = object()
|
self.provider = object()
|
||||||
self.sessions = _FakeSessionManager()
|
self.sessions = _FakeSessionManager()
|
||||||
|
self.runtime_resolver = MagicMock()
|
||||||
|
|
||||||
def llm_runtime(self) -> None:
|
def llm_runtime(self) -> None:
|
||||||
return None
|
return None
|
||||||
@@ -3262,6 +3266,7 @@ def test_gateway_shutdown_event_exits_forever_runtime_tasks(
|
|||||||
self.model = "test-model"
|
self.model = "test-model"
|
||||||
self.provider = object()
|
self.provider = object()
|
||||||
self.sessions = _FakeSessionManager()
|
self.sessions = _FakeSessionManager()
|
||||||
|
self.runtime_resolver = MagicMock()
|
||||||
|
|
||||||
def llm_runtime(self) -> None:
|
def llm_runtime(self) -> None:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -0,0 +1,60 @@
|
|||||||
|
import asyncio
|
||||||
|
from contextlib import suppress
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from watchfiles import Change
|
||||||
|
|
||||||
|
import nanobot.config.watcher as config_watcher
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_watch_config_file_filters_directory_events(
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch,
|
||||||
|
) -> None:
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
other_path = tmp_path / "other.json"
|
||||||
|
seen: dict[str, object] = {}
|
||||||
|
|
||||||
|
async def fake_awatch(*paths, **kwargs):
|
||||||
|
seen["paths"] = paths
|
||||||
|
seen["recursive"] = kwargs["recursive"]
|
||||||
|
watch_filter = kwargs["watch_filter"]
|
||||||
|
assert watch_filter(Change.modified, str(config_path)) is True
|
||||||
|
assert watch_filter(Change.modified, str(other_path)) is False
|
||||||
|
yield {(Change.modified, str(config_path))}
|
||||||
|
|
||||||
|
monkeypatch.setattr(config_watcher, "awatch", fake_awatch)
|
||||||
|
changes: list[None] = []
|
||||||
|
|
||||||
|
await config_watcher.watch_config_file(config_path, lambda: changes.append(None))
|
||||||
|
|
||||||
|
assert seen == {"paths": (tmp_path,), "recursive": False}
|
||||||
|
assert changes == [None]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_watch_config_file_observes_atomic_replace(tmp_path: Path) -> None:
|
||||||
|
config_path = tmp_path / "config.json"
|
||||||
|
config_path.write_text("{}", encoding="utf-8")
|
||||||
|
changed = asyncio.Event()
|
||||||
|
task = asyncio.create_task(
|
||||||
|
config_watcher.watch_config_file(config_path, changed.set)
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
for attempt in range(10):
|
||||||
|
replacement = tmp_path / "config.tmp"
|
||||||
|
replacement.write_text(f'{{"attempt": {attempt}}}', encoding="utf-8")
|
||||||
|
replacement.replace(config_path)
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(changed.wait(), timeout=0.2)
|
||||||
|
break
|
||||||
|
except TimeoutError:
|
||||||
|
continue
|
||||||
|
assert changed.is_set()
|
||||||
|
finally:
|
||||||
|
task.cancel()
|
||||||
|
with suppress(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
Reference in New Issue
Block a user