mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-04 08:28:36 +00:00
refactor(agent): snapshot generation without provider mutation
This commit is contained in:
parent
bd94fefd1a
commit
cb03d2c748
@ -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
|
||||||
|
|||||||
@ -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(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user