[verified] fix: scope model changes to the active profile
This commit is contained in:
@@ -14,6 +14,7 @@ import logging
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from . import protocol as proto
|
||||
@@ -682,7 +683,7 @@ async def models_snapshot(conversation_id: Optional[str] = None) -> Dict[str, An
|
||||
"""Providers + models Hermes currently exposes (credential-aware)."""
|
||||
def _collect() -> Dict[str, Any]:
|
||||
from hermes_cli.config import get_compatible_custom_providers
|
||||
from hermes_cli.model_switch import list_picker_providers
|
||||
from hermes_cli.model_switch_providers import list_picker_providers
|
||||
cfg = _load_cfg()
|
||||
model_cfg = (cfg.get("model") or {}) if isinstance(cfg, dict) else {}
|
||||
current_model = str(model_cfg.get("default", "") or "")
|
||||
@@ -751,9 +752,31 @@ async def set_model(model: str, provider: Optional[str],
|
||||
if not model:
|
||||
return {"ok": False, "code": proto.ERR_BAD_REQUEST,
|
||||
"message": "model is required"}
|
||||
home_token = None
|
||||
secret_token = None
|
||||
try:
|
||||
from agent.secret_scope import (
|
||||
build_profile_secret_scope, reset_secret_scope, set_secret_scope,
|
||||
)
|
||||
from hermes_cli.config import get_compatible_custom_providers
|
||||
from hermes_cli.model_switch import switch_model
|
||||
from hermes_constants import (
|
||||
get_hermes_home, reset_hermes_home_override,
|
||||
set_hermes_home_override,
|
||||
)
|
||||
|
||||
adapter = _current_adapter()
|
||||
profile_home = Path(
|
||||
getattr(adapter, "hermes_home", None) or get_hermes_home())
|
||||
# model.set is a WebSocket control request, not a gateway message turn,
|
||||
# so the runner has not installed this profile's ContextVars for us.
|
||||
# Bind them explicitly before load_config()/switch_model(); to_thread
|
||||
# copies the current context into its worker.
|
||||
home_token = set_hermes_home_override(str(profile_home))
|
||||
secrets = await asyncio.to_thread(
|
||||
build_profile_secret_scope, profile_home)
|
||||
secret_token = set_secret_scope(secrets)
|
||||
|
||||
cfg = _load_cfg()
|
||||
model_cfg = (cfg.get("model") or {}) if isinstance(cfg, dict) else {}
|
||||
result = await asyncio.to_thread(
|
||||
@@ -762,7 +785,7 @@ async def set_model(model: str, provider: Optional[str],
|
||||
str(model_cfg.get("provider", "openrouter") or "openrouter"),
|
||||
str(model_cfg.get("default", "") or ""),
|
||||
str(model_cfg.get("base_url", "") or ""),
|
||||
"", # current_api_key — runtime resolution handles credentials
|
||||
"", # current_api_key — scoped runtime resolution handles credentials
|
||||
False, # is_global → session-scoped when conversation given
|
||||
provider or "",
|
||||
cfg.get("providers") if isinstance(cfg, dict) else None,
|
||||
@@ -772,15 +795,23 @@ async def set_model(model: str, provider: Optional[str],
|
||||
logger.error("[pheby] switch_model failed", exc_info=True)
|
||||
return {"ok": False, "code": proto.ERR_BAD_REQUEST,
|
||||
"message": proto.safe_str(exc, 200)}
|
||||
finally:
|
||||
if secret_token is not None:
|
||||
reset_secret_scope(secret_token)
|
||||
if home_token is not None:
|
||||
reset_hermes_home_override(home_token)
|
||||
|
||||
ok = bool(getattr(result, "success", False))
|
||||
if not ok:
|
||||
error = (getattr(result, "error_message", "") or
|
||||
getattr(result, "error", ""))
|
||||
return {"ok": False, "code": proto.ERR_BAD_REQUEST,
|
||||
"message": proto.safe_str(getattr(result, "error", ""),
|
||||
300)}
|
||||
"message": proto.safe_str(error, 300)}
|
||||
|
||||
resolved_model = getattr(result, "model", model)
|
||||
resolved_provider = getattr(result, "provider", provider or "")
|
||||
resolved_model = (getattr(result, "new_model", "") or
|
||||
getattr(result, "model", "") or model)
|
||||
resolved_provider = (getattr(result, "target_provider", "") or
|
||||
getattr(result, "provider", "") or provider or "")
|
||||
override = {"model": resolved_model}
|
||||
if resolved_provider:
|
||||
override["provider"] = resolved_provider
|
||||
@@ -788,14 +819,36 @@ async def set_model(model: str, provider: Optional[str],
|
||||
store = _session_store()
|
||||
if conversation_id and store is not None:
|
||||
try:
|
||||
session_key = _session_key_for(conversation_id)
|
||||
if await asyncio.to_thread(
|
||||
store.peek_session_id,
|
||||
_session_key_for(conversation_id)) is None:
|
||||
store.peek_session_id, session_key) is None:
|
||||
return {"ok": False, "code": proto.ERR_CONVERSATION_NOT_FOUND,
|
||||
"message": "Conversation not found"}
|
||||
await asyncio.to_thread(store.set_model_override,
|
||||
_session_key_for(conversation_id),
|
||||
override)
|
||||
session_key, override)
|
||||
|
||||
# Match Hermes's native /model commit path: the persisted override
|
||||
# survives restarts, while the richer in-memory override gives the
|
||||
# very next turn its resolved endpoint/key/capabilities. Evict any
|
||||
# cached agent so it cannot answer once more with the old model.
|
||||
runner = _runner()
|
||||
runtime_overrides = getattr(
|
||||
runner, "_session_model_overrides", None)
|
||||
if isinstance(runtime_overrides, dict):
|
||||
runtime_overrides[session_key] = {
|
||||
"model": resolved_model,
|
||||
"provider": resolved_provider,
|
||||
"api_key": getattr(result, "api_key", "") or "",
|
||||
"base_url": getattr(result, "base_url", "") or "",
|
||||
"api_mode": getattr(result, "api_mode", "") or "",
|
||||
"request_overrides": dict(
|
||||
getattr(result, "request_overrides", None) or {}),
|
||||
"capabilities": dict(
|
||||
getattr(result, "runtime_capabilities", None) or {}),
|
||||
}
|
||||
evict = getattr(runner, "_evict_cached_agent", None)
|
||||
if callable(evict):
|
||||
evict(session_key)
|
||||
return {"ok": True, "model": resolved_model,
|
||||
"provider": resolved_provider, "scope": "conversation"}
|
||||
except Exception:
|
||||
|
||||
Reference in New Issue
Block a user