refactor(agent): snapshot generation without provider mutation

This commit is contained in:
chengyongru 2026-07-10 13:59:58 +08:00 committed by Xubin Ren
parent bd94fefd1a
commit cb03d2c748
5 changed files with 33 additions and 7 deletions

View File

@ -173,9 +173,15 @@ class AgentLoop:
return self.tools.tool_names return self.tools.tool_names
def llm_runtime(self) -> LLMRuntime: 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() 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" _RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
_PENDING_USER_TURN_KEY = "pending_user_turn" _PENDING_USER_TURN_KEY = "pending_user_turn"
@ -444,6 +450,8 @@ class AgentLoop:
provider = snapshot.provider provider = snapshot.provider
model = snapshot.model model = snapshot.model
context_window_tokens = snapshot.context_window_tokens context_window_tokens = snapshot.context_window_tokens
if snapshot.generation is not None:
provider.generation = snapshot.generation
old_model = self.model old_model = self.model
self.provider = provider self.provider = provider
self.model = model self.model = model

View File

@ -34,12 +34,12 @@ def build_static_preset_snapshot(
name: str, name: str,
preset: ModelPresetConfig, preset: ModelPresetConfig,
) -> ProviderSnapshot: ) -> ProviderSnapshot:
provider.generation = preset.to_generation_settings()
return ProviderSnapshot( return ProviderSnapshot(
provider=provider, provider=provider,
model=preset.model, model=preset.model,
context_window_tokens=preset.context_window_tokens, context_window_tokens=preset.context_window_tokens,
signature=("model_preset", name, preset.model_dump_json()), signature=("model_preset", name, preset.model_dump_json()),
generation=preset.to_generation_settings(),
) )

View File

@ -6,7 +6,7 @@ from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from nanobot.config.schema import Config, InlineFallbackConfig, ModelPresetConfig, ProviderConfig 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.fallback_provider import FallbackProvider
from nanobot.providers.registry import ProviderSpec, create_dynamic_spec, find_by_name from nanobot.providers.registry import ProviderSpec, create_dynamic_spec, find_by_name
@ -17,6 +17,7 @@ class ProviderSnapshot:
model: str model: str
context_window_tokens: int context_window_tokens: int
signature: tuple[object, ...] signature: tuple[object, ...]
generation: GenerationSettings | None = None
def _resolve_model_preset( def _resolve_model_preset(
@ -268,6 +269,7 @@ def build_provider_snapshot(
model=resolved.model, model=resolved.model,
context_window_tokens=min([resolved.context_window_tokens, *fallback_windows]), context_window_tokens=min([resolved.context_window_tokens, *fallback_windows]),
signature=provider_signature(config, preset=resolved), signature=provider_signature(config, preset=resolved),
generation=resolved.to_generation_settings(),
) )

View File

@ -39,13 +39,18 @@ class LLMRuntime:
) -> LLMRuntime: ) -> LLMRuntime:
"""Capture provider defaults without retaining mutable generation state.""" """Capture provider defaults without retaining mutable generation state."""
generation = provider.generation generation = provider.generation
defaults = GenerationSettings()
return cls( return cls(
provider=provider, provider=provider,
model=model, model=model,
generation=GenerationSettings( generation=GenerationSettings(
temperature=generation.temperature, temperature=getattr(generation, "temperature", defaults.temperature),
max_tokens=generation.max_tokens, max_tokens=getattr(generation, "max_tokens", defaults.max_tokens),
reasoning_effort=generation.reasoning_effort, reasoning_effort=getattr(
generation,
"reasoning_effort",
defaults.reasoning_effort,
),
), ),
context_window_tokens=context_window_tokens, context_window_tokens=context_window_tokens,
model_preset=model_preset, model_preset=model_preset,
@ -83,6 +88,15 @@ def runtime_from_provider_snapshot(
model_preset: str | None = None, model_preset: str | None = None,
) -> LLMRuntime: ) -> LLMRuntime:
"""Convert a provider factory snapshot into the canonical runtime value.""" """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( return LLMRuntime.capture(
snapshot.provider, snapshot.provider,
snapshot.model, snapshot.model,

View File

@ -90,6 +90,8 @@ def test_resolver_resolves_preset_without_mutating_selected_runtime() -> None:
assert resolved.model_preset == "fast" assert resolved.model_preset == "fast"
assert resolver.runtime is initial assert resolver.runtime is initial
assert resolver.model_preset is None 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: def test_resolver_model_override_is_derived_without_default_mutation() -> None: