feat(config): watch runtime configuration changes (#5026)

This commit is contained in:
chengyongru 2026-07-22 13:08:39 +08:00 committed by GitHub
parent 7bf7469d90
commit 3748f664b2
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
11 changed files with 190 additions and 7 deletions

View File

@ -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

View File

@ -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

View File

@ -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(

23
nanobot/config/watcher.py Normal file
View File

@ -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()

View File

@ -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",

View File

@ -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

View File

@ -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)

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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