From 5bd3d1e0afca53a3ce64965082f02521a354d887 Mon Sep 17 00:00:00 2001 From: chengyongru Date: Fri, 10 Jul 2026 14:53:34 +0800 Subject: [PATCH] refactor(agent): centralize runtime default mutations --- nanobot/agent/model_runtime.py | 48 +++++++++++++++++++++ tests/agent/test_model_runtime_resolver.py | 49 ++++++++++++++++++++++ 2 files changed, 97 insertions(+) diff --git a/nanobot/agent/model_runtime.py b/nanobot/agent/model_runtime.py index eacbdcf4..238a6181 100644 --- a/nanobot/agent/model_runtime.py +++ b/nanobot/agent/model_runtime.py @@ -3,6 +3,7 @@ from __future__ import annotations from collections.abc import Callable, Mapping +from dataclasses import replace from nanobot.agent import model_presets as preset_helpers from nanobot.config.schema import Config, ModelPresetConfig @@ -31,6 +32,7 @@ class ModelRuntimeResolver: self._provider_snapshot_loader = provider_snapshot_loader self._preset_snapshot_loader = preset_snapshot_loader self._active_preset = initial_runtime.model_preset + self._tracks_provider_generation = initial_runtime.model_preset is None self._default_selection_signature = preset_helpers.default_selection_signature( initial_runtime.snapshot_signature ) @@ -56,6 +58,7 @@ class ModelRuntimeResolver: """Return the selected runtime, optionally refreshing the default source.""" if refresh: self.refresh() + self._refresh_provider_generation() return self._runtime def resolve_snapshot( @@ -77,6 +80,7 @@ class ModelRuntimeResolver: runtime = self.resolve_snapshot(snapshot, model_preset=model_preset) self._runtime = runtime self._active_preset = model_preset + self._tracks_provider_generation = model_preset is None self._default_selection_signature = preset_helpers.default_selection_signature( runtime.snapshot_signature ) @@ -98,8 +102,51 @@ class ModelRuntimeResolver: runtime = self.resolve_preset(name) self._runtime = runtime self._active_preset = runtime.model_preset + self._tracks_provider_generation = False return runtime + def select_model(self, model: str) -> LLMRuntime: + """Change the default model without reconstructing downstream consumers.""" + if not isinstance(model, str) or not model.strip(): + raise ValueError("model must be a non-empty string") + self._runtime = replace( + self._runtime, + model=model.strip(), + model_preset=None, + ) + self._active_preset = None + return self._runtime + + def select_context_window(self, context_window_tokens: int) -> LLMRuntime: + """Change the default context limit for future admissions.""" + if not isinstance(context_window_tokens, int) or isinstance( + context_window_tokens, + bool, + ): + raise TypeError("context_window_tokens must be an integer") + self._runtime = replace( + self._runtime, + context_window_tokens=context_window_tokens, + ) + return self._runtime + + def _refresh_provider_generation(self) -> LLMRuntime | None: + """Adopt direct provider-default changes only for provider-backed defaults.""" + if not self._tracks_provider_generation: + return None + runtime = self._runtime + captured = LLMRuntime.capture( + runtime.provider, + runtime.model, + context_window_tokens=runtime.context_window_tokens, + model_preset=runtime.model_preset, + snapshot_signature=runtime.snapshot_signature, + ) + if captured.generation == runtime.generation: + return None + self._runtime = replace(runtime, generation=captured.generation) + return self._runtime + def refresh(self) -> LLMRuntime | None: """Refresh configured defaults and return the replacement when changed.""" if self._provider_snapshot_loader is None: @@ -121,6 +168,7 @@ class ModelRuntimeResolver: return None self._runtime = runtime self._active_preset = active_preset + self._tracks_provider_generation = active_preset is None self._default_selection_signature = preset_helpers.default_selection_signature( runtime.snapshot_signature ) diff --git a/tests/agent/test_model_runtime_resolver.py b/tests/agent/test_model_runtime_resolver.py index 4c24d916..14e45607 100644 --- a/tests/agent/test_model_runtime_resolver.py +++ b/tests/agent/test_model_runtime_resolver.py @@ -137,3 +137,52 @@ def test_resolver_refresh_preserves_unchanged_active_preset() -> None: assert refreshed is None assert resolver.runtime.provider is preset_provider assert resolver.model_preset == "fast" + + +def test_resolver_refreshes_provider_generation_for_next_default_turn() -> None: + provider = _provider(temperature=0.2, max_tokens=2048) + resolver = ModelRuntimeResolver(_runtime(provider)) + admitted = resolver.current() + + provider.generation = GenerationSettings(temperature=0.8, max_tokens=512) + refreshed = resolver.current(refresh=True) + + assert admitted.generation == GenerationSettings(0.2, 2048, None) + assert refreshed.generation == GenerationSettings(0.8, 512, None) + + +def test_selected_preset_generation_does_not_fall_back_to_provider_defaults() -> None: + provider = _provider(temperature=0.1, max_tokens=1024) + resolver = ModelRuntimeResolver( + _runtime(provider), + model_presets={ + "creative": ModelPresetConfig( + model="creative-model", + temperature=0.7, + max_tokens=4096, + ) + }, + ) + selected = resolver.select_preset("creative") + + provider.generation = GenerationSettings(temperature=0.9, max_tokens=64) + refreshed = resolver.current(refresh=True) + + assert refreshed is selected + assert refreshed.generation == GenerationSettings(0.7, 4096, None) + + +def test_resolver_mutates_only_its_default_selection() -> None: + initial = _runtime() + resolver = ModelRuntimeResolver(initial) + + selected_model = resolver.select_model("next-model") + selected_window = resolver.select_context_window(65_536) + + assert selected_model.model == "next-model" + assert selected_window.model == "next-model" + assert selected_window.context_window_tokens == 65_536 + assert resolver.runtime is selected_window + assert resolver.model_preset is None + assert initial.model == "base-model" + assert initial.context_window_tokens == 10_000