fix(models): preserve preset rename compatibility
This commit is contained in:
@@ -95,6 +95,7 @@ class ChannelManager:
|
||||
cron_service: CronService | None = None,
|
||||
local_trigger_store: LocalTriggerStore | None = None,
|
||||
webui_runtime_model_name: Callable[[], str | None] | None = None,
|
||||
webui_refresh_runtime_config: Callable[[], None] | None = None,
|
||||
webui_cron_pending_job_ids: Callable[[str], set[str]] | None = None,
|
||||
webui_local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
||||
webui_static_dist: bool = True,
|
||||
@@ -116,6 +117,7 @@ class ChannelManager:
|
||||
self._cron_service = cron_service
|
||||
self._local_trigger_store = local_trigger_store
|
||||
self._webui_runtime_model_name = webui_runtime_model_name
|
||||
self._webui_refresh_runtime_config = webui_refresh_runtime_config
|
||||
self._webui_cron_pending_job_ids = webui_cron_pending_job_ids
|
||||
self._webui_local_trigger_pending_ids = webui_local_trigger_pending_ids
|
||||
self._webui_static_dist = webui_static_dist
|
||||
@@ -183,6 +185,7 @@ class ChannelManager:
|
||||
config_path=self._config_path,
|
||||
disabled_skills=set(self.config.agents.defaults.disabled_skills),
|
||||
runtime_model_name=self._webui_runtime_model_name,
|
||||
refresh_runtime_config=self._webui_refresh_runtime_config,
|
||||
runtime_surface=self._webui_runtime_surface,
|
||||
runtime_capabilities_overrides=self._webui_runtime_capabilities,
|
||||
cron_service=self._cron_service,
|
||||
|
||||
@@ -657,6 +657,9 @@ def _run_gateway(
|
||||
def _webui_runtime_model_name() -> str | None:
|
||||
return agent.model.strip() or None
|
||||
|
||||
def _webui_refresh_runtime_config() -> None:
|
||||
agent.invalidate_runtime_config()
|
||||
|
||||
def _webui_skill_state_action(disabled_skills: set[str]) -> None:
|
||||
config.agents.defaults.disabled_skills = sorted(disabled_skills)
|
||||
agent.context.skills.disabled_skills = set(disabled_skills)
|
||||
@@ -671,6 +674,7 @@ def _run_gateway(
|
||||
cron_service=cron,
|
||||
local_trigger_store=trigger_store,
|
||||
webui_runtime_model_name=_webui_runtime_model_name,
|
||||
webui_refresh_runtime_config=_webui_refresh_runtime_config,
|
||||
webui_cron_pending_job_ids=agent.pending_cron_job_ids_for_session,
|
||||
webui_local_trigger_pending_ids=agent.pending_local_trigger_ids_for_session,
|
||||
webui_static_dist=webui_static_dist,
|
||||
|
||||
@@ -462,18 +462,10 @@ class Config(BaseSettings):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_model_preset(self) -> "Config":
|
||||
names_by_case: dict[str, str] = {}
|
||||
for preset_name in self.model_presets:
|
||||
if preset_name != preset_name.strip() or not preset_name.isprintable():
|
||||
raise ValueError(f"invalid model_preset name {preset_name!r}")
|
||||
normalized = preset_name.casefold()
|
||||
if normalized in names_by_case:
|
||||
raise ValueError(
|
||||
"model_preset names must be unique ignoring case: "
|
||||
f"{names_by_case[normalized]!r} and {preset_name!r}"
|
||||
)
|
||||
names_by_case[normalized] = preset_name
|
||||
if "default" in names_by_case:
|
||||
# Keep persisted names accepted by previous releases loadable. New
|
||||
# names are normalized and checked case-insensitively at mutation
|
||||
# boundaries, where conflicts can be reported without breaking startup.
|
||||
if "default" in self.model_presets:
|
||||
raise ValueError("model_preset name 'default' is reserved for agents.defaults")
|
||||
name = self.agents.defaults.model_preset
|
||||
if name and name != "default" and name not in self.model_presets:
|
||||
|
||||
@@ -26,6 +26,7 @@ from nanobot.runtime_context import (
|
||||
RUNTIME_CONTEXT_HISTORY_META,
|
||||
public_history_message,
|
||||
)
|
||||
from nanobot.session.model_selection import SESSION_MODEL_PRESET_METADATA_KEY
|
||||
from nanobot.utils.helpers import (
|
||||
content_with_media_breadcrumbs,
|
||||
ensure_dir,
|
||||
@@ -1663,6 +1664,47 @@ class SessionManager:
|
||||
self._store.save(session, fsync=fsync)
|
||||
self._remember(session)
|
||||
|
||||
def rename_model_preset(self, old_name: str, new_name: str) -> int:
|
||||
"""Rename a session-scoped model preset across durable and live sessions."""
|
||||
if old_name == new_name:
|
||||
return 0
|
||||
|
||||
cached = dict(self._overflow_cache.items())
|
||||
cached.update(self._cache)
|
||||
keys = set(cached)
|
||||
keys.update(item["key"] for item in self._store.list_sessions())
|
||||
|
||||
changed: list[Session] = []
|
||||
try:
|
||||
for key in sorted(keys):
|
||||
session = cached.get(key) or self._load(key)
|
||||
if (
|
||||
session is None
|
||||
or session.metadata.get(SESSION_MODEL_PRESET_METADATA_KEY) != old_name
|
||||
):
|
||||
continue
|
||||
session.metadata[SESSION_MODEL_PRESET_METADATA_KEY] = new_name
|
||||
changed.append(session)
|
||||
if session.policy.persist:
|
||||
self.save(session, fsync=True)
|
||||
else:
|
||||
self._remember(session)
|
||||
except BaseException:
|
||||
for session in reversed(changed):
|
||||
session.metadata[SESSION_MODEL_PRESET_METADATA_KEY] = old_name
|
||||
try:
|
||||
if session.policy.persist:
|
||||
self.save(session, fsync=True)
|
||||
else:
|
||||
self._remember(session)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to roll back model preset rename for session {}",
|
||||
session.key,
|
||||
)
|
||||
raise
|
||||
return len(changed)
|
||||
|
||||
def flush_all(self) -> int:
|
||||
"""Re-save every cached session with fsync for durable shutdown.
|
||||
|
||||
|
||||
@@ -56,6 +56,7 @@ def build_gateway_services(
|
||||
default_restrict_to_workspace: bool,
|
||||
config_path: Path | None = None,
|
||||
runtime_model_name: Callable[[], str | None] | None,
|
||||
refresh_runtime_config: Callable[[], None] | None = None,
|
||||
runtime_surface: str,
|
||||
runtime_capabilities_overrides: dict[str, Any] | None,
|
||||
disabled_skills: set[str] | None = None,
|
||||
@@ -70,7 +71,15 @@ def build_gateway_services(
|
||||
skill_state_action: Callable[[set[str]], None] | None = None,
|
||||
logger: Any = default_logger,
|
||||
) -> GatewayServices:
|
||||
settings = WebUISettingsServices.create(config_path or get_config_path())
|
||||
settings = WebUISettingsServices.create(
|
||||
config_path or get_config_path(),
|
||||
rename_model_preset=(
|
||||
session_manager.rename_model_preset
|
||||
if session_manager is not None
|
||||
else None
|
||||
),
|
||||
refresh_runtime_config=refresh_runtime_config,
|
||||
)
|
||||
tokens = GatewayTokenStore()
|
||||
ingress = DEFAULT_WEBUI_INGRESS_POLICY
|
||||
minimum_frame_bytes = ingress.minimum_full_policy_frame_bytes()
|
||||
|
||||
@@ -7,7 +7,7 @@ domains; this module preserves the established Python and HTTP-facing seams.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from collections.abc import Callable, Iterable
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
|
||||
@@ -234,14 +234,32 @@ def update_model_configuration(
|
||||
query: QueryParams,
|
||||
*,
|
||||
config_path: Path | None = None,
|
||||
rename_model_preset: Callable[[str, str], int] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
config = _load_settings_config(config_path)
|
||||
if models.update_model_configuration(
|
||||
names_before = set(config.model_presets)
|
||||
changed = models.update_model_configuration(
|
||||
config,
|
||||
query,
|
||||
oauth_status=_oauth_provider_status,
|
||||
):
|
||||
_save_settings_config(config, config_path)
|
||||
)
|
||||
if changed:
|
||||
removed = names_before - set(config.model_presets)
|
||||
added = set(config.model_presets) - names_before
|
||||
rename = (
|
||||
(next(iter(removed)), next(iter(added)))
|
||||
if len(removed) == len(added) == 1
|
||||
else None
|
||||
)
|
||||
if rename is not None and rename_model_preset is not None:
|
||||
rename_model_preset(*rename)
|
||||
try:
|
||||
_save_settings_config(config, config_path)
|
||||
except BaseException:
|
||||
rename_model_preset(rename[1], rename[0])
|
||||
raise
|
||||
else:
|
||||
_save_settings_config(config, config_path)
|
||||
return settings_payload(config_path=config_path)
|
||||
|
||||
|
||||
|
||||
@@ -1661,9 +1661,18 @@ class ModelSettingsHandler:
|
||||
restart_section="runtime",
|
||||
)
|
||||
|
||||
if action == "model-update":
|
||||
payload = self.settings.mutate(
|
||||
operations.update_model,
|
||||
request.query,
|
||||
rename_model_preset=self.settings.rename_model_preset,
|
||||
)
|
||||
if self.settings.refresh_runtime_config is not None:
|
||||
self.settings.refresh_runtime_config()
|
||||
return SettingsRouteResult.success(payload, decorate_restart=True)
|
||||
|
||||
mutation = {
|
||||
"model-create": operations.create_model,
|
||||
"model-update": operations.update_model,
|
||||
"model-delete": operations.delete_model,
|
||||
"models-migrate": operations.migrate_models,
|
||||
"call-order-update": operations.update_call_order,
|
||||
|
||||
@@ -113,12 +113,22 @@ class WebUISettingsServices:
|
||||
|
||||
config: WebUISettingsConfig
|
||||
oauth_flows: WebUIOAuthFlowRegistry
|
||||
rename_model_preset: Callable[[str, str], int] | None = None
|
||||
refresh_runtime_config: Callable[[], None] | None = None
|
||||
|
||||
@classmethod
|
||||
def create(cls, config_path: Path) -> WebUISettingsServices:
|
||||
def create(
|
||||
cls,
|
||||
config_path: Path,
|
||||
*,
|
||||
rename_model_preset: Callable[[str, str], int] | None = None,
|
||||
refresh_runtime_config: Callable[[], None] | None = None,
|
||||
) -> WebUISettingsServices:
|
||||
return cls(
|
||||
config=WebUISettingsConfig(config_path),
|
||||
oauth_flows=WebUIOAuthFlowRegistry(),
|
||||
rename_model_preset=rename_model_preset,
|
||||
refresh_runtime_config=refresh_runtime_config,
|
||||
)
|
||||
|
||||
def read(
|
||||
|
||||
Reference in New Issue
Block a user