From cb03d2c7482fb352326210f30ac72b6a9b407e1d Mon Sep 17 00:00:00 2001 From: chengyongru Date: Fri, 10 Jul 2026 13:59:58 +0800 Subject: [PATCH] refactor(agent): snapshot generation without provider mutation --- nanobot/agent/loop.py | 12 ++++++++++-- nanobot/agent/model_presets.py | 2 +- nanobot/providers/factory.py | 4 +++- nanobot/utils/llm_runtime.py | 20 +++++++++++++++++--- tests/agent/test_model_runtime_resolver.py | 2 ++ 5 files changed, 33 insertions(+), 7 deletions(-) diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index a2dc2707..250db317 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -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 diff --git a/nanobot/agent/model_presets.py b/nanobot/agent/model_presets.py index f5468e84..d6484453 100644 --- a/nanobot/agent/model_presets.py +++ b/nanobot/agent/model_presets.py @@ -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(), ) diff --git a/nanobot/providers/factory.py b/nanobot/providers/factory.py index 756eadc5..8e2543d9 100644 --- a/nanobot/providers/factory.py +++ b/nanobot/providers/factory.py @@ -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(), ) diff --git a/nanobot/utils/llm_runtime.py b/nanobot/utils/llm_runtime.py index 945d2d06..2af2f15a 100644 --- a/nanobot/utils/llm_runtime.py +++ b/nanobot/utils/llm_runtime.py @@ -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, diff --git a/tests/agent/test_model_runtime_resolver.py b/tests/agent/test_model_runtime_resolver.py index b26c0764..4c24d916 100644 --- a/tests/agent/test_model_runtime_resolver.py +++ b/tests/agent/test_model_runtime_resolver.py @@ -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: