diff --git a/gateway/slash_commands_model.py b/gateway/slash_commands_model.py index 3e28c999c9..212090853b 100644 --- a/gateway/slash_commands_model.py +++ b/gateway/slash_commands_model.py @@ -2,15 +2,17 @@ /model, /codex-runtime, /reasoning, /fast, /personality. Split out of ``gateway/slash_commands.py``; bound onto ``GatewayRunner`` through -``GatewaySlashCommandsMixin``. Origin internals are imported lazily (``from gateway.slash_commands -import ...``) inside the bodies to avoid the import cycle. +``GatewaySlashCommandsMixin``. Origin internals are imported lazily inside the bodies to avoid +the import cycle. """ from __future__ import annotations -import logging import asyncio -from typing import Optional +import contextlib +import dataclasses +import logging +from typing import Any, Optional from agent.i18n import t from gateway.platforms.base import MessageEvent @@ -20,13 +22,26 @@ from utils import base_url_host_matches # Log-record parity with gateway/run.py and the origin module. logger = logging.getLogger("gateway.run") +# /fast argument -> (service tier, persisted value, i18n label key; None = value.upper()). +_FAST_SELECTIONS = { + "fast": ("priority", "fast", "gateway.fast.label_fast"), + "on": ("priority", "fast", "gateway.fast.label_fast"), + "normal": (None, "normal", "gateway.fast.label_normal"), + "off": (None, "normal", "gateway.fast.label_normal"), + "auto": ("auto", "auto", None), + "cold": ("cold", "cold", None), +} + +# /reasoning display-toggle arguments -> show_reasoning value. +_REASONING_DISPLAY_TOGGLES = {"show": True, "on": True, "hide": False, "off": False} + def _model_switch_skew_guard() -> Optional[str]: """Refuse a model switch when the gateway is running stale code. A long-lived gateway keeps boot-time modules in memory; if the checkout changed underneath it, - a first-time lazy import on a new code path can crash on a stale cached dependency. Detect the - drift and ask for a restart. Scoped to model switching only (the highest-risk trigger). + a first-time lazy import on a new code path can crash on a stale cached dependency. Scoped to + model switching only (the highest-risk trigger). """ from gateway.code_skew import detect_code_skew @@ -47,12 +62,10 @@ def _model_switch_skew_guard() -> Optional[str]: async def _persist_model_switch_to_config(result, config_path) -> None: """Write-through a resolved /model switch to ``config_path`` (model.default/provider/base_url). - Write-back round-trip: raw read is correct (merged defaults must not be persisted back to the - user's file). A scalar/None ``model:`` is coerced into a dict first — otherwise - ``cfg.setdefault("model", {})`` returns the existing scalar and the next assignment raises - ``TypeError``. Named providers re-resolve base_url/api_mode fresh, so leftovers are cleared - unconditionally; custom providers have no registry entry to re-derive from, so they need an - explicit set-or-clear (a lone ``if base_url:`` leaves stale values). + Raw read is correct (merged defaults must not be persisted back). A scalar/None ``model:`` is + coerced into a dict first or the assignments below raise ``TypeError``. Named providers + re-resolve base_url/api_mode fresh, so leftovers are cleared; custom providers have no registry + entry to re-derive from, so they need an explicit set-or-clear. """ from hermes_cli.config import read_user_config_raw, save_config @@ -85,46 +98,75 @@ async def _persist_model_switch_to_config(result, config_path) -> None: model_cfg["base_url"] = result.base_url elif is_custom_target: model_cfg.pop("base_url", None) - if is_custom_target: - if result.api_mode: - model_cfg["api_mode"] = result.api_mode - else: - model_cfg.pop("api_mode", None) - else: + if not is_custom_target: clear_model_endpoint_credentials(model_cfg, clear_base_url=True) + elif result.api_mode: + model_cfg["api_mode"] = result.api_mode + else: + model_cfg.pop("api_mode", None) save_config(cfg) -def _read_model_command_config(config_path): - """Current (model, provider, base_url, user_providers, custom_providers, excluded) for /model. +@dataclasses.dataclass +class _ModelSwitchContext: + """Everything a /model switch needs beyond the target: current route + persistence policy.""" - Fail-open: any config read error yields the defaults (``provider="openrouter"``). - """ - from gateway.run import _load_gateway_config + session_key: str + source: Any + config_path: Any + persist_global: bool + one_turn: bool = False + restore_snapshot: Optional[dict] = None + current_model: str = "" + current_provider: str = "openrouter" + current_base_url: str = "" + current_api_key: str = "" + user_provs: Any = None + custom_provs: Any = None + excluded_provs: list = dataclasses.field(default_factory=list) - current_model, current_provider, current_base_url = "", "openrouter", "" - user_provs = custom_provs = None - excluded_provs: list = [] - try: - cfg = _load_gateway_config(config_path=config_path) - if cfg: + def read_config(self) -> None: + """Fill the current route from ``config_path``; fail-open to the defaults.""" + from gateway.run import _load_gateway_config + + try: + cfg = _load_gateway_config(config_path=self.config_path) + if not cfg: + return model_cfg = cfg.get("model", {}) if isinstance(model_cfg, dict): - current_model = model_cfg.get("default", "") - current_provider = model_cfg.get("provider", current_provider) - current_base_url = model_cfg.get("base_url", "") - user_provs = cfg.get("providers") + self.current_model = model_cfg.get("default", "") + self.current_provider = model_cfg.get("provider", self.current_provider) + self.current_base_url = model_cfg.get("base_url", "") + self.user_provs = cfg.get("providers") try: from hermes_cli.config import get_compatible_custom_providers - custom_provs = get_compatible_custom_providers(cfg) + self.custom_provs = get_compatible_custom_providers(cfg) except Exception: - custom_provs = cfg.get("custom_providers") - _excl = cfg.get("model_catalog", {}).get("excluded_providers") - if isinstance(_excl, list): - excluded_provs = _excl - except Exception: - pass - return current_model, current_provider, current_base_url, user_provs, custom_provs, excluded_provs + self.custom_provs = cfg.get("custom_providers") + excl = cfg.get("model_catalog", {}).get("excluded_providers") + if isinstance(excl, list): + self.excluded_provs = excl + except Exception: + pass + + def apply_override(self, override: dict) -> None: + """A session /model override supersedes the configured route.""" + if override: + self.current_model = override.get("model", self.current_model) + self.current_provider = override.get("provider", self.current_provider) + self.current_base_url = override.get("base_url", self.current_base_url) + self.current_api_key = override.get("api_key", self.current_api_key) + + def listing_kwargs(self) -> dict: + return dict( + current_provider=self.current_provider, + current_base_url=self.current_base_url, + current_model=self.current_model, + user_providers=self.user_provs, + custom_providers=self.custom_provs, + excluded_providers=self.excluded_provs, + ) def _model_provider_listing_lines(providers) -> list[str]: @@ -135,7 +177,8 @@ def _model_provider_listing_lines(providers) -> list[str]: lines.append(f"**{p['name']}** `--provider {p['slug']}`{tag}:") if p["models"]: model_strs = ", ".join(f"`{m}`" for m in p["models"]) - extra = t("gateway.model.more_models_suffix", count=p["total_models"] - len(p["models"])) if p["total_models"] > len(p["models"]) else "" + hidden = p["total_models"] - len(p["models"]) + extra = t("gateway.model.more_models_suffix", count=hidden) if hidden > 0 else "" lines.append(f" {model_strs}{extra}") elif p.get("api_url"): lines.append(f" `{p['api_url']}`") @@ -143,144 +186,133 @@ def _model_provider_listing_lines(providers) -> list[str]: return lines +async def _configured_display_context() -> tuple[dict, Optional[int]]: + """(model config section, its ``context_length`` as int) for the switch confirmation; fail-open.""" + from gateway.run import _load_gateway_config + + model_cfg: dict = {} + config_ctx = None + with contextlib.suppress(Exception): + model_cfg = _load_gateway_config().get("model", {}) + if isinstance(model_cfg, dict): + raw = model_cfg.get("context_length") + if raw is not None: + config_ctx = int(raw) + if not isinstance(model_cfg, dict): + model_cfg = {} + return model_cfg, config_ctx + + class GatewayModelCommandsMixin: """Model-route slash commands (/model, /codex-runtime, /reasoning, /fast, /personality).""" + # ----------------------------------------------------------------- /model + async def _perform_model_switch( - self, - switch_model, - *, - raw_input: str, - explicit_provider, - session_key: str, - source, - current_model, - current_provider, - current_base_url, - current_api_key, - persist_global: bool, - user_provs, - custom_provs, + self, ctx: _ModelSwitchContext, raw_input: str, explicit_provider, source ): """Resolve a /model switch off-loop. Returns ``(result, None)`` or ``(None, error_text)``.""" from gateway.run import _load_gateway_config + from hermes_cli.model_switch import switch_model skew_error = _model_switch_skew_guard() if skew_error: return None, skew_error - # Offload the switch off the event loop — switch_model() can fall through to a synchronous - # models.dev HTTP fetch (requests.get, 15s timeout) on a cold/expired cache, which freezes - # the gateway otherwise. + # Off the event loop: switch_model() can fall through to a synchronous models.dev HTTP + # fetch (15s timeout) on a cold/expired cache, which would freeze the gateway. result = await asyncio.to_thread( switch_model, raw_input=raw_input, - current_provider=current_provider, - current_model=current_model, - current_base_url=current_base_url, - current_api_key=current_api_key, - is_global=persist_global, + current_provider=ctx.current_provider, + current_model=ctx.current_model, + current_base_url=ctx.current_base_url, + current_api_key=ctx.current_api_key, + is_global=ctx.persist_global, explicit_provider=explicit_provider, - user_providers=user_provs, - custom_providers=custom_provs, + user_providers=ctx.user_provs, + custom_providers=ctx.custom_provs, ) if not result.success: return None, t("gateway.model.error_prefix", error=result.error_message) try: from hermes_cli.context_switch_guard import enrich_model_switch_warnings_for_gateway - # Offload: merge_preflight_compression_warning() calls the sync - # resolve_display_context_length() provider probe ladder — must not run on the loop. + # Off-loop: merge_preflight_compression_warning() runs the sync provider probe ladder. await asyncio.to_thread( enrich_model_switch_warnings_for_gateway, result, self, - session_key=session_key, + session_key=ctx.session_key, source=source, - custom_providers=custom_provs, + custom_providers=ctx.custom_provs, load_gateway_config=_load_gateway_config, ) except Exception as exc: logger.debug("preflight-compression switch warning failed: %s", exc) return result, None - async def _commit_model_switch( - self, - result, - *, - session_key: str, - source, - current_model, - current_base_url, - current_api_key, - custom_provs, - persist_global: bool, - config_path, - one_turn: bool = False, - restore_snapshot=None, - picker: bool = False, - ) -> str: - """Apply a resolved switch (cached agent, session, config) and build the confirmation. + def _switch_cached_agent_model(self, result, ctx: _ModelSwitchContext, picker: bool) -> Optional[str]: + """In-place swap on the cached agent; returns the error reply when it failed. - Shared by the typed ``/model `` path and the picker callback (``picker=True``). + The agent rolls back to the OLD working model/client and re-raises. Abort the commit (DB + persist, session override, cache eviction, config write) so a failed switch is a no-op — + otherwise the next message rebuilds a broken agent from the override. """ - from gateway.run import _load_gateway_config - from hermes_cli.model_switch import format_model_for_display, resolve_display_context_length_async + cached_agent = self._cached_agent_for(ctx.session_key) + if cached_agent is None: + return None + try: + cached_agent.switch_model( + new_model=result.new_model, + new_provider=result.target_provider, + api_key=result.api_key, + base_url=result.base_url, + api_mode=result.api_mode, + capabilities=getattr(result, "runtime_capabilities", None), + ) + except Exception as exc: + logger.warning( + "%s model switch failed for cached agent: %s", "Picker" if picker else "In-place", exc + ) + return t( + "gateway.model.error_prefix", + error=f"Model switch to {result.new_model} failed ({exc}); staying on {ctx.current_model}.", + ) + return None - # If there's a cached agent, update it in-place - cached_agent = self._cached_agent_for(session_key) - if cached_agent is not None: - try: - cached_agent.switch_model( - new_model=result.new_model, - new_provider=result.target_provider, - api_key=result.api_key, - base_url=result.base_url, - api_mode=result.api_mode, - capabilities=getattr(result, "runtime_capabilities", None), - ) - except Exception as exc: - # In-place swap rolled back to the OLD working model/client and re-raised. Abort the - # commit (DB persist, session override, cache eviction, config write) so a failed switch - # is a no-op — otherwise the next message rebuilds a broken agent from the override. - logger.warning( - "%s model switch failed for cached agent: %s", "Picker" if picker else "In-place", exc - ) - return t( - "gateway.model.error_prefix", - error=f"Model switch to {result.new_model} failed ({exc}); staying on {current_model}.", - ) + async def _record_model_switch( + self, result, ctx: _ModelSwitchContext, *, source, one_turn: bool, picker: bool + ) -> None: + """Persist a committed switch: session DB, next-turn note, override map, config write-through.""" + from hermes_cli.model_switch import format_model_for_display - # Persist the new model to the session DB so the dashboard shows the updated model. + # Session DB so the dashboard shows the updated model. _sess_db = getattr(self, "_session_db", None) if _sess_db is not None: try: _sess_entry = await self.async_session_store.get_or_create_session(source) - # Typed path: if this session was auto-reset, consume the flag so the next regular - # message's cleanup does not wipe the model override just stored below. + # Typed path: consume an auto-reset flag so the next regular message's cleanup does + # not wipe the override stored below. if not picker and getattr(_sess_entry, "was_auto_reset", False): _sess_entry.was_auto_reset = False await _sess_db.update_session_model( - _sess_entry.session_id, result.new_model, - provider=result.target_provider, + _sess_entry.session_id, result.new_model, provider=result.target_provider, ) except Exception as exc: logger.debug("Failed to persist model switch to DB: %s", exc) - # Store a note to prepend to the next user message so the model knows about the switch - # (avoids system messages mid-history). Display form strips opaque Palantir RID - # prefixes; the override map below keeps the full ID for the wire. + # Note prepended to the next user message (avoids system messages mid-history). Display + # form strips opaque Palantir RID prefixes; the override map keeps the full ID for the wire. if not hasattr(self, "_pending_model_notes"): self._pending_model_notes = {} - self._pending_model_notes[session_key] = ( - f"[Note: model was just switched from {format_model_for_display(current_model)} to " + self._pending_model_notes[ctx.session_key] = ( + f"[Note: model was just switched from {format_model_for_display(ctx.current_model)} to " f"{format_model_for_display(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.]" ) - - # Store session override so next agent creation uses the new model - self._session_model_overrides[session_key] = { + self._session_model_overrides[ctx.session_key] = { "model": result.new_model, "provider": result.target_provider, "api_key": result.api_key, @@ -292,86 +324,72 @@ class GatewayModelCommandsMixin: 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} + self._pending_one_turn_model_restores[ctx.session_key] = ( + ctx.restore_snapshot or {"had_override": False, "override": None} ) elif not picker and hasattr(self, "_pending_one_turn_model_restores"): - self._pending_one_turn_model_restores.pop(session_key, None) + self._pending_one_turn_model_restores.pop(ctx.session_key, None) - # Write-through the non-secret parts (model/provider/base_url) so the override survives a - # restart; api_key/api_mode are never persisted (re-resolved on rehydration). /model --once is - # EXCLUDED: a one-turn override must not outlive a restart; the pre-once value stays persisted. + # Non-secret write-through so the override survives a restart (api_key/api_mode are + # re-resolved on rehydration). /model --once is EXCLUDED: a one-turn override must not + # outlive a restart. if not one_turn: try: await self.async_session_store.set_model_override( - session_key, self._session_model_overrides[session_key] + ctx.session_key, self._session_model_overrides[ctx.session_key] ) except Exception: logger.debug("Failed to persist session model override", exc_info=True) - # Evict cached agent so the next turn creates a fresh agent from the - # override rather than relying on cache signature mismatch detection. - self._evict_cached_agent(session_key) + # Evict so the next turn builds a fresh agent from the override. + self._evict_cached_agent(ctx.session_key) - # Persist to config (default) unless --session opted out - if persist_global: + if ctx.persist_global: try: - await _persist_model_switch_to_config(result, config_path) + await _persist_model_switch_to_config(result, ctx.config_path) except Exception as e: logger.warning("Failed to persist model switch: %s", e) - # Build confirmation message with full metadata. Display form shortens opaque Palantir - # IDs (ri.language-model-service..*) to their trailing slug. - provider_label = result.provider_label or result.target_provider - lines = [t("gateway.model.switched", model=format_model_for_display(result.new_model))] - lines.append(t("gateway.model.provider_label", provider=provider_label)) + async def _model_switch_confirmation( + self, result, ctx: _ModelSwitchContext, *, one_turn: bool, picker: bool + ) -> str: + """Confirmation text with full metadata (display form shortens opaque Palantir IDs).""" + from hermes_cli.model_switch import format_model_for_display, resolve_display_context_length_async - # Context: always resolve via the provider-aware chain so Codex OAuth, - # Copilot, and Nous-enforced caps win over the raw models.dev entry. + lines = [ + t("gateway.model.switched", model=format_model_for_display(result.new_model)), + t("gateway.model.provider_label", provider=result.provider_label or result.target_provider), + ] + # Context: the provider-aware chain so Codex OAuth, Copilot and Nous-enforced caps win over + # the raw models.dev entry. mi = result.model_info - _sw_config_ctx = None - _sw_model_cfg = {} - try: - _sw_model_cfg = _load_gateway_config().get("model", {}) - if isinstance(_sw_model_cfg, dict): - _sw_raw = _sw_model_cfg.get("context_length") - if _sw_raw is not None: - _sw_config_ctx = int(_sw_raw) - except Exception: - pass - if not isinstance(_sw_model_cfg, dict): - _sw_model_cfg = {} - ctx = await resolve_display_context_length_async( + model_cfg, config_ctx = await _configured_display_context() + ctx_len = await resolve_display_context_length_async( result.new_model, result.target_provider, - base_url=result.base_url or current_base_url or "", - api_key=result.api_key or current_api_key or "", + base_url=result.base_url or ctx.current_base_url or "", + api_key=result.api_key or ctx.current_api_key or "", model_info=mi, - custom_providers=custom_provs, - config_context_length=_sw_config_ctx, - configured_model=_sw_model_cfg.get("default") or _sw_model_cfg.get("model"), - configured_provider=_sw_model_cfg.get("provider"), - configured_base_url=_sw_model_cfg.get("base_url"), + custom_providers=ctx.custom_provs, + config_context_length=config_ctx, + configured_model=model_cfg.get("default") or model_cfg.get("model"), + configured_provider=model_cfg.get("provider"), + configured_base_url=model_cfg.get("base_url"), ) - if ctx: - lines.append(t("gateway.model.context_label", tokens=f"{ctx:,}")) + if ctx_len: + lines.append(t("gateway.model.context_label", tokens=f"{ctx_len:,}")) if mi: if mi.max_output: lines.append(t("gateway.model.max_output_label", tokens=f"{mi.max_output:,}")) lines.append(t("gateway.model.capabilities_label", capabilities=mi.format_capabilities())) - - if not picker: - cache_enabled = ( - (base_url_host_matches(result.base_url or "", "openrouter.ai") and "claude" in result.new_model.lower()) - or result.api_mode == "anthropic_messages" - ) - if cache_enabled: - lines.append(t("gateway.model.prompt_caching_enabled")) - + if not picker and ( + (base_url_host_matches(result.base_url or "", "openrouter.ai") and "claude" in result.new_model.lower()) + or result.api_mode == "anthropic_messages" + ): + lines.append(t("gateway.model.prompt_caching_enabled")) if result.warning_message: lines.append(t("gateway.model.warning_prefix", warning=result.warning_message)) - - if persist_global: + if ctx.persist_global: lines.append(t("gateway.model.saved_global")) elif one_turn: lines.append(" (next turn only — restores after one response)") @@ -379,6 +397,29 @@ class GatewayModelCommandsMixin: lines.append(t("gateway.model.session_only_hint")) return "\n".join(lines) + async def _commit_model_switch( + self, result, ctx: _ModelSwitchContext, *, source, picker: bool = False + ) -> str: + """Apply a resolved switch (cached agent, session, config) and build the confirmation. + + Shared by the typed ``/model `` path and the picker callback (``picker=True``, which + never carries a one-turn override). + """ + one_turn = False if picker else ctx.one_turn + error = self._switch_cached_agent_model(result, ctx, picker) + if error is not None: + return error + await self._record_model_switch(result, ctx, source=source, one_turn=one_turn, picker=picker) + return await self._model_switch_confirmation(result, ctx, one_turn=one_turn, picker=picker) + + async def _switch_and_commit( + self, ctx: _ModelSwitchContext, model_id: str, provider_slug, *, source, picker: bool = False + ) -> str: + result, error = await self._perform_model_switch(ctx, model_id, provider_slug, source) + if error is not None: + return error + return await self._commit_model_switch(result, ctx, source=source, picker=picker) + async def _send_model_picker(self, event: MessageEvent, source, adapter, session_key: str, listing_kwargs: dict, on_model_selected) -> bool: """Send the interactive /model picker; False when nothing was sent (text fallback). @@ -388,8 +429,7 @@ class GatewayModelCommandsMixin: from hermes_cli.model_switch import list_picker_providers try: - # Offload blocking provider-listing (can fall through to a synchronous urllib HTTP fetch - # on a stale cache) off the event loop so the gateway doesn't freeze. See #41289. + # Off-loop: provider listing can fall through to a synchronous HTTP fetch on a stale cache. providers = await asyncio.to_thread( list_picker_providers, max_models=50, include_moa=True, **listing_kwargs ) @@ -408,244 +448,168 @@ class GatewayModelCommandsMixin: ) return bool(result.success) - async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: - """Handle /model command — switch model.""" - from gateway.run import _hermes_home - from hermes_cli.model_switch import ( - switch_model as _switch_model, parse_model_switch_args, - resolve_persist_behavior, - list_authenticated_providers, - ) + async def _model_listing_reply( + self, event: MessageEvent, ctx: _ModelSwitchContext, profile_home + ) -> Optional[str]: + """``/model`` with no args: interactive picker where supported, else the text list.""" + from hermes_cli.model_switch import list_authenticated_providers from hermes_cli.providers import get_label - raw_args = event.get_command_args().strip() - source = event.source - _command_profile_home = None - if getattr(getattr(self, "config", None), "multiplex_profiles", False): - _command_profile_home = self._resolve_profile_home_for_source(source) + listing_kwargs = ctx.listing_kwargs() + adapter = self._adapter_for_source(ctx.source) + if adapter is not None and getattr(type(adapter), "send_model_picker", None) is not None: + async def _on_model_selected(_chat_id: str, model_id: str, provider_slug: str) -> str: + """Perform the model switch and return confirmation text.""" + # The picker callback binds the raw event source (pre-normalization). + if profile_home is None: + return await self._switch_and_commit(ctx, model_id, provider_slug, source=event.source, picker=True) + from gateway.run import _profile_runtime_scope - # Parse --provider, --global, --session, --once, and --refresh flags - # via the shared single-owner parser (hermes_cli.model_switch). - request = parse_model_switch_args(raw_args) - model_input = request.target - explicit_provider = request.explicit_provider - is_global_flag = request.is_global - force_refresh = request.force_refresh - is_session = request.is_session - one_turn = request.is_once - if request.errors: - # Gateway decoration: "❌ " prefix over the canonical error copy. - return f"❌ {request.error_messages()[0]}" - persist_global = resolve_persist_behavior( - is_global_flag, - is_session, - is_once=one_turn, - explicit_provider=explicit_provider, - ) + with _profile_runtime_scope(profile_home): + return await self._switch_and_commit(ctx, model_id, provider_slug, source=event.source, picker=True) - # --refresh: bust the disk cache so the picker shows live data. - if force_refresh: - try: - from hermes_cli.models import clear_provider_models_cache - clear_provider_models_cache() - except Exception: - pass + if await self._send_model_picker(event, ctx.source, adapter, ctx.session_key, listing_kwargs, _on_model_selected): + return None # Picker sent — adapter handles the response - # Read current model/provider from config - config_path = (_command_profile_home or _hermes_home) / "config.yaml" - current_model, current_provider, current_base_url, user_provs, custom_provs, excluded_provs = ( - _read_model_command_config(config_path) - ) - current_api_key = "" + lines = [t("gateway.model.current_label", model=ctx.current_model or "unknown", provider=get_label(ctx.current_provider)), ""] + try: + # Off-loop: provider listing can fall through to a stale-cache HTTP fetch. + providers = await asyncio.to_thread(list_authenticated_providers, max_models=5, **listing_kwargs) + lines.extend(_model_provider_listing_lines(providers)) + except Exception: + pass + lines.append(t("gateway.model.usage_switch_model")) + lines.append(t("gateway.model.usage_switch_provider")) + lines.append(t("gateway.model.usage_persist")) + return "\n".join(lines) - # Check for session override. Normalize the source the same way a normal message turn does - # (Telegram DM topic recovery) before deriving the override key, so the override is stored - # under the key the next message turn reads. - 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) - current_base_url = override.get("base_url", current_base_url) - current_api_key = override.get("api_key", current_api_key) + async def _model_selection_guard_reply( + self, event: MessageEvent, ctx: _ModelSwitchContext, result + ) -> tuple[bool, Optional[str]]: + """Selection-guard confirmation for the typed path (pickers confirm via their own UI). - async def perform_switch(model_id: str, provider_slug, *, src=source): - return await self._perform_model_switch( - _switch_model, - raw_input=model_id, - explicit_provider=provider_slug, - session_key=session_key, - source=src, - current_model=current_model, - current_provider=current_provider, - current_base_url=current_base_url, - current_api_key=current_api_key, - persist_global=persist_global, - user_provs=user_provs, - custom_provs=custom_provs, - ) - - async def commit_switch(result, *, picker: bool = False, src=source) -> str: - """Apply the resolved switch (agent, session, config) and build the reply.""" - return await self._commit_model_switch( - result, - session_key=session_key, - source=src, - current_model=current_model, - current_base_url=current_base_url, - current_api_key=current_api_key, - custom_provs=custom_provs, - persist_global=persist_global, - config_path=config_path, - one_turn=False if picker else one_turn, - restore_snapshot=None if picker else restore_snapshot, - picker=picker, - ) - - async def switch_and_commit(model_id: str, provider_slug, *, picker: bool) -> str: - # The picker callback binds the raw event source (pre-normalization), as it always has. - src = event.source if picker else source - result, error = await perform_switch(model_id, provider_slug, src=src) - if error is not None: - return error - return await commit_switch(result, picker=picker, src=src) - - # No args: show interactive picker (Telegram/Discord) or text list - if not model_input and not explicit_provider: - listing_kwargs = dict( - current_provider=current_provider, - current_base_url=current_base_url, - current_model=current_model, - user_providers=user_provs, - custom_providers=custom_provs, - excluded_providers=excluded_provs, - ) - # Try interactive picker if the platform supports it - adapter = self._adapter_for_source(source) - if adapter is not None and getattr(type(adapter), "send_model_picker", None) is not None: - async def _on_model_selected(_chat_id: str, model_id: str, provider_slug: str) -> str: - """Perform the model switch and return confirmation text.""" - if _command_profile_home is None: - return await switch_and_commit(model_id, provider_slug, picker=True) - from gateway.run import _profile_runtime_scope - - with _profile_runtime_scope(_command_profile_home): - return await switch_and_commit(model_id, provider_slug, picker=True) - - if await self._send_model_picker(event, source, adapter, session_key, listing_kwargs, _on_model_selected): - return None # Picker sent — adapter handles the response - - # Fallback: text list (for platforms without picker or if picker failed) - lines = [t("gateway.model.current_label", model=current_model or "unknown", provider=get_label(current_provider)), ""] - try: - # Offload blocking provider-listing off the event loop so the - # gateway doesn't freeze on a stale-cache HTTP fetch. See #41289. - providers = await asyncio.to_thread(list_authenticated_providers, max_models=5, **listing_kwargs) - lines.extend(_model_provider_listing_lines(providers)) - except Exception: - pass - lines.append(t("gateway.model.usage_switch_model")) - lines.append(t("gateway.model.usage_switch_provider")) - lines.append(t("gateway.model.usage_persist")) - return "\n".join(lines) - - # Perform the switch - result, error = await perform_switch(model_input, explicit_provider) - if error is not None: - return error - - # Selection-guard confirmation for the typed /model path (pickers confirm via their own - # UI). Runs the unified registry (cost + data-policy guards); pricing lookups may hit - # models.dev or a /models endpoint on a cache miss, so run it off the event loop. - _cost_warning = None + Runs the unified registry (cost + data-policy guards) off the event loop — pricing lookups + may hit models.dev or a /models endpoint on a cache miss. Returns ``(fired, reply)``; the + reply may be None when the platform rendered confirm buttons itself. + """ try: from hermes_cli.model_selection_guards import combined_selection_warning - _cost_warning = await asyncio.to_thread( + warning = await asyncio.to_thread( combined_selection_warning, result.new_model, provider=result.target_provider, - base_url=result.base_url or current_base_url or "", - api_key=result.api_key or current_api_key or "", + base_url=result.base_url or ctx.current_base_url or "", + api_key=result.api_key or ctx.current_api_key or "", model_info=result.model_info, ) except Exception: - _cost_warning = None - if _cost_warning is not None: - async def _on_cost_confirm(choice: str) -> str: - if choice == "cancel": - return ( - f"🟡 Model switch cancelled. Current model unchanged " - f"({current_model or 'unknown'})." - ) - # "once" and "always" both proceed — there is no persistent - # opt-out for selection guards (each guarded switch should be - # an explicit decision). - return await commit_switch(result) + warning = None + if warning is None: + return False, None - _p = self._typed_command_prefix_for(event.source.platform) - return await self._request_slash_confirm( - event=event, - command="model", - title=_cost_warning.title, - message=( - f"⚠️ **{_cost_warning.title}**\n\n{_cost_warning.message}\n\n" - f"_Text fallback: reply `{_p}approve` to switch or `{_p}cancel` to keep " - "the current model._" - ), - handler=_on_cost_confirm, - ) + async def _on_cost_confirm(choice: str) -> str: + if choice == "cancel": + return ( + f"🟡 Model switch cancelled. Current model unchanged " + f"({ctx.current_model or 'unknown'})." + ) + # "once" and "always" both proceed — selection guards have no persistent opt-out. + return await self._commit_model_switch(result, ctx, source=ctx.source) - return await commit_switch(result) + _p = self._typed_command_prefix_for(event.source.platform) + return True, await self._request_slash_confirm( + event=event, + command="model", + title=warning.title, + message=( + f"⚠️ **{warning.title}**\n\n{warning.message}\n\n" + f"_Text fallback: reply `{_p}approve` to switch or `{_p}cancel` to keep " + "the current model._" + ), + handler=_on_cost_confirm, + ) + + async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: + """Handle /model command — switch model.""" + from gateway.run import _hermes_home + from hermes_cli.model_switch import parse_model_switch_args, resolve_persist_behavior + + profile_home = None + if getattr(getattr(self, "config", None), "multiplex_profiles", False): + profile_home = self._resolve_profile_home_for_source(event.source) + + # --provider/--global/--session/--once/--refresh via the single-owner parser. + request = parse_model_switch_args(event.get_command_args().strip()) + if request.errors: + # Gateway decoration: "❌ " prefix over the canonical error copy. + return f"❌ {request.error_messages()[0]}" + if request.force_refresh: + # Bust the disk cache so the picker shows live data. + with contextlib.suppress(Exception): + from hermes_cli.models import clear_provider_models_cache + clear_provider_models_cache() + + # Normalize the source the same way a message turn does (Telegram DM topic recovery) before + # deriving the override key, so the override is stored under the key the next turn reads. + source = await asyncio.to_thread(self._normalize_source_for_session_key, event.source) + session_key = self._session_key_for_source(source) + ctx = _ModelSwitchContext( + session_key=session_key, + source=source, + config_path=(profile_home or _hermes_home) / "config.yaml", + persist_global=resolve_persist_behavior( + request.is_global, + request.is_session, + is_once=request.is_once, + explicit_provider=request.explicit_provider, + ), + one_turn=request.is_once, + restore_snapshot=self._snapshot_session_model_override(session_key) if request.is_once else None, + ) + ctx.read_config() + ctx.apply_override(self._session_model_overrides.get(session_key, {})) + + if not request.target and not request.explicit_provider: + return await self._model_listing_reply(event, ctx, profile_home) + + result, error = await self._perform_model_switch(ctx, request.target, request.explicit_provider, source) + if error is not None: + return error + guard_fired, guard_reply = await self._model_selection_guard_reply(event, ctx, result) + if guard_fired: + return guard_reply + return await self._commit_model_switch(result, ctx, source=source) + + # -------------------------------------------------- /codex-runtime, /personality async def _handle_codex_runtime_command(self, event: MessageEvent) -> str: - """Handle /codex-runtime command in the gateway. - - On change the cached agent is evicted so the next message builds a fresh AIAgent with the - new api_mode (avoids prompt-cache invalidation mid-session). - """ + """Handle /codex-runtime; a real change evicts the cached agent so the new api_mode applies + on the next message (avoids prompt-cache invalidation mid-session).""" from hermes_cli import codex_runtime_switch as crs raw_args = event.get_command_args().strip() if event else "" new_value, errors = crs.parse_args(raw_args) if errors: return "❌ " + "\n❌ ".join(errors) - - # Load + persist via the same helpers used for /model and /yolo try: from hermes_cli.config import load_config, save_config except Exception as exc: return f"❌ Could not load config: {exc}" - cfg = load_config() - result = crs.apply( - cfg, - new_value, - persist_callback=(save_config if new_value is not None else None), + load_config(), new_value, persist_callback=(save_config if new_value is not None else None), ) - - # On a real change, evict the cached agent so the new runtime takes - # effect on the next message rather than waiting for cache TTL. if result.success and new_value is not None and result.requires_new_session: try: - session_key = self._session_key_for_source(event.source) - self._evict_cached_agent(session_key) + self._evict_cached_agent(self._session_key_for_source(event.source)) except Exception: logger.debug("could not evict cached agent after codex-runtime change", exc_info=True) - prefix = "✓" if result.success else "✗" return f"{prefix} {result.message}" async def _handle_personality_command(self, event: MessageEvent) -> str: - """Handle /personality command - list or set a personality. - - All resolution/persistence goes through hermes_cli.personality, the single owner of state. - """ + """Handle /personality — list or set a personality (hermes_cli.personality owns the state).""" from gateway.run import _load_gateway_config from hermes_cli.personality import ( active_personality_name, @@ -656,7 +620,6 @@ class GatewayModelCommandsMixin: ) args = event.get_command_args().strip() - try: config = _load_gateway_config() except Exception: @@ -665,16 +628,11 @@ class GatewayModelCommandsMixin: if not args: current = active_personality_name(config) - lines = [t("gateway.personality.header")] - lines.append(t("gateway.personality.none_option")) + lines = [t("gateway.personality.header"), t("gateway.personality.none_option")] for name, prompt in personalities.items(): marker = " ✓" if name == current else "" lines.append( - t( - "gateway.personality.item", - name=f"{name}{marker}", - preview=describe_personality(prompt), - ) + t("gateway.personality.item", name=f"{name}{marker}", preview=describe_personality(prompt)) ) lines.append(t("gateway.personality.usage")) return "\n".join(lines) @@ -684,27 +642,25 @@ class GatewayModelCommandsMixin: except ValueError: available = "`none`, " + ", ".join(f"`{n}`" for n in personalities) return t("gateway.personality.unknown", name=args.lower(), available=available) - - # Persist the selection only — hermes_cli.personality never writes agent.system_prompt (user- - # owned overlay). persist_personality writes get_hermes_home()/config.yaml (the routed profile - # under multiplex) and the next turn re-resolves the prompt from it: no process-global state. + # Persist the selection only — never agent.system_prompt (user-owned overlay). It lands in + # get_hermes_home()/config.yaml (the routed profile under multiplex) and the next turn + # re-resolves the prompt from it: no process-global state. if not persist_personality(name): return t("gateway.personality.save_failed", error="config write failed") - if not name: return t("gateway.personality.cleared") return t("gateway.personality.set_to", name=name) + # ----------------------------------------------------------- /reasoning, /fast + def _save_gateway_config_key(self, key_path: str, value) -> bool: - """Save a dot-separated key to config.yaml (shared by /reasoning, /fast - and their interactive pickers).""" + """Save a dot-separated key to config.yaml (shared by /reasoning, /fast and their pickers).""" from gateway.slash_commands import _nested_dict from gateway.run import _gateway_config_home from hermes_cli.config import read_user_config_raw config_path = _gateway_config_home() / "config.yaml" try: - # Write-back round-trip: raw read is correct (merged defaults must - # not be persisted back to the user's file). + # Raw read: merged defaults must not be persisted back to the user's file. user_config = read_user_config_raw(config_path) *parents, leaf = key_path.split(".") _nested_dict(user_config, *parents)[leaf] = value @@ -714,12 +670,13 @@ class GatewayModelCommandsMixin: logger.error("Failed to save config key %s: %s", key_path, e) return False + def _set_reasoning_override(self, session_key: str, value) -> None: + """Store (or clear with None) the session reasoning override and drop the cached agent.""" + self._set_session_reasoning_override(session_key, value) + self._evict_cached_agent(session_key) + def _apply_reasoning_selection( - self, - session_key: str, - platform_key: str, - value: str, - persist_global: bool = False, + self, session_key: str, platform_key: str, value: str, persist_global: bool = False, ) -> str: """Apply a /reasoning argument (typed or picked) and return the reply. @@ -728,21 +685,12 @@ class GatewayModelCommandsMixin: from hermes_constants import parse_reasoning_effort value = (value or "").strip().lower() - - # Display toggle (per-platform) - if value in {"show", "on"}: - self._show_reasoning = True - self._save_gateway_config_key( - f"display.platforms.{platform_key}.show_reasoning", True - ) - return t("gateway.reasoning.display_set_on", platform=platform_key) - if value in {"hide", "off"}: - self._show_reasoning = False - self._save_gateway_config_key( - f"display.platforms.{platform_key}.show_reasoning", False - ) - return t("gateway.reasoning.display_set_off", platform=platform_key) - + show = _REASONING_DISPLAY_TOGGLES.get(value) + if show is not None: # per-platform display toggle + self._show_reasoning = show + self._save_gateway_config_key(f"display.platforms.{platform_key}.show_reasoning", show) + key = "gateway.reasoning.display_set_on" if show else "gateway.reasoning.display_set_off" + return t(key, platform=platform_key) if value == "reset": if persist_global: return t("gateway.reasoning.reset_global_unsupported") @@ -754,19 +702,14 @@ class GatewayModelCommandsMixin: parsed = parse_reasoning_effort(value) if parsed is None: return t("gateway.reasoning.unknown_arg", arg=value) - self._reasoning_config = parsed if persist_global: if self._save_gateway_config_key("agent.reasoning_effort", value): - self._set_session_reasoning_override(session_key, None) - self._evict_cached_agent(session_key) + self._set_reasoning_override(session_key, None) return t("gateway.reasoning.set_global", effort=value) - self._set_session_reasoning_override(session_key, parsed) - self._evict_cached_agent(session_key) + self._set_reasoning_override(session_key, parsed) return t("gateway.reasoning.set_global_save_failed", effort=value) - - self._set_session_reasoning_override(session_key, parsed) - self._evict_cached_agent(session_key) + self._set_reasoning_override(session_key, parsed) return t("gateway.reasoning.set_session", effort=value) def _reasoning_picker_choices(self, current_effort: str) -> list: @@ -782,12 +725,7 @@ class GatewayModelCommandsMixin: return choices async def _try_send_choice_picker( - self, - event: MessageEvent, - session_key: str, - title: str, - choices: list, - on_choice_selected, + self, event: MessageEvent, session_key: str, title: str, choices: list, on_choice_selected, ) -> bool: """Send an interactive choice picker when the platform supports it. @@ -795,21 +733,16 @@ class GatewayModelCommandsMixin: (``send_choice_picker``); a failed send returns False (text fallback) instead of erroring. """ adapter = self._adapter_for_source(event.source) - has_picker = ( - adapter is not None - and getattr(type(adapter), "send_choice_picker", None) is not None - ) - if not has_picker: + if adapter is None or getattr(type(adapter), "send_choice_picker", None) is None: return False try: - metadata = self._reply_metadata(event) result = await adapter.send_choice_picker( chat_id=event.source.chat_id, title=title, choices=choices, session_key=session_key, on_choice_selected=on_choice_selected, - metadata=metadata, + metadata=self._reply_metadata(event), ) return bool(getattr(result, "success", False)) except Exception as e: @@ -822,83 +755,68 @@ class GatewayModelCommandsMixin: raw_args = event.get_command_args().strip() args, persist_global = self._parse_reasoning_command_args(raw_args) - # Normalize the source (Telegram DM topic recovery) before deriving - # the override key so storage matches the key the next message turn - # reads — same fix as /model (#30479). + # Normalize the source (Telegram DM topic recovery) before deriving the override key so + # storage matches the key the next message turn reads — same as /model. _reasoning_source = await asyncio.to_thread(self._normalize_source_for_session_key, event.source) session_key = self._session_key_for_source(_reasoning_source) self._show_reasoning = self._load_show_reasoning() - # Use the session's effective model (session /model override wins over - # config default) so per-model reasoning_overrides display correctly. + # The session's effective model (session /model override wins over config default) so + # per-model reasoning_overrides display correctly. _session_model = str( ((getattr(self, "_session_model_overrides", {}) or {}).get(session_key) or {}).get("model") or "" ) self._reasoning_config = self._resolve_session_reasoning_config( - source=event.source, - session_key=session_key, - model=_session_model, + source=event.source, session_key=session_key, model=_session_model, ) - - if not raw_args: - # Show current state - rc = self._reasoning_config - if rc is None: - level = t("gateway.reasoning.level_default") - current_effort = "medium" - elif rc.get("enabled") is False: - level = t("gateway.reasoning.level_disabled") - current_effort = "none" - else: - level = rc.get("effort", "medium") - current_effort = level - display_state = ( - t("gateway.reasoning.display_on") - if self._show_reasoning - else t("gateway.reasoning.display_off") - ) - has_session_override = session_key in (getattr(self, "_session_reasoning_overrides", {}) or {}) - scope = ( - t("gateway.reasoning.scope_session") - if has_session_override - else t("gateway.reasoning.scope_global") - ) - - # Interactive picker on platforms that support it (parity with the - # /model picker). Falls through to the text status card otherwise. - _picker_platform_key = _platform_config_key(event.source.platform) - - async def _on_reasoning_choice(_chat_id: str, value: str) -> str: - return self._apply_reasoning_selection( - session_key, _picker_platform_key, value - ) - - picker_sent = await self._try_send_choice_picker( - event, - session_key, - title=t( - "gateway.reasoning.picker_title", - level=level, - scope=scope, - display=display_state, - ), - choices=self._reasoning_picker_choices(current_effort), - on_choice_selected=_on_reasoning_choice, - ) - if picker_sent: - return None # Picker sent — adapter handles the response - - return t( - "gateway.reasoning.status", - level=level, - scope=scope, - display=display_state, - ) - - # Typed argument path — same applier the picker uses. platform_key = _platform_config_key(event.source.platform) - return self._apply_reasoning_selection( - session_key, platform_key, args, persist_global=persist_global + if raw_args: + # Typed argument path — same applier the picker uses. + return self._apply_reasoning_selection(session_key, platform_key, args, persist_global=persist_global) + + rc = self._reasoning_config + if rc is None: + level, current_effort = t("gateway.reasoning.level_default"), "medium" + elif rc.get("enabled") is False: + level, current_effort = t("gateway.reasoning.level_disabled"), "none" + else: + level = current_effort = rc.get("effort", "medium") + display_state = t("gateway.reasoning.display_on") if self._show_reasoning else t("gateway.reasoning.display_off") + has_session_override = session_key in (getattr(self, "_session_reasoning_overrides", {}) or {}) + scope = t("gateway.reasoning.scope_session") if has_session_override else t("gateway.reasoning.scope_global") + + async def _on_reasoning_choice(_chat_id: str, value: str) -> str: + return self._apply_reasoning_selection(session_key, platform_key, value) + + # Interactive picker where supported (parity with /model); else the text status card. + picker_sent = await self._try_send_choice_picker( + event, + session_key, + title=t("gateway.reasoning.picker_title", level=level, scope=scope, display=display_state), + choices=self._reasoning_picker_choices(current_effort), + on_choice_selected=_on_reasoning_choice, ) + if picker_sent: + return None # Picker sent — adapter handles the response + return t("gateway.reasoning.status", level=level, scope=scope, display=display_state) + + def _apply_fast_selection(self, session_key: str, value: str, persist: bool = False) -> str: + """Apply a /fast argument (typed or picked) and return the reply.""" + selection = _FAST_SELECTIONS.get(value) + if selection is None: + return t("gateway.fast.unknown_arg", arg=value) + tier, saved_value, label_key = selection + label = t(label_key) if label_key else value.upper() + self._service_tier = tier + if persist and self._save_gateway_config_key("agent.service_tier", saved_value): + # Global write supersedes any session override. + self._set_session_service_tier_override(session_key, None, clear=True) + self._evict_cached_agent(session_key) + return t("gateway.fast.saved", label=label) + # Session override; also the fallback when the config write failed so the user's choice + # still applies (mirrors /reasoning --global). + self._set_session_service_tier_override(session_key, tier) + self._evict_cached_agent(session_key) + return t("gateway.fast.session_only", label=label) async def _handle_fast_command(self, event: MessageEvent) -> Optional[str]: """Handle /fast — mirror the CLI Priority Processing toggle in gateway chats. @@ -909,73 +827,32 @@ class GatewayModelCommandsMixin: from hermes_cli.models import model_supports_fast_mode raw_args = event.get_command_args().strip().lower() - # Reuse the /reasoning arg parser: strips --global (any position), - # normalizes unicode dashes. + # The /reasoning arg parser strips --global (any position) and normalizes unicode dashes. args, persist_global = self._parse_reasoning_command_args(raw_args) session_key = self._session_key_for_source(event.source) - self._service_tier = self._resolve_session_service_tier( - session_key=session_key - ) + self._service_tier = self._resolve_session_service_tier(session_key=session_key) - user_config = _load_gateway_config() - model = _resolve_gateway_model(user_config) - if not model_supports_fast_mode(model): + if not model_supports_fast_mode(_resolve_gateway_model(_load_gateway_config())): return t("gateway.fast.not_supported") + if args and args != "status": + return self._apply_fast_selection(session_key, args, persist=persist_global) - def _apply_fast_selection(value: str, persist: bool = False) -> str: - """Apply a /fast argument (typed or picked) and return the reply.""" - if value in {"fast", "on"}: - tier = "priority" - saved_value = "fast" - label = t("gateway.fast.label_fast") - elif value in {"normal", "off"}: - tier = None - saved_value = "normal" - label = t("gateway.fast.label_normal") - elif value in {"auto", "cold"}: - tier = saved_value = value - label = value.upper() - else: - return t("gateway.fast.unknown_arg", arg=value) - self._service_tier = tier - if persist: - if self._save_gateway_config_key("agent.service_tier", saved_value): - # Global write supersedes any session override. - self._set_session_service_tier_override( - session_key, None, clear=True - ) - self._evict_cached_agent(session_key) - return t("gateway.fast.saved", label=label) - # Config write failed — fall back to a session override so the - # user's choice still applies (mirrors /reasoning --global). - self._set_session_service_tier_override(session_key, tier) - self._evict_cached_agent(session_key) - return t("gateway.fast.session_only", label=label) - self._set_session_service_tier_override(session_key, tier) - self._evict_cached_agent(session_key) - return t("gateway.fast.session_only", label=label) + mode = "fast" if self._service_tier == "priority" else (self._service_tier or "normal") + status = {"fast": t("gateway.fast.status_fast"), "normal": t("gateway.fast.status_normal")}.get(mode, mode) - if not args or args == "status": - is_fast = self._service_tier == "priority" - mode = "fast" if is_fast else (self._service_tier or "normal") - status = {"fast": t("gateway.fast.status_fast"), "normal": t("gateway.fast.status_normal")}.get(mode, mode) + async def _on_fast_choice(_chat_id: str, value: str) -> str: + return self._apply_fast_selection(session_key, value, persist=persist_global) - async def _on_fast_choice(_chat_id: str, value: str) -> str: - return _apply_fast_selection(value, persist=persist_global) - - picker_sent = await self._try_send_choice_picker( - event, - session_key, - title=t("gateway.fast.picker_title", mode=status), - choices=[ - {"value": v, "label": t(f"gateway.fast.choice_{v}"), "is_current": mode == v} - for v in ("fast", "normal", "auto", "cold") - ], - on_choice_selected=_on_fast_choice, - ) - if picker_sent: - return None # Picker sent — adapter handles the response - - return t("gateway.fast.status", mode=status) - - return _apply_fast_selection(args, persist=persist_global) + picker_sent = await self._try_send_choice_picker( + event, + session_key, + title=t("gateway.fast.picker_title", mode=status), + choices=[ + {"value": v, "label": t(f"gateway.fast.choice_{v}"), "is_current": mode == v} + for v in ("fast", "normal", "auto", "cold") + ], + on_choice_selected=_on_fast_choice, + ) + if picker_sent: + return None # Picker sent — adapter handles the response + return t("gateway.fast.status", mode=status) diff --git a/gateway/slash_commands_session.py b/gateway/slash_commands_session.py index 6283161176..81028372d2 100644 --- a/gateway/slash_commands_session.py +++ b/gateway/slash_commands_session.py @@ -2,16 +2,16 @@ /new, /resume, /sessions, /branch, /title, /save, /undo, /retry, /topic, /compress. Split out of ``gateway/slash_commands.py``; bound onto ``GatewayRunner`` through -``GatewaySlashCommandsMixin``. Origin internals are imported lazily (``from gateway.slash_commands -import ...``) inside the bodies to avoid the import cycle. +``GatewaySlashCommandsMixin``. Origin internals are imported lazily inside the bodies to avoid +the import cycle. """ from __future__ import annotations -import logging import asyncio import contextlib import dataclasses +import logging import os import shlex from typing import Optional, Union @@ -25,19 +25,31 @@ from gateway.session import SessionSource, build_session_key, is_shared_multi_us # Log-record parity with gateway/run.py and the origin module. logger = logging.getLogger("gateway.run") -# Upper bound on the off-loop agent-resource cleanup during a /new or /reset (see -# _handle_reset_command). A stuck teardown must not block the event loop; past this the reset -# proceeds and the cleanup is left to finish (or leak) in its worker thread. +# Upper bound on the off-loop agent-resource cleanup during a /new or /reset. A stuck teardown must +# not block the event loop; past this the reset proceeds and the cleanup finishes (or leaks) in its +# worker thread. _RESET_CLEANUP_TIMEOUT_S = 30.0 +# chat_type values whose session key is per-user (DM-like), incl. the unknown/blank case. +_DM_CHAT_TYPES = {"dm", "direct", "private", ""} + +_BRANCH_COPIED_FIELDS = ( + "content", "tool_calls", "tool_call_id", "finish_reason", "reasoning", "reasoning_content", + "reasoning_details", "codex_reasoning_items", "codex_message_items", "timestamp", +) + + +def _sattr(obj, name: str) -> str: + """``str(getattr(obj, name) or "")`` — normalized identity field for origin comparisons.""" + return str(getattr(obj, name, "") or "") + def _manual_compression_reply_lines(summary: dict, compressor, focus_topic) -> list[str]: """Lines for the manual /compress confirmation, surfacing summariser/aux-model failures. - ``_last_compress_aborted`` = no usable summary, messages unchanged (force=True bypasses any - cooldown). Provider exception text is force-redacted at this UI boundary even when global - redaction is off. A configured aux model that failed and was recovered via main is an info - note so the user can fix their config. + ``_last_compress_aborted`` = no usable summary, messages unchanged. Provider exception text is + force-redacted at this UI boundary even when global redaction is off. A configured aux model + that failed and was recovered via main is an info note so the user can fix their config. """ lines = [f"🗜️ {summary['headline']}"] if focus_topic: @@ -53,13 +65,11 @@ def _manual_compression_reply_lines(summary: dict, compressor, focus_topic) -> l if getattr(compressor, "_last_compress_aborted", False): lines.append(t("gateway.compress.aborted", error=(summary_err or "unknown error"))) elif aux_fail_model: - lines.append( - t( - "gateway.compress.aux_failed", - model=aux_fail_model, - error=(getattr(compressor, "_last_aux_model_failure_error", None) or "unknown error"), - ) - ) + lines.append(t( + "gateway.compress.aux_failed", + model=aux_fail_model, + error=(getattr(compressor, "_last_aux_model_failure_error", None) or "unknown error"), + )) return lines @@ -84,21 +94,12 @@ def _compress_preview_reply(history, partial: bool, keep_last, focus_topic, agg_ def _reset_process_scoped_tool_state() -> None: """Drop env-passthrough and credential-file state at a conversation boundary (best-effort).""" - try: + with contextlib.suppress(Exception): from tools.env_passthrough import clear_env_passthrough clear_env_passthrough() - except Exception: - pass - try: + with contextlib.suppress(Exception): from tools.credential_files import clear_credential_files clear_credential_files() - except Exception: - pass - -_BRANCH_COPIED_FIELDS = ( - "content", "tool_calls", "tool_call_id", "finish_reason", "reasoning", "reasoning_content", - "reasoning_details", "codex_reasoning_items", "codex_message_items", "timestamp", -) def _branch_row(msg: dict) -> dict: @@ -111,14 +112,86 @@ def _branch_row(msg: dict) -> dict: return row +def _strip_resume_name(parts: list[str]) -> str: + """Join the non-flag /resume tokens; strip literal ``<...>``/``[...]``/quotes typed from the + usage hint (mirrors the CLI).""" + name = " ".join(p for p in parts if p not in {"--all", "--cross-room"}).strip() + if len(name) >= 2 and (name[0], name[-1]) in {("<", ">"), ("[", "]"), ('"', '"'), ("'", "'")}: + name = name[1:-1].strip() + return name + + class GatewaySessionCommandsMixin: """Session-transcript slash commands (/new, /resume, /sessions, /branch, /title, /save, /undo, /retry, /topic, /compress).""" + # ------------------------------------------------------------------ /new, /reset + + async def _cleanup_old_agent_for_reset(self, session_key: str) -> None: + """Close the old agent's tool resources (sandboxes, browser daemons, subprocesses) before eviction. + + _cleanup_agent_resources is blocking and this runs ON the event loop (confirm-button click), + so it is offloaded with a bounded timeout. wait_for cancels the await but not the worker + thread — a wedged teardown keeps running (or leaks); the reset proceeds either way. + """ + _old_agent = self._cached_agent_for(session_key) + if _old_agent is None: + return + try: + await asyncio.wait_for( + self._run_in_executor_with_context(self._cleanup_agent_resources, _old_agent), + timeout=_RESET_CLEANUP_TIMEOUT_S, + ) + except asyncio.TimeoutError: + logger.warning( + "Agent resource cleanup for session %s exceeded %ss during " + "/new reset; proceeding with reset (the worker thread is left " + "to finish on its own). (#35994)", + session_key, _RESET_CLEANUP_TIMEOUT_S, + ) + except Exception as cleanup_exc: + logger.warning( + "Agent resource cleanup for session %s failed during /new " + "reset: %s (#35994)", + session_key, cleanup_exc, + ) + + async def _fire_session_reset_hooks( + self, source: SessionSource, session_key: str, old_sid, new_sid + ) -> None: + """Session-boundary hooks: plugin finalize (off-loop + bounded), session:end/reset, on_session_reset.""" + platform_value = source.platform.value if source.platform else "" + # Finalize hooks can block arbitrarily (observability trace exports) and this handler runs + # on the gateway event loop (see GatewayRunner._finalize_session_off_loop). + with contextlib.suppress(Exception): + await self._finalize_session_off_loop( + session_id=old_sid, + platform=platform_value, + reason="new_session", + old_session_id=old_sid, + new_session_id=new_sid, + ) + hook_payload = {"platform": platform_value, "user_id": source.user_id, "session_key": session_key} + await self.hooks.emit("session:end", dict(hook_payload)) + await self.hooks.emit("session:reset", dict(hook_payload)) + + def _invoke_session_reset_lifecycle_hook(self, source: SessionSource, old_sid, new_sid) -> None: + """Plugin on_session_reset hook (new session guaranteed to exist); best-effort.""" + try: + from hermes_cli.lifecycle import invoke_hook as _invoke_hook + _invoke_hook( + "on_session_reset", + session_id=new_sid, + platform=source.platform.value if source.platform else "", + reason="new_session", + old_session_id=old_sid, + new_session_id=new_sid, + ) + except Exception: + pass + async def _handle_reset_command(self, event: MessageEvent) -> Union[str, EphemeralReply]: """Handle /new or /reset command.""" source = event.source - - # Get existing session key session_key = self._session_key_for_source(source) self._invalidate_session_run_generation(session_key, reason="session_reset") # Evict the running-agent slot now that the generation is bumped: the in-flight run's own @@ -126,38 +199,10 @@ class GatewaySessionCommandsMixin: # drops all later messages. Idempotent, so the run's finally calling it again is harmless. self._release_running_agent_state(session_key) - # Snapshot the old entry so on_session_finalize can report the - # expiring session id before reset_session() rotates it. + # Snapshot the old entry so on_session_finalize can report the expiring session id. old_entry = self.session_store._entries.get(session_key) - - # Close the old agent's tool resources (sandboxes, browser daemons, subprocesses) before - # evicting it; getattr-guarded since test fixtures may skip __init__. _cleanup_agent_resources - # is blocking and this handler runs ON the event loop (confirm-button click), so an inline - # call wedges the loop — offload to a worker thread with a bounded timeout. - _old_agent = self._cached_agent_for(session_key) - if _old_agent is not None: - try: - await asyncio.wait_for( - self._run_in_executor_with_context(self._cleanup_agent_resources, _old_agent), - timeout=_RESET_CLEANUP_TIMEOUT_S, - ) - except asyncio.TimeoutError: - # wait_for cancels the await, but the worker thread cannot be cancelled — a wedged - # teardown keeps running (or leaks) for the gateway's lifetime. The reset proceeds. - logger.warning( - "Agent resource cleanup for session %s exceeded %ss during " - "/new reset; proceeding with reset (the worker thread is left " - "to finish on its own). (#35994)", - session_key, _RESET_CLEANUP_TIMEOUT_S, - ) - except Exception as cleanup_exc: - logger.warning( - "Agent resource cleanup for session %s failed during /new " - "reset: %s (#35994)", - session_key, cleanup_exc, - ) + await self._cleanup_old_agent_for_reset(session_key) self._evict_cached_agent(session_key) - # Conversation boundary: clear ALL conversation-scoped per-session state (model/reasoning # overrides, one-turn restores, model notes, last-resolved cache, /queue overflow) + # security state in one funnel call. See _CONVERSATION_SCOPED_STATE in gateway/run.py. @@ -178,85 +223,47 @@ class GatewaySessionCommandsMixin: pass _reset_process_scoped_tool_state() - # Reset the session new_entry = await self.async_session_store.reset_session(session_key) - - # (Conversation-scoped overrides + security state were already - # cleared via _clear_conversation_scope above.) - _old_sid = old_entry.session_id if old_entry else None - platform_value = source.platform.value if source.platform else "" + await self._fire_session_reset_hooks( + source, session_key, _old_sid, new_entry.session_id if new_entry else None + ) - # Fire plugin on_session_finalize hook (session boundary). Off-loop + bounded: finalize - # hooks can block arbitrarily (observability trace exports) and this handler runs on the - # gateway event loop (see GatewayRunner._finalize_session_off_loop). - with contextlib.suppress(Exception): - await self._finalize_session_off_loop( - session_id=_old_sid, - platform=platform_value, - reason="new_session", - old_session_id=_old_sid, - new_session_id=new_entry.session_id if new_entry else None, - ) - - # Emit session:end (session is ending) then session:reset hooks. - hook_payload = {"platform": platform_value, "user_id": source.user_id, "session_key": session_key} - await self.hooks.emit("session:end", dict(hook_payload)) - await self.hooks.emit("session:reset", dict(hook_payload)) - - # Resolve session config info to surface to the user, scoped to the - # profile serving this source so a multiplexed /reset //new banner - # reports the profile's model, not the base config's (#59003). + # Session config info for the user, scoped to the profile serving this source so a + # multiplexed /new banner reports the profile's model, not the base config's. try: - session_info = await asyncio.to_thread( - self._reset_notice_session_info, source - ) + session_info = await asyncio.to_thread(self._reset_notice_session_info, source) except Exception: session_info = "" if new_entry: - header = await asyncio.to_thread(self._telegram_topic_new_header, source) or t("gateway.reset.header_default") + default_header = t("gateway.reset.header_default") else: # No existing session, just create one new_entry = await self.async_session_store.get_or_create_session(source, force_new=True) - header = await asyncio.to_thread(self._telegram_topic_new_header, source) or t("gateway.reset.header_new") + default_header = t("gateway.reset.header_new") + header = await asyncio.to_thread(self._telegram_topic_new_header, source) or default_header - # Set session title if provided with /new _title_arg = event.get_command_args().strip() if _title_arg and self._session_db and new_entry: header = await self._reset_titled_header(header, new_entry.session_id, _title_arg) - # When /new runs inside a Telegram DM topic lane, rewrite the (chat_id, thread_id) → - # session_id binding so the next message uses the freshly-created session. Otherwise the - # binding-lookup at the top of _handle_message_with_agent switches right back to the old one. + # Inside a Telegram DM topic lane, rewrite the (chat_id, thread_id) → session_id binding so + # the next message uses the fresh session instead of switching back to the old one. if await asyncio.to_thread(self._is_telegram_topic_lane, source) and new_entry is not None: try: await asyncio.to_thread(self._record_telegram_topic_binding, source, new_entry) except Exception: logger.debug("Failed to rebind Telegram topic after /new", exc_info=True) - # Fire plugin on_session_reset hook (new session guaranteed to exist) - try: - from hermes_cli.lifecycle import invoke_hook as _invoke_hook - _new_sid = new_entry.session_id if new_entry else None - _invoke_hook( - "on_session_reset", - session_id=_new_sid, - platform=platform_value, - reason="new_session", - old_session_id=_old_sid, - new_session_id=_new_sid, - ) - except Exception: - pass - - # Append a random tip to the reset message + self._invoke_session_reset_lifecycle_hook( + source, _old_sid, new_entry.session_id if new_entry else None + ) try: from hermes_cli.tips import get_random_tip _tip_line = t("gateway.reset.tip", tip=get_random_tip()) except Exception: _tip_line = "" - if session_info: return EphemeralReply(f"{header}\n\n{session_info}{_tip_line}") return EphemeralReply(f"{header}{_tip_line}") @@ -283,15 +290,16 @@ class GatewaySessionCommandsMixin: note = t("gateway.reset.title_empty_untitled") return header + note + # ------------------------------------------------------- origin / ownership guards + def _gateway_session_origin_for_id(self, session_id: str) -> Optional[SessionSource]: """Best-effort origin lookup for gateway session IDs.""" lookup = getattr(type(self.session_store), "lookup_by_session_id", None) if callable(lookup): entry = lookup(self.session_store, session_id) return getattr(entry, "origin", None) if entry is not None else None - - # Test doubles and older stores may not expose the public lookup helper. - # Keep the Matrix resume guard fail-closed if no origin can be resolved. + # Test doubles and older stores may not expose the public lookup helper. Keep the Matrix + # resume guard fail-closed if no origin can be resolved. entries = getattr(self.session_store, "_entries", {}) or {} for entry in entries.values(): if getattr(entry, "session_id", None) == session_id: @@ -300,17 +308,15 @@ class GatewaySessionCommandsMixin: @staticmethod def _same_matrix_room(current: SessionSource, origin: Optional[SessionSource]) -> bool: + # thread_id is part of the session key and Matrix scopes a turn to the current room/thread, + # so a live session in another thread of the SAME room is a DIFFERENT session. Non-threaded + # rooms compare "" == "" unchanged. return ( origin is not None and origin.platform == Platform.MATRIX and current.platform == Platform.MATRIX and origin.chat_id == current.chat_id - # thread_id is part of the session key (build_session_key appends it for every chat - # type when present) and Matrix scopes a turn to the current room/thread, so a live - # session in another thread of the SAME room is a DIFFERENT session: thread A must not - # resume/enumerate a target from thread B. Non-threaded rooms compare "" == "" unchanged. - and str(getattr(current, "thread_id", "") or "") - == str(getattr(origin, "thread_id", "") or "") + and _sattr(current, "thread_id") == _sattr(origin, "thread_id") ) def _same_origin_chat(self, current: SessionSource, origin: Optional[SessionSource]) -> bool: @@ -322,45 +328,28 @@ class GatewaySessionCommandsMixin: """ if origin is None or current is None: return False - if origin.platform != current.platform: + if origin.platform != current.platform or origin.chat_id != current.chat_id: return False - if origin.chat_id != current.chat_id: + # thread_id is part of the session key for every chat type, so threads of the same parent + # chat are DIFFERENT sessions; require thread equality before any sharing logic. + if _sattr(current, "thread_id") != _sattr(origin, "thread_id"): return False - # thread_id is part of the session key for every chat type (build_session_key appends it - # unconditionally), so threads of the same parent chat are DIFFERENT sessions. - # is_shared_multi_user_session only decides sharing WITHIN a thread — require thread equality - # before any sharing logic so a live origin in thread A cannot match a caller in thread B. - if str(getattr(current, "thread_id", "") or "") != str( - getattr(origin, "thread_id", "") or "" - ): - return False - chat_type = (getattr(current, "chat_type", "") or "").lower() - # DM-like chats are always per-user. - if chat_type in {"dm", "direct", "private", ""}: - # chat_id was already required equal above and, when present, IS the DM session key, so - # an equal non-empty chat_id suffices. build_session_key falls back to the participant - # (``user_id_alt or user_id`` — Signal/Feishu key on user_id_alt) only when there is NO - # chat_id; mirror that and fail closed on a missing/different participant so two - # no-chat_id DM origins are never conflated. - if str(getattr(current, "chat_id", "") or ""): + if _sattr(current, "chat_type").lower() in _DM_CHAT_TYPES: + # DM: an equal non-empty chat_id IS the session key. build_session_key falls back to the + # participant (``user_id_alt or user_id`` — Signal/Feishu key on user_id_alt) only when + # there is NO chat_id; mirror that and fail closed on a missing/different participant. + if _sattr(current, "chat_id"): return True cur_pid = str(current.user_id_alt or current.user_id or "") org_pid = str(origin.user_id_alt or origin.user_id or "") return bool(cur_pid) and cur_pid == org_pid - # Non-DM: scope by participant whenever the session key for this source - # is per-user. is_shared_multi_user_session mirrors build_session_key's - # isolation rules exactly, so the guard stays in lock-step with the key. + # Non-DM: a shared key is one session for every participant; a per-user key must compare the + # participant it is built from, failing closed when either side lacks one. if self._is_shared_session_source(current): return True - # Per-user key: compare the participant id the key is actually built - # from (user_id_alt or user_id — Signal/Feishu key on user_id_alt). cur_pid = current.user_id_alt or current.user_id org_pid = origin.user_id_alt or origin.user_id - if cur_pid and org_pid: - return cur_pid == org_pid - # Per-user key but a participant id is missing on one side: cannot prove - # the same owner — fail closed. - return False + return bool(cur_pid and org_pid) and cur_pid == org_pid def _is_shared_session_source(self, source: SessionSource) -> bool: """Whether *source*'s session key is shared by every participant (not per-user). @@ -388,6 +377,45 @@ class GatewaySessionCommandsMixin: except Exception: return False + def _persisted_row_proves_owner(self, source: SessionSource, row: dict) -> bool: + """Whether a persisted (inactive) session *row* provably belongs to *source*'s session key. + + Rows once stored only source + user_id, so a same-user row could belong to a DIFFERENT chat; + the persisted chat/thread origin closes that gap and legacy NULL rows fail closed. The table + has no user_id_alt column, so an alt-keyed (Signal/Feishu) caller can never be proven by + row_uid == caller_uid alone (CWE-639). Same platform alone is never ownership proof. + """ + caller_src = source.platform.value if source.platform else None + row_src = row.get("source") + caller_uid = _sattr(source, "user_id") + if not caller_uid: + return False + row_thread = str(row.get("thread_id") or "") + if not (row_src and caller_src and str(row_src) == str(caller_src) + and row_thread == _sattr(source, "thread_id")): + return False # blank/legacy source cannot prove the platform; other thread = other session + row_uid = str(row.get("user_id") or "") + row_chat = str(row.get("chat_id") or "") + caller_chat = _sattr(source, "chat_id") + caller_keys_on_alt = bool(_sattr(source, "user_id_alt")) + if _sattr(source, "chat_type").lower() in _DM_CHAT_TYPES: + # DMs key on user_id; a no-chat_id DM is keyed PURELY on the participant (alt-keyed + # caller fails closed); with chat_id on both sides equality is the DM key. + if caller_keys_on_alt and not (row_chat and caller_chat): + return False + return bool(row_uid) and row_uid == caller_uid and row_chat == caller_chat + # Non-DM: the key includes chat_id — both sides must carry it and match (a legacy NULL-chat + # row fails closed even when both normalize to ""). + if not (row_chat and caller_chat and row_chat == caller_chat): + return False + # A SHARED group/thread is one session for every participant: same-chat proof suffices (a + # user-id check would block co-members). A per-user key still requires the same owner. + if self._is_shared_session_source(source): + return True + if caller_keys_on_alt: + return False + return bool(row_uid) and row_uid == caller_uid + async def _resume_target_allowed( self, source: SessionSource, target_id: str, allow_override: bool = False ) -> bool: @@ -395,85 +423,23 @@ class GatewaySessionCommandsMixin: Generalizes the Matrix-only room guard to every adapter so a caller cannot bind to another user's/room's session (IDOR). Uses the live origin when the target is active, else the DB - row's source + user_id; the row must PROVE ownership or fail closed. Admin ``--all`` bypasses. + row must PROVE ownership or fail closed. Admin ``--all`` bypasses. """ if allow_override and self._resume_caller_is_admin(source): return True - # Use the live origin only when it resolves to a real SessionSource; a - # store that can't resolve it (or an unexpected lookup error) must not - # silently allow/deny — fall through to the deterministic DB scoping. + # The live origin decides only when it resolves to a real SessionSource; a store that can't + # resolve it (or a lookup error) falls through to the deterministic DB scoping. try: origin = self._gateway_session_origin_for_id(target_id) except Exception: origin = None if isinstance(origin, SessionSource): return self._same_origin_chat(source, origin) - # Inactive/persisted-only: best-effort scope by DB row source + user. try: row = await self._session_db.get_session(target_id) or {} except Exception: return False - caller_src = source.platform.value if source.platform else None - row_src = row.get("source") - if row_src and caller_src and str(row_src) != str(caller_src): - return False # different platform / source - caller_uid = str(getattr(source, "user_id", "") or "") - row_uid = str(row.get("user_id") or "") - # Chat/thread origin recorded at session creation. Rows once stored only source + user_id, - # so a same-user row could belong to a DIFFERENT chat; comparing the persisted origin closes - # that gap. Legacy rows (NULL) fail closed — resume via a live session or an admin override. - caller_chat = str(getattr(source, "chat_id", "") or "") - row_chat = str(row.get("chat_id") or "") - caller_thread = str(getattr(source, "thread_id", "") or "") - row_thread = str(row.get("thread_id") or "") - chat_type = (getattr(source, "chat_type", "") or "").lower() - caller_is_dm = chat_type in {"dm", "direct", "private", ""} - # build_session_key keys the participant on ``user_id_alt or user_id``, but the sessions table - # has no user_id_alt column, so a row cannot prove the canonical participant for an alt-keyed - # (Signal/Feishu) caller: per-user row_uid == caller_uid checks must fail closed (CWE-639). - caller_keys_on_alt = bool(str(getattr(source, "user_id_alt", "") or "")) - if caller_uid: - # Identity-bearing caller: the row must PROVE the same owner AND platform AND chat/thread. - # A blank/legacy source can't prove the platform (row_src above only rejects a *mismatching* - # non-blank one); a different thread is a different session. Any gap fails closed. - origin_ok = ( - bool(row_src) and bool(caller_src) - and str(row_src) == str(caller_src) - and row_thread == caller_thread - ) - if not origin_ok: - return False - if caller_is_dm: - # DMs are keyed on user_id; require the same owner. A no-chat_id DM is keyed PURELY on - # the participant (so an alt-keyed caller fails closed); when both sides carry chat_id, - # equality is the DM key and suffices, and a mismatching chat_id is rejected. - if caller_keys_on_alt and not (bool(row_chat) and bool(caller_chat)): - return False - return ( - bool(row_uid) and row_uid == caller_uid - and row_chat == caller_chat - ) - # Non-DM (group/channel/forum/thread): build_session_key includes chat_id, so a row (or - # caller) with NO chat provenance cannot prove same-chat. Require both non-blank and - # equal — a legacy NULL-chat row fails closed even when both normalize to "". (CWE-639) - if not (bool(row_chat) and bool(caller_chat) and row_chat == caller_chat): - return False - # Same non-DM chat/thread: mirror build_session_key's participant scoping. A SHARED - # group/thread session (group_sessions_per_user=False, or a shared thread) is one session - # for every participant, so the same-chat proof suffices — do NOT also require user-id - # equality (it would block co-members). A per-user session still requires the same owner. - if self._is_shared_session_source(source): - return True - # Per-user non-DM: the session key includes the participant (``user_id_alt or - # user_id``). If the caller keys on user_id_alt, the persisted row (user_id only) cannot - # prove the canonical participant, so fail closed rather than matching on user_id alone. - if caller_keys_on_alt: - return False - return bool(row_uid) and row_uid == caller_uid - # No caller identity: the row carries only source + user_id, so a same-platform row can belong - # to a DIFFERENT chat or user — same platform alone is NOT ownership proof; fail closed - # (CWE-639). Same-chat resume of an ACTIVE session still works via the live-origin branch. - return False + return self._persisted_row_proves_owner(source, row) async def _resume_row_visible( self, source: SessionSource, row: dict, allow_all: bool @@ -482,30 +448,21 @@ class GatewaySessionCommandsMixin: Prevents cross-origin enumeration of session ids/previews via the numbered /resume list; keeps Matrix room-scoping, scopes every other platform to the caller unless admin ``--all``. + Cross-room enumeration is cross-ORIGIN data access, so ``--all`` needs a configured admin. """ - sid = str(row.get("id") or "") - if source.platform == Platform.MATRIX: - # Cross-room enumeration is cross-ORIGIN data access: gate the ``--all`` short-circuit - # behind a real configured admin, exactly like the non-Matrix branch below. - if allow_all and self._resume_caller_is_admin(source): - return True - return self._same_matrix_room(source, self._gateway_session_origin_for_id(sid)) if allow_all and self._resume_caller_is_admin(source): return True + sid = str(row.get("id") or "") + if source.platform == Platform.MATRIX: + return self._same_matrix_room(source, self._gateway_session_origin_for_id(sid)) return await self._resume_target_allowed(source, sid, allow_override=False) + # ------------------------------------------------------------------ /retry, /undo + async def _handle_retry_command(self, event: MessageEvent) -> str: """Handle /retry command - re-send the last user message.""" - source = event.source - session_entry = await self.async_session_store.get_or_create_session(source) - history = await self.async_session_store.load_transcript(session_entry.session_id) - - # Find the last *real* user message. Timeline bookkeeping rows carry role=user + - # display_kind (model_switch / async_delegation_complete / auto_continue / hidden); clients - # never count them as user turns. - last_user_idx = None - # The canonical projection excludes bookkeeping and pure handoffs while - # still recognizing a real ask embedded in a compaction carrier. + # The canonical projection excludes bookkeeping rows (role=user + display_kind) and pure + # handoffs while still recognizing a real ask embedded in a compaction carrier. from agent.context_compressor import ( history_before_user_originated_turn, retryable_user_text, @@ -513,60 +470,47 @@ class GatewaySessionCommandsMixin: user_originated_turn_view, ) - for i in range(len(history) - 1, -1, -1): - msg = history[i] - if user_originated_turn_view(msg) is not None: - last_user_idx = i - break - + source = event.source + session_entry = await self.async_session_store.get_or_create_session(source) + history = await self.async_session_store.load_transcript(session_entry.session_id) + last_user_idx = next( + (i for i in range(len(history) - 1, -1, -1) if user_originated_turn_view(history[i]) is not None), + None, + ) if last_user_idx is None: return t("gateway.retry.no_previous") - # Resolve the live text and the scaffold-preserving prefix before any - # transcript write. Messaging retries cannot reconstruct attachments; - # reject media/unknown content without truncating the session. + # Resolve the live text and the scaffold-preserving prefix before any transcript write. + # Messaging retries cannot reconstruct attachments; reject media/unknown content untouched. try: - truncated, live_view = history_before_user_originated_turn( - history, last_user_idx - ) + truncated, live_view = history_before_user_originated_turn(history, last_user_idx) last_user_msg = retryable_user_text(live_view.get("content")) handoff, _ = split_user_originated_turn(history[last_user_idx]) except ValueError as exc: return f"Cannot retry that message safely: {exc}" if handoff is not None: - # A composite carrier is one physical row containing both the retained summary and the - # live ask. Let the carrier-aware rewind archive that row/tail and insert its pure - # scaffold atomically. + # A composite carrier is one physical row holding both the retained summary and the live + # ask: the carrier-aware rewind archives that row/tail and inserts its pure scaffold + # atomically, reselecting the latest carrier on the same snapshot so a concurrent newer + # turn can never be removed while stale text is resent. try: rewind_result = await self.async_session_store.rewind_session( - session_entry.session_id, - 1, - require_retryable_composite=True, + session_entry.session_id, 1, require_retryable_composite=True, ) except ValueError as exc: return f"Cannot retry that message safely: {exc}" if rewind_result is None: return "Retry failed; transcript was not changed." - # The store reselects and validates the latest carrier on the same - # snapshot used by the atomic rewind. A concurrent newer turn can - # therefore never be removed while this handler resends stale text. last_user_msg = rewind_result["target_text"] - else: - # After in-place compaction the pre-compaction transcript lives on as - # active=0/compacted=1 rows under this session id. active_only preserves that archive; a - # separate existence probe could fail open or race with the write. - if not await self.async_session_store.rewrite_transcript( - session_entry.session_id, - truncated, - active_only=True, - reject_active_turn_lease=True, - ): - return "Retry failed; transcript was not changed." - # Reset stored token count — transcript was truncated - session_entry.last_prompt_tokens = 0 + # After in-place compaction the pre-compaction transcript lives on as active=0/compacted=1 + # rows under this id; active_only preserves that archive. + elif not await self.async_session_store.rewrite_transcript( + session_entry.session_id, truncated, active_only=True, reject_active_turn_lease=True, + ): + return "Retry failed; transcript was not changed." + session_entry.last_prompt_tokens = 0 # transcript was truncated - # Re-send by creating a fake text event with the old message retry_event = MessageEvent( text=last_user_msg, message_type=MessageType.TEXT, @@ -574,8 +518,6 @@ class GatewaySessionCommandsMixin: raw_message=event.raw_message, channel_prompt=event.channel_prompt, ) - - # Let the normal message handler process it return await self._handle_message(retry_event) async def _handle_undo_command(self, event: MessageEvent) -> str: @@ -584,31 +526,21 @@ class GatewaySessionCommandsMixin: from the active-only transcript (gateway analogue of the CLI's history surgery). """ source = event.source - - # Parse optional turn count: "/undo" → 1, "/undo 3" → 3. n = 1 raw_args = event.get_command_args().strip() if raw_args: try: - n = int(raw_args.split()[0]) + n = max(1, int(raw_args.split()[0])) except (ValueError, IndexError): return t("gateway.undo.invalid_count", arg=raw_args.split()[0]) - if n < 1: - n = 1 session_entry = await self.async_session_store.get_or_create_session(source) result = await self.async_session_store.rewind_session(session_entry.session_id, n) - if result is None: return t("gateway.undo.nothing") - - # Reset stored token count — transcript was truncated. - session_entry.last_prompt_tokens = 0 - # Evict the cached agent so the next turn rebuilds from the active-only - # transcript and memory providers refresh their per-session caches. + session_entry.last_prompt_tokens = 0 # transcript was truncated try: - session_key = build_session_key(source) - self._evict_cached_agent(session_key) + self._evict_cached_agent(build_session_key(source)) except Exception as e: logger.debug("undo: cached-agent eviction skipped: %s", e) @@ -621,6 +553,8 @@ class GatewaySessionCommandsMixin: preview=preview, ) + # --------------------------------------------------------------------- /compress + async def _handle_compress_command(self, event: MessageEvent) -> str: """Profile-scoping wrapper around manual /compress. @@ -665,14 +599,11 @@ class GatewaySessionCommandsMixin: count_before = getattr(compressor, "compression_count", 0) try: await self._run_in_executor_with_context( - lambda: agent._compress_context( - [], "", force=True, - ) + lambda: agent._compress_context([], "", force=True) ) except Exception as exc: return t("gateway.compress.failed", error=exc) - count_after = getattr(compressor, "compression_count", 0) - if count_after > count_before: + if getattr(compressor, "compression_count", 0) > count_before: return ( "🗜️ Codex app-server thread compacted (thread/compact). " "The transcript mirror is unchanged by design — the " @@ -689,27 +620,17 @@ class GatewaySessionCommandsMixin: Optional ``/compress <focus>`` tells the summariser what to preserve, discarding the rest. """ + from hermes_cli.partial_compress import extract_compress_flags, parse_partial_compress_args + source = event.source session_entry = await self.async_session_store.get_or_create_session(source) history = await self.async_session_store.load_transcript(session_entry.session_id) - if not history or len(history) < 4: return t("gateway.compress.not_enough") - # Parse args: either a focus topic (full compress) or the - # boundary-aware "here [N]" form (partial compress). - from hermes_cli.partial_compress import ( - extract_compress_flags, - parse_partial_compress_args, - rejoin_compressed_head_and_tail, - split_history_for_partial_compress, - ) - from agent.conversation_compression import ( - finalize_context_engine_compression_notification, - ) + # Strip --preview/--dry-run/--aggressive before positional parsing so the flags coexist with + # the boundary-aware "here [N]" (partial) and focus-topic (full) forms. _raw_args = (event.get_command_args() or "").strip() - # Strip --preview/--dry-run/--aggressive before positional parsing - # so the flags coexist with 'here [N]' / focus-topic forms. _raw_args, _preview, _aggressive = extract_compress_flags(_raw_args) partial, keep_last, focus_topic = parse_partial_compress_args(_raw_args) @@ -720,146 +641,115 @@ class GatewaySessionCommandsMixin: _agg_note = t("gateway.compress.aggressive_unsupported") if not _preview: return _agg_note - if _preview: return _compress_preview_reply(history, partial, keep_last, focus_topic, _agg_note) - try: - from agent.manual_compression_feedback import summarize_manual_compression - from agent.model_metadata import estimate_request_tokens_rough - from gateway.run import _platform_config_key - - session_key = self._session_key_for_source(source) - # Preserve the platform + stable gateway session identity of a normal turn so external - # context engines bind this agent to the original conversation, not a default "cli" host. - platform_key = ( - _platform_config_key(source.platform) if source.platform else None + return await self._run_manual_compression( + source, session_entry, history, partial, keep_last, focus_topic ) - model, runtime_kwargs = self._resolve_session_agent_runtime( - source=source, - session_key=session_key, - ) - if str(runtime_kwargs.get("api_mode") or "").lower() == "codex_app_server": - # codex app-server: the model's context is the server-side thread owned by the LIVE - # cached agent; a temporary agent has none (and finally-eviction would destroy the - # real context). Compact the live thread and KEEP the agent cached; no mirror fallback. - return await self._compress_codex_app_server_session( - session_key, session_entry.session_id - ) - if not runtime_kwargs.get("api_key"): - return t("gateway.compress.no_provider") - - # Pass the FULL transcript (tool results included), like auto-compress: user/assistant- - # only starves tool-result pruning and can trip the protect-first/last early-return. - msgs = [ - m for m in history - if m.get("role") in {"user", "assistant", "tool"} - ] - - # Boundary-aware split: only the head is summarized; the most recent `keep_last` - # exchanges are preserved verbatim. The split snaps the tail to a user-turn start so the - # rejoined transcript keeps role alternation valid. - tail: list = [] - head = msgs - if partial: - head, tail = split_history_for_partial_compress(msgs, keep_last) - if not tail: - # Degenerate split — fall back to full compression. - partial = False - head = msgs - - # Bind the temporary compression agent to the source's platform + stable gateway session - # key. Assign directly (not setdefault: a resolver value would be a stale placeholder, - # and it avoids duplicate-kwarg TypeError); platform only when known so None -> "cli" holds. - if platform_key is not None: - runtime_kwargs["platform"] = platform_key - runtime_kwargs["gateway_session_key"] = session_key - - tmp_agent = await self._build_manual_compression_agent( - session_entry.session_id, model, runtime_kwargs - ) - try: - # Estimate with system prompt + tool schemas included so the figure reflects real - # request pressure, not a transcript-only underestimate. Must be computed after - # tmp_agent is built so _cached_system_prompt/tools are populated. - _sys_prompt = getattr(tmp_agent, "_cached_system_prompt", "") or "" - _tools = getattr(tmp_agent, "tools", None) or None - approx_tokens = estimate_request_tokens_rough( - msgs, system_prompt=_sys_prompt, tools=_tools - ) - - compressor = tmp_agent.context_compressor - if not compressor.has_content_to_compress(head): - return t("gateway.compress.nothing_to_do") - - # Not a bare run_in_executor: the profile secret scope is a contextvar and the - # default-executor hop would drop it, making the compressor's aux-client credential - # resolution fail closed under multiplexing. - compressed, _ = await self._run_in_executor_with_context( - lambda: tmp_agent._compress_context( - head, - "", - approx_tokens=approx_tokens, - focus_topic=focus_topic, - force=True, - defer_context_engine_notification=True, - ) - ) - - # If _compress_context returned unchanged because a concurrent compression lock is - # held, tell the user clearly instead of showing the misleading "No changes from - # compression" no-op text. - _lock_skipped = getattr(tmp_agent, "_compression_skipped_due_to_lock", None) - if _lock_skipped is True or isinstance(_lock_skipped, str): - from agent.manual_compression_feedback import ( - describe_compression_lock_skip, - ) - return describe_compression_lock_skip(_lock_skipped) - - if partial and tail: - compressed = rejoin_compressed_head_and_tail(compressed, tail) - - await self._persist_manual_compression(tmp_agent, session_entry, source, compressed) - finalize_context_engine_compression_notification( - tmp_agent, - committed=True, - ) - new_tokens = estimate_request_tokens_rough( - compressed, system_prompt=_sys_prompt, tools=_tools - ) - summary = summarize_manual_compression( - msgs, - compressed, - approx_tokens, - new_tokens, - compression_state=compressor, - ) - finally: - finalize_context_engine_compression_notification( - tmp_agent, - committed=False, - ) - # Evict cached agent so next turn rebuilds system prompt - # from current files (SOUL.md, memory, etc.). - self._evict_cached_agent(session_key) - # Off-loop + bounded: temporary-agent teardown can block on - # subprocess/network/SQLite work. - await self._cleanup_agent_resources_off_loop( - tmp_agent, context="manual compression" - ) - return "\n".join(_manual_compression_reply_lines(summary, compressor, focus_topic)) except Exception as e: logger.warning("Manual compress failed: %s", e) return t("gateway.compress.failed", error=e) + async def _run_manual_compression( + self, source, session_entry, history: list, partial: bool, keep_last, focus_topic + ) -> str: + """Build a temporary agent, compress the transcript, persist, and describe the outcome.""" + from agent.conversation_compression import finalize_context_engine_compression_notification + from agent.manual_compression_feedback import summarize_manual_compression + from agent.model_metadata import estimate_request_tokens_rough + from gateway.run import _platform_config_key + from hermes_cli.partial_compress import ( + rejoin_compressed_head_and_tail, + split_history_for_partial_compress, + ) + + session_key = self._session_key_for_source(source) + # Preserve the platform + stable gateway session identity of a normal turn so external + # context engines bind this agent to the original conversation, not a default "cli" host. + platform_key = _platform_config_key(source.platform) if source.platform else None + model, runtime_kwargs = self._resolve_session_agent_runtime(source=source, session_key=session_key) + if str(runtime_kwargs.get("api_mode") or "").lower() == "codex_app_server": + # The model's context is the server-side thread owned by the LIVE cached agent; a + # temporary agent has none (and finally-eviction would destroy the real context). + return await self._compress_codex_app_server_session(session_key, session_entry.session_id) + if not runtime_kwargs.get("api_key"): + return t("gateway.compress.no_provider") + + # FULL transcript (tool results included), like auto-compress: user/assistant-only starves + # tool-result pruning and can trip the protect-first/last early-return. + msgs = [m for m in history if m.get("role") in {"user", "assistant", "tool"}] + # Boundary-aware split: only the head is summarized; the most recent `keep_last` exchanges + # stay verbatim. The split snaps the tail to a user-turn start so role alternation holds. + head, tail = msgs, [] + if partial: + head, tail = split_history_for_partial_compress(msgs, keep_last) + if not tail: # degenerate split — fall back to full compression + partial, head = False, msgs + + # Assign directly (not setdefault: a resolver value would be a stale placeholder, and it + # avoids duplicate-kwarg TypeError); platform only when known so None -> "cli" holds. + if platform_key is not None: + runtime_kwargs["platform"] = platform_key + runtime_kwargs["gateway_session_key"] = session_key + + tmp_agent = await self._build_manual_compression_agent(session_entry.session_id, model, runtime_kwargs) + try: + # Estimate with system prompt + tool schemas so the figure reflects real request + # pressure; computed after tmp_agent is built so _cached_system_prompt/tools exist. + _sys_prompt = getattr(tmp_agent, "_cached_system_prompt", "") or "" + _tools = getattr(tmp_agent, "tools", None) or None + approx_tokens = estimate_request_tokens_rough(msgs, system_prompt=_sys_prompt, tools=_tools) + + compressor = tmp_agent.context_compressor + if not compressor.has_content_to_compress(head): + return t("gateway.compress.nothing_to_do") + + # Not a bare run_in_executor: the profile secret scope is a contextvar and the default- + # executor hop would drop it, failing the aux-client credential resolution closed. + compressed, _ = await self._run_in_executor_with_context( + lambda: tmp_agent._compress_context( + head, + "", + approx_tokens=approx_tokens, + focus_topic=focus_topic, + force=True, + defer_context_engine_notification=True, + ) + ) + # Unchanged because a concurrent compression lock is held: say so instead of the + # misleading "No changes from compression" no-op text. + _lock_skipped = getattr(tmp_agent, "_compression_skipped_due_to_lock", None) + if _lock_skipped is True or isinstance(_lock_skipped, str): + from agent.manual_compression_feedback import describe_compression_lock_skip + return describe_compression_lock_skip(_lock_skipped) + + if partial and tail: + compressed = rejoin_compressed_head_and_tail(compressed, tail) + await self._persist_manual_compression(tmp_agent, session_entry, source, compressed) + finalize_context_engine_compression_notification(tmp_agent, committed=True) + new_tokens = estimate_request_tokens_rough(compressed, system_prompt=_sys_prompt, tools=_tools) + summary = summarize_manual_compression( + msgs, compressed, approx_tokens, new_tokens, compression_state=compressor, + ) + finally: + finalize_context_engine_compression_notification(tmp_agent, committed=False) + # Evict so the next turn rebuilds the system prompt from current files (SOUL.md, memory). + self._evict_cached_agent(session_key) + # Off-loop + bounded: temporary-agent teardown can block on subprocess/network/SQLite. + await self._cleanup_agent_resources_off_loop(tmp_agent, context="manual compression") + return "\n".join(_manual_compression_reply_lines(summary, compressor, focus_topic)) + async def _build_manual_compression_agent(self, session_id: str, model, runtime_kwargs: dict): """Build the throwaway AIAgent that performs a manual /compress rewrite of *session_id*.""" from run_agent import AIAgent from gateway.run import _GATEWAY_HYGIENE_PLATFORM, _seed_hygiene_system_prompt + from hermes_cli.config import load_config as _load_cfg + from utils import is_truthy_value as _is_truthy - # The manual compression helper runs outside the live session's fully initialized prompt - # environment and _compress_context may persist its cached system prompt — restore the - # exact live-session prompt so provider blocks are retained. + # This helper runs outside the live session's fully initialized prompt environment and + # _compress_context may persist its cached system prompt — restore the exact live-session + # prompt so provider blocks are retained. session_row = None get_session = getattr(self._session_db, "get_session", None) if callable(get_session): @@ -875,12 +765,9 @@ class GatewaySessionCommandsMixin: exc_info=True, ) - # This agent performs a lossy rewrite. When compression.checkpoint_required is on, the - # memory provider must be loaded so _compress_context() can write the pre-compression - # checkpoint; otherwise keep the historical fast path (no provider init). - from hermes_cli.config import load_config as _load_cfg - from utils import is_truthy_value as _is_truthy - + # Lossy rewrite: when compression.checkpoint_required is on, the memory provider must be + # loaded so _compress_context() can write the pre-compression checkpoint; otherwise keep + # the historical fast path (no provider init). _checkpoint_required = _is_truthy( ((_load_cfg() or {}).get("compression") or {}).get("checkpoint_required"), default=False, @@ -897,12 +784,12 @@ class GatewaySessionCommandsMixin: ) _seed_hygiene_system_prompt(tmp_agent, session_row) # Keep the real source platform during construction so external context engines bind - # correctly. If compression has to rebuild the prompt, stamp that provider-less fallback - # as stale for the next real gateway turn. + # correctly; if compression has to rebuild the prompt, stamp that provider-less fallback as + # stale for the next real gateway turn. tmp_agent.platform = _GATEWAY_HYGIENE_PLATFORM tmp_agent._print_fn = lambda *a, **kw: None - # Prevent close() from ending the newly rotated session — the gateway session entry now - # points at the new id and must remain open for the next user turn. + # close() must not end the newly rotated session — the gateway entry now points at the new + # id and must remain open for the next user turn. tmp_agent._end_session_on_close = False return tmp_agent @@ -910,13 +797,11 @@ class GatewaySessionCommandsMixin: """Commit a manual /compress result to the session store. _compress_context either rotated (new continuation id — write compressed messages into the - NEW session so the original stays searchable) or compacted in place (compression.in_place: - same id, transcript replaced). Persist BEFORE repointing the live session: repoint first + - failed DB write would leave the entry on an empty session while reporting success; a failed - write is fatal so old history stays reachable. Only rewrite when rotation produced a NEW id: - in-place compaction already archived + inserted rows and rewrite_transcript() - (active_only=False) would DELETE the archived turns; an unchanged id without in-place means - rotation FAILED and a rewrite would leave only the summary. + NEW session so the original stays searchable) or compacted in place (same id, transcript + replaced). Persist BEFORE repointing the live session: a failed write is fatal so old history + stays reachable. Only rewrite when rotation produced a NEW id: in-place compaction already + archived + inserted rows and rewrite_transcript() (active_only=False) would DELETE the + archived turns; an unchanged id without in-place means rotation FAILED. """ new_session_id = tmp_agent.session_id if new_session_id != session_entry.session_id: @@ -937,9 +822,11 @@ class GatewaySessionCommandsMixin: "preserving original transcript instead of overwriting " "it (#44794)." ) - # Reset stored token count — transcript changed, old value is stale + # Transcript changed — the stored token count is stale. await self.async_session_store.update_session(session_entry.session_key, last_prompt_tokens=0) + # ------------------------------------------------------------------------ /topic + async def _handle_topic_command(self, event: MessageEvent, args: str = "") -> str: """Handle /topic for Telegram DM user-managed topic sessions.""" source = event.source @@ -948,9 +835,8 @@ class GatewaySessionCommandsMixin: if not self._session_db: return self._session_db_unavailable_reply() - # Authorization: /topic activates multi-session mode and mutates SQLite side tables. - # Unauthorized senders (not in allowlist) must not be able to do that. Gateway routes - # already authorize the message before reaching here, but defense in depth. + # /topic activates multi-session mode and mutates SQLite side tables; routes already + # authorized the message, but defense in depth against unauthorized (non-allowlist) senders. auth_fn = getattr(self, "_is_user_authorized", None) if callable(auth_fn): try: @@ -960,15 +846,10 @@ class GatewaySessionCommandsMixin: logger.debug("Topic auth check failed", exc_info=True) args = event.get_command_args().strip() - - # /topic help — inline usage without leaving the bot. if args.lower() in {"help", "?", "-h", "--help"}: return self._telegram_topic_help_text() - - # /topic off — clean disable path so users don't have to edit the DB. if args.lower() in {"off", "disable", "stop"}: return await self._disable_telegram_topic_mode_for_chat(source) - if args: if not source.thread_id: return t("gateway.topic.restore_needs_topic") @@ -976,22 +857,23 @@ class GatewaySessionCommandsMixin: capabilities = await self._get_telegram_topic_capabilities(source) if capabilities.get("checked"): + blocked_key = None if capabilities.get("has_topics_enabled") is False: - # Debounce the BotFather screenshot: don't re-send on every - # /topic while threads are still disabled. + blocked_key = "gateway.topic.topics_disabled" + elif capabilities.get("allows_users_to_create_topics") is False: + blocked_key = "gateway.topic.topics_user_disallowed" + if blocked_key: + # Debounce the BotFather screenshot: don't re-send on every /topic while blocked. if self._should_send_telegram_capability_hint(source): await self._send_telegram_topic_setup_image(source) - return t("gateway.topic.topics_disabled") - if capabilities.get("allows_users_to_create_topics") is False: - if self._should_send_telegram_capability_hint(source): - await self._send_telegram_topic_setup_image(source) - return t("gateway.topic.topics_user_disallowed") + return t(blocked_key) + profile_name = self._telegram_topic_profile_name(source) try: await self._session_db.enable_telegram_topic_mode( chat_id=str(source.chat_id), user_id=str(source.user_id), - profile_name=self._telegram_topic_profile_name(source), + profile_name=profile_name, has_topics_enabled=capabilities.get("has_topics_enabled"), allows_users_to_create_topics=capabilities.get("allows_users_to_create_topics"), ) @@ -1001,36 +883,34 @@ class GatewaySessionCommandsMixin: if not source.thread_id: await self._ensure_telegram_system_topic(source) - - if source.thread_id: - try: - binding = await self._session_db.get_telegram_topic_binding( - chat_id=str(source.chat_id), - thread_id=str(source.thread_id), - profile_name=self._telegram_topic_profile_name(source), - ) - except Exception: - logger.debug("Failed to read Telegram topic binding", exc_info=True) - binding = None - if binding: - session_id = str(binding.get("session_id") or "") - title = None - try: - title = await self._session_db.get_session_title(session_id) - except Exception: - title = None - session_label = title or t("gateway.topic.untitled_session") - return t( - "gateway.topic.bound_status", - label=session_label, - session_id=session_id, - ) + return await self._telegram_topic_root_status_message(source) + try: + binding = await self._session_db.get_telegram_topic_binding( + chat_id=str(source.chat_id), + thread_id=str(source.thread_id), + profile_name=profile_name, + ) + except Exception: + logger.debug("Failed to read Telegram topic binding", exc_info=True) + binding = None + if not binding: return t("gateway.topic.thread_ready") + session_id = str(binding.get("session_id") or "") + try: + title = await self._session_db.get_session_title(session_id) + except Exception: + title = None + return t( + "gateway.topic.bound_status", + label=title or t("gateway.topic.untitled_session"), + session_id=session_id, + ) - return await self._telegram_topic_root_status_message(source) + # ------------------------------------------------------------------ /save, /title async def _handle_save_command(self, event: MessageEvent) -> str: """Handle /save — export the current session and send it as a document.""" + import tempfile from hermes_cli.session_export import ( SAVE_USAGE, default_save_filename, @@ -1039,15 +919,11 @@ class GatewaySessionCommandsMixin: ) parts = event.get_command_args().split() + redact = bool(parts) and parts[-1].lower() in ("redact", "--redact") + if redact: + parts = parts[:-1] if not parts: return SAVE_USAGE - redact = False - if parts[-1].lower() in ("redact", "--redact"): - redact = True - parts = parts[:-1] - if not parts: - return SAVE_USAGE - try: fmt = normalize_save_format(parts[0]) except ValueError as e: @@ -1056,82 +932,66 @@ class GatewaySessionCommandsMixin: source = event.source session_entry = await self.async_session_store.get_or_create_session(source) session_id = session_entry.session_id - if not self._session_db: return "Session database not available." + # The filename is echoed to the platform only — never trust path separators from chat input. filename = parts[1] if len(parts) > 1 else default_save_filename(session_id, fmt) - # The filename is echoed to the platform only — never trust path - # separators from chat input. filename = os.path.basename(filename) or default_save_filename(session_id, fmt) - # self._session_db is an AsyncSessionDB — every forwarded call is - # offloaded to a thread and must be awaited. + # self._session_db is an AsyncSessionDB — every forwarded call is awaited. export_data = await self._session_db.export_session(session_id) if not export_data: return f"No stored messages found for this session ({session_id})." - if redact: from hermes_cli.session_export_md import redact_session_data - export_data = redact_session_data(export_data) - import tempfile - temp_dir = tempfile.mkdtemp(prefix="hermes_save_") temp_path = os.path.join(temp_dir, filename) try: - # Off-loop: rendering a long session and writing it to disk are CPU/disk-bound and scale - # with transcript size (multi-MB for long sessions). Inline they stall every other chat - # on the gateway event loop (Pattern A). One thread hop covers both. + # Off-loop: rendering a long session and writing it are CPU/disk-bound and scale with + # transcript size (multi-MB); inline they stall every other chat on the event loop. def _render_and_write() -> None: rendered = render_session_for_save(export_data, fmt) with open(temp_path, "w", encoding="utf-8") as f: f.write(rendered) await asyncio.to_thread(_render_and_write) - adapter = self.get_adapter(source.platform) - if adapter: - await adapter.send_document( - chat_id=source.chat_id, - file_path=temp_path, - caption=f"Session export: {filename}", - file_name=filename, - ) - return "Export complete." - return "Platform adapter not found to send the document." + if not adapter: + return "Platform adapter not found to send the document." + await adapter.send_document( + chat_id=source.chat_id, + file_path=temp_path, + caption=f"Session export: {filename}", + file_name=filename, + ) + return "Export complete." except Exception as e: logger.warning("Session /save failed: %s", e) return f"Error exporting session: {e}" finally: - try: + with contextlib.suppress(Exception): os.remove(temp_path) os.rmdir(temp_dir) - except Exception: - pass async def _handle_title_command(self, event: MessageEvent) -> str: """Handle /title command — set or show the current session's title.""" source = event.source session_entry = await self.async_session_store.get_or_create_session(source) session_id = session_entry.session_id - if not self._session_db: return self._session_db_unavailable_reply() - # Ensure session exists in SQLite DB (it may only exist in session_store - # if this is the first command in a new session) - existing_title = await self._session_db.get_session_title(session_id) - if existing_title is None: - # Session doesn't exist in DB yet — create it + # The session may only exist in session_store so far (first command in a new session). + if await self._session_db.get_session_title(session_id) is None: try: await self._session_db.create_session( session_id=session_id, source=source.platform.value if source.platform else "unknown", user_id=source.user_id, - # Persist the messaging origin so a later /resume of this - # titled-but-now-inactive session can prove it belongs to the - # caller's chat/thread (IDOR scoping). + # Persist the messaging origin so a later /resume of this titled-but-inactive + # session can prove it belongs to the caller's chat/thread (IDOR scoping). chat_id=source.chat_id, chat_type=source.chat_type, thread_id=source.thread_id, @@ -1140,98 +1000,51 @@ class GatewaySessionCommandsMixin: pass # Session might already exist, ignore errors title_arg = event.get_command_args().strip() - if title_arg: - # Sanitize the title before setting - try: - from hermes_state import SessionDB - sanitized = SessionDB.sanitize_title(title_arg) - except ValueError as e: - return t("gateway.shared.warn_passthrough", error=e) - if not sanitized: - return t("gateway.title.empty_after_clean") - # Set the title - try: - if await self._session_db.set_session_title(session_id, sanitized): - # Propagate the user-chosen title to the visible Telegram forum topic name too. - # Auto-generated titles already rename the topic; without this, /title only - # updated the DB title and the topic kept its auto-assigned name. - schedule_rename = getattr( - self, "_schedule_telegram_topic_title_rename", None - ) - if callable(schedule_rename): - try: - await asyncio.to_thread(schedule_rename, source, session_id, sanitized) - except Exception: - logger.debug( - "Failed to rename Telegram topic from /title", - exc_info=True, - ) - return t("gateway.title.set_to", title=sanitized) - else: - return t("gateway.title.not_found") - except ValueError as e: - return t("gateway.shared.warn_passthrough", error=e) - else: - # Show the current title and session ID + if not title_arg: title = await self._session_db.get_session_title(session_id) if title: return t("gateway.title.current_with_title", session_id=session_id, title=title) - else: - return t("gateway.title.current_no_title", session_id=session_id) - - async def _handle_resume_command(self, event: MessageEvent) -> str: - """Handle /resume command — list or switch to a previous session.""" - if not self._session_db: - return self._session_db_unavailable_reply() - - source = await asyncio.to_thread( - self._normalize_source_for_session_key, event.source - ) - session_key = self._session_key_for_source(source) - raw_args = event.get_command_args().strip() + return t("gateway.title.current_no_title", session_id=session_id) try: - parts = shlex.split(raw_args) - except ValueError as exc: - return t("gateway.resume.parse_error", error=exc) - allow_all = "--all" in parts - allow_cross_room = "--cross-room" in parts - name = " ".join(p for p in parts if p not in {"--all", "--cross-room"}).strip() - - # Strip common outer brackets/quotes users may type literally from the - # usage hint (e.g. ``/resume <abc123>``). Mirrors the CLI behavior. - if len(name) >= 2 and ( - (name[0] == "<" and name[-1] == ">") - or (name[0] == "[" and name[-1] == "]") - or (name[0] == '"' and name[-1] == '"') - or (name[0] == "'" and name[-1] == "'") - ): - name = name[1:-1].strip() - - async def _list_titled_sessions() -> list[dict]: - """Titled sessions visible to the caller (origin-scoped unless admin ``--all``).""" - user_source = source.platform.value if source.platform else None - widen = allow_all and self._resume_caller_is_admin(source) - sessions = await self._session_db.list_sessions_rich( - source=user_source, - session_key=None if widen else session_key, - limit=10, - ) - titled = [s for s in sessions if s.get("title")][:10] - return [s for s in titled if await self._resume_row_visible(source, s, allow_all)] - - if not name: - # List recent titled sessions for this user/platform + from hermes_state import SessionDB + sanitized = SessionDB.sanitize_title(title_arg) + except ValueError as e: + return t("gateway.shared.warn_passthrough", error=e) + if not sanitized: + return t("gateway.title.empty_after_clean") + try: + if not await self._session_db.set_session_title(session_id, sanitized): + return t("gateway.title.not_found") + except ValueError as e: + return t("gateway.shared.warn_passthrough", error=e) + # Propagate the user-chosen title to the visible Telegram forum topic name too (auto- + # generated titles already rename the topic). + schedule_rename = getattr(self, "_schedule_telegram_topic_title_rename", None) + if callable(schedule_rename): try: - titled = await _list_titled_sessions() - return self._resume_listing_reply(source, titled, allow_all) - except Exception as e: - logger.debug("Failed to list titled sessions: %s", e) - return t("gateway.resume.list_failed", error=e) + await asyncio.to_thread(schedule_rename, source, session_id, sanitized) + except Exception: + logger.debug("Failed to rename Telegram topic from /title", exc_info=True) + return t("gateway.title.set_to", title=sanitized) - # Resolve a numbered choice or a title to a session ID. + # -------------------------------------------------------------- /resume, /sessions + + async def _list_titled_sessions(self, source, session_key: str, allow_all: bool) -> list[dict]: + """Titled sessions visible to the caller (origin-scoped unless admin ``--all``).""" + widen = allow_all and self._resume_caller_is_admin(source) + sessions = await self._session_db.list_sessions_rich( + source=source.platform.value if source.platform else None, + session_key=None if widen else session_key, + limit=10, + ) + titled = [s for s in sessions if s.get("title")][:10] + return [s for s in titled if await self._resume_row_visible(source, s, allow_all)] + + async def _resolve_resume_target(self, source, session_key: str, name: str, allow_all: bool): + """``(target_id, display_name)`` for a numbered choice, session id or title; else an error reply.""" if name.isdigit(): try: - titled = await _list_titled_sessions() + titled = await self._list_titled_sessions(source, session_key, allow_all) except Exception as e: logger.debug("Failed to list titled sessions for numeric resume: %s", e) return t("gateway.resume.list_failed", error=e) @@ -1239,75 +1052,94 @@ class GatewaySessionCommandsMixin: if index < 1 or index > len(titled): return t("gateway.resume.out_of_range", index=index) target = titled[index - 1] - target_id = target.get("id") - name = target.get("title") or name + target_id, name = target.get("id"), target.get("title") or name else: - # Try direct session ID lookup first (so `/resume <session_id>` - # works in the gateway, not just `/resume <title>`). + # Direct session ID lookup first so `/resume <session_id>` works, not just `<title>`. session = await self._session_db.get_session(name) - if session: - target_id = session["id"] - else: - target_id = await self._session_db.resolve_session_by_title(name) + target_id = session["id"] if session else await self._session_db.resolve_session_by_title(name) if not target_id: return t("gateway.resume.not_found", name=name) - # Compression creates child continuations that hold the live transcript. - # Follow that chain so gateway /resume matches CLI behavior (#15000). + # Compression creates child continuations that hold the live transcript; follow the chain + # so gateway /resume matches CLI behavior. try: target_id = await self._session_db.resolve_resume_session_id(target_id) except Exception as e: logger.debug("Failed to resolve resume continuation for %s: %s", target_id, e) + return target_id, name + async def _resume_access_denied_reply( + self, source, target_id: str, name: str, allow_all: bool, allow_cross_room: bool + ) -> Optional[str]: + """IDOR guard: a session id/title is a routing handle, not authority. Bind /resume to the + caller's own room (Matrix) or platform/user/chat (every other adapter).""" if source.platform == Platform.MATRIX: target_origin = self._gateway_session_origin_for_id(target_id) - if not self._same_matrix_room(source, target_origin) and not allow_cross_room: - if target_origin is None: - return t("gateway.resume.matrix_blocked_no_origin", name=name) - return t( - "gateway.resume.matrix_blocked_other_room", - room=target_origin.chat_name or target_origin.chat_id, - name=name, - ) - elif not await self._resume_target_allowed( - source, target_id, allow_override=(allow_all or allow_cross_room) - ): - # IDOR guard: a session id/title is a routing handle, not authority. Bind /resume to the - # caller's own platform/user/chat on every non-Matrix adapter so one user can't attach - # to another's persisted transcript. - return t("gateway.resume.blocked_not_owner", name=name) + if self._same_matrix_room(source, target_origin) or allow_cross_room: + return None + if target_origin is None: + return t("gateway.resume.matrix_blocked_no_origin", name=name) + return t( + "gateway.resume.matrix_blocked_other_room", + room=target_origin.chat_name or target_origin.chat_id, + name=name, + ) + if await self._resume_target_allowed(source, target_id, allow_override=(allow_all or allow_cross_room)): + return None + return t("gateway.resume.blocked_not_owner", name=name) + + async def _handle_resume_command(self, event: MessageEvent) -> str: + """Handle /resume command — list or switch to a previous session.""" + if not self._session_db: + return self._session_db_unavailable_reply() + + source = await asyncio.to_thread(self._normalize_source_for_session_key, event.source) + session_key = self._session_key_for_source(source) + try: + parts = shlex.split(event.get_command_args().strip()) + except ValueError as exc: + return t("gateway.resume.parse_error", error=exc) + allow_all = "--all" in parts + allow_cross_room = "--cross-room" in parts + name = _strip_resume_name(parts) + + if not name: + try: + titled = await self._list_titled_sessions(source, session_key, allow_all) + return self._resume_listing_reply(source, titled, allow_all) + except Exception as e: + logger.debug("Failed to list titled sessions: %s", e) + return t("gateway.resume.list_failed", error=e) + + resolved = await self._resolve_resume_target(source, session_key, name, allow_all) + if isinstance(resolved, str): + return resolved + target_id, name = resolved + denied = await self._resume_access_denied_reply(source, target_id, name, allow_all, allow_cross_room) + if denied is not None: + return denied - # Check if already on that session current_entry = await self.async_session_store.get_or_create_session(source) if current_entry.session_id == target_id: return t("gateway.resume.already_on", name=name) - # Clear any running agent for this session key self._release_running_agent_state(session_key) - - # Switch the session entry to point at the old session new_entry = await self.async_session_store.switch_session(session_key, target_id) if not new_entry: return t("gateway.resume.switch_failed") - # Conversation boundary: clear ALL conversation-scoped per-session state (model/reasoning - # overrides #10702, one-turn restores, model notes, last-resolved cache #58403, /queue - # overflow) + security state in one funnel call. + # overrides, one-turn restores, model notes, last-resolved cache, /queue overflow) + + # security state in one funnel call. self._clear_conversation_scope(session_key, reason="resume") - - # Evict any cached agent for this session so the next message rebuilds with the correct - # session_id end-to-end — mirrors /branch and /reset. Otherwise the cached AIAgent (and its - # memory provider, which cached _session_id at initialize()) keeps writing to the wrong session. + # Evict so the next message rebuilds with the correct session_id end-to-end — otherwise the + # cached AIAgent (and its memory provider, which cached _session_id at initialize()) keeps + # writing to the wrong session. Mirrors /branch and /reset. self._evict_cached_agent(session_key) - # Get the title for confirmation title = await self._session_db.get_session_title(target_id) or name - - # Count messages for context history = await self.async_session_store.load_transcript(target_id) msg_count = len([m for m in history if m.get("role") == "user"]) if history else 0 - msg_part = f" ({msg_count} message{'s' if msg_count != 1 else ''})" if msg_count else "" - if source.platform == Platform.MATRIX and allow_cross_room: + msg_part = f" ({msg_count} message{'s' if msg_count != 1 else ''})" if msg_count else "" return t( "gateway.resume.matrix_cross_room_success", title=title, @@ -1359,33 +1191,25 @@ class GatewaySessionCommandsMixin: query_session_listing, ) - raw_args = event.get_command_args().strip() try: include_all, include_unnamed, target, search_query = ( - parse_session_listing_args(raw_args) + parse_session_listing_args(event.get_command_args().strip()) ) except ValueError as exc: return t("gateway.resume.parse_error", error=exc) - if search_query == "": return "Usage: `/sessions search <query>`" - if target: resume_event = dataclasses.replace(event, text=f"/resume {target}") return await self._handle_resume_command(resume_event) - source = await asyncio.to_thread( - self._normalize_source_for_session_key, event.source - ) + source = await asyncio.to_thread(self._normalize_source_for_session_key, event.source) session_key = self._session_key_for_source(source) - - # A cross-origin listing (`/sessions all`) is honored only for an admin, mirroring the - # `/resume --all` override. `all` is just a parsed user argument; ungated, any caller could - # enumerate other origins' session ids/titles/previews — the enumeration half of the IDOR. + # A cross-origin listing (`/sessions all`) is honored only for an admin, mirroring + # `/resume --all`; ungated, any caller could enumerate other origins' session ids/titles/ + # previews — the enumeration half of the IDOR. A non-admin `all` gets explicit feedback + # rather than a silently scoped list that reads as "my session vanished". cross_origin = include_all and self._resume_caller_is_admin(source) - # Don't silently no-op a requested widening: a non-admin `/sessions all` - # used to render the same scoped list with zero feedback, which reads - # as "my session vanished" (community report, Aug 2026). scope_notice = None if include_all and not cross_origin: scope_notice = ( @@ -1403,83 +1227,64 @@ class GatewaySessionCommandsMixin: include_all_sources=cross_origin, include_unnamed=include_unnamed, search_query=search_query, - # Search filters at SQL level, so over-fetch before the visibility - # cut: origin-invisible matches would otherwise consume the page. + # Search filters at SQL level, so over-fetch before the visibility cut: origin-invisible + # matches would otherwise consume the page. limit=50 if search_query else 10, exclude_sources=["tool"], ) if not cross_origin: - # Scope the listing to the caller's own origin on every adapter so - # session ids/previews from other users/rooms aren't enumerable. - rows = [ - row for row in rows - if await self._resume_row_visible(source, row, allow_all=False) - ] + # Scope to the caller's own origin so other users'/rooms' ids/previews aren't enumerable. + rows = [row for row in rows if await self._resume_row_visible(source, row, allow_all=False)] rows = rows[:10] if search_query: title = f"Sessions matching “{search_query}”" else: title = "Sessions" if include_unnamed else "Named Sessions" return format_gateway_session_listing( - rows, - include_source=cross_origin, - title=title, - notice=scope_notice, + rows, include_source=cross_origin, title=title, notice=scope_notice, ) + # ----------------------------------------------------------------------- /branch + async def _handle_branch_command(self, event: MessageEvent) -> str: """Handle /branch [name] — fork the current session into a new independent copy so the user can explore a different approach without losing the original. """ + import json as _json import uuid as _uuid + from datetime import datetime as _dt if not self._session_db: return self._session_db_unavailable_reply() source = event.source session_key = self._session_key_for_source(source) - - # Load the current session and its transcript current_entry = await self.async_session_store.get_or_create_session(source) history = await self.async_session_store.load_transcript(current_entry.session_id) if not history: return t("gateway.branch.no_conversation") - branch_name = event.get_command_args().strip() - - # Generate the new session ID - from datetime import datetime as _dt - now = _dt.now() - timestamp_str = now.strftime("%Y%m%d_%H%M%S") - short_uuid = _uuid.uuid4().hex[:6] - new_session_id = f"{timestamp_str}_{short_uuid}" - - # Determine branch title - if branch_name: - branch_title = branch_name - else: + new_session_id = f"{_dt.now().strftime('%Y%m%d_%H%M%S')}_{_uuid.uuid4().hex[:6]}" + branch_title = event.get_command_args().strip() + if not branch_title: current_title = await self._session_db.get_session_title(current_entry.session_id) - base = current_title or "branch" - branch_title = await self._session_db.get_next_title_in_lineage(base) - + branch_title = await self._session_db.get_next_title_in_lineage(current_title or "branch") parent_session_id = current_entry.session_id # Serialize the parent's full origin (same shape as the reset path's db_create_kwargs in - # gateway/session.py, #82633) so the branch row carries complete identity from birth. Prefer - # the live entry's origin (it may hold richer metadata than the triggering event's source). + # gateway/session.py) so the branch row carries complete identity from birth. Prefer the + # live entry's origin (it may hold richer metadata than the triggering event's source). _branch_origin = current_entry.origin or source _branch_origin_json = None if _branch_origin is not None: - try: - import json as _json - + with contextlib.suppress(Exception): _branch_origin_json = _json.dumps(_branch_origin.to_dict()) - except Exception: - _branch_origin_json = None - # Create the new session with parent link. Persist a stable ``_branched_from`` marker in - # model_config so list_sessions_rich() keeps the branch visible in /resume and /sessions - # even after the parent is reopened and re-ended with a different end_reason. + # Persist a stable ``_branched_from`` marker in model_config so list_sessions_rich() keeps + # the branch visible in /resume and /sessions even after the parent is reopened and + # re-ended with a different end_reason. Forward ALL gateway routing columns at CREATE time: + # otherwise they're NULL until switch_session() records the peer, and a crash in between + # leaves the branch unroutable (by chat/thread lookup and by /resume's IDOR guard). try: await self._session_db.create_session( session_id=new_session_id, @@ -1487,12 +1292,6 @@ class GatewaySessionCommandsMixin: model=(self.config.get("model", {}) or {}).get("default") if isinstance(self.config, dict) else None, model_config={"_branched_from": parent_session_id}, parent_session_id=parent_session_id, - # Forward ALL gateway routing columns at CREATE time: otherwise they're NULL until - # switch_session() calls _record_gateway_session_peer(), and a crash in between (each - # append_message is best-effort) leaves the branch unroutable — by chat/thread lookup - # and by /resume's IDOR guard. user_id feeds the full-peer-tuple fallback lookup; - # origin_json/display_name complete the identity (same shape as session.py's reset - # path) so state.db consumers see a fully formed row with no backfill gap. user_id=source.user_id, session_key=session_key, chat_id=source.chat_id, @@ -1505,27 +1304,19 @@ class GatewaySessionCommandsMixin: logger.error("Failed to create branch session: %s", e) return t("gateway.branch.create_failed", error=e) - # Copy conversation history to the new session in bounded-chunk transactions: one txn per - # row was the removed write-amplification pattern, and a history can be hundreds of rows. - # Best-effort like the old loop — a failed copy still yields a usable (partial) branch. - try: + # Bounded-chunk transactions (a history can be hundreds of rows). Best-effort: a failed + # copy still yields a usable (partial) branch. + with contextlib.suppress(Exception): await self._session_db.append_messages_batch( new_session_id, [_branch_row(msg) for msg in history], chunk_rows=500, ) - except Exception: - pass # Best-effort copy - - # Set title with contextlib.suppress(Exception): await self._session_db.set_session_title(new_session_id, branch_title) - # Switch the session store entry to the new session new_entry = await self.async_session_store.switch_session(session_key, new_session_id) if not new_entry: return t("gateway.branch.switch_failed") self._clear_session_boundary_security_state(session_key) - - # Evict any cached agent for this session self._evict_cached_agent(session_key) msg_count = len([m for m in history if m.get("role") == "user"])