refactor(agent): snapshot generation without provider mutation
This commit is contained in:
+10
-2
@@ -173,9 +173,15 @@ class AgentLoop:
|
||||
return self.tools.tool_names
|
||||
|
||||
def llm_runtime(self) -> LLMRuntime:
|
||||
"""Return the current provider/model pair owned by this loop."""
|
||||
"""Capture the current provider/model settings owned by this loop."""
|
||||
self._refresh_provider_snapshot()
|
||||
return LLMRuntime(self.provider, self.model)
|
||||
return LLMRuntime.capture(
|
||||
self.provider,
|
||||
self.model,
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
model_preset=self.model_preset,
|
||||
snapshot_signature=self._provider_signature,
|
||||
)
|
||||
|
||||
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||
_PENDING_USER_TURN_KEY = "pending_user_turn"
|
||||
@@ -444,6 +450,8 @@ class AgentLoop:
|
||||
provider = snapshot.provider
|
||||
model = snapshot.model
|
||||
context_window_tokens = snapshot.context_window_tokens
|
||||
if snapshot.generation is not None:
|
||||
provider.generation = snapshot.generation
|
||||
old_model = self.model
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
|
||||
@@ -34,12 +34,12 @@ def build_static_preset_snapshot(
|
||||
name: str,
|
||||
preset: ModelPresetConfig,
|
||||
) -> ProviderSnapshot:
|
||||
provider.generation = preset.to_generation_settings()
|
||||
return ProviderSnapshot(
|
||||
provider=provider,
|
||||
model=preset.model,
|
||||
context_window_tokens=preset.context_window_tokens,
|
||||
signature=("model_preset", name, preset.model_dump_json()),
|
||||
generation=preset.to_generation_settings(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from nanobot.config.schema import Config, InlineFallbackConfig, ModelPresetConfig, ProviderConfig
|
||||
from nanobot.providers.base import LLMProvider
|
||||
from nanobot.providers.base import GenerationSettings, LLMProvider
|
||||
from nanobot.providers.fallback_provider import FallbackProvider
|
||||
from nanobot.providers.registry import ProviderSpec, create_dynamic_spec, find_by_name
|
||||
|
||||
@@ -17,6 +17,7 @@ class ProviderSnapshot:
|
||||
model: str
|
||||
context_window_tokens: int
|
||||
signature: tuple[object, ...]
|
||||
generation: GenerationSettings | None = None
|
||||
|
||||
|
||||
def _resolve_model_preset(
|
||||
@@ -268,6 +269,7 @@ def build_provider_snapshot(
|
||||
model=resolved.model,
|
||||
context_window_tokens=min([resolved.context_window_tokens, *fallback_windows]),
|
||||
signature=provider_signature(config, preset=resolved),
|
||||
generation=resolved.to_generation_settings(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -39,13 +39,18 @@ class LLMRuntime:
|
||||
) -> LLMRuntime:
|
||||
"""Capture provider defaults without retaining mutable generation state."""
|
||||
generation = provider.generation
|
||||
defaults = GenerationSettings()
|
||||
return cls(
|
||||
provider=provider,
|
||||
model=model,
|
||||
generation=GenerationSettings(
|
||||
temperature=generation.temperature,
|
||||
max_tokens=generation.max_tokens,
|
||||
reasoning_effort=generation.reasoning_effort,
|
||||
temperature=getattr(generation, "temperature", defaults.temperature),
|
||||
max_tokens=getattr(generation, "max_tokens", defaults.max_tokens),
|
||||
reasoning_effort=getattr(
|
||||
generation,
|
||||
"reasoning_effort",
|
||||
defaults.reasoning_effort,
|
||||
),
|
||||
),
|
||||
context_window_tokens=context_window_tokens,
|
||||
model_preset=model_preset,
|
||||
@@ -83,6 +88,15 @@ def runtime_from_provider_snapshot(
|
||||
model_preset: str | None = None,
|
||||
) -> LLMRuntime:
|
||||
"""Convert a provider factory snapshot into the canonical runtime value."""
|
||||
if snapshot.generation is not None:
|
||||
return LLMRuntime(
|
||||
provider=snapshot.provider,
|
||||
model=snapshot.model,
|
||||
generation=snapshot.generation,
|
||||
context_window_tokens=snapshot.context_window_tokens,
|
||||
model_preset=model_preset,
|
||||
snapshot_signature=snapshot.signature,
|
||||
)
|
||||
return LLMRuntime.capture(
|
||||
snapshot.provider,
|
||||
snapshot.model,
|
||||
|
||||
@@ -90,6 +90,8 @@ def test_resolver_resolves_preset_without_mutating_selected_runtime() -> None:
|
||||
assert resolved.model_preset == "fast"
|
||||
assert resolver.runtime is initial
|
||||
assert resolver.model_preset is None
|
||||
assert initial.provider.generation == GenerationSettings(0.1, 1024, None)
|
||||
assert resolved.generation == GenerationSettings(0.5, 512, None)
|
||||
|
||||
|
||||
def test_resolver_model_override_is_derived_without_default_mutation() -> None:
|
||||
|
||||
Reference in New Issue
Block a user