refactor(agent): introduce immutable model runtime resolver
This commit is contained in:
@@ -0,0 +1,154 @@
|
|||||||
|
"""Public resolution boundary for default and overridden LLM runtimes."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Callable, Mapping
|
||||||
|
|
||||||
|
from nanobot.agent import model_presets as preset_helpers
|
||||||
|
from nanobot.config.schema import Config, ModelPresetConfig
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot, build_provider_snapshot
|
||||||
|
from nanobot.utils.llm_runtime import LLMRuntime, runtime_from_provider_snapshot
|
||||||
|
|
||||||
|
|
||||||
|
class ModelRuntimeResolver:
|
||||||
|
"""Own model selection and resolve it to immutable execution values.
|
||||||
|
|
||||||
|
The resolver is deliberately independent of ``AgentLoop``. Command, SDK,
|
||||||
|
and tool admission layers can depend on this public service without reading
|
||||||
|
or mutating private loop state.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
initial_runtime: LLMRuntime,
|
||||||
|
*,
|
||||||
|
model_presets: Mapping[str, ModelPresetConfig] | None = None,
|
||||||
|
provider_snapshot_loader: Callable[[], ProviderSnapshot] | None = None,
|
||||||
|
preset_snapshot_loader: preset_helpers.PresetSnapshotLoader | None = None,
|
||||||
|
) -> None:
|
||||||
|
self._runtime = initial_runtime
|
||||||
|
self._model_presets = dict(model_presets or {})
|
||||||
|
self._provider_snapshot_loader = provider_snapshot_loader
|
||||||
|
self._preset_snapshot_loader = preset_snapshot_loader
|
||||||
|
self._active_preset = initial_runtime.model_preset
|
||||||
|
self._default_selection_signature = preset_helpers.default_selection_signature(
|
||||||
|
initial_runtime.snapshot_signature
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def runtime(self) -> LLMRuntime:
|
||||||
|
"""Return the current immutable default without refreshing configuration."""
|
||||||
|
return self._runtime
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_presets(self) -> Mapping[str, ModelPresetConfig]:
|
||||||
|
return self._model_presets
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_preset(self) -> str | None:
|
||||||
|
return self._active_preset
|
||||||
|
|
||||||
|
@property
|
||||||
|
def provider_signature(self) -> tuple[object, ...] | None:
|
||||||
|
return self._runtime.snapshot_signature
|
||||||
|
|
||||||
|
def current(self, *, refresh: bool = False) -> LLMRuntime:
|
||||||
|
"""Return the selected runtime, optionally refreshing the default source."""
|
||||||
|
if refresh:
|
||||||
|
self.refresh()
|
||||||
|
return self._runtime
|
||||||
|
|
||||||
|
def resolve_snapshot(
|
||||||
|
self,
|
||||||
|
snapshot: ProviderSnapshot,
|
||||||
|
*,
|
||||||
|
model_preset: str | None = None,
|
||||||
|
) -> LLMRuntime:
|
||||||
|
"""Resolve a factory snapshot without changing the selected default."""
|
||||||
|
return runtime_from_provider_snapshot(snapshot, model_preset=model_preset)
|
||||||
|
|
||||||
|
def adopt_snapshot(
|
||||||
|
self,
|
||||||
|
snapshot: ProviderSnapshot,
|
||||||
|
*,
|
||||||
|
model_preset: str | None = None,
|
||||||
|
) -> LLMRuntime:
|
||||||
|
"""Select a snapshot as the default for future turns."""
|
||||||
|
runtime = self.resolve_snapshot(snapshot, model_preset=model_preset)
|
||||||
|
self._runtime = runtime
|
||||||
|
self._active_preset = model_preset
|
||||||
|
self._default_selection_signature = preset_helpers.default_selection_signature(
|
||||||
|
runtime.snapshot_signature
|
||||||
|
)
|
||||||
|
return runtime
|
||||||
|
|
||||||
|
def resolve_preset(self, name: str | None) -> LLMRuntime:
|
||||||
|
"""Resolve a named preset without changing the selected default."""
|
||||||
|
normalized = preset_helpers.normalize_preset_name(name, self._model_presets)
|
||||||
|
snapshot = preset_helpers.build_runtime_preset_snapshot(
|
||||||
|
name=normalized,
|
||||||
|
presets=self._model_presets,
|
||||||
|
provider=self._runtime.provider,
|
||||||
|
loader=self._preset_snapshot_loader,
|
||||||
|
)
|
||||||
|
return self.resolve_snapshot(snapshot, model_preset=normalized)
|
||||||
|
|
||||||
|
def select_preset(self, name: str | None) -> LLMRuntime:
|
||||||
|
"""Select a named preset as the default for future turns."""
|
||||||
|
runtime = self.resolve_preset(name)
|
||||||
|
self._runtime = runtime
|
||||||
|
self._active_preset = runtime.model_preset
|
||||||
|
return runtime
|
||||||
|
|
||||||
|
def refresh(self) -> LLMRuntime | None:
|
||||||
|
"""Refresh configured defaults and return the replacement when changed."""
|
||||||
|
if self._provider_snapshot_loader is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
snapshot = self._provider_snapshot_loader()
|
||||||
|
default_selection = preset_helpers.default_selection_signature(snapshot.signature)
|
||||||
|
active_preset = self._active_preset
|
||||||
|
if active_preset and self._default_selection_signature in (None, default_selection):
|
||||||
|
self._default_selection_signature = default_selection
|
||||||
|
runtime = self.resolve_preset(active_preset)
|
||||||
|
else:
|
||||||
|
active_preset = None
|
||||||
|
self._active_preset = None
|
||||||
|
self._default_selection_signature = default_selection
|
||||||
|
runtime = self.resolve_snapshot(snapshot)
|
||||||
|
|
||||||
|
if runtime.snapshot_signature == self._runtime.snapshot_signature:
|
||||||
|
return None
|
||||||
|
self._runtime = runtime
|
||||||
|
self._active_preset = active_preset
|
||||||
|
self._default_selection_signature = preset_helpers.default_selection_signature(
|
||||||
|
runtime.snapshot_signature
|
||||||
|
)
|
||||||
|
return runtime
|
||||||
|
|
||||||
|
def resolve_override(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
model: str | None,
|
||||||
|
model_preset: str | None,
|
||||||
|
config: Config | None = None,
|
||||||
|
) -> LLMRuntime | None:
|
||||||
|
"""Resolve an SDK-style per-run override without mutating the default."""
|
||||||
|
if model is not None and model_preset is not None:
|
||||||
|
raise ValueError("model and model_preset are mutually exclusive")
|
||||||
|
if model_preset is not None:
|
||||||
|
return self.resolve_preset(model_preset)
|
||||||
|
if model is None:
|
||||||
|
return None
|
||||||
|
if config is None:
|
||||||
|
return LLMRuntime(
|
||||||
|
provider=self._runtime.provider,
|
||||||
|
model=model,
|
||||||
|
generation=self._runtime.generation,
|
||||||
|
context_window_tokens=self._runtime.context_window_tokens,
|
||||||
|
snapshot_signature=("model_override", model),
|
||||||
|
)
|
||||||
|
|
||||||
|
base = config.resolve_preset(self._active_preset)
|
||||||
|
preset = base.model_copy(update={"model": model, "provider": "auto"})
|
||||||
|
return self.resolve_snapshot(build_provider_snapshot(config, preset=preset))
|
||||||
@@ -1,13 +1,92 @@
|
|||||||
"""Small helpers for passing the active LLM provider/model together."""
|
"""Immutable execution settings for one LLM turn."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, replace
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from nanobot.providers.base import LLMProvider
|
from nanobot.providers.base import GenerationSettings, LLMProvider
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class LLMRuntime:
|
class LLMRuntime:
|
||||||
|
"""One captured provider/model configuration used for an entire execution.
|
||||||
|
|
||||||
|
The provider itself is stateful, but all mutable selection and generation
|
||||||
|
values are copied into this frozen value. Consumers must use these fields
|
||||||
|
instead of consulting ``provider.generation`` after admission.
|
||||||
|
"""
|
||||||
|
|
||||||
provider: LLMProvider
|
provider: LLMProvider
|
||||||
model: str
|
model: str
|
||||||
|
generation: GenerationSettings
|
||||||
|
context_window_tokens: int
|
||||||
|
model_preset: str | None = None
|
||||||
|
snapshot_signature: tuple[object, ...] | None = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def capture(
|
||||||
|
cls,
|
||||||
|
provider: LLMProvider,
|
||||||
|
model: str,
|
||||||
|
*,
|
||||||
|
context_window_tokens: int,
|
||||||
|
model_preset: str | None = None,
|
||||||
|
snapshot_signature: tuple[object, ...] | None = None,
|
||||||
|
) -> LLMRuntime:
|
||||||
|
"""Capture provider defaults without retaining mutable generation state."""
|
||||||
|
generation = provider.generation
|
||||||
|
return cls(
|
||||||
|
provider=provider,
|
||||||
|
model=model,
|
||||||
|
generation=GenerationSettings(
|
||||||
|
temperature=generation.temperature,
|
||||||
|
max_tokens=generation.max_tokens,
|
||||||
|
reasoning_effort=generation.reasoning_effort,
|
||||||
|
),
|
||||||
|
context_window_tokens=context_window_tokens,
|
||||||
|
model_preset=model_preset,
|
||||||
|
snapshot_signature=snapshot_signature,
|
||||||
|
)
|
||||||
|
|
||||||
|
def with_generation_overrides(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
temperature: float | None = None,
|
||||||
|
max_tokens: int | None = None,
|
||||||
|
reasoning_effort: str | None = None,
|
||||||
|
) -> LLMRuntime:
|
||||||
|
"""Return a derived runtime for explicit per-run generation overrides."""
|
||||||
|
generation = self.generation
|
||||||
|
return replace(
|
||||||
|
self,
|
||||||
|
generation=GenerationSettings(
|
||||||
|
temperature=(
|
||||||
|
generation.temperature if temperature is None else temperature
|
||||||
|
),
|
||||||
|
max_tokens=generation.max_tokens if max_tokens is None else max_tokens,
|
||||||
|
reasoning_effort=(
|
||||||
|
generation.reasoning_effort
|
||||||
|
if reasoning_effort is None
|
||||||
|
else reasoning_effort
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def runtime_from_provider_snapshot(
|
||||||
|
snapshot: ProviderSnapshot,
|
||||||
|
*,
|
||||||
|
model_preset: str | None = None,
|
||||||
|
) -> LLMRuntime:
|
||||||
|
"""Convert a provider factory snapshot into the canonical runtime value."""
|
||||||
|
return LLMRuntime.capture(
|
||||||
|
snapshot.provider,
|
||||||
|
snapshot.model,
|
||||||
|
context_window_tokens=snapshot.context_window_tokens,
|
||||||
|
model_preset=model_preset,
|
||||||
|
snapshot_signature=snapshot.signature,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,137 @@
|
|||||||
|
from dataclasses import FrozenInstanceError
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.agent.model_runtime import ModelRuntimeResolver
|
||||||
|
from nanobot.config.schema import ModelPresetConfig
|
||||||
|
from nanobot.providers.base import GenerationSettings
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
|
from nanobot.utils.llm_runtime import LLMRuntime, runtime_from_provider_snapshot
|
||||||
|
|
||||||
|
|
||||||
|
def _provider(
|
||||||
|
*,
|
||||||
|
temperature: float = 0.1,
|
||||||
|
max_tokens: int = 1024,
|
||||||
|
reasoning_effort: str | None = None,
|
||||||
|
) -> MagicMock:
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.generation = GenerationSettings(
|
||||||
|
temperature=temperature,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
reasoning_effort=reasoning_effort,
|
||||||
|
)
|
||||||
|
return provider
|
||||||
|
|
||||||
|
|
||||||
|
def _runtime(provider: MagicMock | None = None) -> LLMRuntime:
|
||||||
|
return LLMRuntime.capture(
|
||||||
|
provider or _provider(),
|
||||||
|
"base-model",
|
||||||
|
context_window_tokens=10_000,
|
||||||
|
snapshot_signature=("base-model", "auto"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_runtime_captures_generation_and_is_immutable() -> None:
|
||||||
|
provider = _provider(temperature=0.2, max_tokens=2048, reasoning_effort="low")
|
||||||
|
runtime = _runtime(provider)
|
||||||
|
|
||||||
|
provider.generation = GenerationSettings(temperature=0.9, max_tokens=99)
|
||||||
|
|
||||||
|
assert runtime.generation == GenerationSettings(0.2, 2048, "low")
|
||||||
|
with pytest.raises(FrozenInstanceError):
|
||||||
|
runtime.model = "changed" # type: ignore[misc]
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_snapshot_has_one_canonical_runtime_conversion() -> None:
|
||||||
|
provider = _provider(temperature=0.3, max_tokens=4096)
|
||||||
|
snapshot = ProviderSnapshot(
|
||||||
|
provider=provider,
|
||||||
|
model="snapshot-model",
|
||||||
|
context_window_tokens=32_768,
|
||||||
|
signature=("snapshot-model", "openai"),
|
||||||
|
)
|
||||||
|
|
||||||
|
runtime = runtime_from_provider_snapshot(snapshot, model_preset="fast")
|
||||||
|
|
||||||
|
assert runtime.provider is provider
|
||||||
|
assert runtime.model == "snapshot-model"
|
||||||
|
assert runtime.generation == GenerationSettings(0.3, 4096, None)
|
||||||
|
assert runtime.context_window_tokens == 32_768
|
||||||
|
assert runtime.model_preset == "fast"
|
||||||
|
assert runtime.snapshot_signature == snapshot.signature
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolver_resolves_preset_without_mutating_selected_runtime() -> None:
|
||||||
|
initial = _runtime()
|
||||||
|
preset_provider = _provider(temperature=0.5, max_tokens=512)
|
||||||
|
preset = ModelPresetConfig(
|
||||||
|
model="fast-model",
|
||||||
|
temperature=0.5,
|
||||||
|
max_tokens=512,
|
||||||
|
context_window_tokens=8192,
|
||||||
|
)
|
||||||
|
resolver = ModelRuntimeResolver(
|
||||||
|
initial,
|
||||||
|
model_presets={"fast": preset},
|
||||||
|
preset_snapshot_loader=lambda name: ProviderSnapshot(
|
||||||
|
provider=preset_provider,
|
||||||
|
model=preset.model,
|
||||||
|
context_window_tokens=preset.context_window_tokens,
|
||||||
|
signature=(name, preset.model),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
resolved = resolver.resolve_preset("fast")
|
||||||
|
|
||||||
|
assert resolved.model == "fast-model"
|
||||||
|
assert resolved.model_preset == "fast"
|
||||||
|
assert resolver.runtime is initial
|
||||||
|
assert resolver.model_preset is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolver_model_override_is_derived_without_default_mutation() -> None:
|
||||||
|
initial = _runtime()
|
||||||
|
resolver = ModelRuntimeResolver(initial)
|
||||||
|
|
||||||
|
override = resolver.resolve_override(
|
||||||
|
model="override-model",
|
||||||
|
model_preset=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert override is not None
|
||||||
|
assert override.model == "override-model"
|
||||||
|
assert override.provider is initial.provider
|
||||||
|
assert override.generation is initial.generation
|
||||||
|
assert resolver.runtime is initial
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolver_refresh_preserves_unchanged_active_preset() -> None:
|
||||||
|
initial = _runtime()
|
||||||
|
preset = ModelPresetConfig(model="fast-model")
|
||||||
|
preset_provider = _provider()
|
||||||
|
resolver = ModelRuntimeResolver(
|
||||||
|
initial,
|
||||||
|
model_presets={"fast": preset},
|
||||||
|
provider_snapshot_loader=lambda: ProviderSnapshot(
|
||||||
|
provider=_provider(),
|
||||||
|
model="base-model",
|
||||||
|
context_window_tokens=10_000,
|
||||||
|
signature=("base-model", "auto", "refreshed"),
|
||||||
|
),
|
||||||
|
preset_snapshot_loader=lambda _name: ProviderSnapshot(
|
||||||
|
provider=preset_provider,
|
||||||
|
model="fast-model",
|
||||||
|
context_window_tokens=20_000,
|
||||||
|
signature=("fast-model", "auto", "refreshed"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
resolver.select_preset("fast")
|
||||||
|
|
||||||
|
refreshed = resolver.refresh()
|
||||||
|
|
||||||
|
assert refreshed is None
|
||||||
|
assert resolver.runtime.provider is preset_provider
|
||||||
|
assert resolver.model_preset == "fast"
|
||||||
Reference in New Issue
Block a user