mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-07 21:08:34 +03:00
fix(agent): rebuild provider on preset switch
This commit is contained in:
@@ -452,7 +452,9 @@ class AgentLoop:
|
|||||||
context_window_tokens = extra.pop("context_window_tokens", None) or resolved.context_window_tokens
|
context_window_tokens = extra.pop("context_window_tokens", None) or resolved.context_window_tokens
|
||||||
provider_snapshot_loader = extra.pop("provider_snapshot_loader", None)
|
provider_snapshot_loader = extra.pop("provider_snapshot_loader", None)
|
||||||
preset_snapshot_loader = extra.pop("preset_snapshot_loader", None)
|
preset_snapshot_loader = extra.pop("preset_snapshot_loader", None)
|
||||||
if preset_snapshot_loader is None and explicit_provider is None:
|
if preset_snapshot_loader is None and (
|
||||||
|
explicit_provider is None or provider_snapshot_loader is not None
|
||||||
|
):
|
||||||
preset_snapshot_loader = preset_helpers.make_preset_snapshot_loader(
|
preset_snapshot_loader = preset_helpers.make_preset_snapshot_loader(
|
||||||
config,
|
config,
|
||||||
provider_snapshot_loader,
|
provider_snapshot_loader,
|
||||||
|
|||||||
@@ -321,3 +321,44 @@ def test_settings_context_window_refreshes_runtime_state(
|
|||||||
assert payload["requires_restart"] is False
|
assert payload["requires_restart"] is False
|
||||||
assert loop.context_window_tokens == 262_144
|
assert loop.context_window_tokens == 262_144
|
||||||
assert loop.llm_runtime().context_window_tokens == 262_144
|
assert loop.llm_runtime().context_window_tokens == 262_144
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_config_uses_snapshot_loader_for_preset_switch_with_injected_provider(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
config = Config()
|
||||||
|
config.agents.defaults.workspace = str(tmp_path / "workspace")
|
||||||
|
config.agents.defaults.model_preset = "default"
|
||||||
|
config.model_presets = {
|
||||||
|
"default": ModelPresetConfig(model="base-model", provider="openai"),
|
||||||
|
"fast": ModelPresetConfig(model="fast-model", provider="deepseek"),
|
||||||
|
}
|
||||||
|
initial_provider = _provider("base-model")
|
||||||
|
default_provider = _provider("base-model")
|
||||||
|
fast_provider = _provider("fast-model")
|
||||||
|
loaded_presets: list[str | None] = []
|
||||||
|
|
||||||
|
def loader(*, preset_name: str | None = None) -> ProviderSnapshot:
|
||||||
|
loaded_presets.append(preset_name)
|
||||||
|
provider = fast_provider if preset_name == "fast" else default_provider
|
||||||
|
model = "fast-model" if preset_name == "fast" else "base-model"
|
||||||
|
return ProviderSnapshot(
|
||||||
|
provider=provider,
|
||||||
|
model=model,
|
||||||
|
context_window_tokens=32_768,
|
||||||
|
signature=(model, preset_name),
|
||||||
|
)
|
||||||
|
|
||||||
|
loop = AgentLoop.from_config(
|
||||||
|
config,
|
||||||
|
provider=initial_provider,
|
||||||
|
provider_snapshot_loader=loader,
|
||||||
|
)
|
||||||
|
|
||||||
|
runtime = loop.runtime_resolver.resolve_preset("fast")
|
||||||
|
|
||||||
|
assert runtime.provider is fast_provider
|
||||||
|
assert runtime.model == "fast-model"
|
||||||
|
assert runtime.model_preset == "fast"
|
||||||
|
assert loop.provider is default_provider
|
||||||
|
assert loaded_presets == ["default", "fast"]
|
||||||
|
|||||||
Reference in New Issue
Block a user