From 3f84b7a16334ede95e68100d5cdacdc33ab96d7e Mon Sep 17 00:00:00 2001 From: deusyu Date: Sat, 18 Jul 2026 12:59:40 -0700 Subject: [PATCH] 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). --- cli.py | 109 ++++++++++-- gateway/run.py | 42 ++++- gateway/slash_commands.py | 39 ++++- hermes_cli/model_switch.py | 92 +++++++--- .../gateway/test_model_switch_persistence.py | 41 +++++ tests/hermes_cli/test_cli_model_once.py | 142 +++++++++++++++ .../test_model_switch_once_flags.py | 22 +++ tests/test_tui_gateway_server.py | 165 ++++++++++++++++++ tui_gateway/server.py | 94 ++++++++-- 9 files changed, 683 insertions(+), 63 deletions(-) create mode 100644 tests/hermes_cli/test_cli_model_once.py create mode 100644 tests/hermes_cli/test_model_switch_once_flags.py diff --git a/cli.py b/cli.py index 23e3e97ebb43..f2c9773aef74 100644 --- a/cli.py +++ b/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 — switch model (persists by default) + /model --once — switch for the next turn only /model --session — switch for this session only /model --global — switch and persist (explicit) /model --provider — switch provider + model /model --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 switch model (persists)") + _cprint(" /model --once switch for the next turn only") _cprint(" /model --session switch for this session only") _cprint(" /model --provider 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 diff --git a/gateway/run.py b/gateway/run.py index ec256d590d09..aa4c3471a2b0 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -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: 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: diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index 34c786b4893e..7aa4aac73eb5 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -1454,6 +1454,7 @@ class GatewaySlashCommandsMixin: Supports: /model — interactive picker (Telegram/Discord) or text list /model — switch model (persists by default) + /model --once — switch for the next turn only /model --session — switch for this session only /model --global — switch and persist (explicit) /model --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")) diff --git a/hermes_cli/model_switch.py b/hermes_cli/model_switch.py index f83089b83853..fcb669171627 100644 --- a/hermes_cli/model_switch.py +++ b/hermes_cli/model_switch.py @@ -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 + # 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 `` 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: diff --git a/tests/gateway/test_model_switch_persistence.py b/tests/gateway/test_model_switch_persistence.py index 29adf19e6f8e..25f4a0349611 100644 --- a/tests/gateway/test_model_switch_persistence.py +++ b/tests/gateway/test_model_switch_persistence.py @@ -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 diff --git a/tests/hermes_cli/test_cli_model_once.py b/tests/hermes_cli/test_cli_model_once.py new file mode 100644 index 000000000000..45a6ffd99ea2 --- /dev/null +++ b/tests/hermes_cli/test_cli_model_once.py @@ -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 == [] diff --git a/tests/hermes_cli/test_model_switch_once_flags.py b/tests/hermes_cli/test_model_switch_once_flags.py new file mode 100644 index 000000000000..9f6b2012b92d --- /dev/null +++ b/tests/hermes_cli/test_model_switch_once_flags.py @@ -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 diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index 3c7386c8aabe..5c6dc79765c9 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -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, diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 523dada178a0..039b658ecdde 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -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: