refactor(agent): centralize runtime default mutations
This commit is contained in:
@@ -3,6 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from collections.abc import Callable, Mapping
|
from collections.abc import Callable, Mapping
|
||||||
|
from dataclasses import replace
|
||||||
|
|
||||||
from nanobot.agent import model_presets as preset_helpers
|
from nanobot.agent import model_presets as preset_helpers
|
||||||
from nanobot.config.schema import Config, ModelPresetConfig
|
from nanobot.config.schema import Config, ModelPresetConfig
|
||||||
@@ -31,6 +32,7 @@ class ModelRuntimeResolver:
|
|||||||
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._active_preset = initial_runtime.model_preset
|
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(
|
self._default_selection_signature = preset_helpers.default_selection_signature(
|
||||||
initial_runtime.snapshot_signature
|
initial_runtime.snapshot_signature
|
||||||
)
|
)
|
||||||
@@ -56,6 +58,7 @@ class ModelRuntimeResolver:
|
|||||||
"""Return the selected runtime, optionally refreshing the default source."""
|
"""Return the selected runtime, optionally refreshing the default source."""
|
||||||
if refresh:
|
if refresh:
|
||||||
self.refresh()
|
self.refresh()
|
||||||
|
self._refresh_provider_generation()
|
||||||
return self._runtime
|
return self._runtime
|
||||||
|
|
||||||
def resolve_snapshot(
|
def resolve_snapshot(
|
||||||
@@ -77,6 +80,7 @@ class ModelRuntimeResolver:
|
|||||||
runtime = self.resolve_snapshot(snapshot, model_preset=model_preset)
|
runtime = self.resolve_snapshot(snapshot, model_preset=model_preset)
|
||||||
self._runtime = runtime
|
self._runtime = runtime
|
||||||
self._active_preset = model_preset
|
self._active_preset = model_preset
|
||||||
|
self._tracks_provider_generation = model_preset is None
|
||||||
self._default_selection_signature = preset_helpers.default_selection_signature(
|
self._default_selection_signature = preset_helpers.default_selection_signature(
|
||||||
runtime.snapshot_signature
|
runtime.snapshot_signature
|
||||||
)
|
)
|
||||||
@@ -98,8 +102,51 @@ class ModelRuntimeResolver:
|
|||||||
runtime = self.resolve_preset(name)
|
runtime = self.resolve_preset(name)
|
||||||
self._runtime = runtime
|
self._runtime = runtime
|
||||||
self._active_preset = runtime.model_preset
|
self._active_preset = runtime.model_preset
|
||||||
|
self._tracks_provider_generation = False
|
||||||
return runtime
|
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:
|
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:
|
||||||
@@ -121,6 +168,7 @@ class ModelRuntimeResolver:
|
|||||||
return None
|
return None
|
||||||
self._runtime = runtime
|
self._runtime = runtime
|
||||||
self._active_preset = active_preset
|
self._active_preset = active_preset
|
||||||
|
self._tracks_provider_generation = active_preset is None
|
||||||
self._default_selection_signature = preset_helpers.default_selection_signature(
|
self._default_selection_signature = preset_helpers.default_selection_signature(
|
||||||
runtime.snapshot_signature
|
runtime.snapshot_signature
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -137,3 +137,52 @@ def test_resolver_refresh_preserves_unchanged_active_preset() -> None:
|
|||||||
assert refreshed is None
|
assert refreshed is None
|
||||||
assert resolver.runtime.provider is preset_provider
|
assert resolver.runtime.provider is preset_provider
|
||||||
assert resolver.model_preset == "fast"
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user