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:
deusyu 2026-07-18 12:59:40 -07:00 committed by Teknium
parent 7ab95b4c9c
commit 3f84b7a163
9 changed files with 683 additions and 63 deletions

109
cli.py
View file

@ -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

View file

@ -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:

View file

@ -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"))

View file

@ -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:

View file

@ -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

View 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 == []

View 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

View file

@ -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,

View file

@ -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: