Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
109 changes: 97 additions & 12 deletions cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
pass

import logging
import copy
import os
import shutil
import sys
Expand Down Expand Up @@ -6163,7 +6164,7 @@ def _ensure_tirith_security(self) -> None:
except Exception:
pass


def _show_security_advisories(self):
"""Show a startup banner if any unacked security advisories match.

Expand Down Expand Up @@ -7835,6 +7836,67 @@ def _close_model_picker(self) -> None:
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,
Expand Down Expand Up @@ -8068,17 +8130,19 @@ def _handle_model_switch(self, cmd_original: str):
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
Expand All @@ -8087,19 +8151,25 @@ def _handle_model_switch(self, cmd_original: str):
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
Expand Down Expand Up @@ -8148,6 +8218,7 @@ def _handle_model_switch(self, cmd_original: str):
_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")
Expand Down Expand Up @@ -8200,6 +8271,7 @@ def _handle_model_switch(self, cmd_original: str):
# 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 = {
Expand Down Expand Up @@ -8254,8 +8326,13 @@ def _handle_model_switch(self, cmd_original: str):
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
Expand Down Expand Up @@ -8305,6 +8382,8 @@ def _handle_model_switch(self, cmd_original: str):
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)")

Expand Down Expand Up @@ -11809,6 +11888,10 @@ def run_agent():
_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,
Expand Down Expand Up @@ -11838,6 +11921,8 @@ def run_agent():
"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
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
deusyu
# PR #29923 salvage
42 changes: 40 additions & 2 deletions gateway/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -3180,6 +3181,7 @@ def __init__(self, config: Optional[GatewayConfig] = None):
# 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]] = {}
Expand Down Expand Up @@ -8178,6 +8180,7 @@ async def _session_expiry_watcher(self, interval: int = 300):
# 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)
Expand Down Expand Up @@ -11176,6 +11179,7 @@ async def _do_undo():
# 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
Expand Down Expand Up @@ -11206,6 +11210,18 @@ def _restore_moa_one_shot(self, event: "MessageEvent", quick_key: str) -> None:
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,
*,
Expand Down Expand Up @@ -11802,6 +11818,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str, run_g
# 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)
Expand Down Expand Up @@ -12912,6 +12929,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str, run_g
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)
Expand Down Expand Up @@ -17107,6 +17125,26 @@ def _apply_session_model_override(
)
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)
Expand Down Expand Up @@ -20393,7 +20431,7 @@ def _approval_notify_sync(approval_data: dict) -> None:
"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
Expand Down Expand Up @@ -20431,7 +20469,7 @@ def _approval_notify_sync(approval_data: dict) -> None:
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:
Expand Down
Loading
Loading