mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-21 16:18:55 +00:00
feat: add /model --once one-turn model override (#29914)
Adds --once to /model across CLI, TUI, and gateway: switch model for the next turn only, restoring the previous model in a finally block so success, exception, and interrupt all revert. Parsing extends parse_model_flags_detailed(); resolve_persist_behavior() treats --once as a persistence opt-out; --global + --once is rejected. Salvaged from PR #29923 (image-generation lane split to #59815 per review; conflict resolution against current main by the maintainers).
This commit is contained in:
parent
7ab95b4c9c
commit
3f84b7a163
9 changed files with 683 additions and 63 deletions
109
cli.py
109
cli.py
|
|
@ -24,6 +24,7 @@ except ModuleNotFoundError:
|
|||
pass
|
||||
|
||||
import logging
|
||||
import copy
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
|
|
@ -6163,7 +6164,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
def _show_security_advisories(self):
|
||||
"""Show a startup banner if any unacked security advisories match.
|
||||
|
||||
|
|
@ -7835,6 +7836,67 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
self._restore_modal_input_snapshot()
|
||||
self._invalidate(min_interval=0.0)
|
||||
|
||||
def _snapshot_model_runtime(self) -> dict:
|
||||
"""Capture current CLI and agent model runtime for one-turn restore."""
|
||||
agent = getattr(self, "agent", None)
|
||||
return {
|
||||
"model": self.model,
|
||||
"provider": self.provider,
|
||||
"requested_provider": self.requested_provider,
|
||||
"_explicit_api_key": getattr(self, "_explicit_api_key", None),
|
||||
"_explicit_base_url": getattr(self, "_explicit_base_url", None),
|
||||
"api_key": self.api_key,
|
||||
"base_url": self.base_url,
|
||||
"api_mode": self.api_mode,
|
||||
"agent_primary_runtime": copy.deepcopy(
|
||||
getattr(agent, "_primary_runtime", None)
|
||||
) if agent is not None else None,
|
||||
}
|
||||
|
||||
def _restore_model_runtime_snapshot(self, snapshot: dict | None) -> None:
|
||||
"""Restore a model runtime captured before a one-turn override."""
|
||||
if not snapshot:
|
||||
return
|
||||
for key in (
|
||||
"model",
|
||||
"provider",
|
||||
"requested_provider",
|
||||
"_explicit_api_key",
|
||||
"_explicit_base_url",
|
||||
"api_key",
|
||||
"base_url",
|
||||
"api_mode",
|
||||
):
|
||||
if key in snapshot:
|
||||
setattr(self, key, snapshot.get(key))
|
||||
|
||||
agent = getattr(self, "agent", None)
|
||||
if agent is None:
|
||||
return
|
||||
|
||||
primary = snapshot.get("agent_primary_runtime")
|
||||
if primary and hasattr(agent, "_restore_primary_runtime"):
|
||||
try:
|
||||
agent._primary_runtime = copy.deepcopy(primary)
|
||||
agent._fallback_activated = True
|
||||
agent._rate_limited_until = 0
|
||||
if agent._restore_primary_runtime():
|
||||
return
|
||||
except Exception:
|
||||
logger.debug("CLI one-turn model restore via primary runtime failed", exc_info=True)
|
||||
|
||||
if hasattr(agent, "switch_model"):
|
||||
try:
|
||||
agent.switch_model(
|
||||
new_model=snapshot.get("model", ""),
|
||||
new_provider=snapshot.get("provider", ""),
|
||||
api_key=snapshot.get("api_key", ""),
|
||||
base_url=snapshot.get("base_url", ""),
|
||||
api_mode=snapshot.get("api_mode", ""),
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("CLI one-turn model restore failed: %s", exc)
|
||||
|
||||
@staticmethod
|
||||
def _compute_model_picker_viewport(
|
||||
selected: int,
|
||||
|
|
@ -8068,17 +8130,19 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
Supports:
|
||||
/model — show current model + usage hints
|
||||
/model <name> — switch model (persists by default)
|
||||
/model <name> --once — switch for the next turn only
|
||||
/model <name> --session — switch for this session only
|
||||
/model <name> --global — switch and persist (explicit)
|
||||
/model <name> --provider <provider> — switch provider + model
|
||||
/model --provider <provider> — switch to provider, auto-detect model
|
||||
|
||||
Persistence defaults to on (``model.persist_switch_by_default`` in
|
||||
config.yaml, default True). Use ``--session`` for a one-off switch.
|
||||
config.yaml, default True). Use ``--session`` for this CLI session or
|
||||
``--once`` for the next turn only.
|
||||
"""
|
||||
from hermes_cli.model_switch import (
|
||||
switch_model,
|
||||
parse_model_flags,
|
||||
parse_model_flags_detailed,
|
||||
resolve_persist_behavior,
|
||||
)
|
||||
from hermes_cli.providers import get_label
|
||||
|
|
@ -8087,19 +8151,25 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
parts = cmd_original.split(None, 1) # split off '/model'
|
||||
raw_args = parts[1].strip() if len(parts) > 1 else ""
|
||||
|
||||
# Parse --provider, --global, --session, and --refresh flags
|
||||
(
|
||||
model_input,
|
||||
explicit_provider,
|
||||
is_global_flag,
|
||||
force_refresh,
|
||||
is_session,
|
||||
) = parse_model_flags(raw_args)
|
||||
# Parse --provider, --global, --session, --once, and --refresh flags
|
||||
parsed_flags = parse_model_flags_detailed(raw_args)
|
||||
model_input = parsed_flags.model_input
|
||||
explicit_provider = parsed_flags.explicit_provider
|
||||
is_global_flag = parsed_flags.is_global
|
||||
force_refresh = parsed_flags.force_refresh
|
||||
is_session = parsed_flags.is_session
|
||||
one_turn = parsed_flags.is_once
|
||||
if is_global_flag and one_turn:
|
||||
_cprint(" ✗ /model --once cannot be combined with --global")
|
||||
return
|
||||
if one_turn and not model_input and not explicit_provider:
|
||||
_cprint(" ✗ /model --once requires a model or provider.")
|
||||
return
|
||||
# Resolve the effective persistence once: --session overrides the
|
||||
# config-gated default, --global forces persist, otherwise defer to
|
||||
# model.persist_switch_by_default (defaults to True so /model survives
|
||||
# across sessions).
|
||||
persist_global = resolve_persist_behavior(is_global_flag, is_session)
|
||||
persist_global = resolve_persist_behavior(is_global_flag, is_session, is_once=one_turn)
|
||||
|
||||
# --refresh: wipe the on-disk picker cache before building the
|
||||
# provider list. Forces a live re-fetch of every authed provider's
|
||||
|
|
@ -8148,6 +8218,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
_cprint(" No authenticated providers found.")
|
||||
_cprint("")
|
||||
_cprint(" /model <name> switch model (persists)")
|
||||
_cprint(" /model <name> --once switch for the next turn only")
|
||||
_cprint(" /model <name> --session switch for this session only")
|
||||
_cprint(" /model --provider <slug> switch provider")
|
||||
_cprint(" /model --refresh re-fetch live model lists")
|
||||
|
|
@ -8200,6 +8271,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
# Update requested_provider so _ensure_runtime_credentials() doesn't
|
||||
# overwrite the switch on the next turn (it re-resolves from this).
|
||||
old_model = self.model
|
||||
_one_turn_restore_snapshot = self._snapshot_model_runtime() if one_turn else None
|
||||
# Snapshot CLI-level fields before mutation so a failed in-place swap
|
||||
# rolls the whole CLI back to the old working model (#50163).
|
||||
_cli_snapshot = {
|
||||
|
|
@ -8254,8 +8326,13 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
self._pending_model_switch_note = (
|
||||
f"[Note: model was just switched from {old_model} to {result.new_model} "
|
||||
f"via {result.provider_label or result.target_provider}. "
|
||||
f"{'This override applies to the next turn only. ' if one_turn else ''}"
|
||||
f"Adjust your self-identification accordingly.]"
|
||||
)
|
||||
if one_turn:
|
||||
self._pending_one_turn_model_restore = _one_turn_restore_snapshot
|
||||
else:
|
||||
self._pending_one_turn_model_restore = None
|
||||
|
||||
# Display confirmation with full metadata
|
||||
provider_label = result.provider_label or result.target_provider
|
||||
|
|
@ -8305,6 +8382,8 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
save_config_value("model.base_url", result.base_url or None)
|
||||
save_config_value("model.api_mode", result.api_mode or None)
|
||||
_cprint(" Saved to config.yaml")
|
||||
elif one_turn:
|
||||
_cprint(" (next turn only — restores after one response)")
|
||||
else:
|
||||
_cprint(" (session only — add --global to persist)")
|
||||
|
||||
|
|
@ -11809,6 +11888,10 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
_persist_clean_user_message = (
|
||||
message if (_voice_prefix or agent_message != message) else None
|
||||
)
|
||||
_one_turn_model_restore = getattr(
|
||||
self, "_pending_one_turn_model_restore", None
|
||||
)
|
||||
self._pending_one_turn_model_restore = None
|
||||
try:
|
||||
result = self.agent.run_conversation(
|
||||
user_message=agent_message,
|
||||
|
|
@ -11838,6 +11921,8 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
|||
"error": _summary,
|
||||
}
|
||||
finally:
|
||||
if _one_turn_model_restore:
|
||||
self._restore_model_runtime_snapshot(_one_turn_model_restore)
|
||||
# Surface any credit notices queued during the turn (cold-start
|
||||
# seed / per-turn capture) now that the response is done — printing
|
||||
# at this boundary paints cleanly above the prompt instead of being
|
||||
|
|
|
|||
|
|
@ -2988,6 +2988,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
_profile_failed_platforms: Optional[Dict[str, Dict[Platform, asyncio.Task]]] = None
|
||||
_systemd_watchdog: Optional[Any] = None
|
||||
_session_model_overrides: Dict[str, Dict[str, str]] = {}
|
||||
_pending_one_turn_model_restores: Dict[str, Dict[str, Any]] = {}
|
||||
_session_reasoning_overrides: Dict[str, Dict[str, Any]] = {}
|
||||
_startup_restore_in_progress: bool = False
|
||||
# Loop-liveness heartbeat / shutdown-watchdog handles (#66892). Class-level
|
||||
|
|
@ -3180,6 +3181,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
# Per-session model overrides from /model command.
|
||||
# Key: session_key, Value: dict with model/provider/api_key/base_url/api_mode
|
||||
self._session_model_overrides: Dict[str, Dict[str, str]] = {}
|
||||
self._pending_one_turn_model_restores: Dict[str, Dict[str, Any]] = {}
|
||||
# Per-session reasoning effort overrides from /reasoning.
|
||||
# Key: session_key, Value: parsed reasoning config dict.
|
||||
self._session_reasoning_overrides: Dict[str, Dict[str, Any]] = {}
|
||||
|
|
@ -8178,6 +8180,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
# its agent from these overrides. Only true session
|
||||
# finalization, /new, and /reset clear them.)
|
||||
self._session_model_overrides.pop(key, None)
|
||||
self._pending_one_turn_model_restores.pop(key, None)
|
||||
self._set_session_reasoning_override(key, None)
|
||||
if hasattr(self, "_pending_model_notes"):
|
||||
self._pending_model_notes.pop(key, None)
|
||||
|
|
@ -11176,6 +11179,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
# Putting it in finally guarantees the revert on success, exception,
|
||||
# and interrupt alike.
|
||||
self._restore_moa_one_shot(event, _quick_key)
|
||||
self._restore_pending_one_turn_model_override(_quick_key)
|
||||
# Unconditional release covers every exit path. _release_running_agent_state
|
||||
# is idempotent (pop-on-absent is harmless) and, called without a
|
||||
# run_generation guard, always clears the slot regardless of which
|
||||
|
|
@ -11206,6 +11210,18 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
def _restore_pending_one_turn_model_override(self, session_key: str) -> None:
|
||||
"""Restore a per-session model override after ``/model --once`` runs."""
|
||||
if not session_key:
|
||||
return
|
||||
try:
|
||||
snapshot = self._pending_one_turn_model_restores.pop(session_key, None)
|
||||
if not snapshot:
|
||||
return
|
||||
self._restore_session_model_override(session_key, snapshot)
|
||||
except Exception:
|
||||
logger.debug("Failed to restore one-turn model override", exc_info=True)
|
||||
|
||||
async def _prepare_inbound_message_text(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -11802,6 +11818,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
# inherit the previous conversation's model/reasoning overrides
|
||||
# or a queued "/model switched" note.
|
||||
self._session_model_overrides.pop(session_key, None)
|
||||
self._pending_one_turn_model_restores.pop(session_key, None)
|
||||
self._set_session_reasoning_override(session_key, None)
|
||||
if hasattr(self, "_pending_model_notes"):
|
||||
self._pending_model_notes.pop(session_key, None)
|
||||
|
|
@ -12912,6 +12929,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
new_entry = await self.async_session_store.reset_session(session_key)
|
||||
self._evict_cached_agent(session_key)
|
||||
self._session_model_overrides.pop(session_key, None)
|
||||
self._pending_one_turn_model_restores.pop(session_key, None)
|
||||
self._set_session_reasoning_override(session_key, None)
|
||||
if hasattr(self, "_pending_model_notes"):
|
||||
self._pending_model_notes.pop(session_key, None)
|
||||
|
|
@ -17107,6 +17125,26 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
)
|
||||
return model, runtime_kwargs
|
||||
|
||||
def _snapshot_session_model_override(self, session_key: str) -> dict:
|
||||
"""Capture a gateway session override before a one-turn switch."""
|
||||
override = self._session_model_overrides.get(session_key)
|
||||
return {
|
||||
"had_override": override is not None,
|
||||
"override": dict(override) if override is not None else None,
|
||||
}
|
||||
|
||||
def _restore_session_model_override(self, session_key: str, snapshot: dict) -> None:
|
||||
"""Restore the session override captured before a one-turn switch."""
|
||||
if not session_key:
|
||||
return
|
||||
if snapshot.get("had_override"):
|
||||
self._session_model_overrides[session_key] = dict(
|
||||
snapshot.get("override") or {}
|
||||
)
|
||||
else:
|
||||
self._session_model_overrides.pop(session_key, None)
|
||||
self._evict_cached_agent(session_key)
|
||||
|
||||
def _is_intentional_model_switch(self, session_key: str, agent_model: str) -> bool:
|
||||
"""Return True if *agent_model* matches an active /model session override."""
|
||||
override = self._session_model_overrides.get(session_key)
|
||||
|
|
@ -20393,7 +20431,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
"model": _resolved_model,
|
||||
"context_length": _context_length,
|
||||
}
|
||||
|
||||
|
||||
# Scan tool results for MEDIA:<path> tags that need to be delivered
|
||||
# as native audio/file attachments. The TTS tool embeds MEDIA: tags
|
||||
# in its JSON response, but the model's final text reply usually
|
||||
|
|
@ -20431,7 +20469,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
|
|||
if has_voice_directive:
|
||||
unique_tags.insert(0, "[[audio_as_voice]]")
|
||||
final_response = final_response + "\n" + "\n".join(unique_tags)
|
||||
|
||||
|
||||
# Auto-generate session title after first exchange (non-blocking)
|
||||
if final_response and self._session_db:
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1454,6 +1454,7 @@ class GatewaySlashCommandsMixin:
|
|||
Supports:
|
||||
/model — interactive picker (Telegram/Discord) or text list
|
||||
/model <name> — switch model (persists by default)
|
||||
/model <name> --once — switch for the next turn only
|
||||
/model <name> --session — switch for this session only
|
||||
/model <name> --global — switch and persist (explicit)
|
||||
/model <name> --provider <provider> — switch provider + model
|
||||
|
|
@ -1462,7 +1463,7 @@ class GatewaySlashCommandsMixin:
|
|||
from gateway.run import _hermes_home, _load_gateway_config
|
||||
import yaml
|
||||
from hermes_cli.model_switch import (
|
||||
switch_model as _switch_model, parse_model_flags,
|
||||
switch_model as _switch_model, parse_model_flags_detailed,
|
||||
resolve_persist_behavior,
|
||||
list_authenticated_providers,
|
||||
list_picker_providers,
|
||||
|
|
@ -1477,15 +1478,23 @@ class GatewaySlashCommandsMixin:
|
|||
self, "_resolve_profile_home_for_source"
|
||||
)(source)
|
||||
|
||||
# Parse --provider, --global, --session, and --refresh flags
|
||||
(
|
||||
model_input,
|
||||
explicit_provider,
|
||||
# Parse --provider, --global, --session, --once, and --refresh flags
|
||||
parsed_flags = parse_model_flags_detailed(raw_args)
|
||||
model_input = parsed_flags.model_input
|
||||
explicit_provider = parsed_flags.explicit_provider
|
||||
is_global_flag = parsed_flags.is_global
|
||||
force_refresh = parsed_flags.force_refresh
|
||||
is_session = parsed_flags.is_session
|
||||
one_turn = parsed_flags.is_once
|
||||
if is_global_flag and one_turn:
|
||||
return "❌ /model --once cannot be combined with --global"
|
||||
if one_turn and not model_input and not explicit_provider:
|
||||
return "❌ /model --once requires a model or provider."
|
||||
persist_global = resolve_persist_behavior(
|
||||
is_global_flag,
|
||||
force_refresh,
|
||||
is_session,
|
||||
) = parse_model_flags(raw_args)
|
||||
persist_global = resolve_persist_behavior(is_global_flag, is_session)
|
||||
is_once=one_turn,
|
||||
)
|
||||
|
||||
# --refresh: bust the disk cache so the picker shows live data.
|
||||
if force_refresh:
|
||||
|
|
@ -1528,6 +1537,9 @@ class GatewaySlashCommandsMixin:
|
|||
source = await asyncio.to_thread(self._normalize_source_for_session_key, source)
|
||||
session_key = self._session_key_for_source(source)
|
||||
override = self._session_model_overrides.get(session_key, {})
|
||||
restore_snapshot = (
|
||||
self._snapshot_session_model_override(session_key) if one_turn else None
|
||||
)
|
||||
if override:
|
||||
current_model = override.get("model", current_model)
|
||||
current_provider = override.get("provider", current_provider)
|
||||
|
|
@ -1948,6 +1960,7 @@ class GatewaySlashCommandsMixin:
|
|||
self._pending_model_notes[session_key] = (
|
||||
f"[Note: model was just switched from {current_model} to {result.new_model} "
|
||||
f"via {result.provider_label or result.target_provider}. "
|
||||
f"{'This override applies to the next turn only. ' if one_turn else ''}"
|
||||
f"Adjust your self-identification accordingly.]"
|
||||
)
|
||||
|
||||
|
|
@ -1959,6 +1972,14 @@ class GatewaySlashCommandsMixin:
|
|||
"base_url": result.base_url,
|
||||
"api_mode": result.api_mode,
|
||||
}
|
||||
if one_turn:
|
||||
if not hasattr(self, "_pending_one_turn_model_restores"):
|
||||
self._pending_one_turn_model_restores = {}
|
||||
self._pending_one_turn_model_restores[session_key] = (
|
||||
restore_snapshot or {"had_override": False, "override": None}
|
||||
)
|
||||
elif hasattr(self, "_pending_one_turn_model_restores"):
|
||||
self._pending_one_turn_model_restores.pop(session_key, None)
|
||||
|
||||
# Write-through the non-secret parts (model/provider/base_url) to
|
||||
# the session store so the override survives a gateway restart.
|
||||
|
|
@ -2070,6 +2091,8 @@ class GatewaySlashCommandsMixin:
|
|||
|
||||
if persist_global:
|
||||
lines.append(t("gateway.model.saved_global"))
|
||||
elif one_turn:
|
||||
lines.append(" (next turn only — restores after one response)")
|
||||
else:
|
||||
lines.append(t("gateway.model.session_only_hint"))
|
||||
|
||||
|
|
|
|||
|
|
@ -349,14 +349,28 @@ class ModelSwitchResult:
|
|||
capabilities: Optional[ModelCapabilities] = None
|
||||
model_info: Optional[ModelInfo] = None
|
||||
is_global: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelFlagParseResult:
|
||||
"""Parsed flags for a /model command."""
|
||||
|
||||
model_input: str
|
||||
explicit_provider: str = ""
|
||||
is_global: bool = False
|
||||
force_refresh: bool = False
|
||||
is_session: bool = False
|
||||
is_once: bool = False
|
||||
# ---------------------------------------------------------------------------
|
||||
# Flag parsing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def parse_model_flags(raw_args: str) -> tuple[str, str, bool, bool, bool]:
|
||||
"""Parse --provider, --global, --session, and --refresh flags from /model command args.
|
||||
def parse_model_flags_detailed(raw_args: str) -> ModelFlagParseResult:
|
||||
"""Parse flags from /model command args.
|
||||
|
||||
Returns ``(model_input, explicit_provider, is_global, force_refresh, is_session)``.
|
||||
Returns a :class:`ModelFlagParseResult`. ``--once`` is intentionally
|
||||
parsed here but interpreted by each caller because each frontend has its
|
||||
own live-session restore hook.
|
||||
|
||||
``is_global`` and ``is_session`` are independent flag presences; the
|
||||
*effective* persistence decision is resolved by
|
||||
|
|
@ -368,6 +382,7 @@ def parse_model_flags(raw_args: str) -> tuple[str, str, bool, bool, bool]:
|
|||
"sonnet" -> ("sonnet", "", False, False, False)
|
||||
"sonnet --global" -> ("sonnet", "", True, False, False)
|
||||
"sonnet --session" -> ("sonnet", "", False, False, True)
|
||||
"sonnet --once" -> is_once=True
|
||||
"sonnet --provider anthropic" -> ("sonnet", "anthropic", False, False, False)
|
||||
"--provider my-ollama" -> ("", "my-ollama", False, False, False)
|
||||
"--refresh" -> ("", "", False, True, False)
|
||||
|
|
@ -377,33 +392,32 @@ def parse_model_flags(raw_args: str) -> tuple[str, str, bool, bool, bool]:
|
|||
explicit_provider = ""
|
||||
force_refresh = False
|
||||
is_session = False
|
||||
is_once = False
|
||||
|
||||
# Normalize Unicode dashes (Telegram/iOS auto-converts -- to em/en dash)
|
||||
# A single Unicode dash before a flag keyword becomes "--"
|
||||
import re as _re
|
||||
raw_args = _re.sub(r'[\u2012\u2013\u2014\u2015](provider|global|session|refresh)', r'--\1', raw_args)
|
||||
raw_args = _re.sub(r'[\u2012\u2013\u2014\u2015](provider|global|session|refresh|once)', r'--\1', raw_args)
|
||||
|
||||
# Extract --global
|
||||
if "--global" in raw_args:
|
||||
is_global = True
|
||||
raw_args = raw_args.replace("--global", "").strip()
|
||||
|
||||
# Extract --session (explicit session-only; overrides the persist default)
|
||||
if "--session" in raw_args:
|
||||
is_session = True
|
||||
raw_args = raw_args.replace("--session", "").strip()
|
||||
|
||||
# Extract --refresh (bust the model picker disk cache before listing)
|
||||
if "--refresh" in raw_args:
|
||||
force_refresh = True
|
||||
raw_args = raw_args.replace("--refresh", "").strip()
|
||||
|
||||
# Extract --provider <name>
|
||||
# Keep this hand-rolled because model IDs may contain colons/slashes and
|
||||
# the historical parser did not require shell quoting.
|
||||
parts = raw_args.split()
|
||||
i = 0
|
||||
filtered: list[str] = []
|
||||
while i < len(parts):
|
||||
if parts[i] == "--provider" and i + 1 < len(parts):
|
||||
if parts[i] == "--global":
|
||||
is_global = True
|
||||
i += 1
|
||||
elif parts[i] == "--session":
|
||||
is_session = True
|
||||
i += 1
|
||||
elif parts[i] == "--refresh":
|
||||
force_refresh = True
|
||||
i += 1
|
||||
elif parts[i] == "--once":
|
||||
is_once = True
|
||||
i += 1
|
||||
elif parts[i] == "--provider" and i + 1 < len(parts):
|
||||
explicit_provider = parts[i + 1]
|
||||
i += 2
|
||||
else:
|
||||
|
|
@ -411,17 +425,41 @@ def parse_model_flags(raw_args: str) -> tuple[str, str, bool, bool, bool]:
|
|||
i += 1
|
||||
|
||||
model_input = " ".join(filtered).strip()
|
||||
return (model_input, explicit_provider, is_global, force_refresh, is_session)
|
||||
return ModelFlagParseResult(
|
||||
model_input=model_input,
|
||||
explicit_provider=explicit_provider,
|
||||
is_global=is_global,
|
||||
force_refresh=force_refresh,
|
||||
is_session=is_session,
|
||||
is_once=is_once,
|
||||
)
|
||||
|
||||
|
||||
def resolve_persist_behavior(is_global: bool, is_session: bool) -> bool:
|
||||
def parse_model_flags(raw_args: str) -> tuple[str, str, bool, bool, bool]:
|
||||
"""Parse legacy /model flags and return the historical 5-tuple.
|
||||
|
||||
New call sites that care about ``--once`` should use
|
||||
:func:`parse_model_flags_detailed`.
|
||||
"""
|
||||
parsed = parse_model_flags_detailed(raw_args)
|
||||
return (
|
||||
parsed.model_input,
|
||||
parsed.explicit_provider,
|
||||
parsed.is_global,
|
||||
parsed.force_refresh,
|
||||
parsed.is_session,
|
||||
)
|
||||
|
||||
|
||||
def resolve_persist_behavior(is_global: bool, is_session: bool, is_once: bool = False) -> bool:
|
||||
"""Decide whether a ``/model`` switch should persist to ``config.yaml``.
|
||||
|
||||
Resolution order:
|
||||
|
||||
1. ``--session`` explicitly opts out → ``False`` (this session only).
|
||||
2. ``--global`` explicitly opts in → ``True``.
|
||||
3. Otherwise defer to ``model.persist_switch_by_default`` in
|
||||
1. ``--once`` explicitly opts out → ``False`` (next turn only).
|
||||
2. ``--session`` explicitly opts out → ``False`` (this session only).
|
||||
3. ``--global`` explicitly opts in → ``True``.
|
||||
4. Otherwise defer to ``model.persist_switch_by_default`` in
|
||||
``config.yaml`` (defaults to ``True``, so a plain ``/model <name>``
|
||||
survives across sessions — the behavior users expect).
|
||||
|
||||
|
|
@ -429,6 +467,8 @@ def resolve_persist_behavior(is_global: bool, is_session: bool) -> bool:
|
|||
flat string rather than a dict, in which case the built-in default
|
||||
(``True``) applies.
|
||||
"""
|
||||
if is_once:
|
||||
return False
|
||||
if is_session:
|
||||
return False
|
||||
if is_global:
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ def _make_runner():
|
|||
runner._voice_mode = {}
|
||||
runner.hooks = SimpleNamespace(emit=AsyncMock(), loaded_hooks=False)
|
||||
runner._session_model_overrides = {}
|
||||
runner._pending_one_turn_model_restores = {}
|
||||
runner._pending_model_notes = {}
|
||||
runner._background_tasks = set()
|
||||
runner._running_agents = {}
|
||||
|
|
@ -242,3 +243,43 @@ class TestIsIntentionalModelSwitch:
|
|||
}
|
||||
|
||||
assert runner._is_intentional_model_switch(sk, "gpt-5.4") is False
|
||||
|
||||
|
||||
class TestOneTurnModelOverrideRestore:
|
||||
"""Verify gateway one-turn overrides restore previous session state."""
|
||||
|
||||
def test_restores_previous_override(self):
|
||||
runner = _make_runner()
|
||||
sk = build_session_key(_make_source())
|
||||
previous = {
|
||||
"model": "old/model",
|
||||
"provider": "openrouter",
|
||||
"api_key": "old-key",
|
||||
"base_url": "https://openrouter.ai/api/v1",
|
||||
"api_mode": "chat_completions",
|
||||
}
|
||||
runner._session_model_overrides[sk] = previous
|
||||
|
||||
snapshot = runner._snapshot_session_model_override(sk)
|
||||
runner._session_model_overrides[sk] = {
|
||||
"model": "temp/model",
|
||||
"provider": "anthropic",
|
||||
}
|
||||
|
||||
runner._restore_session_model_override(sk, snapshot)
|
||||
|
||||
assert runner._session_model_overrides[sk] == previous
|
||||
|
||||
def test_restores_absent_override_by_clearing(self):
|
||||
runner = _make_runner()
|
||||
sk = build_session_key(_make_source())
|
||||
|
||||
snapshot = runner._snapshot_session_model_override(sk)
|
||||
runner._session_model_overrides[sk] = {
|
||||
"model": "temp/model",
|
||||
"provider": "anthropic",
|
||||
}
|
||||
|
||||
runner._restore_session_model_override(sk, snapshot)
|
||||
|
||||
assert sk not in runner._session_model_overrides
|
||||
|
|
|
|||
142
tests/hermes_cli/test_cli_model_once.py
Normal file
142
tests/hermes_cli/test_cli_model_once.py
Normal file
|
|
@ -0,0 +1,142 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
from hermes_cli.model_switch import ModelSwitchResult
|
||||
|
||||
|
||||
class _FakeAgent:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.model = "old/model"
|
||||
self.provider = "openrouter"
|
||||
|
||||
def switch_model(self, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
self.model = kwargs["new_model"]
|
||||
self.provider = kwargs["new_provider"]
|
||||
|
||||
|
||||
class _StubCLI:
|
||||
model = "old/model"
|
||||
provider = "openrouter"
|
||||
requested_provider = "openrouter"
|
||||
api_key = "sk-old"
|
||||
_explicit_api_key = "sk-old"
|
||||
base_url = "https://openrouter.ai/api/v1"
|
||||
_explicit_base_url = "https://openrouter.ai/api/v1"
|
||||
api_mode = "chat_completions"
|
||||
agent = None
|
||||
_pending_model_switch_note = None
|
||||
_pending_one_turn_model_restore = None
|
||||
|
||||
def _confirm_expensive_model_switch(self, result):
|
||||
return True
|
||||
|
||||
|
||||
def test_cli_model_once_records_restore_and_does_not_persist(monkeypatch):
|
||||
import cli as cli_mod
|
||||
|
||||
stub = _StubCLI()
|
||||
stub.agent = _FakeAgent()
|
||||
stub._snapshot_model_runtime = cli_mod.HermesCLI._snapshot_model_runtime.__get__(stub)
|
||||
printed = []
|
||||
|
||||
monkeypatch.setattr(cli_mod, "_cprint", lambda s, *a, **k: printed.append(str(s)))
|
||||
monkeypatch.setattr(cli_mod, "save_config_value", lambda *a, **k: (_ for _ in ()).throw(AssertionError("should not persist")))
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.inventory.load_picker_context",
|
||||
lambda: SimpleNamespace(
|
||||
user_providers=None,
|
||||
custom_providers=None,
|
||||
with_overrides=lambda **_: SimpleNamespace(user_providers=None, custom_providers=None),
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.model_switch.switch_model",
|
||||
lambda **_: ModelSwitchResult(
|
||||
success=True,
|
||||
new_model="claude-sonnet-4.6",
|
||||
target_provider="anthropic",
|
||||
api_key="sk-ant",
|
||||
base_url="https://api.anthropic.com",
|
||||
api_mode="anthropic_messages",
|
||||
provider_label="Anthropic",
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr("hermes_cli.model_switch.resolve_display_context_length", lambda *a, **k: None)
|
||||
|
||||
cli_mod.HermesCLI._handle_model_switch(
|
||||
stub,
|
||||
"/model claude-sonnet-4.6 --provider anthropic --once",
|
||||
)
|
||||
|
||||
assert stub.model == "claude-sonnet-4.6"
|
||||
assert stub.provider == "anthropic"
|
||||
assert stub.agent.calls[-1]["new_model"] == "claude-sonnet-4.6"
|
||||
assert stub._pending_one_turn_model_restore["model"] == "old/model"
|
||||
assert "next turn only" in printed[-1]
|
||||
|
||||
|
||||
def test_cli_restore_model_runtime_snapshot_restores_agent():
|
||||
import cli as cli_mod
|
||||
|
||||
stub = _StubCLI()
|
||||
stub.agent = _FakeAgent()
|
||||
snapshot = {
|
||||
"model": "old/model",
|
||||
"provider": "openrouter",
|
||||
"requested_provider": "openrouter",
|
||||
"api_key": "sk-old",
|
||||
"explicit_api_key": "sk-old",
|
||||
"base_url": "https://openrouter.ai/api/v1",
|
||||
"explicit_base_url": "https://openrouter.ai/api/v1",
|
||||
"api_mode": "chat_completions",
|
||||
}
|
||||
|
||||
cli_mod.HermesCLI._restore_model_runtime_snapshot(stub, snapshot)
|
||||
|
||||
assert stub.model == "old/model"
|
||||
assert stub.provider == "openrouter"
|
||||
assert stub.agent.calls[-1]["new_model"] == "old/model"
|
||||
|
||||
|
||||
def test_cli_restore_model_runtime_prefers_primary_runtime():
|
||||
import cli as cli_mod
|
||||
|
||||
class Agent(_FakeAgent):
|
||||
_primary_runtime = None
|
||||
_rate_limited_until = 123
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.model = "temp/model"
|
||||
self.provider = "anthropic"
|
||||
|
||||
def _restore_primary_runtime(self):
|
||||
self.model = self._primary_runtime["model"]
|
||||
self.provider = self._primary_runtime["provider"]
|
||||
return True
|
||||
|
||||
stub = _StubCLI()
|
||||
stub.agent = Agent()
|
||||
snapshot = {
|
||||
"model": "old/model",
|
||||
"provider": "openrouter",
|
||||
"requested_provider": "openrouter",
|
||||
"api_key": "sk-old",
|
||||
"explicit_api_key": "sk-old",
|
||||
"base_url": "",
|
||||
"explicit_base_url": "",
|
||||
"api_mode": "chat_completions",
|
||||
"agent_primary_runtime": {
|
||||
"model": "old/model",
|
||||
"provider": "openrouter",
|
||||
},
|
||||
}
|
||||
|
||||
cli_mod.HermesCLI._restore_model_runtime_snapshot(stub, snapshot)
|
||||
|
||||
assert stub.agent.model == "old/model"
|
||||
assert stub.agent.provider == "openrouter"
|
||||
assert stub.agent.calls == []
|
||||
22
tests/hermes_cli/test_model_switch_once_flags.py
Normal file
22
tests/hermes_cli/test_model_switch_once_flags.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
from hermes_cli.model_switch import parse_model_flags, parse_model_flags_detailed
|
||||
|
||||
|
||||
def test_parse_model_flags_detailed_supports_once():
|
||||
parsed = parse_model_flags_detailed("sonnet --provider anthropic --once")
|
||||
|
||||
assert parsed.model_input == "sonnet"
|
||||
assert parsed.explicit_provider == "anthropic"
|
||||
assert parsed.is_global is False
|
||||
assert parsed.force_refresh is False
|
||||
assert parsed.is_session is False
|
||||
assert parsed.is_once is True
|
||||
|
||||
|
||||
def test_parse_model_flags_legacy_wrapper_strips_once():
|
||||
model_input, provider, is_global, force_refresh, is_session = parse_model_flags("sonnet --once")
|
||||
|
||||
assert model_input == "sonnet"
|
||||
assert provider == ""
|
||||
assert is_global is False
|
||||
assert force_refresh is False
|
||||
assert is_session is False
|
||||
|
|
@ -5182,6 +5182,170 @@ def test_config_set_model_switches_agent_without_touching_env(monkeypatch):
|
|||
server._sessions.clear()
|
||||
|
||||
|
||||
def test_config_set_model_once_keeps_env_and_records_restore(monkeypatch):
|
||||
class Agent:
|
||||
model = "old/model"
|
||||
provider = "openrouter"
|
||||
base_url = "https://openrouter.ai/api/v1"
|
||||
api_key = "sk-old"
|
||||
api_mode = "chat_completions"
|
||||
|
||||
def switch_model(self, **kwargs):
|
||||
self.model = kwargs["new_model"]
|
||||
self.provider = kwargs["new_provider"]
|
||||
self.api_key = kwargs["api_key"]
|
||||
self.base_url = kwargs["base_url"]
|
||||
self.api_mode = kwargs["api_mode"]
|
||||
|
||||
result = types.SimpleNamespace(
|
||||
success=True,
|
||||
new_model="claude-sonnet-4.6",
|
||||
target_provider="anthropic",
|
||||
api_key="sk-ant",
|
||||
base_url="https://api.anthropic.com",
|
||||
api_mode="anthropic_messages",
|
||||
warning_message="",
|
||||
)
|
||||
seen = {}
|
||||
agent = Agent()
|
||||
session = _session(agent=agent)
|
||||
server._sessions["sid"] = session
|
||||
monkeypatch.setenv("HERMES_INFERENCE_PROVIDER", "openrouter")
|
||||
monkeypatch.setenv("HERMES_MODEL", "old/model")
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.model_switch.switch_model",
|
||||
lambda **kwargs: seen.update(kwargs) or result,
|
||||
)
|
||||
monkeypatch.setattr(server, "_restart_slash_worker", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(server, "_emit", lambda *args, **kwargs: None)
|
||||
|
||||
try:
|
||||
resp = server.handle_request(
|
||||
{
|
||||
"id": "1",
|
||||
"method": "config.set",
|
||||
"params": {
|
||||
"session_id": "sid",
|
||||
"key": "model",
|
||||
"value": "claude-sonnet-4.6 --provider anthropic --once",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert resp["result"]["scope"] == "once"
|
||||
assert seen["is_global"] is False
|
||||
assert agent.model == "claude-sonnet-4.6"
|
||||
assert session["one_turn_model_restore"]["model"] == "old/model"
|
||||
assert os.environ["HERMES_INFERENCE_PROVIDER"] == "openrouter"
|
||||
assert os.environ["HERMES_MODEL"] == "old/model"
|
||||
finally:
|
||||
server._sessions.clear()
|
||||
|
||||
|
||||
def test_config_set_model_once_requires_live_session(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.model_switch.switch_model",
|
||||
lambda **_: (_ for _ in ()).throw(AssertionError("switch should not run")),
|
||||
)
|
||||
|
||||
resp = server.handle_request(
|
||||
{
|
||||
"id": "1",
|
||||
"method": "config.set",
|
||||
"params": {
|
||||
"key": "model",
|
||||
"value": "claude-sonnet-4.6 --provider anthropic --once",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert resp["error"]["code"] == 5001
|
||||
assert "/model --once requires a live session" in resp["error"]["message"]
|
||||
|
||||
|
||||
def test_config_set_model_session_switch_clears_pending_once_restore(monkeypatch):
|
||||
class Agent:
|
||||
model = "temp/model"
|
||||
provider = "anthropic"
|
||||
base_url = "https://api.anthropic.com"
|
||||
api_key = "sk-temp"
|
||||
api_mode = "anthropic_messages"
|
||||
|
||||
def switch_model(self, **kwargs):
|
||||
self.model = kwargs["new_model"]
|
||||
self.provider = kwargs["new_provider"]
|
||||
self.api_key = kwargs["api_key"]
|
||||
self.base_url = kwargs["base_url"]
|
||||
self.api_mode = kwargs["api_mode"]
|
||||
|
||||
result = types.SimpleNamespace(
|
||||
success=True,
|
||||
new_model="new/model",
|
||||
target_provider="openrouter",
|
||||
api_key="sk-new",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
api_mode="chat_completions",
|
||||
warning_message="",
|
||||
)
|
||||
session = _session(agent=Agent())
|
||||
session["one_turn_model_restore"] = {"model": "old/model"}
|
||||
server._sessions["sid"] = session
|
||||
monkeypatch.setattr("hermes_cli.model_switch.switch_model", lambda **_kwargs: result)
|
||||
monkeypatch.setattr(server, "_restart_slash_worker", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(server, "_emit", lambda *args, **kwargs: None)
|
||||
|
||||
try:
|
||||
resp = server.handle_request(
|
||||
{
|
||||
"id": "1",
|
||||
"method": "config.set",
|
||||
"params": {
|
||||
"session_id": "sid",
|
||||
"key": "model",
|
||||
"value": "new/model --provider openrouter --session",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert resp["result"]["scope"] == "session"
|
||||
assert "one_turn_model_restore" not in session
|
||||
finally:
|
||||
server._sessions.clear()
|
||||
|
||||
|
||||
def test_restore_agent_model_runtime_falls_back_to_switch_model():
|
||||
class Agent:
|
||||
model = "temp/model"
|
||||
provider = "anthropic"
|
||||
base_url = "https://api.anthropic.com"
|
||||
api_key = "sk-temp"
|
||||
api_mode = "anthropic_messages"
|
||||
|
||||
def switch_model(self, **kwargs):
|
||||
self.model = kwargs["new_model"]
|
||||
self.provider = kwargs["new_provider"]
|
||||
self.api_key = kwargs["api_key"]
|
||||
self.base_url = kwargs["base_url"]
|
||||
self.api_mode = kwargs["api_mode"]
|
||||
|
||||
agent = Agent()
|
||||
|
||||
server._restore_agent_model_runtime(
|
||||
agent,
|
||||
{
|
||||
"model": "old/model",
|
||||
"provider": "openrouter",
|
||||
"api_key": "sk-old",
|
||||
"base_url": "https://openrouter.ai/api/v1",
|
||||
"api_mode": "chat_completions",
|
||||
},
|
||||
)
|
||||
|
||||
assert agent.model == "old/model"
|
||||
assert agent.provider == "openrouter"
|
||||
assert agent.base_url == "https://openrouter.ai/api/v1"
|
||||
|
||||
|
||||
def test_config_set_personality_rejects_unknown_name(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
server,
|
||||
|
|
@ -8342,6 +8506,7 @@ def test_browser_manage_connect_defaults_to_loopback(monkeypatch):
|
|||
|
||||
def test_browser_manage_connect_default_local_reports_launch_hint(monkeypatch):
|
||||
monkeypatch.delenv("BROWSER_CDP_URL", raising=False)
|
||||
monkeypatch.setattr("platform.system", lambda: "Linux")
|
||||
emitted: list[tuple[str, dict]] = []
|
||||
monkeypatch.setattr(
|
||||
server,
|
||||
|
|
|
|||
|
|
@ -3133,6 +3133,42 @@ def _persist_model_switch(result) -> None:
|
|||
save_config_value("model.base_url", None)
|
||||
|
||||
|
||||
def _snapshot_agent_model_runtime(agent) -> dict:
|
||||
"""Capture the current agent model runtime for a one-turn restore."""
|
||||
return {
|
||||
"model": getattr(agent, "model", ""),
|
||||
"provider": getattr(agent, "provider", ""),
|
||||
"api_key": getattr(agent, "api_key", ""),
|
||||
"base_url": getattr(agent, "base_url", ""),
|
||||
"api_mode": getattr(agent, "api_mode", ""),
|
||||
"primary_runtime": copy.deepcopy(getattr(agent, "_primary_runtime", None)),
|
||||
}
|
||||
|
||||
|
||||
def _restore_agent_model_runtime(agent, snapshot: dict | None) -> None:
|
||||
"""Restore an agent model runtime captured before a one-turn override."""
|
||||
if not snapshot or agent is None:
|
||||
return
|
||||
primary = snapshot.get("primary_runtime")
|
||||
if primary and hasattr(agent, "_restore_primary_runtime"):
|
||||
try:
|
||||
agent._primary_runtime = copy.deepcopy(primary)
|
||||
agent._fallback_activated = True
|
||||
agent._rate_limited_until = 0
|
||||
if agent._restore_primary_runtime():
|
||||
return
|
||||
except Exception:
|
||||
logger.debug("TUI one-turn model restore via primary runtime failed", exc_info=True)
|
||||
if hasattr(agent, "switch_model"):
|
||||
agent.switch_model(
|
||||
new_model=snapshot.get("model", ""),
|
||||
new_provider=snapshot.get("provider", ""),
|
||||
api_key=snapshot.get("api_key", ""),
|
||||
base_url=snapshot.get("base_url", ""),
|
||||
api_mode=snapshot.get("api_mode", ""),
|
||||
)
|
||||
|
||||
|
||||
def _apply_model_switch(
|
||||
sid: str,
|
||||
session: dict,
|
||||
|
|
@ -3140,34 +3176,44 @@ def _apply_model_switch(
|
|||
*,
|
||||
confirm_expensive_model: bool = False,
|
||||
pin_session_override: bool = True,
|
||||
parsed_flags: tuple[str, str, bool, bool, bool] | None = None,
|
||||
parsed_flags: Any | None = None,
|
||||
persist_override: bool | None = None,
|
||||
) -> dict:
|
||||
from hermes_cli.model_switch import (
|
||||
parse_model_flags,
|
||||
parse_model_flags_detailed,
|
||||
resolve_persist_behavior,
|
||||
switch_model,
|
||||
)
|
||||
from hermes_cli.runtime_provider import resolve_runtime_provider
|
||||
|
||||
if parsed_flags is None:
|
||||
parsed_flags = parse_model_flags(raw_input)
|
||||
(
|
||||
model_input,
|
||||
explicit_provider,
|
||||
is_global_flag,
|
||||
_force_refresh,
|
||||
is_session,
|
||||
) = parsed_flags
|
||||
parsed_flags = parse_model_flags_detailed(raw_input)
|
||||
if hasattr(parsed_flags, "model_input"):
|
||||
model_input = parsed_flags.model_input
|
||||
explicit_provider = parsed_flags.explicit_provider
|
||||
is_global_flag = parsed_flags.is_global
|
||||
is_session = parsed_flags.is_session
|
||||
one_turn = parsed_flags.is_once
|
||||
else:
|
||||
model_input, explicit_provider, is_global_flag, _force_refresh, is_session = parsed_flags
|
||||
one_turn = False
|
||||
if is_global_flag and one_turn:
|
||||
raise ValueError("/model --once cannot be combined with --global")
|
||||
persist_global = (
|
||||
persist_override
|
||||
if persist_override is not None
|
||||
else resolve_persist_behavior(is_global_flag, is_session)
|
||||
else resolve_persist_behavior(
|
||||
is_global_flag,
|
||||
is_session,
|
||||
is_once=one_turn,
|
||||
)
|
||||
)
|
||||
if not model_input:
|
||||
raise ValueError("model value required")
|
||||
|
||||
agent = session.get("agent")
|
||||
if one_turn and not agent:
|
||||
raise ValueError("/model --once requires a live session")
|
||||
if agent:
|
||||
current_provider = getattr(agent, "provider", "") or ""
|
||||
current_model = getattr(agent, "model", "") or ""
|
||||
|
|
@ -3197,6 +3243,7 @@ def _apply_model_switch(
|
|||
# endpoints (e.g. "ollama-launch") and validate against saved model lists.
|
||||
user_provs = None
|
||||
custom_provs = None
|
||||
cfg = None
|
||||
try:
|
||||
from hermes_cli.config import get_compatible_custom_providers, load_config
|
||||
|
||||
|
|
@ -3220,6 +3267,8 @@ def _apply_model_switch(
|
|||
if not result.success:
|
||||
raise ValueError(result.error_message or "model switch failed")
|
||||
|
||||
restore_snapshot = _snapshot_agent_model_runtime(agent) if (one_turn and agent) else None
|
||||
|
||||
if agent:
|
||||
try:
|
||||
from hermes_cli.context_switch_guard import merge_preflight_compression_warning
|
||||
|
|
@ -3292,6 +3341,10 @@ def _apply_model_switch(
|
|||
session, model=result.new_model, provider=result.target_provider
|
||||
)
|
||||
_emit("session.info", sid, _session_info(agent, session))
|
||||
if one_turn:
|
||||
session["one_turn_model_restore"] = restore_snapshot
|
||||
else:
|
||||
session.pop("one_turn_model_restore", None)
|
||||
|
||||
# Record the switch as a PER-SESSION override so a later rebuild of THIS
|
||||
# session (e.g. /new via _reset_session_agent, or resume) re-derives the
|
||||
|
|
@ -3306,7 +3359,7 @@ def _apply_model_switch(
|
|||
# contamination bug). agent.switch_model() above already mutated the right
|
||||
# agent in place; the override dict makes that choice survive a rebuild
|
||||
# without touching shared process state.
|
||||
if pin_session_override and isinstance(session, dict):
|
||||
if pin_session_override and isinstance(session, dict) and not one_turn:
|
||||
session["model_override"] = {
|
||||
"model": result.new_model,
|
||||
"provider": result.target_provider,
|
||||
|
|
@ -3320,6 +3373,7 @@ def _apply_model_switch(
|
|||
"value": result.new_model,
|
||||
"warning": result.warning_message or "",
|
||||
"confirm_required": False,
|
||||
"scope": "once" if one_turn else ("global" if persist_global else "session"),
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -9781,6 +9835,7 @@ def _run_prompt_submit(rid, sid: str, session: dict, text: Any) -> None:
|
|||
session_tokens = []
|
||||
home_token = None # per-turn HERMES_HOME override for a resumed remote profile
|
||||
goal_followup = None # set by the post-turn goal hook below
|
||||
one_turn_restore = session.pop("one_turn_model_restore", None)
|
||||
try:
|
||||
from tools.approval import (
|
||||
reset_current_session_key,
|
||||
|
|
@ -10190,6 +10245,14 @@ def _run_prompt_submit(rid, sid: str, session: dict, text: Any) -> None:
|
|||
)
|
||||
_emit("error", sid, {"message": str(e)})
|
||||
finally:
|
||||
if one_turn_restore:
|
||||
try:
|
||||
_restore_agent_model_runtime(agent, one_turn_restore)
|
||||
_restart_slash_worker(sid, session)
|
||||
_persist_live_session_runtime(session)
|
||||
_persist_live_session_system_prompt(session)
|
||||
except Exception:
|
||||
logger.debug("TUI one-turn model restore failed", exc_info=True)
|
||||
try:
|
||||
if approval_token is not None:
|
||||
reset_current_session_key(approval_token)
|
||||
|
|
@ -11154,10 +11217,10 @@ def _(rid, params: dict) -> dict:
|
|||
4009,
|
||||
"session busy — /interrupt the current turn before switching models",
|
||||
)
|
||||
from hermes_cli.model_switch import parse_model_flags
|
||||
from hermes_cli.model_switch import parse_model_flags_detailed
|
||||
|
||||
parsed_flags = parse_model_flags(value)
|
||||
_model_input, explicit_provider, _persist_global, _force_refresh, _is_session = parsed_flags
|
||||
parsed_flags = parse_model_flags_detailed(value)
|
||||
explicit_provider = parsed_flags.explicit_provider
|
||||
if session.get("agent") is None and not explicit_provider.strip():
|
||||
session_id = params.get("session_id", "")
|
||||
_start_agent_build(session_id, session)
|
||||
|
|
@ -11192,6 +11255,7 @@ def _(rid, params: dict) -> dict:
|
|||
"warning": result["warning"],
|
||||
"confirm_required": result.get("confirm_required", False),
|
||||
"confirm_message": result.get("confirm_message", ""),
|
||||
"scope": result.get("scope", "session"),
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue