diff --git a/tests/tui_gateway/test_hosted_room_peer_http.py b/tests/tui_gateway/test_hosted_room_peer_http.py index 3f578607ff..c43e85e287 100644 --- a/tests/tui_gateway/test_hosted_room_peer_http.py +++ b/tests/tui_gateway/test_hosted_room_peer_http.py @@ -463,7 +463,6 @@ def test_invalid_room_dispatch_http_403_is_definitively_not_admitted(monkeypatch assert caught.value.error_code == "room_capability_catalog_changed" assert caught.value.not_admitted is True assert caught.value.ambiguous is False - assert caught.value.needs_capability_refresh is True def test_capability_mismatch_requires_reauthorization_without_retry(tmp_path): @@ -525,7 +524,6 @@ def test_peer_http_error_body_is_never_exposed_or_logged(monkeypatch, caplog): assert caught.value.status_code == 500 assert hostile not in str(caught.value) - assert caught.value.error_message is None assert hostile not in caplog.text diff --git a/tui_gateway/agent_callbacks.py b/tui_gateway/agent_callbacks.py index 21592409a2..ea30f137c8 100644 --- a/tui_gateway/agent_callbacks.py +++ b/tui_gateway/agent_callbacks.py @@ -1,9 +1,6 @@ -"""Agent callback wiring: child-session live mirror, per-session agent callbacks, -personality overlay, background/preview agent kwargs, agent reset. - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so they reference server.py globals bare. -""" +"""Agent callback wiring: child-session live mirror, per-session agent callbacks, personality +overlay, background/preview agent kwargs, agent reset. Bodies are rebound onto server.py's +globals at install time (method_ctx.bind_module), so they reference server.py globals bare.""" from __future__ import annotations @@ -12,18 +9,18 @@ import threading from .method_ctx import bind_module -# Child-session live mirror: a delegated child's activity reaches the gateway only -# as relayed ``subagent.*`` events on the PARENT sid, so a window opened on the -# child's own session would sit silent until the run persists. Translate them into -# the native stream events emitted on the CHILD sid (write_json routes by sid). +# Child-session live mirror: a delegated child's activity reaches the gateway only as +# relayed ``subagent.*`` events on the PARENT sid; translate them into native stream +# events on the CHILD sid (write_json routes by sid) so its own window is not silent. _child_mirrors: dict[str, dict] = {} _child_mirrors_lock = threading.Lock() -# Child session ids with a run in flight (refreshed per relayed event, popped on -# complete) so a lazy watch resume reports running=true during a silent long tool. +# Child sids with a run in flight (refreshed per relayed event, popped on complete) so a +# lazy watch resume reports running=true during a silent long tool. _active_child_runs: dict[str, float] = {} -# Anything quiet this long lost its completion event (callback raised, parent -# crashed) — don't pin "running". +# Anything quiet this long lost its completion event — don't pin "running". _CHILD_RUN_STALE_S = 3600.0 +_CHILD_DELTA_EVENTS = {"subagent.thinking": "reasoning.delta", "subagent.text": "message.delta", + "subagent.start": "message.delta"} def _child_run_active(child_key: str) -> bool: @@ -35,15 +32,13 @@ def _mirror_subagent_to_child(event_type: str, payload: dict) -> None: child_key = str(payload.get("child_session_id") or "") if not child_key: return - # Liveness registry first: accurate with no window open, so one opened mid-run - # immediately knows the child is busy. + # Liveness registry first: accurate with no window open (one opened mid-run knows busy). if event_type == "subagent.complete": _active_child_runs.pop(child_key, None) else: _active_child_runs[child_key] = time.time() - # Mirror only into a live watch session NOT upgraded to a full agent: an - # upgraded one owns a real native stream and mirroring would interleave two - # turns on one sid. Either way drop state so a reopened window starts fresh. + # Mirror only into a live watch session NOT upgraded to a full agent (an upgraded one owns + # a real native stream). Either way drop state so a reopened window starts fresh. live = _find_live_session_by_key(child_key) if live is None or live[1].get("agent") is not None: with _child_mirrors_lock: @@ -51,19 +46,16 @@ def _mirror_subagent_to_child(event_type: str, payload: dict) -> None: return csid = live[0] text = str(payload.get("text") or "") - # thinking/text/start (the child's goal, as a one-time header) are plain deltas. - delta = {"subagent.thinking": "reasoning.delta", "subagent.text": "message.delta", - "subagent.start": "message.delta"} with _child_mirrors_lock: st = _child_mirrors.setdefault(child_key, {"seq": 0, "open_tool": None, "started": False}) if not st["started"]: st["started"] = True _emit("message.start", csid) - if event_type in delta: + # thinking/text/start (the child's goal, as a one-time header) are plain deltas. + if event_type in _CHILD_DELTA_EVENTS: if text: - if event_type == "subagent.start": - text = f"{text}\n" - _emit(delta[event_type], csid, {"text": text}) + _emit(_CHILD_DELTA_EVENTS[event_type], csid, + {"text": f"{text}\n" if event_type == "subagent.start" else text}) return if event_type not in ("subagent.tool", "subagent.complete"): return @@ -85,11 +77,11 @@ def _mirror_subagent_to_child(event_type: str, payload: dict) -> None: def _agent_cbs(sid: str) -> dict: def _read_block(event: str, timeout: int): - # read_terminal / read_preview (desktop GUI): blocking bridge like clarify; the - # preview read gets longer since a URL tab extracts text from a live page. + # read_terminal / read_preview (desktop GUI): blocking bridge like clarify; the preview + # read gets longer since a URL tab extracts text from a live page. return lambda start=None, count=None: _block( - event, sid, {k: v for k, v in (("start", start), ("count", count)) if v is not None}, timeout=timeout - ) + event, sid, {k: v for k, v in (("start", start), ("count", count)) if v is not None}, + timeout=timeout) callbacks = { "tool_start_callback": lambda tc_id, name, args: _on_tool_start(sid, tc_id, name, args), @@ -98,49 +90,42 @@ def _agent_cbs(sid: str) -> dict: sid, event_type, name, preview, args, **kwargs), "tool_gen_callback": lambda name: _tool_progress_enabled(sid) and _emit("tool.generating", sid, {"name": name}), "thinking_callback": lambda text: _emit("thinking.delta", sid, {"text": text}), - # Affection reaction (ily / <3 / good bot) → hearts; core-detected so TUI and desktop share it. + # Affection reaction (ily / <3 / good bot) → hearts; core-detected so TUI/desktop share it. "reaction_callback": lambda kind: _emit("reaction", sid, {"kind": kind}), "reasoning_callback": lambda text: _emit( - "reasoning.delta", sid, {"text": text, **({"verbose": True} if _session_verbose(sid) else {})} - ), + "reasoning.delta", sid, {"text": text, **({"verbose": True} if _session_verbose(sid) else {})}), "status_callback": lambda kind, text=None: _status_update(sid, str(kind), None if text is None else str(text)), - # Credits/notice spine: AgentNotice → notification.show; recovery clear → notification.clear. + # Credits/notice spine: AgentNotice → notification.show; recovery → notification.clear. "notice_callback": lambda n: _emit( "notification.show", sid, - {"text": n.text, "level": n.level, "kind": n.kind, "ttl_ms": n.ttl_ms, "key": n.key, "id": n.id}, - ), + {"text": n.text, "level": n.level, "kind": n.kind, "ttl_ms": n.ttl_ms, "key": n.key, "id": n.id}), "notice_clear_callback": lambda key: _emit("notification.clear", sid, {"key": key}), "clarify_callback": lambda q, c, multi_select=False, questions=None: ( _clarify_block(sid, q, c, multi_select=multi_select, questions=questions)), "read_terminal_callback": _read_block("terminal.read.request", 30), "read_preview_callback": _read_block("preview.read.request", 45), - # drive_preview / annotate_preview (desktop GUI): renderer drives the preview webview and - # answers with outcome + refreshed element inventory; same budget as the preview read it ends with. + # drive_preview / annotate_preview (desktop GUI): same budget as the preview read it ends with. "drive_preview_callback": lambda payload: _block("preview.act.request", sid, dict(payload), timeout=45), # read_window_below (desktop GUI): main process enumerates native windows. "read_window_below_callback": lambda: _block("window.read.request", sid, {}, timeout=30), - # setup_mcp (desktop GUI): consent card + install/enable/OAuth. Long timeout on purpose (typing - # an API key, browser OAuth); like clarify, timeout returns "unanswered" and a late answer is tolerated. + # setup_mcp (desktop GUI): consent card + install/enable/OAuth; long timeout on purpose + # (typing an API key, browser OAuth) and, like clarify, a late answer is tolerated. "setup_mcp_callback": lambda server, action, reason: _block( - "mcp.setup.request", sid, {"server": server, "action": action, "reason": reason}, timeout=600 - ), + "mcp.setup.request", sid, {"server": server, "action": action, "reason": reason}, timeout=600), # tour (desktop GUI): renderer drives driver.js and answers tour.respond. "tour_callback": lambda payload: _tour_request(sid, payload)} - # Interim assistant commentary (text alongside tool calls). Gated on - # display.interim_assistant_messages (default true); _run_prompt_submit overwrites - # it per turn and clears it in its finally so a stale closure can't fire. + # Interim assistant commentary (text alongside tool calls), gated on display.interim_assistant_ + # messages; _run_prompt_submit overwrites it per turn and clears it so a stale closure can't fire. if _load_interim_assistant_messages(): callbacks["interim_assistant_callback"] = lambda text, *, already_streamed=False: _emit( "message.interim", sid, {"text": str(text), "already_streamed": bool(already_streamed)}) - return callbacks def _apply_project_workspace(task_id: str, path: str, _name: str = "") -> None: - """Intentional workspace move from the project_* tools: re-anchor the live - session's cwd and push session.info so the desktop follows. This is the ONLY - auto-cwd path — driven by an explicit tool call, never a terminal `cd`.""" + """Intentional workspace move from the project_* tools: re-anchor the live session's cwd + and push session.info. The ONLY auto-cwd path — an explicit tool call, never a `cd`.""" if not path: return # task_id is the durable session_key; _sessions (and desktop event routing) key by sid. @@ -149,24 +134,19 @@ def _apply_project_workspace(task_id: str, path: str, _name: str = "") -> None: sid, session = (key, _sessions[key]) if key in _sessions else next( ((s, c) for s, c in _sessions.items() if c.get("session_key") == key or getattr(c.get("agent"), "session_id", None) == key), - ("", None), - ) - if session is None: - return + ("", None)) resolved = os.path.abspath(os.path.expanduser(str(path))) - if not os.path.isdir(resolved): + if session is None or not os.path.isdir(resolved): return - session["cwd"] = resolved - session["explicit_cwd"] = True - session["cwd_from_settle"] = False # explicit switch supersedes a settle-adopted cwd + # explicit switch supersedes a settle-adopted cwd + session.update(cwd=resolved, explicit_cwd=True, cwd_from_settle=False) _register_session_cwd(session) _persist_session_cwd_and_schedule_git_meta(session, resolved) try: agent = session.get("agent") info = _session_info(agent, session) if agent is not None else { "cwd": resolved, "branch": _git_branch_for_cwd(resolved), - "project": _project_info_for_cwd(resolved), "lazy": True, - } + "project": _project_info_for_cwd(resolved), "lazy": True} _emit("session.info", sid, info) except Exception: logger.debug("failed to emit session.info after project workspace move", exc_info=True) @@ -176,19 +156,17 @@ def _wire_callbacks(sid: str): from tools.terminal_tool import set_sudo_password_callback from tools.skills_tool import set_secret_capture_callback from tools.project_tools import set_project_workspace_callback - set_sudo_password_callback(lambda: _block("sudo.request", sid, {}, timeout=120)) - set_project_workspace_callback(_apply_project_workspace) def secret_cb(env_var, prompt, metadata=None): - pl = {"prompt": prompt, "env_var": env_var} - if metadata: - pl["metadata"] = metadata + pl = {"prompt": prompt, "env_var": env_var, **({"metadata": metadata} if metadata else {})} val = _block("secret.request", sid, pl) if not val: return {"success": True, "stored_as": env_var, "validated": False, "skipped": True, "message": "skipped"} from hermes_cli.config import save_env_value_secure return {**save_env_value_secure(env_var, val), "skipped": False, "message": "ok"} + set_sudo_password_callback(lambda: _block("sudo.request", sid, {}, timeout=120)) + set_project_workspace_callback(_apply_project_workspace) set_secret_capture_callback(secret_cb) @@ -199,12 +177,10 @@ def _available_personalities(cfg: dict | None = None) -> dict: def _validate_personality(value: str, cfg: dict | None = None) -> tuple[str, str]: - """Resolve a requested personality to (name, prompt) or raise ValueError. Same - contract as hermes_cli.personality.resolve_personality, but goes through the - module-level _available_personalities so tests keep a single patch point.""" + """(name, prompt) for a requested personality or ValueError; like resolve_personality but + via the module-level _available_personalities so tests keep a single patch point.""" from hermes_cli.personality import normalize_personality_name, render_personality_prompt - name = normalize_personality_name(value) - if not name: + if not (name := normalize_personality_name(value)): return "", "" personalities = _available_personalities(cfg) if name not in personalities: @@ -221,16 +197,13 @@ def _prompt_text(value) -> str: def _apply_personality_to_session( sid: str, session: dict, new_prompt: str, personality: str = "") -> tuple[bool, dict | None]: - """Apply a personality change to a live session without resetting history: the - ephemeral system prompt is updated in place (appended at API-call time, so - prompt-cache hits survive) plus a pivot marker so the model stops pattern-matching - its earlier tone. Returns (history_reset=False, info).""" + """Apply a personality change without resetting history: the ephemeral system prompt is + updated in place (appended at API-call time, so prompt-cache hits survive) plus a pivot + marker so the model stops pattern-matching its earlier tone. Returns (False, info).""" if not session: return False, None session["personality"] = personality - - agent = session.get("agent") - if not agent: + if not (agent := session.get("agent")): return False, None agent.ephemeral_system_prompt = new_prompt or None marker = ( @@ -239,12 +212,10 @@ def _apply_personality_to_session( f"accordingly: {new_prompt}]" if new_prompt else "[System: The user has cleared the personality overlay. " - "From this point forward, respond in your normal default style.]" - ) - # Like the model-switch marker: role=user so strict providers accept it - # mid-conversation, but `display_kind` keeps it out of the - # `truncate_before_user_ordinal` addressing space (untagged, every rewind would - # land one turn early and `replace_messages` hard-delete the difference). + "From this point forward, respond in your normal default style.]") + # Like the model-switch marker: role=user so strict providers accept it mid-conversation, + # but `display_kind` keeps it out of the `truncate_before_user_ordinal` addressing space + # (untagged, every rewind would land one turn early and hard-delete the difference). with session["history_lock"]: session["history"].append({"role": "user", "content": marker, "display_kind": "personality_switch"}) session["history_version"] = int(session.get("history_version", 0)) + 1 @@ -256,8 +227,7 @@ def _apply_personality_to_session( def _cfg_max_turns(cfg: dict, default: int) -> int: from hermes_cli.config import resolve_turn_limit as _resolve_turn_limit # Env override wins; resolve_turn_limit makes "none"/"unlimited"/0 first-class spellings. - env_val = os.environ.get("HERMES_TUI_MAX_TURNS") - if env_val: + if env_val := os.environ.get("HERMES_TUI_MAX_TURNS"): return _resolve_turn_limit(env_val, default=default) raw = (cfg.get("agent") or {}).get("max_turns") if raw is None: @@ -267,12 +237,7 @@ def _cfg_max_turns(cfg: dict, default: int) -> int: def _parse_tui_skills_env() -> list[str]: raw = os.environ.get("HERMES_TUI_SKILLS", "") - skills: list[str] = [] - for part in raw.replace("\n", ",").split(","): - item = part.strip() - if item and item not in skills: - skills.append(item) - return skills + return list(dict.fromkeys(p.strip() for p in raw.replace("\n", ",").split(",") if p.strip())) def _load_fallback_model(): @@ -282,40 +247,33 @@ def _load_fallback_model(): return get_fallback_chain(_load_cfg()) -def _agent_fallback_model(agent): - """Return an agent's fallback chain without rehydrating deliberately empty chains.""" - if hasattr(agent, "_fallback_chain"): - return agent._fallback_chain or [] - return agent._fallback_model if hasattr(agent, "_fallback_model") else _load_fallback_model() - - def _background_agent_kwargs(agent, task_id: str) -> dict: cfg = _load_cfg() def g(name, default=None): return getattr(agent, name, default) - kwargs = {k: g(k) or None for k in ( - "base_url", "api_key", "provider", "api_mode", "acp_command", "acp_args", - "ephemeral_system_prompt")} - kwargs.update({k: g(k) for k in ( - "providers_allowed", "providers_ignored", "providers_order", "provider_sort", - "provider_data_collection", "openrouter_min_coding_score")}) - kwargs.update( - model=g("model") or _resolve_model(), - max_iterations=_cfg_max_turns(cfg, 25), - # Detached tasks declare platform="tui" (no UI sid for renderer-routed - # events), so resolve toolsets against it — never GUI schema they can't use. - enabled_toolsets=g("enabled_toolsets") or _load_enabled_toolsets("tui"), - quiet_mode=True, verbose_logging=False, - provider_require_parameters=g("provider_require_parameters", False), - session_id=task_id, - reasoning_config=g("reasoning_config") or _load_reasoning_config(str(g("model", "") or "")), - service_tier=g("service_tier") or _load_service_tier(), - request_overrides=dict(g("request_overrides", {}) or {}), - platform="tui", session_db=_get_db(), fallback_model=_agent_fallback_model(agent), - ) - return kwargs + # Don't rehydrate a deliberately empty fallback chain. + if hasattr(agent, "_fallback_chain"): + fallback = agent._fallback_chain or [] + else: + fallback = (agent._fallback_model if hasattr(agent, "_fallback_model") + else _load_fallback_model()) + # Detached tasks declare platform="tui" (no UI sid for renderer-routed events), so resolve + # toolsets against it — never GUI schema they can't use. + return { + **{k: g(k) or None for k in ("base_url", "api_key", "provider", "api_mode", "acp_command", + "acp_args", "ephemeral_system_prompt")}, + **{k: g(k) for k in ("providers_allowed", "providers_ignored", "providers_order", "provider_sort", + "provider_data_collection", "openrouter_min_coding_score")}, + "model": g("model") or _resolve_model(), "max_iterations": _cfg_max_turns(cfg, 25), + "enabled_toolsets": g("enabled_toolsets") or _load_enabled_toolsets("tui"), + "quiet_mode": True, "verbose_logging": False, + "provider_require_parameters": g("provider_require_parameters", False), "session_id": task_id, + "reasoning_config": g("reasoning_config") or _load_reasoning_config(str(g("model", "") or "")), + "service_tier": g("service_tier") or _load_service_tier(), + "request_overrides": dict(g("request_overrides", {}) or {}), + "platform": "tui", "session_db": _get_db(), "fallback_model": fallback} def _ephemeral_preview_agent_kwargs(agent, task_id: str) -> dict: @@ -323,13 +281,10 @@ def _ephemeral_preview_agent_kwargs(agent, task_id: str) -> dict: "enabled_toolsets": ["terminal", "file"], "session_db": None, "skip_memory": True} -_PREVIEW_HISTORY_ROLES = ("user", "assistant", "tool", "system") - - def _preview_restart_history(session: dict, max_messages: int = 24, max_tool_chars: int = 1200) -> list[dict]: - """Distill recent parent history for the ephemeral preview-restart agent (else it - guesses app/server/cwd/port from the bare URL). Keeps the last ``max_messages`` - (always back to the last user turn); tool results truncated to ``max_tool_chars``.""" + """Distill recent parent history for the ephemeral preview-restart agent (else it guesses + app/cwd/port from the bare URL): last ``max_messages`` back to the last user turn, tool + results truncated to ``max_tool_chars``.""" try: with session["history_lock"]: history = list(session.get("history") or []) @@ -337,14 +292,13 @@ def _preview_restart_history(session: dict, max_messages: int = 24, max_tool_cha history = list(session.get("history") or []) if not history: return [] + last_user = next((i for i in range(len(history) - 1, -1, -1) if history[i].get("role") == "user"), None) start = max(0, len(history) - max_messages) - for idx in range(len(history) - 1, -1, -1): - if history[idx].get("role") == "user": - start = min(start, idx) - break + if last_user is not None: + start = min(start, last_user) trimmed: list[dict] = [] for msg in history[start:]: - if not isinstance(msg, dict) or msg.get("role") not in _PREVIEW_HISTORY_ROLES: + if not isinstance(msg, dict) or msg.get("role") not in ("user", "assistant", "tool", "system"): continue copy = {k: v for k, v in msg.items() if k != "reasoning"} content = copy.get("content") @@ -358,7 +312,7 @@ def _preview_tool_result_preview(name: str, result: str) -> str: try: data = json.loads(result) except Exception: - return "" + data = None if not isinstance(data, dict): return "" if name == "terminal": @@ -375,8 +329,7 @@ def _preview_restart_callbacks(parent: str, task_id: str) -> dict: started_at: dict[str, float] = {} def progress(message: str, level: str = "info") -> None: - text = str(message or "").strip() - if text: + if text := str(message or "").strip(): _emit("preview.restart.progress", parent, {"task_id": task_id, "level": level, "text": text}) def tool_start(tool_call_id: str, name: str, args: dict) -> None: @@ -391,27 +344,22 @@ def _preview_restart_callbacks(parent: str, task_id: str) -> dict: progress(summary + (f"\n{output}" if output else "")) def tool_progress(event_type: str, name: str | None = None, preview: str | None = None, **_kwargs) -> None: - if preview: - progress(str(preview)) - elif name: - progress(f"{event_type.replace('.', ' ')}: {name}") + if preview or name: + progress(str(preview) if preview else f"{event_type.replace('.', ' ')}: {name}") return { "tool_start_callback": tool_start, "tool_complete_callback": tool_complete, "tool_progress_callback": tool_progress, "tool_gen_callback": lambda name: progress(f"Preparing {name}"), - "status_callback": lambda kind, text=None: progress(text if text is not None else kind), - } + "status_callback": lambda kind, text=None: progress(text if text is not None else kind)} def _reset_session_agent(sid: str, session: dict) -> dict: tokens = _set_session_context(session["session_key"]) try: - # /new is a full conversation boundary: session-scoped runtime overrides - # (/model, /reasoning, /fast) do NOT carry forward — the fresh agent - # re-derives them from config.yaml, and the pins are cleared so a rebuild - # can't resurrect them. Global process state is never touched (see the - # cross-session-contamination note in _apply_model_switch). + # /new is a full conversation boundary: session-scoped runtime overrides (/model, + # /reasoning, /fast) do NOT carry forward and the pins are cleared so a rebuild can't + # resurrect them. Global process state is never touched (see _apply_model_switch). for k in ("model_override", "create_reasoning_override", "create_service_tier_override", "one_turn_model_restore"): session.pop(k, None) new_agent = _make_agent( @@ -425,8 +373,7 @@ def _reset_session_agent(sid: str, session: dict) -> dict: queued_prompt=None, _queued_prompt_generation=int(session.get("_queued_prompt_generation", 0)) + 1, edit_snapshots={}, image_counter=0, running=False, show_reasoning=_load_show_reasoning(), - tool_progress_mode=_load_tool_progress_mode(), tool_started_at={}, - ) + tool_progress_mode=_load_tool_progress_mode(), tool_started_at={}) session.pop("queued_prompts", None) with session["history_lock"]: session["history"] = [] diff --git a/tui_gateway/change_watcher.py b/tui_gateway/change_watcher.py index 395cee5a31..9cb3dfa698 100644 --- a/tui_gateway/change_watcher.py +++ b/tui_gateway/change_watcher.py @@ -1,8 +1,6 @@ -"""Skin + config-change watcher: signatures for skin/pet/cron/sessions/platforms/pairing/bot-relay state and the broadcast loop that pushes *.changed events. - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so they reference server.py globals bare. -""" +"""Skin + config-change watcher: on-disk signatures for skin/pet/cron/sessions/platforms/ +pairing/bot-relay state and the broadcast loop that pushes *.changed events. Bodies are +rebound onto server.py's globals at install time (method_ctx.bind_module).""" from __future__ import annotations @@ -14,12 +12,11 @@ _registry = HandlerRegistry() def resolve_skin() -> dict: try: from hermes_cli.skin_engine import init_skin_from_config, get_active_skin - init_skin_from_config(_load_cfg()) skin = get_active_skin() + # light/dark are paired palettes: the TUI prefers the block matching terminal polarity. return { "name": skin.name, "colors": skin.colors, - # Paired palettes: the TUI prefers the block matching terminal polarity. "light_colors": skin.light_colors, "dark_colors": skin.dark_colors, "branding": skin.branding, "banner_logo": skin.banner_logo, "banner_hero": skin.banner_hero, "tool_prefix": skin.tool_prefix, @@ -47,10 +44,13 @@ def _watcher_mtime_ns(path: Path): return None +def _home_mtime_ns(*parts: str): + return _watcher_mtime_ns(_watcher_home().joinpath(*parts)) + + def _newest_mtime_ns(paths) -> int | None: """Max ``st_mtime_ns`` across ``paths`` (unstat-able ignored); None when none stat'ed.""" - mtimes = (_watcher_mtime_ns(p) for p in paths) - return max((m for m in mtimes if m is not None), default=None) + return max((m for m in map(_watcher_mtime_ns, paths) if m is not None), default=None) def _skin_sig() -> tuple[str, float | None]: @@ -58,10 +58,9 @@ def _skin_sig() -> tuple[str, float | None]: their name moves; a user skin's mtime lets an in-place color edit repaint too.""" name = str((_load_cfg().get("display") or {}).get("skin") or "default") try: - mtime: float | None = (_watcher_home() / "skins" / f"{name}.yaml").stat().st_mtime + return name, (_watcher_home() / "skins" / f"{name}.yaml").stat().st_mtime except OSError: - mtime = None - return name, mtime + return name, None def _note_skin_broadcast() -> None: @@ -75,14 +74,11 @@ def _broadcast_skin_if_changed() -> None: """Emit ``skin.changed`` when the active skin moved, via the SAME live path as ``/skin`` so every surface repaints. The check is a dict lookup + one stat.""" global _last_skin_sig - try: - sig = _skin_sig() - except Exception: - return - if sig == _last_skin_sig: - return - _last_skin_sig = sig with contextlib.suppress(Exception): + sig = _skin_sig() + if sig == _last_skin_sig: + return + _last_skin_sig = sig _broadcast_global_event("skin.changed", resolve_skin()) @@ -99,55 +95,39 @@ def _pet_sig() -> tuple: if not pet_cfg or not is_truthy_value(pet_cfg.get("enabled"), default=False): return ("off",) try: - active = _active_pet() - if not active: - return ("off",) - pet, scale = active - return (pet.slug, _pet_sheet_revision(pet.spritesheet), scale) + if active := _active_pet(): + pet, scale = active + return (pet.slug, _pet_sheet_revision(pet.spritesheet), scale) except Exception: # noqa: BLE001 - cosmetic, never break the watcher - return ("off",) + pass + return ("off",) def _pet_changed_payload() -> dict: """``pet.info.meta``-shaped payload so the renderer can decide whether to refetch sprites.""" try: - active = _active_pet() - if not active: - return {"enabled": False} - pet, scale = active - return { - "enabled": True, "slug": pet.slug, "displayName": pet.display_name, "scale": scale, - "spritesheetRevision": _pet_sheet_revision(pet.spritesheet)} + if active := _active_pet(): + pet, scale = active + return {"enabled": True, "slug": pet.slug, "displayName": pet.display_name, + "scale": scale, "spritesheetRevision": _pet_sheet_revision(pet.spritesheet)} except Exception: # noqa: BLE001 - cosmetic, never break the watcher - return {"enabled": False} - - -def _cron_sig(): - """mtime of cron/jobs.json — moves on edits AND scheduler tick bookkeeping.""" - return _watcher_mtime_ns(_watcher_home() / "cron" / "jobs.json") + pass + return {"enabled": False} def _sessions_sig(): - """Newest mtime across state.db + WAL: the one thing messaging-gateway turns and - cron runs (which never touch this gateway's transports) all move. Served sibling - profile homes are probed too, else a routed profile's Bot Chat never refreshes.""" + """Newest mtime across state.db + WAL: the one thing messaging-gateway turns and cron runs + all move. Served sibling profile homes are probed too, else a routed Bot Chat never refreshes.""" return _newest_mtime_ns( root / name for root in (_watcher_home(), *_served_profile_homes) - for name in ("state.db", "state.db-wal") - ) - - -def _platforms_sig(): - """mtime of gateway_state.json — where the messaging gateway persists platform - connect/disconnect/health, i.e. the Messaging page's status-changed signal.""" - return _watcher_mtime_ns(_watcher_home() / "gateway_state.json") + for name in ("state.db", "state.db-wal")) def _pairing_sig(): """Newest mtime across every profile's pairing ledgers (legacy ``pairing/`` and - ``platforms/pairing/``). Pending codes are written by the gateway process, so the - files are the only shared signal; a pairing request moves nothing in gateway_state.json.""" + ``platforms/pairing/``): the gateway process writes pending codes, so the files are the only + shared signal (a pairing request moves nothing in gateway_state.json).""" home = _watcher_home() roots = [home / "pairing", home / "platforms" / "pairing"] with contextlib.suppress(OSError): @@ -168,29 +148,27 @@ _bot_relay_outbox_seen = 0 def _bot_relay_outbox_sig(): - """Newest mtime across pending bot-relay outbox envelopes (monotone). Written by - the AGENT process, so the files are the only shared signal; the Desktop reacts - to ``bot_relay.outbox.pending`` with an immediate debounced drain.""" + """Newest mtime across pending bot-relay outbox envelopes (monotone). Written by the AGENT + process, so the files are the only shared signal; the Desktop reacts with a debounced drain.""" global _bot_relay_outbox_seen home = _watcher_home() root = home.parent.parent if home.parent.name == "profiles" else home - newest = 0 with contextlib.suppress(OSError): for entry in (root / "bot_relay" / "outbox").iterdir(): if entry.name.endswith(".json"): - newest = max(newest, _watcher_mtime_ns(entry) or 0) - if newest > _bot_relay_outbox_seen: - _bot_relay_outbox_seen = newest + _bot_relay_outbox_seen = max(_bot_relay_outbox_seen, _watcher_mtime_ns(entry) or 0) return _bot_relay_outbox_seen or None -# event → (check interval, signature fn, payload fn). Signatures are stat-cheap; -# the interval keeps pricier probes (pet resolves the sheet off disk) off the 0.5s tick. +# event → (check interval, signature fn, payload fn). Signatures are stat-cheap; the interval +# keeps pricier probes (pet resolves the sheet off disk) off the 0.5s tick. cron/jobs.json +# moves on edits AND scheduler ticks; gateway_state.json is where the messaging gateway +# persists platform connect/disconnect/health (the Messaging page's status signal). _CHANGE_WATCHES: dict[str, tuple[float, Any, Any]] = { "pet.changed": (2.0, _pet_sig, _pet_changed_payload), - "cron.changed": (1.0, _cron_sig, lambda: {}), + "cron.changed": (1.0, lambda: _home_mtime_ns("cron", "jobs.json"), lambda: {}), "sessions.changed": (0.5, _sessions_sig, lambda: {}), - "platforms.changed": (2.0, _platforms_sig, lambda: {}), + "platforms.changed": (2.0, lambda: _home_mtime_ns("gateway_state.json"), lambda: {}), "pairing.changed": (2.0, _pairing_sig, lambda: {}), # 1s so a queued DM envelope reaches the Desktop's push-triggered drain fast. "bot_relay.outbox.pending": (1.0, _bot_relay_outbox_sig, lambda: {})} @@ -220,9 +198,9 @@ def _broadcast_watched_changes(now: float | None = None) -> None: if event not in _change_sigs: _change_sigs[event] = sig continue + floor = _CHANGE_BROADCAST_FLOOR_S.get(event, 0.0) if sig == _change_sigs[event]: continue - floor = _CHANGE_BROADCAST_FLOOR_S.get(event, 0.0) if floor and now - _change_broadcast_at.get(event, -floor) < floor: continue # floored: old signature stays so it re-fires when the window opens _change_sigs[event] = sig @@ -235,9 +213,8 @@ _skin_watcher_started = False def _ensure_skin_watcher() -> None: - """Start the process's one change watcher (named for its original skin-only - duty): cheap on-disk signatures → broadcast events, so skin/pet/cron/cross-process - changes go live everywhere within seconds without client polling. Idempotent.""" + """Start the process's one change watcher (named for its original skin-only duty): cheap + on-disk signatures → broadcast events, so changes go live without client polling. Idempotent.""" global _skin_watcher_started if _skin_watcher_started: return @@ -249,7 +226,6 @@ def _ensure_skin_watcher() -> None: time.sleep(0.5) _broadcast_skin_if_changed() _broadcast_watched_changes() - threading.Thread(target=_loop, name="hermes-change-watcher", daemon=True).start() diff --git a/tui_gateway/compute_host.py b/tui_gateway/compute_host.py index 0f3d49cf40..52b5370339 100644 --- a/tui_gateway/compute_host.py +++ b/tui_gateway/compute_host.py @@ -1,8 +1,5 @@ -"""Persistent dashboard compute-host process. - -The long-lived child that owns live AIAgent objects when ``dashboard.turn_isolation`` -is enabled; frames are line-JSON over stdin/stdout. -""" +"""Persistent dashboard compute-host child: owns live AIAgent objects when +``dashboard.turn_isolation`` is enabled; frames are line-JSON over stdin/stdout.""" from __future__ import annotations @@ -43,16 +40,18 @@ class _HostTransport: return None -# Slice of ``ComputeHost.shutdown``'s budget held back for the post-drain finalize. -# ``HostSupervisor._terminate_pid`` SIGKILLs the host ``_SHUTDOWN_TIMEOUT_SECS`` (10s, -# same as ``shutdown``'s default ``wait``) after SIGTERM, so a drain allowed to consume -# the whole budget would leave the flush racing that kill and persist nothing at all. +# Slice of ``ComputeHost.shutdown``'s budget held back for the post-drain finalize: the +# supervisor SIGKILLs the host 10s (= default ``wait``) after SIGTERM, so a drain that ate +# the whole budget would leave the flush racing that kill and persist nothing. _FLUSH_RESERVE_SECS = 1.0 +# Fallback control.error text when a routed server method returns an error without a message. +_CONTROL_FAILURES = { + "session.save": "session save failed", "session.compress": "session compression failed"} + class ComputeHost: - # frame ``type`` -> handler method name (resolved per call so instance - # monkeypatches of a handler still take effect). + # frame ``type`` -> handler method name (resolved per call so monkeypatches take effect). _FRAME_HANDLERS: dict[str, str] = { "turn.start": "_handle_turn_start", "interrupt": "_handle_interrupt", "respond": "_handle_respond", "reload_mcp": "_handle_reload_mcp", @@ -70,8 +69,7 @@ class ComputeHost: self._boot_id = uuid.uuid4().hex self._progress_counter = 0 self._progress_lock = threading.Lock() - # Future -> the ``sid`` whose turn it is running. ``shutdown`` needs to know - # *whose* turn is still live so it can leave those sessions unfinalized. + # Future -> the ``sid`` whose turn it runs; ``shutdown`` leaves live sids unfinalized. self._turn_futures: dict[concurrent.futures.Future, str] = {} self._turn_futures_lock = threading.Lock() self._transport = _HostTransport(self.emit) @@ -101,36 +99,25 @@ class ComputeHost: def shutdown(self, *, reason: str = "shutdown", wait: float = 10.0) -> None: """Drain in-flight turns, then finalize every session. - Order matters: ``_finalize_session`` is a one-shot latch, so finalizing before - the drain would spend the flush's single chance mid-turn, fire - ``on_session_end(interrupted=True)`` against a running session and release the - active-session lease under a live turn. ``_FLUSH_RESERVE_SECS`` (never more than - half of ``wait``) is withheld from the drain so the flush still runs when turns - outlast the window. Sessions whose turn is *still running* at the deadline are - excluded from the flush (``_executor.shutdown`` does not join them): finalizing - one mid-turn would leave it un-finalizable with its lease released, whereas - leaving it unfinalized keeps it recoverable. ``server._shutdown_sessions`` - (atexit) may still re-finalize skipped sessions on the SIGTERM / stdin_closed - paths; the orphan path (``os._exit``) bypasses atexit. + ``_finalize_session`` is a one-shot latch, so finalizing before the drain would spend + it mid-turn and release the lease. ``_FLUSH_RESERVE_SECS`` (at most half of ``wait``) + is withheld from the drain so the flush still runs when turns outlast the window. + Sessions still running at the deadline are skipped (unfinalized keeps them + recoverable; atexit ``server._shutdown_sessions`` may re-finalize them). """ self._closed.set() budget = max(0.0, wait) deadline = time.monotonic() + budget - min(_FLUSH_RESERVE_SECS, budget / 2.0) while True: remaining = deadline - time.monotonic() - if remaining <= 0: + if remaining <= 0 or not self._live_turns(): break - with self._turn_futures_lock: - pending = [f for f in self._turn_futures if not f.done()] - if not pending: - break - # Bounded by ``remaining``: a flat sleep would overshoot the deadline and - # eat the reserve it protects (all of it for small ``wait``). + # Bounded by ``remaining``: a flat sleep would eat the reserve it protects. time.sleep(min(0.05, remaining)) with self._turn_futures_lock: live_sids = {sid for f, sid in self._turn_futures.items() if sid and not f.done()} self.flush_all_sessions(reason=reason, skip_sids=live_sids) - self._executor.shutdown(wait=False, cancel_futures=True) + self.close() def flush_all_sessions( self, *, reason: str = "shutdown", skip_sids: Collection[str] | None = None) -> None: @@ -153,19 +140,16 @@ class ComputeHost: self.emit({ "type": "error", "request_id": frame.get("request_id"), "message": f"unknown frame type: {kind}"}) - return - getattr(self, handler)(frame) + else: + getattr(self, handler)(frame) def _handle_shutdown(self, frame: dict[str, Any]) -> None: self.emit({"type": "shutdown.ack", "request_id": frame.get("request_id")}) - # Explicit supervisor/test shutdown is a clean child-process close; - # SIGTERM and orphan paths are the durability flush paths. - self._closed.set() - self._executor.shutdown(wait=False, cancel_futures=True) + # Explicit shutdown is a clean close; SIGTERM and orphan paths do the durability flush. + self.close() def _track_turn_future(self, future: concurrent.futures.Future, sid: str) -> None: - """Register an in-flight turn against its session; the done callback must pop - under the lock or the mapping grows for the host's life.""" + """Track an in-flight turn; the done callback pops it or the map grows forever.""" with self._turn_futures_lock: self._turn_futures[future] = sid future.add_done_callback(self._untrack_turn_future) @@ -178,41 +162,44 @@ class ComputeHost: future = self._executor.submit(self._run_real_turn, dict(frame)) self._track_turn_future(future, str(frame.get("sid") or "")) - def _handle_interrupt(self, frame: dict[str, Any]) -> None: + def _guarded( + self, frame: dict[str, Any], error_kind: str, body: Callable, *, + on_error: Callable[[str], None] | None = None, **error_extra: Any) -> None: + """Run ``body(server, sid, request_id)``; any exception becomes an ``error_kind`` reply.""" sid = str(frame.get("sid") or "") request_id = frame.get("request_id") try: from tui_gateway import server + body(server, sid, request_id) + except Exception as exc: + if on_error is not None: + on_error(sid) + self._reply(error_kind, sid, request_id, **error_extra, message=str(exc)) + + def _handle_interrupt(self, frame: dict[str, Any]) -> None: + def body(server: Any, sid: str, request_id: Any) -> None: session = server._sessions.get(sid) if session is None: self._reply("interrupt.ack", sid, request_id, applied=False) return - # In the child, `_session_uses_compute_host()` is false, so the shared helper - # interrupts the local agent and releases this process's pending clarify - # Event; the parent only has a metadata mirror and cannot. + # In the child the shared helper interrupts the local agent and releases this + # process's pending clarify Event (the parent only has a metadata mirror). server._interrupt_session_turn(sid, session) self._reply("interrupt.ack", sid, request_id, applied=True, applied_ns=now_ns()) - except Exception as exc: - self._reply("interrupt.ack", sid, request_id, applied=False, message=str(exc)) + self._guarded(frame, "interrupt.ack", body, applied=False) def _handle_respond(self, frame: dict[str, Any]) -> None: """Resolve an interactive request in the host-owned pending registry.""" - sid = str(frame.get("sid") or "") - request_id = frame.get("request_id") - try: - from tui_gateway import server - if sid not in server._sessions: - self._reply("respond.error", sid, request_id, message="session not found") - return + def body(server: Any, sid: str, request_id: Any) -> None: params = frame.get("params") - if not isinstance(params, dict): - self._reply( - "respond.error", sid, request_id, message="response params must be an object") + error = ("session not found" if sid not in server._sessions + else None if isinstance(params, dict) else "response params must be an object") + if error: + self._reply("respond.error", sid, request_id, message=error) return response = server._methods["clarify.respond"](request_id, params) self._reply("respond.ack", sid, request_id, response=response) - except Exception as exc: - self._reply("respond.error", sid, request_id, message=str(exc)) + self._guarded(frame, "respond.error", body) def _run_real_turn(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") @@ -223,7 +210,8 @@ class ComputeHost: try: from tui_gateway import server session = self._ensure_server_session(server, frame) - text = frame.get("text") if "text" in frame else frame.get("prompt", "") + text = frame["text"] if "text" in frame else frame.get("prompt", "") + inflight = frame["text"] if "text" in frame else frame.get("prompt") with session["history_lock"]: queued_gen = frame.get("queued_prompt_generation") current_gen = int(session.get("_queued_prompt_generation", 0)) @@ -233,11 +221,8 @@ class ComputeHost: if session.get("running"): self._reply("turn.error", sid, request_id, message="session busy") return - session["running"] = True - session["_turn_cancel_requested"] = False - session["last_active"] = time.time() - server._start_inflight_turn( - session, frame.get("text") if "text" in frame else frame.get("prompt")) + session.update(running=True, _turn_cancel_requested=False, last_active=time.time()) + server._start_inflight_turn(session, inflight) self._reply("turn.started", sid, request_id, started_ns=now_ns()) with contextlib.suppress(Exception): server._ensure_session_db_row(session) @@ -252,16 +237,14 @@ class ComputeHost: if run_thread is not None and hasattr(run_thread, "join"): run_thread.join() with session["history_lock"]: - history_version = int(session.get("history_version", 0)) - message_count = len(session.get("history") or []) + meta = _history_meta(session) interrupted = bool(session.get("_turn_cancel_requested")) - session_key = str(session.get("session_key") or "") session_info = server._session_info(session.get("agent"), session) - self._bump_progress() + with self._progress_lock: + self._progress_counter += 1 self._reply( - "turn.end", sid, request_id, history_version=history_version, - session_key=session_key, message_count=message_count, interrupted=interrupted, - ended_ns=now_ns(), session_info=session_info, session_info_emitted=True) + "turn.end", sid, request_id, **meta, interrupted=interrupted, ended_ns=now_ns(), + session_info=session_info, session_info_emitted=True) except Exception as exc: with contextlib.suppress(Exception): from tui_gateway import server @@ -274,25 +257,27 @@ class ComputeHost: def _ensure_server_session(self, server: Any, frame: dict[str, Any]) -> dict: sid = str(frame.get("sid") or "") - key = str(frame.get("session_key") or sid) session = server._sessions.get(sid) if session is not None: session["transport"] = self._transport if frame.get("cols") is not None: session["cols"] = int(frame.get("cols") or 80) - if frame.get("cwd"): - session["cwd"] = str(frame.get("cwd")) - if frame.get("profile_home"): - session["profile_home"] = str(frame.get("profile_home")) - if isinstance(frame.get("attached_images"), list): - session["attached_images"] = list(frame.get("attached_images") or []) - return session + for key in ("cwd", "profile_home"): + if frame.get(key): + session[key] = str(frame[key]) + else: + session = self._build_server_session(server, frame, sid) + if isinstance(frame.get("attached_images"), list): + session["attached_images"] = list(frame.get("attached_images") or []) + return session + + def _build_server_session(self, server: Any, frame: dict[str, Any], sid: str) -> dict: + """Build the agent under the frame's profile scope and register the session.""" + key = str(frame.get("session_key") or sid) history = frame.get("history") if isinstance(frame.get("history"), list) else [] profile_home = str(frame.get("profile_home") or "") - session_db = None + session_db = home_token = secret_token = None owns_db = False - home_token = None - secret_token = None try: if profile_home: from hermes_constants import set_hermes_home_override @@ -300,10 +285,8 @@ class ComputeHost: from hermes_state import get_shared_session_db home_token = set_hermes_home_override(profile_home) secret_token = set_secret_scope(build_profile_secret_scope(Path(profile_home))) - # DEDICATED handle — ours only until _make_agent succeeds; after that the - # agent (registered in server._sessions[sid] via _init_session or the - # fallback dict below) owns it. A RAISING _make_agent is the one path - # where nothing takes it, hence ``owns_db``. + # DEDICATED handle — ours only until _make_agent succeeds, then the agent owns + # it. A RAISING _make_agent is the one path where nothing takes it (``owns_db``). session_db = get_shared_session_db(Path(profile_home) / "state.db") owns_db = True agent = server._make_agent( @@ -338,9 +321,8 @@ class ComputeHost: finally: reset_transport(token) except Exception: - # If _init_session's side machinery (slash worker, approval notify) is - # unavailable, keep a minimal host-owned session rather than failing the - # turn after the expensive agent build succeeded. + # _init_session's side machinery (slash worker, approval notify) unavailable: keep a + # minimal host-owned session rather than failing after the expensive agent build. server._sessions[sid] = { "agent": agent, "session_key": key, "history": list(history), "history_lock": threading.Lock(), @@ -356,117 +338,88 @@ class ComputeHost: session = server._sessions[sid] session["transport"] = self._transport session["profile_home"] = profile_home or session.get("profile_home") - if isinstance(frame.get("attached_images"), list): - session["attached_images"] = list(frame.get("attached_images") or []) if frame.get("model_override") is not None: session["model_override"] = frame.get("model_override") return session def _handle_reload_mcp(self, frame: dict[str, Any]) -> None: - sid = str(frame.get("sid") or "") - request_id = frame.get("request_id") - try: - from tui_gateway import server + def body(server: Any, sid: str, request_id: Any) -> None: resp = server.handle_request({ "id": request_id, "method": "reload.mcp", "params": {"session_id": sid, "confirm": True}}) self._reply("reload_mcp.ack", sid, request_id, response=resp) - except Exception as exc: - self._reply("control.error", sid, request_id, message=str(exc)) + self._guarded(frame, "control.error", body) def _handle_control(self, frame: dict[str, Any]) -> None: - sid = str(frame.get("sid") or "") - request_id = frame.get("request_id") route_name = str(frame.get("route_name") or "") - def _error(message: str) -> None: - self._reply("control.error", sid, request_id, message=message) - - def _ack(**extra: Any) -> None: - self._reply("control.ack", sid, request_id, route_name=route_name, **extra) - - def _call_method(name: str, params: dict[str, Any], failure: str) -> dict | None: - """Run a server method; emit control.error and return None on error.""" - response = server._methods[name](request_id, params) - if "error" in response: - _error(str(response["error"].get("message") or failure)) - return None - return response - - def _history_meta() -> dict[str, Any]: - """Ack metadata read under ``history_lock`` (caller holds it).""" - return { - "session_key": str(session.get("session_key") or ""), - "history_version": int(session.get("history_version", 0)), - "message_count": len(session.get("history") or [])} - try: - from tui_gateway import server + def body(server: Any, sid: str, request_id: Any) -> None: route = MUTATOR_ROUTE_TABLE.get(route_name) - if route is None: - _error(f"unclassified route: {route_name}") - return session = server._sessions.get(sid) - if session is None: - _error("session not found") - return - if route == "idle-gated" and session.get("running"): - _error("session busy") - return - if route_name == "reload.mcp": + error = (f"unclassified route: {route_name}" if route is None + else "session not found" if session is None + else "session busy" if route == "idle-gated" and session.get("running") + else None) + if error: + self._reply("control.error", sid, request_id, message=error) + elif route_name == "reload.mcp": self._handle_reload_mcp({**frame, "type": "reload_mcp"}) - return - if route_name == "session.save": - response = _call_method("session.save", {"session_id": sid}, "session save failed") - if response is not None: - _ack(result=response.get("result") or {}) - return - if route_name == "session.compress": - focus_topic = str(frame.get("command") or "").removeprefix("/compress").strip() - params = {"session_id": sid} - if focus_topic: - params["focus_topic"] = focus_topic - response = _call_method("session.compress", params, "session compression failed") - if response is None: - return - with session["history_lock"]: - meta = _history_meta() - _ack( - result=response.get("result") or {}, **meta, - session_info=server._session_info(session.get("agent"), session)) - return - command = str(frame.get("command") or "") - output = server._mirror_slash_side_effects(sid, session, command) if command else "" - with session["history_lock"]: - messages = server._history_to_messages(list(session.get("history") or [])) - meta = _history_meta() - _ack( - output=output, session_key=meta["session_key"], - history_version=meta["history_version"], message_count=meta["message_count"], - messages=messages, session_info=server._session_info(session.get("agent"), session)) - except Exception as exc: + else: + ack = self._control_ack(server, frame, session) + if "error" in ack: + self._reply("control.error", sid, request_id, message=ack["error"]) + else: + self._reply("control.ack", sid, request_id, route_name=route_name, **ack) + + def on_error(sid: str) -> None: if route_name in {"session.compress", "slash.compress"}: - # The compress mirror defers the context-engine boundary notification until - # the host commits. If anything raises between queueing and finalize (e.g. - # building the ack's session_info), discard the pending notification so it - # can't fire against a rejected boundary on a later compress. finalize is - # exactly-once, so this is a no-op if the mirror already emitted it. + # The compress mirror defers the context-engine boundary notification until the + # host commits; discard it so it can't fire against a rejected boundary later + # (finalize is exactly-once, so a no-op if the mirror already emitted it). with contextlib.suppress(Exception): from tui_gateway import server as _server from agent.conversation_compression import ( - finalize_context_engine_compression_notification) + finalize_context_engine_compression_notification as _finalize) _agent = (_server._sessions.get(sid) or {}).get("agent") if _agent is not None: - finalize_context_engine_compression_notification(_agent, committed=False) - _error(str(exc)) + _finalize(_agent, committed=False) + self._guarded(frame, "control.error", body, on_error=on_error) - def _bump_progress(self) -> None: - with self._progress_lock: - self._progress_counter += 1 + def _control_ack(self, server: Any, frame: dict[str, Any], session: dict) -> dict: + """control.ack payload for one classified route, or ``{"error": message}``.""" + sid = str(frame.get("sid") or "") + route_name = str(frame.get("route_name") or "") + command = str(frame.get("command") or "") + if route_name in {"session.save", "session.compress"}: + params = {"session_id": sid} + if route_name == "session.compress": + focus_topic = command.removeprefix("/compress").strip() + if focus_topic: + params["focus_topic"] = focus_topic + response = server._methods[route_name](frame.get("request_id"), params) + if "error" in response: + failure = _CONTROL_FAILURES[route_name] + return {"error": str(response["error"].get("message") or failure)} + ack = {"result": response.get("result") or {}} + if route_name == "session.save": + return ack + with session["history_lock"]: + ack.update(_history_meta(session)) + else: + output = server._mirror_slash_side_effects(sid, session, command) if command else "" + with session["history_lock"]: + messages = server._history_to_messages(list(session.get("history") or [])) + ack = {"output": output, **_history_meta(session), "messages": messages} + ack["session_info"] = server._session_info(session.get("agent"), session) + return ack + + def _live_turns(self) -> list[concurrent.futures.Future]: + with self._turn_futures_lock: + return [f for f in self._turn_futures if not f.done()] def _heartbeat_loop(self) -> None: while not self._closed.wait(self._heartbeat_secs): - with self._turn_futures_lock: - active_turns = sum(1 for f in self._turn_futures if not f.done()) + active_turns = len(self._live_turns()) with self._progress_lock: counter = self._progress_counter self.emit({ @@ -482,6 +435,14 @@ class ComputeHost: os._exit(0) +def _history_meta(session: dict) -> dict[str, Any]: + """Transcript identity for turn.end / control.ack frames; caller holds history_lock.""" + return { + "session_key": str(session.get("session_key") or ""), + "history_version": int(session.get("history_version", 0)), + "message_count": len(session.get("history") or [])} + + def _rss_mb(pid: int) -> float: try: out = subprocess.check_output( @@ -547,8 +508,7 @@ def run_host(stdin: Any = None, stdout: Any = None) -> None: def main(argv: list[str] | None = None) -> int: - parser = argparse.ArgumentParser(description="Dashboard compute-host process") - parser.parse_args(argv) + argparse.ArgumentParser(description="Dashboard compute-host process").parse_args(argv) run_host() return 0 diff --git a/tui_gateway/compute_host_bridge.py b/tui_gateway/compute_host_bridge.py index 4754a14efd..5a5a0398d9 100644 --- a/tui_gateway/compute_host_bridge.py +++ b/tui_gateway/compute_host_bridge.py @@ -1,9 +1,6 @@ -"""Compute-host (turn isolation) bridge: relay prompts/controls to the child process -and mirror its metadata/clarify/compress acks back into the session. - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so they reference server.py globals bare. -""" +"""Compute-host (turn isolation) bridge: relay prompts/controls to the child process and +mirror its metadata/clarify/compress acks back into the session. Bodies are rebound onto +server.py's globals at install time (method_ctx.bind_module), so they use them bare.""" from __future__ import annotations @@ -14,7 +11,6 @@ from .method_ctx import HandlerRegistry, bind_module _registry = HandlerRegistry() - _compute_host_supervisor = None _compute_host_supervisor_lock = threading.Lock() # Cap on how long session.compress blocks its RPC on the compute host. Must stay @@ -23,23 +19,18 @@ _compute_host_supervisor_lock = threading.Lock() _COMPUTE_HOST_COMPRESS_WAIT_CAP_SECS = 630.0 -def _inside_compute_host_child() -> bool: - return os.environ.get("HERMES_COMPUTE_HOST_CHILD") == "1" - - def _turn_isolation_enabled(cfg: dict | None = None) -> bool: - if _inside_compute_host_child(): + if os.environ.get("HERMES_COMPUTE_HOST_CHILD") == "1": return False return bool((cfg or _load_dashboard_process_isolation_config()).get("turn_isolation")) def _session_uses_compute_host(session: dict, cfg: dict | None = None) -> bool: - if not _turn_isolation_enabled(cfg): - return False # Routes lazy sessions whose AIAgent was never built in-process; already-built # sessions keep the in-process path unless a prior isolated turn marked host ownership. - return bool(session.get("_compute_host_active")) or ( - session.get("agent") is None and session.get("agent_ready") is not None) + return _turn_isolation_enabled(cfg) and ( + bool(session.get("_compute_host_active")) + or (session.get("agent") is None and session.get("agent_ready") is not None)) def _get_compute_host_supervisor(cfg: dict | None = None): @@ -61,8 +52,7 @@ def _compute_host_turn_frame( with session["history_lock"]: history = list(session.get("history", [])) history_version = int(session.get("history_version", 0)) - attached_images = list( - image_paths if image_paths is not None else session.get("attached_images", [])) + attached_images = list(image_paths if image_paths is not None else session.get("attached_images", [])) return { "type": "turn.start", "sid": sid, "request_id": rid, "session_key": session.get("session_key") or sid, "text": text, @@ -97,8 +87,8 @@ def _compute_host_adopt_frame_meta(session: dict, frame: dict) -> None: session["session_key"] = str(frame.get("session_key")) if frame.get("history_version") is not None: with contextlib.suppress(Exception): - session["history_version"] = max( - int(session.get("history_version", 0)), int(frame.get("history_version") or 0)) + session["history_version"] = max(int(session.get("history_version", 0)), + int(frame.get("history_version") or 0)) def _relay_compute_host_rpc(message: dict) -> bool: @@ -110,7 +100,7 @@ def _relay_compute_host_rpc(message: dict) -> bool: payload = params.get("payload") request_id = payload.get("request_id") if isinstance(payload, dict) else None if session is not None and request_id: - with session.get("history_lock", threading.Lock()): + with _history_lock(session): if kind == "clarify.request": session["_compute_host_pending_clarify"] = dict(payload) elif _pending_clarify_matches(session, request_id): @@ -118,6 +108,10 @@ def _relay_compute_host_rpc(message: dict) -> bool: return write_json(message) +def _history_lock(session: dict): + return session.get("history_lock", threading.Lock()) + + def _pending_clarify_matches(session: dict, request_id) -> bool: """Whether ``session``'s mirrored pending clarify is ``request_id``. Caller holds history_lock.""" @@ -127,32 +121,26 @@ def _pending_clarify_matches(session: dict, request_id) -> bool: def _compute_host_clarify_session(request_id: str) -> tuple[str, dict] | None: """Find the parent mirror for one host-owned clarify request.""" - if not request_id: - return None - for sid, session in list(_sessions.items()): - with session.get("history_lock", threading.Lock()): + for sid, session in list(_sessions.items()) if request_id else (): + with _history_lock(session): if _pending_clarify_matches(session, request_id): return sid, session return None -def _update_compute_host_clarify_snapshot( - sid: str, session: dict, params: dict, result: dict) -> None: +def _update_compute_host_clarify_snapshot(sid: str, session: dict, params: dict, result: dict) -> None: """Keep reconnect snapshots accurate while a batch clarify is answered.""" request_id = str(params.get("request_id") or "") - with session.get("history_lock", threading.Lock()): + question_id = str(params.get("question_id") or "") + with _history_lock(session): if not _pending_clarify_matches(session, request_id): return pending = session["_compute_host_pending_clarify"] - expired = result.get("status") == "expired" - if expired or not result.get("remaining") and not params.get("question_id"): + if result.get("status") == "expired" or not result.get("remaining") and not question_id: session.pop("_compute_host_pending_clarify", None) - return - question_id = str(params.get("question_id") or "") - if question_id and isinstance(result.get("remaining"), list): - answers = dict(pending.get("answers") or {}) - answers[question_id] = str(params.get("answer") or "") - pending["answers"] = answers + elif question_id and isinstance(result.get("remaining"), list): + pending["answers"] = {**(pending.get("answers") or {}), + question_id: str(params.get("answer") or "")} if not result["remaining"]: session.pop("_compute_host_pending_clarify", None) @@ -160,11 +148,9 @@ def _update_compute_host_clarify_snapshot( def _respond_compute_host_clarify(rid: str, params: dict) -> dict | None: """Proxy a clarify answer into the process that owns its pending Event.""" located = _compute_host_clarify_session(str(params.get("request_id") or "")) - if located is None: + if located is None or not _session_uses_compute_host(located[1]): return None sid, session = located - if not _session_uses_compute_host(session): - return None try: ack = _get_compute_host_supervisor().respond(sid, params) except Exception as exc: @@ -176,9 +162,8 @@ def _respond_compute_host_clarify(rid: str, params: dict) -> dict | None: return _err(rid, 5019, "compute-host clarify response returned an invalid response") if "error" in response: error = response["error"] if isinstance(response["error"], dict) else {} - return _err( - rid, int(error.get("code") or 5000), - str(error.get("message") or "clarify response failed")) + return _err(rid, int(error.get("code") or 5000), + str(error.get("message") or "clarify response failed")) result = response.get("result") if not isinstance(result, dict): return _err(rid, 5019, "compute-host clarify response returned an invalid result") @@ -191,7 +176,7 @@ def _apply_compute_host_metadata_mirror(session: dict, frame: dict | None) -> No writer of live agent/history state, and UI reads must not build a second agent.""" if not isinstance(frame, dict): return - with session.get("history_lock", threading.Lock()): + with _history_lock(session): _compute_host_adopt_frame_meta(session, frame) if frame.get("message_count") is not None: with contextlib.suppress(Exception): @@ -223,17 +208,15 @@ def _submit_prompt_to_compute_host( rid: str, sid: str, session: dict, text: Any, image_paths: list[str] | None = None, queued_prompt_generation: int | None = None, display_kind: str | None = None) -> dict: cfg = _load_dashboard_process_isolation_config() - frame = _compute_host_turn_frame( - rid, sid, session, text, image_paths=image_paths, - queued_prompt_generation=queued_prompt_generation, display_kind=display_kind) + frame = _compute_host_turn_frame(rid, sid, session, text, image_paths=image_paths, + queued_prompt_generation=queued_prompt_generation, + display_kind=display_kind) def _complete(done: dict) -> None: - # submit_turn reports a synchronous pipe failure via the callback before - # re-raising; leave the session untouched so prompt.submit can fail open - # to the in-process path without a duplicate terminal error. - if done.get("reason") == "send_failed": - return - _on_compute_host_turn_done(rid, sid, session, done) + # submit_turn reports a synchronous pipe failure via the callback before re-raising; + # leave the session untouched so prompt.submit can fail open to the in-process path. + if done.get("reason") != "send_failed": + _on_compute_host_turn_done(rid, sid, session, done) try: _get_compute_host_supervisor(cfg).submit_turn(frame, on_complete=_complete) except Exception as exc: @@ -271,29 +254,22 @@ def _compute_host_compress_wait_seconds(cfg: dict | None = None) -> float: return float(min(max(ceiling + 30.0, 120.0), _COMPUTE_HOST_COMPRESS_WAIT_CAP_SECS)) -def _announce_compute_host_compress_done(sid: str, session: dict, ack: dict) -> None: - """Mirror a compress ack and push the ``session.info`` + ``compacted`` edges the - in-process /compress path emits, so a client whose RPC wait expired still learns.""" - _apply_compute_host_metadata_mirror(session, ack) - _emit("session.info", sid, _compute_host_session_info(session)) - _status_update(sid, "compacted", "✓ Context compression complete") - - -def _adopt_late_compute_host_compress_ack( - sid: str, session: dict, ack: dict, *, route_name: str) -> None: - """Adopt a compress ack that arrived after its RPC waiter answered ``pending``: the - only place the rotated session_key / history_version / mirror can land and the - client's only signal. A late ``control.error`` goes out via ``error``.""" +def _adopt_late_compute_host_compress_ack(sid: str, session: dict, ack: dict, *, route_name: str) -> None: + """Adopt a compress ack that arrived after its RPC waiter answered ``pending``: the only place + the rotated session_key / history_version / mirror can land and the client's only signal + (the same ``session.info`` + ``compacted`` edges the in-process /compress path emits). A late + ``control.error`` goes out via ``error``.""" with _sessions_lock: - live = _sessions.get(sid) - if live is not session: - return + if _sessions.get(sid) is not session: + return if not isinstance(ack, dict) or ack.get("type") in {"control.error", "error"}: message = str((ack or {}).get("message") or f"compute-host {route_name} failed") _emit("error", sid, {"message": f"compression failed: {message}"}) _status_update(sid, "ready") return - _announce_compute_host_compress_done(sid, session, ack) + _apply_compute_host_metadata_mirror(session, ack) + _emit("session.info", sid, _compute_host_session_info(session)) + _status_update(sid, "compacted", "✓ Context compression complete") def register(server) -> None: diff --git a/tui_gateway/entry.py b/tui_gateway/entry.py index 33d83df476..209c48cbf9 100644 --- a/tui_gateway/entry.py +++ b/tui_gateway/entry.py @@ -1,9 +1,8 @@ import os import sys -# Stop a ``utils/`` (or ``proxy/``, ``ui/``) package in the launch directory from -# shadowing Hermes's own top-level modules. ``hermes_bootstrap`` lives at the repo -# root (its name can't collide with a user package), so importing it first is safe. +# Stop a ``utils/``-style package in the launch directory from shadowing Hermes's own +# top-level modules; ``hermes_bootstrap``'s name can't collide, so importing it first is safe. import hermes_bootstrap hermes_bootstrap.harden_import_path() @@ -29,15 +28,13 @@ logger = logging.getLogger(__name__) # Discovery thread spawned by THIS module; None when delegated to the shared owner in # hermes_cli.mcp_startup (current path). The wait/in-flight/join helpers consult both. _mcp_discovery_thread = None -# Set once MCP servers are found configured, so wait_for_mcp_discovery can re-invoke -# the idempotent spawn on later builds (retry-after-zero-connected) without a config -# re-probe — non-MCP sessions never pay the tools.mcp_tool import per build. +# Set once MCP servers are found configured so wait_for_mcp_discovery can re-invoke the +# idempotent spawn on later builds without a config re-probe. _mcp_discovery_enabled = False def _install_sidecar_publisher() -> None: - """Mirror every dispatcher emit to the dashboard sidebar via WS when - `HERMES_TUI_SIDECAR_URL` is set (best-effort: a dropped WS falls back to stdio-only).""" + """Mirror every dispatcher emit to the dashboard sidebar via WS when set (best-effort).""" url = os.environ.get("HERMES_TUI_SIDECAR_URL") if not url: return @@ -46,8 +43,7 @@ def _install_sidecar_publisher() -> None: # Grace for orderly shutdown before ``os._exit(0)`` so a worker wedged mid-flush can't -# strand the process; ``HERMES_TUI_GATEWAY_SHUTDOWN_GRACE_S`` overrides (a longer grace -# also means a longer wait on a real deadlock). +# strand the process; ``HERMES_TUI_GATEWAY_SHUTDOWN_GRACE_S`` overrides. _DEFAULT_SHUTDOWN_GRACE_S = 1.0 @@ -56,10 +52,6 @@ def _shutdown_grace_seconds() -> float: return value if value > 0 else _DEFAULT_SHUTDOWN_GRACE_S -def _stamp() -> str: - return time.strftime("%Y-%m-%d %H:%M:%S") - - def _mcp_startup_call(name: str, *args, default=None, log=None, **kwargs): """Call ``hermes_cli.mcp_startup.`` (lazy import); ``default`` on any failure, optionally logged as ``(level, message)``.""" @@ -88,10 +80,9 @@ def _append_crash_log(header: str, dump=None) -> None: def _log_signal(signum: int, frame) -> None: - """Capture WHICH thread and WHERE a termination signal hit us, then exit. - ``sys.exit(0)`` alone raced the worker pool (a thread holding ``_stdout_lock`` - mid-flush blocks interpreter shutdown), so: log all thread stacks, give the - configured grace to drain, then ``os._exit(0)``.""" + """Capture WHICH thread and WHERE a termination signal hit us, then exit. ``sys.exit(0)`` + alone raced the worker pool (a thread holding ``_stdout_lock`` mid-flush blocks interpreter + shutdown), so: log all thread stacks, give the configured grace to drain, then ``os._exit``.""" # SIGPIPE/SIGHUP don't exist on Windows — only look up attributes present. names = {int(sig): attr for attr in ("SIGPIPE", "SIGTERM", "SIGHUP", "SIGINT", "SIGBREAK") if (sig := getattr(signal, attr, None)) is not None} @@ -106,41 +97,35 @@ def _log_signal(signum: int, frame) -> None: f.write(f"\n--- thread {th.name} (id={tid}) ---\n") f.write("".join(traceback.format_stack(sys._current_frames().get(tid)))) - _append_crash_log(f"{name} received · {_stamp()}", _dump) + _append_crash_log(f"{name} received · {time.strftime('%Y-%m-%d %H:%M:%S')}", _dump) print(f"[gateway-signal] {name}", file=sys.stderr, flush=True) - - # ``os._exit`` skips atexit but breaks the mid-flush deadlock; the crash log - # + stderr line above are the forensic trail. + # ``os._exit`` skips atexit but breaks the mid-flush deadlock; the crash log is the trail. timer = threading.Timer(_shutdown_grace_seconds(), lambda: os._exit(0)) timer.daemon = True timer.start() - # atexit (_shutdown_sessions) can be blocked past the grace window by a worker - # holding the GIL/_stdout_lock; finalize explicitly so unpersisted messages reach - # state.db before the hard-exit timer fires. + # atexit (_shutdown_sessions) can be blocked past the grace window by a worker holding + # the GIL/_stdout_lock; finalize explicitly so unpersisted messages reach state.db first. with suppress(Exception): from tui_gateway.server import _shutdown_sessions _shutdown_sessions() - # Unwind the main thread so atexit + finalisers run inside the grace window; - # the daemon timer is the safety net if that unwind hangs. + # Unwind the main thread so atexit + finalisers run; the daemon timer is the safety net. sys.exit(0) def _install_signal(signame, handler): - """Install a signal handler if legal here: signal.signal() raises off the main - thread (Desktop build path: server._build imports entry from a worker), and - Windows lacks SIGPIPE/SIGHUP — both are skipped. Handlers are process-global.""" + """Install a signal handler if legal here: signal.signal() raises off the main thread + (Desktop build path imports entry from a worker) and Windows lacks SIGPIPE/SIGHUP.""" sig = getattr(signal, signame, None) if sig is None or threading.current_thread() is not threading.main_thread(): return - # Off the main thread despite the check, or handler rejected by the platform. - with suppress(ValueError, OSError, RuntimeError): + with suppress(ValueError, OSError, RuntimeError): # platform rejected the handler signal.signal(sig, handler) -# SIGPIPE: ignore, don't exit — SIG_DFL killed the process silently whenever a -# *background* thread (TTS, beep) wrote to a pipe the TUI had gone quiet on; ignoring -# lets write_json see BrokenPipeError and exit cleanly via _log_exit. Terminal signals -# route through _log_signal so kills/hangups are diagnosable (SIGBREAK = Windows SIGHUP). +# SIGPIPE: ignore, don't exit — SIG_DFL killed the process silently whenever a background +# thread wrote to a pipe the TUI had gone quiet on; ignoring lets write_json see +# BrokenPipeError and exit via _log_exit. Terminal signals route through _log_signal so +# kills/hangups are diagnosable (SIGBREAK = Windows SIGHUP). _install_signal("SIGPIPE", signal.SIG_IGN) _install_signal("SIGTERM", _log_signal) if hasattr(signal, "SIGHUP"): @@ -151,26 +136,23 @@ _install_signal("SIGINT", signal.SIG_IGN) def _log_exit(reason: str) -> None: - """Record why the gateway exits: every path collapses into a silent sys.exit(0), - and without this trail the TUI can't tell WHICH broken pipe triggered it.""" - _append_crash_log(f"gateway exit · {_stamp()} · reason={reason}") + """Record why the gateway exits (every path is a silent sys.exit(0) otherwise).""" + _append_crash_log(f"gateway exit · {time.strftime('%Y-%m-%d %H:%M:%S')} · reason={reason}") print(f"[gateway-exit] {reason}", file=sys.stderr, flush=True) def wait_for_mcp_discovery(timeout: "float | None" = None) -> None: - """Block until background MCP discovery finishes, up to the resolved bound - (``mcp_discovery_timeout`` from config; ``timeout`` overrides). The agent snapshots - its tool list ONCE at build time, so this bounded join lets already-spawning - servers land without re-introducing the startup hang.""" + """Block until background MCP discovery finishes, up to the resolved bound (config + ``mcp_discovery_timeout``; ``timeout`` overrides). The agent snapshots its tool list ONCE + at build time, so this bounded join lets already-spawning servers land.""" thread = _mcp_discovery_thread if thread is not None and thread.is_alive(): fallback = timeout if timeout is not None else 0.75 bound = _mcp_startup_call("_resolve_discovery_timeout", timeout, default=fallback) thread.join(timeout=bound) return - # Shared-owner path: re-invoke the idempotent spawn first so a zero-connected run - # gets its retry instead of latching the process MCP-less. Runs under the CALLER's - # profile context (agent build binds the session profile's HERMES_HOME first). + # Shared-owner path: re-invoke the idempotent spawn first so a zero-connected run gets + # its retry instead of latching the process MCP-less (runs under the CALLER's profile). if not _mcp_discovery_enabled: return _spawn_discovery(("debug", "TUI MCP discovery retry-spawn failed")) @@ -178,10 +160,8 @@ def wait_for_mcp_discovery(timeout: "float | None" = None) -> None: def mcp_discovery_in_flight() -> bool: - """True if ANY background MCP discovery thread is still running. Two owners by - surface (stdio thread here, ``hermes_cli.mcp_startup`` for desktop/dashboard); - the late-refresh scheduler calls this regardless of surface, so it MUST consult - both or slow MCP servers' tools never surface on desktop.""" + """True if ANY background MCP discovery thread is still running: the late-refresh + scheduler calls this regardless of surface, so it MUST consult both owners.""" thread = _mcp_discovery_thread if thread is not None and thread.is_alive(): return True @@ -189,9 +169,8 @@ def mcp_discovery_in_flight() -> bool: def join_mcp_discovery(timeout: float | None = None) -> bool: - """Join both discovery owners; True once neither is alive. Unlike - ``wait_for_mcp_discovery`` this accepts an unbounded wait (off-critical-path - late-refresh waiter); ``timeout`` bounds EACH join, entry thread first.""" + """Join both discovery owners; True once neither is alive. Accepts an unbounded wait + (off-critical-path late-refresh waiter); ``timeout`` bounds EACH join, entry thread first.""" entry_done = True thread = _mcp_discovery_thread if thread is not None: @@ -211,13 +190,10 @@ def _has_configured_mcp_servers() -> bool: def ensure_mcp_discovery_started() -> None: - """Start background MCP discovery for the current profile context, once. - ``main()`` calls this for stdio; WS/Desktop skip ``main()``, so - ``server._start_agent_build`` also calls it AFTER binding the session profile's - HERMES_HOME so discovery reads the SELECTED profile's ``mcp_servers``. MCP - registration is process-global: the FIRST profile to build an agent wins.""" + """Start background MCP discovery for the current profile context, once. ``main()`` calls + this for stdio; ``server._start_agent_build`` also calls it AFTER binding the session + profile's HERMES_HOME. MCP registration is process-global: the FIRST profile wins.""" global _mcp_discovery_enabled - if not _has_configured_mcp_servers(): return _mcp_discovery_enabled = True @@ -233,40 +209,30 @@ def _write_or_exit(payload: dict, reason: str) -> None: def main(): _install_sidecar_publisher() - # Heartbeat row lets the orphan sweep tell "live but idle" from "truly orphaned"; - # it must run BEFORE the sweep. The sweep is once-per-process and config-gated. + # The heartbeat row lets the orphan sweep tell "live but idle" from "truly orphaned", + # so it must start BEFORE the sweep. for start, what in ( - (server._start_backend_heartbeat_refresher, "backend heartbeat refresher start"), - (server._schedule_startup_orphan_sweep, "startup orphan sweep scheduling"), - ): + (server._start_backend_heartbeat_refresher, "backend heartbeat refresher start"), + (server._schedule_startup_orphan_sweep, "startup orphan sweep scheduling")): try: start() except Exception: logger.warning("%s failed", what, exc_info=True) - # Backgrounded so a dead MCP server (~7s of retries) can't freeze startup; - # _make_agent briefly joins it. The config gate keeps the MCP SDK import off - # the no-mcp_servers path. + # Backgrounded so a dead MCP server can't freeze startup; _make_agent briefly joins it. ensure_mcp_discovery_started() + # change_events: clients demote legacy polls; replay_epoch: WS restart detection. _write_or_exit({ - "jsonrpc": "2.0", - "method": "event", - "params": { - "type": "gateway.ready", - # change_events: clients demote legacy polls (see tui_gateway/ws.py). - # replay_epoch: WS restart detection; the stdio TUI ignores it. - "payload": { - "skin": resolve_skin(), "change_events": True, "replay_epoch": replay_epoch(), - }, - }, - }, "startup write failed (broken stdout pipe before first event)") + "jsonrpc": "2.0", "method": "event", + "params": {"type": "gateway.ready", "payload": { + "skin": resolve_skin(), "change_events": True, "replay_epoch": replay_epoch()}}}, + "startup write failed (broken stdout pipe before first event)") # Live-apply skins Hermes activates mid-conversation. server._ensure_skin_watcher() - # Warm the /model picker's provider-models cache in this idle window, else the - # first /model open blocks on serial /v1/models fetches. Fire-and-forget. + # Warm the /model picker's provider-models cache in this idle window (fire-and-forget). try: from hermes_cli.model_switch import prewarm_picker_cache_async prewarm_picker_cache_async() @@ -280,11 +246,9 @@ def main(): if not handle_spurious_eof(_recovery_times, _log_exit): break continue - line = raw.strip() if not line: continue - try: req = json.loads(line) except json.JSONDecodeError: diff --git a/tui_gateway/git_probe.py b/tui_gateway/git_probe.py index 5930e90374..875385c7d7 100644 --- a/tui_gateway/git_probe.py +++ b/tui_gateway/git_probe.py @@ -1,11 +1,8 @@ """Git working-tree probing for the gateway: run git, resolve repo roots, fold linked worktrees. - Probing runs where the gateway runs (covers remote backends). Roots go through a thread-safe -single-flight cache so concurrent identical probes share one ``git`` spawn. Positives are cached -for the process lifetime; negatives (not a repo / deleted dir) only for ``_NEG_TTL`` — -``build_tree`` resolves a cwd once *per session*, so hundreds of non-git cwds would otherwise -re-spawn ``git`` on every sidebar open, while the TTL keeps a fresh ``git init`` re-probable. -""" +single-flight cache so concurrent identical probes share one ``git`` spawn: positives live for the +process, negatives (not a repo / deleted dir) for ``_NEG_TTL`` — hundreds of non-git session cwds +would otherwise re-spawn ``git`` on every sidebar open, while the TTL keeps ``git init`` re-probable.""" from __future__ import annotations @@ -19,19 +16,14 @@ from hermes_cli._subprocess_compat import bounded_git_probe _GIT_TIMEOUT = 1.5 _WARM_WORKERS = 8 -# "Not a git repo" TTL: short enough that a fresh `git init` shows within seconds, -# long enough to collapse a tree build's hundreds of redundant probes. -_NEG_TTL = 30.0 +_NEG_TTL = 30.0 # "not a git repo" TTL: a fresh `git init` shows within seconds def run_git(cwd: str, *args: str) -> str: - """``git -C `` → stripped stdout, or ``""`` on any failure. - - ``bounded_git_probe`` bounds post-kill cleanup on Windows — a plain ``subprocess.run(timeout)`` - deadlocked Desktop readiness when a killed git left a suspended descendant holding the pipes. - """ - # `git -C` on a missing dir can only fail, at the price of a fork; deleted worktrees - # dominate a long session history's cwds, so the stat pays off. + """``git -C `` → stripped stdout, or ``""`` on any failure. ``bounded_git_probe`` + bounds post-kill cleanup on Windows (a killed git's suspended descendant held the pipes).""" + # A missing dir can only fail at the price of a fork; deleted worktrees dominate a long + # session history's cwds, so the stat pays off. if not cwd or not os.path.isdir(cwd): return "" return bounded_git_probe(["git", "-C", cwd, *args], timeout=_GIT_TIMEOUT) @@ -60,8 +52,7 @@ class _RootCache: def resolve(self, key: str, probe) -> str: while True: with self._lock: - hit = self._roots.get(key) - if hit: + if hit := self._roots.get(key): return hit expiry = self._neg.get(key) if expiry is not None: @@ -72,8 +63,7 @@ class _RootCache: leader = gate is None if leader: gate = self._inflight[key] = threading.Event() - if not leader: - # Another thread is probing this key — wait, then re-read. + if not leader: # another thread is probing this key — wait, then re-read gate.wait(timeout=_GIT_TIMEOUT + 0.5) continue value = "" @@ -100,21 +90,16 @@ def invalidate() -> None: def repo_root(cwd: str) -> str: """Top-level git repo root for ``cwd`` (``""`` when not a repo).""" - if not cwd: - return "" - return _cache.resolve(cwd, lambda: run_git(cwd, "rev-parse", "--show-toplevel")) + return _cache.resolve(cwd, lambda: run_git(cwd, "rev-parse", "--show-toplevel")) if cwd else "" def common_repo_root(cwd: str) -> str: - """The MAIN (common) repo root for ``cwd``, folding linked worktrees. - - ``--show-toplevel`` returns a linked worktree's OWN root; the parent of the shared - ``--git-common-dir`` is the one true root (fallback: the toplevel root). Normalized to git's - forward-slash spelling so it compares equal to :func:`repo_root` — with native ``\\`` on - Windows the main checkout was misread as a linked worktree and the sidebar rendered it twice. - """ - # Not a repo: nothing to fold. Checking the (warmed, negative-cached) toplevel first spares - # every non-repo cwd a second `git` spawn the parallel warm can't absorb. + """The MAIN (common) repo root for ``cwd``, folding linked worktrees: ``--show-toplevel`` is a + linked worktree's OWN root; the parent of the shared ``--git-common-dir`` is the one true root + (fallback: toplevel). Normalized to git's forward-slash spelling so it compares equal to + :func:`repo_root` (native ``\\`` on Windows made the main checkout look like a worktree).""" + # Checking the (warmed, negative-cached) toplevel first spares every non-repo cwd a second + # `git` spawn the parallel warm can't absorb. if not cwd or not repo_root(cwd): return "" @@ -142,10 +127,8 @@ def warm_roots(cwds: Iterable[str], max_workers: int = _WARM_WORKERS) -> None: """Pre-resolve many cwds' roots in parallel (bounded) so a cold first paint doesn't serialize one git spawn per session cwd; results land in the cache.""" pending = sorted({(cwd or "").strip() for cwd in cwds} - {""}) - if not pending: - return if len(pending) == 1: resolve(pending[0]) - return - with ThreadPoolExecutor(max_workers=min(max_workers, len(pending))) as pool: - list(pool.map(resolve, pending)) + elif pending: + with ThreadPoolExecutor(max_workers=min(max_workers, len(pending))) as pool: + list(pool.map(resolve, pending)) diff --git a/tui_gateway/host_supervisor.py b/tui_gateway/host_supervisor.py index dddb4401b9..cd19373672 100644 --- a/tui_gateway/host_supervisor.py +++ b/tui_gateway/host_supervisor.py @@ -1,9 +1,6 @@ -"""Supervisor for the dashboard compute-host child process. - -When ``dashboard.turn_isolation`` is enabled, agent turns move behind one persistent -``python -m tui_gateway.compute_host`` child so compute-heavy agent threads do not -contend with the serving process' event loop for the same GIL. -""" +"""Supervisor for the dashboard compute-host child: with ``dashboard.turn_isolation`` +agent turns run in one persistent ``python -m tui_gateway.compute_host`` child so heavy +agent threads do not contend with the serving process' event loop for the GIL.""" from __future__ import annotations @@ -26,7 +23,6 @@ from hermes_constants import get_hermes_home from tools.environments.local import hermes_subprocess_env logger = logging.getLogger(__name__) -_Thread = threading.Thread MUTATOR_ROUTE_TABLE: dict[str, str] = { "prompt.submit": "turn-path", "session.interrupt": "turn-path", "reload.mcp": "run-concurrent", @@ -39,9 +35,8 @@ MUTATOR_ROUTE_TABLE: dict[str, str] = { _REGISTRY_NAME = "dashboard-compute-host.json" _RESPAWN_WINDOW_SECS = 300.0 _SHUTDOWN_TIMEOUT_SECS = 10.0 -# Late control-ack handlers: a compress that outlives its RPC waiter can run for the -# full compression ceiling plus a stall-fallback retry, so keep registrations well -# past that — but bounded. +# Late control-ack handlers: a compress that outlives its RPC waiter can run for the full +# compression ceiling plus a stall-fallback retry, so keep registrations past that — bounded. _LATE_CONTROL_TTL_SECS = 1800.0 _LATE_CONTROL_MAX = 64 # Host frames whose ``request_id`` resolves a pending/late control waiter. @@ -52,10 +47,9 @@ _CONTROL_REPLY_TYPES = frozenset({ def append_log_record(path: str | Path, record: str) -> None: """Append one log record using O_APPEND and exactly one os.write call.""" - p = Path(path) - p.parent.mkdir(parents=True, exist_ok=True) + Path(path).parent.mkdir(parents=True, exist_ok=True) text = record if record.endswith("\n") else f"{record}\n" - fd = os.open(str(p), os.O_WRONLY | os.O_CREAT | os.O_APPEND, 0o600) + fd = os.open(str(path), os.O_WRONLY | os.O_CREAT | os.O_APPEND, 0o600) try: os.write(fd, text.encode("utf-8", errors="replace")) finally: @@ -68,31 +62,45 @@ def _repo_root() -> Path: def _check_output(argv: list[str], **kwargs: Any) -> str: """Stripped stdout of a short subprocess, or ``""`` on any failure.""" - try: + with contextlib.suppress(Exception): return subprocess.check_output( argv, text=True, encoding="utf-8", errors="replace", stderr=subprocess.DEVNULL, timeout=2, **kwargs).strip() - except Exception: - return "" + return "" def _build_sha() -> str: - """Current checkout's HEAD sha, or ``"unknown"``. Shared with ``compute_host`` so - the hello handshake and the supervisor's expectation agree byte-for-byte.""" + """HEAD sha or ``"unknown"``; shared with ``compute_host`` so the hello handshake agrees.""" return _check_output(["git", "rev-parse", "HEAD"], cwd=str(_repo_root())) or "unknown" +def _call_logged(cb: Callable[[dict], None], frame: dict, failure: str) -> None: + """Invoke a host-frame callback; a raising callback is logged, never propagated.""" + try: + cb(frame) + except Exception: + logger.exception(failure) + + def _pid_alive(pid: int) -> bool: if pid <= 0: return False try: os.kill(pid, 0) return True + except Exception as exc: + return isinstance(exc, PermissionError) + + +def _signal_pid(pid: int, sig: int, label: str) -> bool: + """Send ``sig``; False when the pid is gone or the signal failed (logged).""" + try: + os.kill(pid, sig) + return True except ProcessLookupError: return False - except PermissionError: - return True except Exception: + logger.debug("failed to %s compute host pid=%s", label, pid, exc_info=True) return False @@ -128,10 +136,9 @@ class HostSupervisor: self.rpc_sink = rpc_sink or (lambda _obj: None) self.respawn_max = max(0, int(respawn_max)) self.heartbeat_secs = max(1, int(heartbeat_secs)) - self.expected_build_sha = ( - expected_build_sha if expected_build_sha is not None else _build_sha()) + self.expected_build_sha = _build_sha() if expected_build_sha is None else expected_build_sha self.expected_hermes_home = ( - expected_hermes_home if expected_hermes_home is not None else str(get_hermes_home())) + str(get_hermes_home()) if expected_hermes_home is None else expected_hermes_home) self._lock = threading.RLock() self._proc: subprocess.Popen[str] | None = None self._hello_event = threading.Event() @@ -141,9 +148,8 @@ class HostSupervisor: self._restart_times: list[float] = [] self._pending_turns: dict[str, tuple[str, Callable[[dict], None] | None]] = {} self._pending_controls: dict[str, queue.Queue[dict]] = {} - # request_id -> (registered_at, handler) for control waiters that timed out - # while their host work still runs; without it the eventual control.ack - # matched no queue and was silently dropped. + # request_id -> (registered_at, handler) for control waiters that timed out while their + # host work still runs, so the eventual control.ack is not silently dropped. self._late_control_handlers: dict[str, tuple[float, Callable[[dict], None]]] = {} self._stderr_tail: list[str] = [] self._last_progress_counter = 0 @@ -155,10 +161,6 @@ class HostSupervisor: proc = self._proc return int(proc.pid or 0) if proc is not None else 0 - @property - def hello(self) -> dict[str, Any]: - return dict(self._hello) - def is_running(self) -> bool: proc = self._proc return proc is not None and proc.poll() is None and not self._stopped_respawning @@ -193,25 +195,24 @@ class HostSupervisor: except FileNotFoundError: return "none" except Exception: - self._remove_registry() - return "invalid-registry" + data = None try: - pid = int(data.get("host_pid") or 0) + pid = int((data or {}).get("host_pid") or 0) except Exception: pid = 0 - if pid <= 0 or not _pid_alive(pid): - self._remove_registry() - return "not-running" - if not self._pid_matches_compute_host(pid): - # PID was reused by another process. Never signal it. - self._remove_registry() - return "pid-reuse-ignored" - self._terminate_pid(pid, timeout=_SHUTDOWN_TIMEOUT_SECS) + if data is None: + outcome = "invalid-registry" + elif pid <= 0 or not _pid_alive(pid): + outcome = "not-running" + elif not self._pid_matches_compute_host(pid): + outcome = "pid-reuse-ignored" # PID reused by another process: never signal it + else: + self._terminate_pid(pid, timeout=_SHUTDOWN_TIMEOUT_SECS) + outcome = "terminated" self._remove_registry() - return "terminated" + return outcome - def submit_turn( - self, frame: dict[str, Any], *, on_complete: Callable[[dict], None] | None = None) -> str: + def submit_turn(self, frame: dict[str, Any], *, on_complete: Callable[[dict], None] | None = None) -> str: self.start() request_id = str(frame.get("request_id") or uuid.uuid4().hex) sid = str(frame.get("sid") or "") @@ -224,16 +225,15 @@ class HostSupervisor: with self._lock: self._pending_turns.pop(request_id, None) if on_complete is not None: - on_complete({ - "type": "turn.error", "sid": sid, "request_id": request_id, - "reason": "send_failed", "message": str(exc)}) + on_complete({"type": "turn.error", "sid": sid, "request_id": request_id, + "reason": "send_failed", "message": str(exc)}) raise return request_id def interrupt(self, sid: str, *, request_id: str | None = None) -> None: self.start() - self._send_frame({ - "type": "interrupt", "sid": sid, "request_id": request_id or uuid.uuid4().hex}) + self._send_frame( + {"type": "interrupt", "sid": sid, "request_id": request_id or uuid.uuid4().hex}) def _await_reply(self, frame: dict[str, Any], request_id: str, timeout: float) -> dict: """Send ``frame`` and block for the host reply carrying ``request_id``.""" @@ -255,29 +255,24 @@ class HostSupervisor: return self._await_reply(frame, request_id, timeout) def reload_mcp(self, sid: str, *, request_id: str | None = None) -> dict: - return self.control( - sid, route_name="reload.mcp", wait=True, - payload={ - "type": "reload_mcp", "sid": sid, "request_id": request_id or uuid.uuid4().hex}) + payload = {"type": "reload_mcp", "sid": sid, "request_id": request_id or uuid.uuid4().hex} + return self.control(sid, route_name="reload.mcp", wait=True, payload=payload) def control( self, sid: str, *, route_name: str, payload: dict[str, Any] | None = None, wait: bool = True, timeout: float = 30.0, on_late_ack: Callable[[dict], None] | None = None, ) -> dict: - """Send a control frame; with ``wait`` block up to ``timeout`` for its ack. - - ``on_late_ack`` (only with ``wait``) keeps the request adoptable after the - waiter gives up: the host's eventual ``control.ack``/``control.error``/``error`` - for this ``request_id`` fires the handler once instead of being dropped. - Bounded by ``_LATE_CONTROL_TTL_SECS`` / ``_LATE_CONTROL_MAX``. - """ + """Send a control frame; with ``wait`` block up to ``timeout`` for its ack. ``on_late_ack`` + (only with ``wait``) keeps the request adoptable after the waiter gives up: the host's + eventual ``control.ack``/``control.error``/``error`` fires it once (bounded by + ``_LATE_CONTROL_TTL_SECS``/``_MAX``) instead of being dropped.""" if route_name not in MUTATOR_ROUTE_TABLE: raise ValueError(f"unclassified host mutator route: {route_name}") self.start() - request_id = str((payload or {}).get("request_id") or uuid.uuid4().hex) - frame = { - "type": "control", **(payload or {}), "sid": sid, "route_name": route_name, - "request_id": request_id} + payload = payload or {} + request_id = str(payload.get("request_id") or uuid.uuid4().hex) + frame = {"type": "control", **payload, "sid": sid, "route_name": route_name, + "request_id": request_id} if not wait: self._send_frame(frame) return {"status": "sent", "request_id": request_id} @@ -288,13 +283,11 @@ class HostSupervisor: self._register_late_control_handler(request_id, on_late_ack) raise - def _register_late_control_handler( - self, request_id: str, handler: Callable[[dict], None]) -> None: + def _register_late_control_handler(self, request_id: str, handler: Callable[[dict], None]) -> None: now = time.monotonic() with self._lock: handlers = self._late_control_handlers - expired = [r for r, (at, _cb) in handlers.items() if now - at > _LATE_CONTROL_TTL_SECS] - for rid in expired: + for rid in [r for r, (at, _cb) in handlers.items() if now - at > _LATE_CONTROL_TTL_SECS]: handlers.pop(rid, None) while len(handlers) >= _LATE_CONTROL_MAX: handlers.pop(min(handlers, key=lambda rid: handlers[rid][0]), None) @@ -307,41 +300,30 @@ class HostSupervisor: if q is not None: with contextlib.suppress(queue.Full): q.put_nowait(frame) - return - if late is None: - return - try: - late[1](frame) - except Exception: - logger.exception( - "compute host late control ack handler failed (request_id=%s)", request_id) + elif late is not None: + _call_logged(late[1], frame, f"compute host late control ack handler failed (request_id={request_id})") def _spawn_locked(self, *, reason: str) -> None: if self._stopped_respawning: raise RuntimeError("compute host respawn disabled after crash loop") self._hello_event.clear() self._hello = {} - env = hermes_subprocess_env(inherit_credentials=True) - env.update(os.environ) - if self.env: - env.update(self.env) + env = {**hermes_subprocess_env(inherit_credentials=True), **os.environ, **(self.env or {})} env["HERMES_COMPUTE_HOST_HEARTBEAT_SECS"] = str(self.heartbeat_secs) root = str(_repo_root()) env.setdefault("PYTHONPATH", root) if root not in env["PYTHONPATH"].split(os.pathsep): env["PYTHONPATH"] = root + os.pathsep + env["PYTHONPATH"] + # Lossy UTF-8 decode: a locale-mismatched byte must not raise inside the drain threads. proc = subprocess.Popen( self.argv, cwd=str(self.cwd), env=env, stdin=subprocess.PIPE, stdout=subprocess.PIPE, - stderr=subprocess.PIPE, text=True, - # Lossy UTF-8 decode: a locale-mismatched byte must not raise inside the - # drain threads and kill the supervisor. - encoding="utf-8", errors="replace", bufsize=1, start_new_session=True) + stderr=subprocess.PIPE, text=True, encoding="utf-8", errors="replace", bufsize=1, + start_new_session=True) self._proc = proc - for target, name in ( - (self._drain_stdout, "compute-host-stdout"), - (self._drain_stderr, "compute-host-stderr"), (self._wait_for_exit, "compute-host-wait"), - ): - _Thread(target=target, args=(proc,), name=name, daemon=True).start() + for target, name in ((self._drain_stdout, "compute-host-stdout"), + (self._drain_stderr, "compute-host-stderr"), + (self._wait_for_exit, "compute-host-wait")): + threading.Thread(target=target, args=(proc,), name=name, daemon=True).start() if not self._hello_event.wait(timeout=10.0): self._terminate_process(proc) raise RuntimeError(f"compute host did not send hello; stderr={self._stderr_tail[-5:]}") @@ -365,10 +347,9 @@ class HostSupervisor: def _persist_registry(self) -> None: self.registry_path.parent.mkdir(parents=True, exist_ok=True) tmp = self.registry_path.with_suffix(self.registry_path.suffix + ".tmp") - payload = { - "host_pid": self.pid, "boot_id": self._hello.get("boot_id") or "", - "build_sha": self._hello.get("build_sha") or "", "started_at": time.time(), - "argv": self.argv} + payload = {"host_pid": self.pid, "boot_id": self._hello.get("boot_id") or "", + "build_sha": self._hello.get("build_sha") or "", "started_at": time.time(), + "argv": self.argv} tmp.write_text(json.dumps(payload, sort_keys=True), encoding="utf-8") tmp.replace(self.registry_path) @@ -400,48 +381,29 @@ class HostSupervisor: def _drain_stderr(self, proc: subprocess.Popen[str]) -> None: assert proc.stderr is not None for raw in proc.stderr: - text = raw.rstrip("\n") - if text: + if text := raw.rstrip("\n"): self._stderr_tail = (self._stderr_tail + [text])[-80:] logger.warning("compute host stderr: %s", text) def _handle_host_frame(self, frame: dict[str, Any]) -> None: ftype = str(frame.get("type") or "") - if ftype in _CONTROL_REPLY_TYPES or (ftype == "error" and frame.get("request_id")): - self._deliver_control_frame(str(frame.get("request_id") or ""), frame) - return - handler = self._HOST_FRAME_HANDLERS.get(ftype) - if handler is not None: - getattr(self, handler)(frame) - - # host frame ``type`` -> handler method name (see also _CONTROL_REPLY_TYPES). - _HOST_FRAME_HANDLERS: dict[str, str] = { - "hello": "_on_hello", "hb": "_on_heartbeat", "rpc": "_on_rpc", "turn.end": "_complete_turn", - "turn.error": "_complete_turn"} - - def _on_hello(self, frame: dict[str, Any]) -> None: - self._hello = dict(frame) - self._hello_event.set() - - def _on_heartbeat(self, frame: dict[str, Any]) -> None: - self._last_progress_counter = int( - frame.get("progress_counter") or self._last_progress_counter) - logger.debug("compute host heartbeat: %s", frame) - - def _on_rpc(self, frame: dict[str, Any]) -> None: - message = frame.get("message") - if isinstance(message, dict): - self.rpc_sink(message) - - def _complete_turn(self, frame: dict[str, Any]) -> None: request_id = str(frame.get("request_id") or "") - with self._lock: - pending = self._pending_turns.pop(request_id, None) - if pending is not None and pending[1] is not None: - try: - pending[1](frame) - except Exception: - logger.exception("compute host turn completion callback failed") + if ftype in _CONTROL_REPLY_TYPES or (ftype == "error" and request_id): + self._deliver_control_frame(request_id, frame) + elif ftype == "hello": + self._hello = dict(frame) + self._hello_event.set() + elif ftype == "hb": + self._last_progress_counter = int(frame.get("progress_counter") or self._last_progress_counter) + logger.debug("compute host heartbeat: %s", frame) + elif ftype == "rpc": + if isinstance(frame.get("message"), dict): + self.rpc_sink(frame["message"]) + elif ftype in ("turn.end", "turn.error"): + with self._lock: + pending = self._pending_turns.pop(request_id, None) + if pending is not None and pending[1] is not None: + _call_logged(pending[1], frame, "compute host turn completion callback failed") def _wait_for_exit(self, proc: subprocess.Popen[str]) -> None: code = proc.wait() @@ -459,32 +421,21 @@ class HostSupervisor: with self._lock: pending = self._pending_turns self._pending_turns = {} + failure = {"reason": reason, "message": message} for request_id, (sid, cb) in pending.items(): - self.rpc_sink({ - "jsonrpc": "2.0", - "method": "event", - "params": { - "type": "error", "session_id": sid, - "payload": {"message": message, "reason": reason}}}) + self.rpc_sink({"jsonrpc": "2.0", "method": "event", + "params": {"type": "error", "session_id": sid, "payload": dict(failure)}}) if cb is not None: - try: - cb({ - "type": "turn.error", "sid": sid, "request_id": request_id, - "reason": reason, "message": message}) - except Exception: - logger.exception("compute host error callback failed") - # A crashed host never emits the late acks timed-out control waiters still - # expect; fail them too so the client's "still running" notice can't hang. + frame = {"type": "turn.error", "sid": sid, "request_id": request_id, **failure} + _call_logged(cb, frame, "compute host error callback failed") + # A crashed host never emits the late acks timed-out control waiters still expect; fail + # them too so the client's "still running" notice can't hang. with self._lock: late = self._late_control_handlers self._late_control_handlers = {} for request_id, (_registered_at, handler) in late.items(): - try: - handler({ - "type": "control.error", "request_id": request_id, "reason": reason, - "message": message}) - except Exception: - logger.exception("compute host late control error handler failed") + frame = {"type": "control.error", "request_id": request_id, **failure} + _call_logged(handler, frame, "compute host late control error handler failed") def _maybe_respawn_after_crash(self) -> None: now = time.monotonic() @@ -508,42 +459,30 @@ class HostSupervisor: self._spawn_locked(reason="crash") except Exception: logger.exception("compute host respawn failed") - _Thread(target=_respawn, name="compute-host-respawn", daemon=True).start() + threading.Thread(target=_respawn, name="compute-host-respawn", daemon=True).start() + _pid_matches_compute_host = staticmethod(is_compute_host_identity) def _terminate_pid(self, pid: int, *, timeout: float = _SHUTDOWN_TIMEOUT_SECS) -> None: - try: - os.kill(pid, signal.SIGTERM) - except ProcessLookupError: - return - except Exception: - logger.debug("failed to SIGTERM compute host pid=%s", pid, exc_info=True) + if not _signal_pid(pid, signal.SIGTERM, "SIGTERM"): return deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - if not _pid_alive(pid): + while _pid_alive(pid): + if time.monotonic() >= deadline: + _signal_pid(pid, signal.SIGKILL, "SIGKILL") return time.sleep(0.05) - try: - os.kill(pid, signal.SIGKILL) - except ProcessLookupError: - return - except Exception: - logger.debug("failed to SIGKILL compute host pid=%s", pid, exc_info=True) def _terminate_process(self, proc: subprocess.Popen[str]) -> None: if proc.poll() is not None: return - try: + with contextlib.suppress(Exception): proc.terminate() proc.wait(timeout=_SHUTDOWN_TIMEOUT_SECS) return - except Exception: - pass - with contextlib.suppress(Exception): - proc.kill() - with contextlib.suppress(Exception): - proc.wait(timeout=2) + for step in (proc.kill, lambda: proc.wait(timeout=2)): + with contextlib.suppress(Exception): + step() __all__ = ["MUTATOR_ROUTE_TABLE", "HostSupervisor", "append_log_record", "is_compute_host_identity"] diff --git a/tui_gateway/hosted_room_driver.py b/tui_gateway/hosted_room_driver.py index 3e89c3d879..1867c02f4b 100644 --- a/tui_gateway/hosted_room_driver.py +++ b/tui_gateway/hosted_room_driver.py @@ -1,12 +1,10 @@ """Runtime adapter for gateway-owned hosted room turns. -The durable state machine lives in :mod:`gateway.hosted_room_driver`. This module -owns the process-local worker and a small injected session adapter (seven methods); -it never imports the gateway server or constructs agents. One bounded supervisor -schedules independent room workers: profile turn locks still serialize Bots sharing -a profile, while a room waiting for approval cannot stall unrelated rooms. Hosted -member sessions reuse ``Group: `` so a local-to-hosted migration keeps one -canonical transcript instead of forking a second conversation. +The durable state machine lives in :mod:`gateway.hosted_room_driver`; this module owns the +process-local worker and an injected session adapter, never the gateway server or agents. +One bounded supervisor schedules independent room workers: profile turn locks serialize Bots +sharing a profile while a room waiting for approval cannot stall unrelated rooms. Member +sessions reuse ``Group: `` so a local-to-hosted migration keeps one transcript. """ from __future__ import annotations @@ -33,36 +31,28 @@ _STOP_PENDING = "stop retry remains pending: {exc}" class InternalSessionRPC(Protocol): - """Normalized in-process session operations required by the room driver.""" + """Normalized in-process session operations required by the room driver. - def resolve_exact(self, *, profile: str, title: str, source: str) -> Mapping[str, Any] | None: - """Return the exact titled session under ``profile``, if it exists.""" - - def create(self, *, profile: str, title: str, source: str) -> Mapping[str, Any]: - """Create a session without model or provider overrides.""" - - def resume(self, *, profile: str, session_id: str, source: str) -> Mapping[str, Any]: - """Resume the canonical room session.""" + ``submit`` durably reports one fenced turn's terminal result via ``on_terminal``; + ``interrupt`` acts only while the current turn still matches ``expected_task_id``. + """ + def resolve_exact( + self, *, profile: str, title: str, source: str) -> Mapping[str, Any] | None: ... + def create(self, *, profile: str, title: str, source: str) -> Mapping[str, Any]: ... + def resume(self, *, profile: str, session_id: str, source: str) -> Mapping[str, Any]: ... def submit( self, *, profile: str, session_id: str, prompt: str, source: str, task: state.TaskIdentity, execution_generation: int, on_terminal: Callable[[Mapping[str, Any]], None], - ) -> Mapping[str, Any]: - """Submit one fenced room turn and durably report its terminal result.""" - - def history(self, *, profile: str, session_id: str, source: str) -> Sequence[Mapping[str, Any]]: - """Return normalized session messages in durable order.""" - - def info(self, *, profile: str, session_id: str, source: str) -> Mapping[str, Any]: - """Return normalized live status for a session.""" - + ) -> Mapping[str, Any]: ... + def history( + self, *, profile: str, session_id: str, source: str) -> Sequence[Mapping[str, Any]]: ... + def info(self, *, profile: str, session_id: str, source: str) -> Mapping[str, Any]: ... def interrupt( - self, *, profile: str, session_id: str, source: str, expected_task_id: str - ) -> Mapping[str, Any] | None: - """Interrupt only when the current turn still matches the expected task.""" + self, *, profile: str, session_id: str, source: str, expected_task_id: str, + ) -> Mapping[str, Any] | None: ... -# Resolve the local or peer session transport for one durable room task. MemberTransportResolver = Callable[["HostedRoomBinding", Mapping[str, Any]], InternalSessionRPC] @@ -97,9 +87,8 @@ def _session_kw(profile: str, session_id: str) -> dict[str, str]: def _fences(task: Mapping[str, Any]) -> dict[str, int]: - return { - "expected_execution_generation": int(task["execution_generation"]), - "expected_cancel_generation": int(task["cancel_generation"])} + return {"expected_execution_generation": int(task["execution_generation"]), + "expected_cancel_generation": int(task["cancel_generation"])} class HostedRoomRuntime: @@ -119,19 +108,16 @@ class HostedRoomRuntime: indeterminate_defer_seconds: float = 60.0, max_concurrent_rooms: int = 4, unavailable_retry_min_seconds: float = 1.0, unavailable_retry_max_seconds: float = 30.0, process_generation: str | None = None) -> None: - for name, value in ( - ("lease_ttl_seconds", lease_ttl_seconds), - ("poll_interval_seconds", poll_interval_seconds), - ("active_poll_interval_seconds", active_poll_interval_seconds), - ("turn_timeout_seconds", turn_timeout_seconds), - ("indeterminate_defer_seconds", indeterminate_defer_seconds)): + positive = dict( + lease_ttl_seconds=lease_ttl_seconds, poll_interval_seconds=poll_interval_seconds, + active_poll_interval_seconds=active_poll_interval_seconds, + turn_timeout_seconds=turn_timeout_seconds, + indeterminate_defer_seconds=indeterminate_defer_seconds) + for name, value in positive.items(): if value <= 0: raise ValueError(f"{name} must be positive") - if ( - not isinstance(max_concurrent_rooms, int) - or isinstance(max_concurrent_rooms, bool) - or max_concurrent_rooms < 1 - ): + if (not isinstance(max_concurrent_rooms, int) or isinstance(max_concurrent_rooms, bool) + or max_concurrent_rooms < 1): raise ValueError("max_concurrent_rooms must be a positive integer") if not 0 < unavailable_retry_min_seconds <= unavailable_retry_max_seconds: raise ValueError("unavailable retry bounds are invalid") @@ -141,22 +127,17 @@ class HostedRoomRuntime: self.rpc, self.transport_resolver, self.turn_lock = rpc, transport_resolver, turn_lock self.prepare_room, self.publish_terminal = prepare_room, publish_terminal self.pending_action, self.clock = pending_action, clock - self.lease_ttl_seconds = float(lease_ttl_seconds) - self.poll_interval_seconds = float(poll_interval_seconds) - self.active_poll_interval_seconds = float(active_poll_interval_seconds) - self.turn_timeout_seconds = float(turn_timeout_seconds) - self.indeterminate_defer_seconds = float(indeterminate_defer_seconds) + for name, value in positive.items(): + setattr(self, name, float(value)) self.max_concurrent_rooms = max_concurrent_rooms - self.unavailable_retry_min_seconds = float(unavailable_retry_min_seconds) - self.unavailable_retry_max_seconds = float(unavailable_retry_max_seconds) + self.unavailable_retry_min_seconds, self.unavailable_retry_max_seconds = ( + float(unavailable_retry_min_seconds), float(unavailable_retry_max_seconds)) self.process_generation = process_generation or uuid.uuid4().hex self._rooms_provider: Callable[[], Iterable[HostedRoomBinding]] = ( - cast(Callable[[], Iterable[HostedRoomBinding]], rooms) - if callable(rooms) + cast(Callable[[], Iterable[HostedRoomBinding]], rooms) if callable(rooms) else (lambda bindings=tuple(rooms): bindings)) - self._stop, self._wake = threading.Event(), threading.Event() - self._thread: threading.Thread | None = None + self._thread = self._last_error = None self._room_threads: dict[str, threading.Thread] = {} self._rooms_needing_reschedule: set[str] = set() self._leases: dict[str, state.DriverLease] = {} @@ -165,10 +146,8 @@ class HostedRoomRuntime: self._ambiguous_rooms: dict[str, float] = {} self._unavailable_route_retries: dict[tuple[str, str], dict[str, float]] = {} self._blocked_rooms: set[str] = set() - self._status_lock = threading.Lock() - self._current_tasks: dict[str, state.TaskIdentity] = {} + self._status_lock, self._current_tasks = threading.Lock(), {} self._room_schedule_cursor, self._cycles = 0, 0 - self._last_error: str | None = None # ------------------------------------------------------------------ lifecycle def start(self) -> None: @@ -196,14 +175,13 @@ class HostedRoomRuntime: room_threads = tuple(self._room_threads.values()) for room_thread in room_threads: room_thread.join(max(0.0, deadline - time.monotonic())) - return not thread.is_alive() and all(not t.is_alive() for t in room_threads) + return not any(t.is_alive() for t in (thread, *room_threads)) def wakeup(self) -> None: """Wake the worker after task admission or a room-state change.""" with self._status_lock: - # Rooms that still own a worker slot are revisited once that thread - # exits, closing the race between terminal publication/route repair - # and the longer idle fallback without busy-looping idle rooms. + # Rooms still owning a worker slot are revisited once that thread exits, closing + # the terminal-publication/route-repair race without busy-looping idle rooms. self._rooms_needing_reschedule.update(self._room_threads) self._wake.set() @@ -213,24 +191,19 @@ class HostedRoomRuntime: thread = self._thread current_tasks = tuple(self._current_tasks.values()) return { - "running": bool(thread and thread.is_alive()), - "stopping": self._stop.is_set(), + "running": bool(thread and thread.is_alive()), "stopping": self._stop.is_set(), "process_generation": self.process_generation, "current_task": current_tasks[0] if current_tasks else None, - "current_tasks": current_tasks, - "leased_rooms": tuple(sorted(self._leases)), + "current_tasks": current_tasks, "leased_rooms": tuple(sorted(self._leases)), "blocked_rooms": tuple(sorted(self._blocked_rooms)), - "last_error": self._last_error, - "cycles": self._cycles} + "last_error": self._last_error, "cycles": self._cycles} # ------------------------------------------------------------------ public ops def cancel(self, identity: state.TaskIdentity, *, cancel_id: str) -> dict[str, Any]: """Persist a stop intent, then commit cancellation after acknowledgement. - The worker thread transitions tasks concurrently (queued -> running -> - terminal), so the status read is only a routing hint: every fast-path - failure caused by a concurrent transition re-reads and re-routes instead - of surfacing a transient `InvalidTaskTransitionError`/`StaleTaskError`. + The worker transitions tasks concurrently, so the status read is only a routing + hint: a fast-path fence failure re-reads and re-routes instead of surfacing it. """ for _ in range(_CANCEL_ROUTE_RETRIES): before = state.get_task(self.db_path, identity) @@ -239,34 +212,27 @@ class HostedRoomRuntime: if before["status"] in state.TERMINAL_STATUSES: raise state.InvalidTaskTransitionError( f"cannot cancel task in state '{before['status']}'") - fenced = dict( - cancel_id=cancel_id, expected_cancel_generation=before["cancel_generation"], - clock=self.clock) - if before["status"] in {"queued", "deferred"}: - try: - cancelled = state.cancel_task(self.db_path, identity, **fenced) - except (state.InvalidTaskTransitionError, state.StaleTaskError): - continue # lost the race with the worker; re-route - self.wakeup() - return cancelled + direct = before["status"] in {"queued", "deferred"} try: - stopping = state.begin_task_cancel(self.db_path, identity, **fenced) + result = (state.cancel_task if direct else state.begin_task_cancel)( + self.db_path, identity, cancel_id=cancel_id, + expected_cancel_generation=before["cancel_generation"], clock=self.clock) except (state.InvalidTaskTransitionError, state.StaleTaskError): - continue # settled or re-queued mid-flight; re-route - binding = self._binding_for_room(identity.room_id) - try: - if binding is not None: - lease = self._ensure_lease(binding) - if self._peer_stop_acknowledged(binding, stopping) or ( - not self._settle_stopping_completion(binding, stopping, lease) - and self._interrupt_stopping_task(binding, stopping)): - self._complete_cancel(stopping, cancel_id=cancel_id) - except Exception as exc: - self._record_error(f"stop remains pending: {exc}") + continue # lost the race with the worker (settled or re-queued); re-route + if not direct: + binding = self._binding_for_room(identity.room_id) + try: + if binding is not None: + lease = self._ensure_lease(binding) + if self._peer_stop_acknowledged(binding, result) or ( + not self._settle_stopping_completion(binding, result, lease) + and self._interrupt_stopping_task(binding, result)): + self._complete_cancel(result, cancel_id=cancel_id) + except Exception as exc: + self._record_error(f"stop remains pending: {exc}") self.wakeup() - return state.get_task(self.db_path, identity) - # Exhausted routing retries under sustained contention: surface the - # live status honestly rather than a transient transition error. + return result if direct else state.get_task(self.db_path, identity) + # Routing retries exhausted under contention: surface the live status honestly. final = state.get_task(self.db_path, identity) if final["status"] == "cancelled": return final @@ -285,20 +251,21 @@ class HostedRoomRuntime: lease = self._ensure_lease(binding) if task["status"] == "deferred": return self._requeue(state.requeue_deferred_task, task, lease, identity.room_id) - # Explicit Retry may resume the exact stored session; an automatic - # abandoned-attempt scan remains non-resuming for local sessions. + # Explicit Retry may resume the exact stored session; the automatic abandoned-attempt + # scan stays non-resuming for local sessions. inspection = self._inspect_recovery_session(binding, task) if inspection.terminal is not None: return self._resolve_indeterminate(binding, task, lease, inspection.terminal) if inspection.status == "cancelled": - return self._resolve_remote_cancel(binding, task, lease) + return self._fenced( + state.resolve_indeterminate_cancellation, binding, task, lease, + cancel_id=f"remote-cancel:{task['execution_generation']}") if inspection.active: self._set_blocked(identity.room_id, True) raise state.InvalidTaskTransitionError( "cannot retry while the original task attempt is still active") return self._requeue(state.requeue_indeterminate_task, task, lease, identity.room_id) - # ------------------------------------------------------------------ durable-state helpers def _publish(self, binding: HostedRoomBinding, task: dict[str, Any]) -> dict[str, Any]: if self.publish_terminal is not None: self.publish_terminal(binding, task) @@ -328,9 +295,9 @@ class HostedRoomRuntime: def _complete_cancel( self, task: Mapping[str, Any], *, cancel_id: str | None = None) -> dict[str, Any]: return state.complete_task_cancel( - self.db_path, task["identity"], + self.db_path, task["identity"], clock=self.clock, cancel_id=task["cancel_id"] if cancel_id is None else cancel_id, - expected_cancel_generation=task["cancel_generation"], clock=self.clock) + expected_cancel_generation=task["cancel_generation"]) def _resolve_indeterminate( self, binding: HostedRoomBinding, task: Mapping[str, Any], lease: state.DriverLease, @@ -339,19 +306,6 @@ class HostedRoomRuntime: state.resolve_indeterminate_task, binding, task, lease, publish=publish, **asdict(terminal)) - def _resolve_remote_cancel( - self, binding: HostedRoomBinding, task: Mapping[str, Any], lease: state.DriverLease, - *, publish: bool = True) -> dict[str, Any]: - return self._fenced( - state.resolve_indeterminate_cancellation, binding, task, lease, publish=publish, - cancel_id=f"remote-cancel:{task['execution_generation']}") - - def _settle_stopping( - self, binding: HostedRoomBinding, task: Mapping[str, Any], lease: state.DriverLease, - **terminal: Any) -> dict[str, Any]: - """Settle a stopping task from ``terminal`` (settlement_id/status/result + fences).""" - return self._fenced(state.settle_stopping_task, binding, task, lease, **terminal) - def _finish_stop( self, binding: HostedRoomBinding, task: Mapping[str, Any], lease: state.DriverLease ) -> bool: @@ -363,13 +317,11 @@ class HostedRoomRuntime: return True return False - # ------------------------------------------------------------------ session probes def _resume_exact( self, transport: InternalSessionRPC, room_id: str, profile: str) -> str | None: - """Resolve the canonical room session and return its resumed runtime id. + """Resume the canonical room session and return its runtime id (None when absent). - Callers must use the returned id (not the stored one) for every - subsequent history/info probe; resume may hand back a different id. + Probes must use the returned id, not the stored one: resume may hand back another. """ session = self._resolve_or_create(transport, profile, room_id, create=False) return None if session is None else _session_id(session) @@ -377,10 +329,8 @@ class HostedRoomRuntime: def _open_session( self, binding: HostedRoomBinding, task: Mapping[str, Any], *, peer_only: bool = False ) -> tuple[InternalSessionRPC | None, str | None, str | None]: - """Return ``(transport, profile, resumed session id)``; the id is None when unusable. - - A missing transport (or the local one when ``peer_only``) yields no session id. - """ + """Return ``(transport, profile, resumed session id)``; no id when the transport is + missing (or local under ``peer_only``) or the session is absent.""" transport = self._transport_for(binding, task) if transport is None or (peer_only and transport is self.rpc): return transport, None, None @@ -395,43 +345,34 @@ class HostedRoomRuntime: transport.history(**_session_kw(profile, session_id)), task["identity"], int(task["execution_generation"])) - @staticmethod - def _info_acknowledges_peer_cancel(info: Mapping[str, Any], task: Mapping[str, Any]) -> bool: - """Accept only one exact peer task attempt's terminal Stop receipt.""" + def _peer_stop_acknowledged(self, binding: HostedRoomBinding, task: Mapping[str, Any]) -> bool: + """Probe a peer's exact durable terminal Stop receipt before reading history.""" + transport, profile, session_id = self._open_session(binding, task, peer_only=True) + if session_id is None: + return False + info = transport.info(**_session_kw(profile, session_id)) return ( not _info_active(info) and str(info.get("status") or "") in _STOP_ACK_STATUSES and str(info.get("task_id") or "") == task["identity"].task_id and int(info.get("execution_generation") or 0) == int(task["execution_generation"])) - def _peer_stop_acknowledged(self, binding: HostedRoomBinding, task: Mapping[str, Any]) -> bool: - """Probe a peer's exact durable terminal status before reading history.""" - transport, profile, session_id = self._open_session(binding, task, peer_only=True) - if session_id is None: - return False - return self._info_acknowledges_peer_cancel( - transport.info(**_session_kw(profile, session_id)), task) - def _interrupt_stopping_task(self, binding: HostedRoomBinding, task: Mapping[str, Any]) -> bool: transport, profile, session_id = self._open_session(binding, task) if session_id is None: - # A local accepted turn cannot survive without its canonical - # session: an authoritative absence is a safe Stop acknowledgement - # (resolution errors raise). A peer remains uncertain instead. + # A local turn cannot survive without its canonical session, so an authoritative + # absence is a safe Stop acknowledgement (errors raise); a peer stays uncertain. return transport is not None and transport is self.rpc info = transport.info(**_session_kw(profile, session_id)) if not _info_active(info): - # History was checked immediately before this probe. An exact - # session that is no longer active cannot keep executing, and after - # a restart its process-local task marker is expected to be absent. + # History was checked just before this probe: an inactive exact session cannot + # keep executing, and after a restart its process-local task marker is absent. return True if not _info_is_active_for(info, task["identity"], require_exact=True): return False result = transport.interrupt( **_session_kw(profile, session_id), expected_task_id=task["identity"].task_id) - if result is None: - return False - return ( + return result is not None and ( result.get("interrupted") is True or str(result.get("status") or "") in _STOP_ACK_STATUSES) @@ -445,32 +386,28 @@ class HostedRoomRuntime: receipt = self._terminal_from_history(transport, profile, session_id, task) if receipt is None: return False - self._settle_stopping(binding, task, lease, **asdict(receipt)) + self._fenced(state.settle_stopping_task, binding, task, lease, **asdict(receipt)) return True def _report_pending_action( self, task: Mapping[str, Any], *, session_id: str, info: Mapping[str, Any]) -> None: if self.pending_action is None: return - approval = info.get("pending_approval") or info.get("approval") - action = None + approval, action = info.get("pending_approval") or info.get("approval"), None if isinstance(approval, Mapping): choices = [c for c in approval.get("choices") or () if c in {"once", "deny"}] safe_approval = {**approval, "choices": choices or ["once", "deny"]} action = { - "kind": "approval", - "task_id": task["identity"].task_id, + "kind": "approval", "task_id": task["identity"].task_id, "execution_generation": int(task["execution_generation"]), - "run_id": info.get("run_id"), - "session_id": session_id, - "request_id": safe_approval.get("request_id"), - "approval": safe_approval} + "run_id": info.get("run_id"), "session_id": session_id, + "request_id": safe_approval.get("request_id"), "approval": safe_approval} self.pending_action(task["identity"].room_id, _member_id(task), action) def _retry_stopping_tasks(self, binding: HostedRoomBinding, lease: state.DriverLease) -> bool: - for task in state.list_tasks(self.db_path, room_id=binding.room_id, status="stopping"): + for task in self._tasks(binding, "stopping"): try: - lease = self._renew_lease_if_needed(binding, lease) + lease = self._renew_lease_if_needed(lease) if self._peer_stop_acknowledged(binding, task): self._complete_cancel(task) continue @@ -485,8 +422,7 @@ class HostedRoomRuntime: def _worker_loop(self) -> None: try: while not self._stop.is_set(): - # Clear before work so a write racing the cycle remains set and - # causes an immediate follow-up pass rather than being lost. + # Clear before work so a write racing the cycle forces a follow-up pass. self._wake.clear() try: self._run_cycle() @@ -516,9 +452,7 @@ class HostedRoomRuntime: return with self._status_lock: self._room_threads = { - room_id: thread - for room_id, thread in self._room_threads.items() - if thread.is_alive()} + room_id: t for room_id, t in self._room_threads.items() if t.is_alive()} available = self.max_concurrent_rooms - len(self._room_threads) active_rooms = set(self._room_threads) if available <= 0: @@ -546,16 +480,15 @@ class HostedRoomRuntime: try: self._process_room(binding) except state.LeaseHeldError: - return + pass except Exception as exc: if isinstance(exc, (state.RoomUnavailableError, state.StaleLeaseError)): self._drop_lease(binding.room_id) self._set_blocked(binding.room_id, False) self._record_error(f"room {binding.room_id}: {exc}") finally: - current = threading.current_thread() with self._status_lock: - if self._room_threads.get(binding.room_id) is current: + if self._room_threads.get(binding.room_id) is threading.current_thread(): self._room_threads.pop(binding.room_id, None) should_wake = binding.room_id in self._rooms_needing_reschedule self._rooms_needing_reschedule.discard(binding.room_id) @@ -568,24 +501,25 @@ class HostedRoomRuntime: self._inspect_abandoned_attempts(binding) deferred_until = self._ambiguous_rooms.get(binding.room_id) if deferred_until is not None: - running = state.list_tasks(self.db_path, room_id=binding.room_id, status="running") - if running and self.clock() < deferred_until: + if self._tasks(binding, "running") and self.clock() < deferred_until: return self._ambiguous_rooms.pop(binding.room_id, None) lease = self._ensure_lease(binding) - recovery_key = (lease.room_id, lease.lease_generation) - if recovery_key not in self._recovered_leases: + if (lease.room_id, lease.lease_generation) not in self._recovered_leases: state.recover_room(self.db_path, lease, clock=self.clock) - self._recovered_leases.add(recovery_key) + self._recovered_leases.add((lease.room_id, lease.lease_generation)) if self._retry_stopping_tasks(binding, lease): self._set_blocked(binding.room_id, True) return if self._reconcile_indeterminate(binding, lease): return - for task in state.list_tasks(self.db_path, room_id=binding.room_id, status="queued"): - if self._stop.is_set() or self._route_retry_is_deferred(task): + for task in self._tasks(binding, "queued"): + retry = self._unavailable_route_retries.get( + (task["identity"].room_id, _member_id(task))) + if self._stop.is_set() or ( + retry is not None and self.clock() < retry["next_attempt_at"]): return - lease = self._renew_lease_if_needed(binding, lease) + lease = self._renew_lease_if_needed(lease) attempt = state.start_task( self.db_path, task["identity"], lease, expected_cancel_generation=task["cancel_generation"], clock=self.clock) @@ -594,17 +528,8 @@ class HostedRoomRuntime: if current["status"] not in state.TERMINAL_STATUSES: return - # ------------------------------------------------------------------ route retry backoff - @staticmethod - def _route_retry_key(task: Mapping[str, Any]) -> tuple[str, str]: - return task["identity"].room_id, _member_id(task) - - def _route_retry_is_deferred(self, task: Mapping[str, Any]) -> bool: - retry = self._unavailable_route_retries.get(self._route_retry_key(task)) - return retry is not None and self.clock() < retry["next_attempt_at"] - def _defer_unavailable_route(self, task: Mapping[str, Any]) -> float: - key = self._route_retry_key(task) + key = (task["identity"].room_id, _member_id(task)) previous = self._unavailable_route_retries.get(key) lo, hi = self.unavailable_retry_min_seconds, self.unavailable_retry_max_seconds delay = lo if previous is None else min(hi, max(lo, previous["delay"] * 2)) @@ -612,35 +537,27 @@ class HostedRoomRuntime: "delay": delay, "next_attempt_at": self.clock() + delay} return delay - def _clear_unavailable_route_retry(self, task: Mapping[str, Any]) -> None: - self._unavailable_route_retries.pop(self._route_retry_key(task), None) - # ------------------------------------------------------------------ leases def _ensure_lease(self, binding: HostedRoomBinding) -> state.DriverLease: with self._status_lock: current = self._leases.get(binding.room_id) if current is not None: try: - return self._renew_lease_if_needed(binding, current) + return self._renew_lease_if_needed(current) except state.StaleLeaseError: self._drop_lease(binding.room_id) - lease = state.acquire_lease( self.db_path, room_id=binding.room_id, gateway_id=binding.gateway_id, authority_epoch=binding.authority_epoch, process_generation=self.process_generation, ttl_seconds=self.lease_ttl_seconds, clock=self.clock) with self._status_lock: self._leases[binding.room_id] = lease - self._recovered_leases = { - key for key in self._recovered_leases if key[0] != binding.room_id} + self._recovered_leases = {k for k in self._recovered_leases if k[0] != binding.room_id} return lease def _renew_lease_if_needed( - self, binding: HostedRoomBinding, lease: state.DriverLease, *, force: bool = False - ) -> state.DriverLease: - del binding - renew_at = lease.expires_at - (self.lease_ttl_seconds / 2) - if not force and self.clock() < renew_at: + self, lease: state.DriverLease, *, force: bool = False) -> state.DriverLease: + if not force and self.clock() < lease.expires_at - (self.lease_ttl_seconds / 2): return lease renewed = state.renew_lease( self.db_path, lease, ttl_seconds=self.lease_ttl_seconds, clock=self.clock) @@ -662,25 +579,23 @@ class HostedRoomRuntime: def _execute_attempt( self, binding: HostedRoomBinding, task: Mapping[str, Any], attempt: state.TaskAttempt ) -> None: - profile = task["payload"]["target_profile"] + profile, submit_attempted = task["payload"]["target_profile"], False transport = self._transport_for(binding, task) - submit_attempted = False with self._status_lock: self._current_tasks[binding.room_id] = attempt.identity try: with self.turn_lock(profile): session = self._resolve_or_create(transport, profile, binding.room_id) - # An in-process submit should fail before admission or return - # after it, but an unexpected exception at that boundary is - # still ambiguous. Never terminalize it as a proven failure. - submit_attempted = True - session_id = _session_id(session) + # A submit should fail before admission or return after it; an unexpected + # exception at that boundary is ambiguous, never a proven failure. + submit_attempted, session_id = True, _session_id(session) deadline_monotonic = time.monotonic() + self.turn_timeout_seconds transport.submit( **_session_kw(profile, session_id), prompt=task["payload"]["prompt"], task=attempt.identity, execution_generation=attempt.execution_generation, on_terminal=lambda receipt: self._on_terminal(binding, attempt, receipt)) - self._clear_unavailable_route_retry(task) + self._unavailable_route_retries.pop( + (task["identity"].room_id, _member_id(task)), None) receipt = self._wait_for_terminal( binding, profile=profile, session_id=session_id, attempt=attempt, transport=transport, deadline_monotonic=deadline_monotonic) @@ -689,34 +604,30 @@ class HostedRoomRuntime: state.settle_task(self.db_path, attempt, **asdict(receipt), clock=self.clock) except (state.StaleLeaseError, state.StaleTaskError) as exc: self._drop_lease(binding.room_id) - self._record_error(f"task {attempt.identity.task_id} fenced: {exc}") + self._record_task_error(attempt, f"fenced: {exc}") except Exception as exc: if submit_attempted and bool(getattr(exc, "not_admitted", False)): try: state.requeue_not_admitted_task(self.db_path, attempt, clock=self.clock) except (state.StaleLeaseError, state.StaleTaskError) as fence_exc: self._mark_ambiguous(binding, attempt) - self._record_error( - f"task {attempt.identity.task_id} not-admitted proof lost " - f"its fence: {fence_exc}") + self._record_task_error( + attempt, f"not-admitted proof lost its fence: {fence_exc}") else: delay = self._defer_unavailable_route(task) - self._record_error( - f"task {attempt.identity.task_id} was not admitted; " - f"queued for retry in {delay:g}s") + self._record_task_error( + attempt, f"was not admitted; queued for retry in {delay:g}s") elif submit_attempted: self._mark_ambiguous(binding, attempt) - self._record_error( - f"task {attempt.identity.task_id} observation failed after submit: {exc}") + self._record_task_error(attempt, f"observation failed after submit: {exc}") else: self._settle_failure_if_current(attempt, exc) finally: with self._status_lock: self._current_tasks.pop(binding.room_id, None) - # The task may have published a reply, deferred a member, or - # exposed the next turn while this room thread still occupied - # its slot. Schedule exactly one immediate follow-up after the - # thread leaves; idle room scans never set this marker. + # The task may have published a reply or exposed the next turn while this + # thread held its slot: schedule exactly one follow-up after it leaves + # (idle room scans never set this marker). self._rooms_needing_reschedule.add(binding.room_id) def _mark_ambiguous(self, binding: HostedRoomBinding, attempt: state.TaskAttempt) -> None: @@ -737,23 +648,23 @@ class HostedRoomRuntime: or f"reply:{attempt.identity.task_id}:{attempt.execution_generation}", result=_bounded_terminal_result(receipt)) try: - settled = state.settle_task(self.db_path, attempt, **asdict(terminal), clock=self.clock) - self._publish(binding, settled) + self._publish( + binding, + state.settle_task(self.db_path, attempt, **asdict(terminal), clock=self.clock)) except state.StaleTaskError: with suppress(state.StaleLeaseError, state.StaleTaskError): current = state.get_task(self.db_path, attempt.identity) if current["status"] == "stopping": - self._settle_stopping( - binding, current, attempt.lease, **asdict(terminal), + self._fenced( + state.settle_stopping_task, binding, current, attempt.lease, + **asdict(terminal), expected_execution_generation=attempt.execution_generation) except state.StaleLeaseError: - # Cancellation, disband, or authority transfer won the durable - # race. The model result is intentionally discarded; never turn a - # correct fence into a worker thread exception. + # Cancellation, disband, or authority transfer won the durable race: the model + # result is discarded rather than turning a correct fence into a thread exception. pass except state.DriverStateError as exc: - # A malformed terminal receipt must not escape the callback and - # hold the profile lock until the deadline. + # A malformed receipt must not escape the callback and hold the profile lock. self._settle_failure_if_current( attempt, RuntimeError(f"terminal result could not be committed: {exc}")) self.wakeup() @@ -769,7 +680,7 @@ class HostedRoomRuntime: return None if task["status"] == "stopping": try: - lease = self._renew_lease_if_needed(binding, lease) + lease = self._renew_lease_if_needed(lease) if self._finish_stop(binding, task, lease): return None except Exception as exc: @@ -780,7 +691,7 @@ class HostedRoomRuntime: if time.monotonic() >= deadline_monotonic: self._expire_attempt_deadline(binding, task, lease) return None - lease = self._renew_lease_if_needed(binding, lease) + lease = self._renew_lease_if_needed(lease) receipt = self._terminal_from_history(transport, profile, session_id, task) if receipt is not None: return receipt @@ -791,19 +702,14 @@ class HostedRoomRuntime: self._wake.clear() return None - # ------------------------------------------------------------------ deadline stops - @staticmethod - def _is_deadline_stop(task: Mapping[str, Any]) -> bool: - return str(task.get("cancel_id") or "").startswith("deadline:") - def _complete_acknowledged_stop( self, binding: HostedRoomBinding, task: Mapping[str, Any], lease: state.DriverLease ) -> dict[str, Any]: """Terminalize an acknowledged Stop: deadline stops publish an explicit failure.""" - if not self._is_deadline_stop(task): + if not str(task.get("cancel_id") or "").startswith("deadline:"): return self._complete_cancel(task) - return self._settle_stopping( - binding, task, lease, + return self._fenced( + state.settle_stopping_task, binding, task, lease, settlement_id=f"deadline:{int(task['execution_generation'])}", status="failed", result={ "error": "This Group Chat turn exceeded its configured time limit and was stopped.", @@ -816,23 +722,22 @@ class HostedRoomRuntime: """Fence, stop, and terminalize one exact attempt at its deadline.""" if task["status"] == "running": task = state.begin_task_cancel( - self.db_path, task["identity"], + self.db_path, task["identity"], clock=self.clock, cancel_id=f"deadline:{int(task['execution_generation'])}", - expected_cancel_generation=int(task["cancel_generation"]), clock=self.clock) + expected_cancel_generation=int(task["cancel_generation"])) elif task["status"] != "stopping": return # A user Stop that won the race keeps its own cancellation semantics. - if not self._is_deadline_stop(task): + if not str(task.get("cancel_id") or "").startswith("deadline:"): return - lease = self._renew_lease_if_needed(binding, lease, force=True) - if self._finish_stop(binding, task, lease): - return - self._record_error( - f"task {task['identity'].task_id} exceeded its deadline; stop remains pending") + lease = self._renew_lease_if_needed(lease, force=True) + if not self._finish_stop(binding, task, lease): + self._record_error( + f"task {task['identity'].task_id} exceeded its deadline; stop remains pending") # ------------------------------------------------------------------ recovery def _inspect_abandoned_attempts(self, binding: HostedRoomBinding) -> None: - for task in state.list_tasks(self.db_path, room_id=binding.room_id, status="running"): + for task in self._tasks(binding, "running"): if task["run_process_generation"] == self.process_generation: continue inspection = ( @@ -842,8 +747,7 @@ class HostedRoomRuntime: if inspection.terminal is not None: self._harvest_previous_attempt(binding, task, inspection.terminal) elif inspection.active: - # The prior session still owns the turn. Do not contend for its - # lease or submit a duplicate prompt. + # The prior session still owns the turn: no lease contention, no duplicate prompt. raise state.LeaseHeldError("recovered session turn is still active") def _inspect_session( @@ -851,9 +755,9 @@ class HostedRoomRuntime: *, read_history: bool) -> _RecoveryInspection: """Probe one resolved session: optional terminal receipt from history, then live info.""" profile = task["payload"]["target_profile"] - receipt = None - if read_history: - receipt = self._terminal_from_history(transport, profile, session_id, task) + receipt = ( + self._terminal_from_history(transport, profile, session_id, task) + if read_history else None) info = transport.info(**_session_kw(profile, session_id)) self._report_pending_action(task, session_id=session_id, info=info) return _RecoveryInspection( @@ -862,8 +766,7 @@ class HostedRoomRuntime: def _inspect_recovery_session( self, binding: HostedRoomBinding, task: Mapping[str, Any]) -> _RecoveryInspection: - profile = task["payload"]["target_profile"] - transport = self._transport_for(binding, task) + profile, transport = task["payload"]["target_profile"], self._transport_for(binding, task) with self.turn_lock(profile): session_id = self._resume_exact(transport, task["identity"].room_id, profile) if session_id is None: @@ -871,13 +774,11 @@ class HostedRoomRuntime: return self._inspect_session(transport, task, session_id, read_history=True) def _inspect_local_recovery_session(self, task: Mapping[str, Any]) -> _RecoveryInspection: - """Check only live process state before explicit local recovery. + """Check only live process state (no resume, no history) before explicit local recovery. - A restart loses the in-process terminal callback identity. Session - history is a display projection that cannot prove which durable task - attempt authored a row, so never hydrate or infer completion from it. - An inactive abandoned attempt remains indeterminate until the user - explicitly retries it under a new fenced generation. + A restart loses the in-process terminal callback identity and history cannot prove + which attempt authored a row, so an inactive abandoned attempt stays indeterminate + until the user retries it under a new fenced generation. """ profile = task["payload"]["target_profile"] with self.turn_lock(profile): @@ -890,14 +791,14 @@ class HostedRoomRuntime: def _reconcile_indeterminate( self, binding: HostedRoomBinding, lease: state.DriverLease) -> bool: - unresolved = state.list_tasks(self.db_path, room_id=binding.room_id, status="indeterminate") + unresolved = self._tasks(binding, "indeterminate") if not unresolved: self._set_blocked(binding.room_id, False) return False inspected = self._inspected_indeterminate_attempts for task in unresolved: - generation = int(task["execution_generation"]) - attempt_key = (binding.room_id, task["identity"].task_id, generation) + attempt_key = ( + binding.room_id, task["identity"].task_id, int(task["execution_generation"])) is_local = self._transport_for(binding, task) is self.rpc if is_local and attempt_key not in inspected: inspection = self._inspect_local_recovery_session(task) @@ -909,12 +810,9 @@ class HostedRoomRuntime: if inspection.active: self._set_blocked(binding.room_id, True) return True - deferred_at = float( - task.get("indeterminate_at") - or task.get("updated_at") - or task.get("created_at") + deadline = self.indeterminate_defer_seconds + float( + task.get("indeterminate_at") or task.get("updated_at") or task.get("created_at") or self.clock()) - deadline = deferred_at + self.indeterminate_defer_seconds inspection = _NO_INSPECTION if attempt_key not in inspected or self.clock() >= deadline: try: @@ -926,7 +824,9 @@ class HostedRoomRuntime: inspected.add(attempt_key) if inspection.status == "cancelled": # Remote-probe resolutions are not republished here. - self._resolve_remote_cancel(binding, task, lease, publish=False) + self._fenced( + state.resolve_indeterminate_cancellation, binding, task, lease, publish=False, + cancel_id=f"remote-cancel:{task['execution_generation']}") inspected.discard(attempt_key) continue if inspection.terminal is not None: @@ -948,21 +848,21 @@ class HostedRoomRuntime: self, binding: HostedRoomBinding, task: Mapping[str, Any], receipt: _TerminalReceipt ) -> None: previous_attempt = state.TaskAttempt( - identity=task["identity"], + identity=task["identity"], execution_generation=task["execution_generation"], + cancel_generation=task["cancel_generation"], lease=state.DriverLease( room_id=binding.room_id, gateway_id=task["run_gateway_id"], authority_epoch=binding.authority_epoch, process_generation=task["run_process_generation"], - lease_generation=task["run_lease_generation"], expires_at=0.0), - execution_generation=task["execution_generation"], - cancel_generation=task["cancel_generation"]) - # Once the previous proof has expired there is deliberately no "trust - # this historical output" escape hatch; fenced recovery leaves the task - # indeterminate for explicit user action. + lease_generation=task["run_lease_generation"], expires_at=0.0)) + # Once the previous proof has expired there is deliberately no "trust this historical + # output" escape hatch; fenced recovery leaves the task indeterminate for the user. with suppress(state.StaleLeaseError, state.StaleTaskError): state.settle_task(self.db_path, previous_attempt, **asdict(receipt), clock=self.clock) - # ------------------------------------------------------------------ misc + def _tasks(self, binding: HostedRoomBinding, status: str) -> list[dict[str, Any]]: + return state.list_tasks(self.db_path, room_id=binding.room_id, status=status) + def _binding_for_room(self, room_id: str) -> HostedRoomBinding | None: return next((b for b in self._rooms_provider() if b.room_id == room_id), None) @@ -991,7 +891,10 @@ class HostedRoomRuntime: self.db_path, attempt, settlement_id=f"failure:{attempt.identity.task_id}:{attempt.execution_generation}", status="failed", result={"error": str(exc)}, clock=self.clock) - self._record_error(f"task {attempt.identity.task_id} failed: {exc}") + self._record_task_error(attempt, f"failed: {exc}") + + def _record_task_error(self, attempt: state.TaskAttempt, message: str) -> None: + self._record_error(f"task {attempt.identity.task_id} {message}") def _record_error(self, message: str) -> None: with self._status_lock: @@ -1004,8 +907,8 @@ def room_session_title(room_id: str) -> str: def _member_id(task: Mapping[str, Any]) -> str: - payload = task.get("payload") or {} - return str(payload.get("target_member_id") or payload.get("target_profile") or "") + p = task.get("payload") or {} + return str(p.get("target_member_id") or p.get("target_profile") or "") def _session_id(session: Mapping[str, Any]) -> str: @@ -1016,12 +919,10 @@ def _session_id(session: Mapping[str, Any]) -> str: def _truncate_utf8(value: Any, *, max_bytes: int) -> tuple[str, bool]: - text = str(value or "") - encoded = text.encode("utf-8") + text, encoded = str(value or ""), str(value or "").encode("utf-8") if len(encoded) <= max_bytes: return text, False - suffix = _TERMINAL_TRUNCATION_NOTICE.encode("utf-8") - prefix = encoded[: max(0, max_bytes - len(suffix))] + prefix = encoded[: max(0, max_bytes - len(_TERMINAL_TRUNCATION_NOTICE.encode("utf-8")))] while prefix: try: return prefix.decode("utf-8") + _TERMINAL_TRUNCATION_NOTICE, True @@ -1034,8 +935,7 @@ def _bounded_terminal_result(receipt: Mapping[str, Any]) -> dict[str, Any]: text, truncated = _truncate_utf8(receipt.get("text", ""), max_bytes=MAX_TERMINAL_TEXT_BYTES) error, error_truncated = _truncate_utf8(receipt.get("error", ""), max_bytes=4096) return { - "message_id": receipt.get("message_id"), - "text": text, + "message_id": receipt.get("message_id"), "text": text, **({"error": error} if error else {}), **({"truncated": True} if truncated or error_truncated else {})} @@ -1048,15 +948,13 @@ def _find_terminal_receipt( if ( message.get("task_id") != identity.task_id or message.get("execution_generation") != execution_generation - or message.get("role") != "assistant" - or status not in {"settled", "failed"}): + or message.get("role") != "assistant" or status not in {"settled", "failed"}): continue receipt_id = message.get("message_id") if not isinstance(receipt_id, str) or not receipt_id: receipt_id = f"reply:{identity.task_id}:{execution_generation}" return _TerminalReceipt( - status=cast(state.TerminalStatus, status), - settlement_id=receipt_id, + status=cast(state.TerminalStatus, status), settlement_id=receipt_id, result=_bounded_terminal_result( {"message_id": receipt_id, "text": message.get("content", "")})) return None @@ -1068,9 +966,5 @@ def _info_active(info: Mapping[str, Any]) -> bool: def _info_is_active_for( info: Mapping[str, Any], identity: state.TaskIdentity, *, require_exact: bool = False) -> bool: - if not _info_active(info): - return False - active_task_id = info.get("task_id") - if require_exact: - return active_task_id == identity.task_id - return active_task_id in {None, identity.task_id} + accepted = (identity.task_id,) if require_exact else (None, identity.task_id) + return _info_active(info) and info.get("task_id") in accepted diff --git a/tui_gateway/hosted_room_peer_http.py b/tui_gateway/hosted_room_peer_http.py index e3af7e3377..416b726226 100644 --- a/tui_gateway/hosted_room_peer_http.py +++ b/tui_gateway/hosted_room_peer_http.py @@ -25,13 +25,11 @@ logger = logging.getLogger(__name__) _NOT_ADMITTED_ERRNOS = frozenset( - value - for name in ("ECONNREFUSED", "ENETDOWN", "ENETUNREACH", "EHOSTDOWN", "EHOSTUNREACH") + value for name in ("ECONNREFUSED", "ENETDOWN", "ENETUNREACH", "EHOSTDOWN", "EHOSTUNREACH") if (value := getattr(errno, name, None)) is not None) _ERROR_CODE_RE = re.compile(r"^[a-z][a-z0-9_]{0,63}$") -# A replay page may legitimately contain many bounded 64 KiB room events. Keep -# enough room for the largest normal page while preventing peer-sized responses -# from scaling memory use without limit. +# A replay page may legitimately hold many bounded 64 KiB room events; cap it so a peer-sized +# response cannot scale memory use without limit. MAX_PEER_RESPONSE_BYTES = 16 * 1024 * 1024 MAX_PEER_ERROR_RESPONSE_BYTES = 16 * 1024 _PEER_RESPONSE_CHUNK_BYTES = 64 * 1024 @@ -43,8 +41,8 @@ _RECEIPT_SCOPE_FIELDS = ( _TERMINAL_RUN_STATES = frozenset({"completed", "failed", "interrupted", "cancelled"}) _ACTIVE_RUN_STATES = frozenset({"queued", "running", "waiting_for_approval", "stopping"}) _KNOWN_RUN_STATES = _TERMINAL_RUN_STATES | _ACTIVE_RUN_STATES -# Older target gateways wrap these conditions inside the generic dispatch -# error; normalize locally until their wire code becomes specific. +_RUN_STATUS_KEYS = ("run_id", "status", "output", "error", "approval", "last_event") +# Older target gateways wrap these inside the generic dispatch error; normalize locally. _LEGACY_DISPATCH_MESSAGE_CODES = ( ("room grant", "invalid_room_grant"), ("capability catalog changed", "room_capability_catalog_changed"), @@ -75,18 +73,9 @@ class _PeerResponseDeadlineExceeded(TimeoutError): """A peer response exceeded the request's monotonic wall-clock budget.""" -def _content_length(response: Any) -> int | None: - try: - value = int(response.headers.get("Content-Length")) - except (AttributeError, TypeError, ValueError): - return None - return value if value >= 0 else None - - def _set_response_socket_timeout(response: Any, remaining: float) -> None: """Best-effort urllib socket timeout tightened to the remaining budget.""" - frontier = [response] - seen: set[int] = set() + frontier, seen = [response], set() for _depth in range(5): next_frontier = [] for value in frontier: @@ -103,12 +92,14 @@ def _set_response_socket_timeout(response: Any, remaining: float) -> None: def _read_bounded_response(response: Any, *, max_bytes: int, deadline: float) -> bytes: - declared = _content_length(response) - if declared is not None and declared > max_bytes: + try: + declared = int(response.headers.get("Content-Length")) + except (AttributeError, TypeError, ValueError): + declared = -1 + if declared > max_bytes: raise _PeerResponseTooLarge reader = getattr(response, "read1", None) - if not callable(reader): - reader = response.read + reader = reader if callable(reader) else response.read body = bytearray() while len(body) <= max_bytes: remaining = deadline - time.monotonic() @@ -131,14 +122,25 @@ def _read_bounded_response(response: Any, *, max_bytes: int, deadline: float) -> raise _PeerResponseTooLarge +def _read_body(response: Any, *, max_bytes: int, deadline: float, kind: str, **flags: Any) -> str: + """Read a bounded body as text; budget overruns become classified ``PeerRunsHTTPError``.""" + try: + return _read_bounded_response( + response, max_bytes=max_bytes, deadline=deadline).decode("utf-8", "replace") + except _PeerResponseTooLarge as exc: + raise PeerRunsHTTPError(_BUDGET_MESSAGES["size"].format(kind=kind), **flags) from exc + except _PeerResponseDeadlineExceeded as exc: + raise PeerRunsHTTPError( + _BUDGET_MESSAGES["time"].format(kind=kind), retryable=True, **flags) from exc + + def _is_proven_pre_admission_failure(exc: BaseException) -> bool: """Return whether no HTTP connection could have carried the request.""" reason: Any = exc while isinstance(reason, urllib.error.URLError): reason = reason.reason - if isinstance(reason, socket.gaierror): - return True - return isinstance(reason, OSError) and reason.errno in _NOT_ADMITTED_ERRNOS + return isinstance(reason, socket.gaierror) or ( + isinstance(reason, OSError) and reason.errno in _NOT_ADMITTED_ERRNOS) def _valid_code(code: Any) -> str | None: @@ -165,33 +167,18 @@ def _response_error_code(detail: str) -> str | None: return _valid_code(payload.get("code")) -def _http_error_message(method: str, path: str, status: int, error_code: str | None) -> str: - renewal = status in {401, 403} and error_code in _GRANT_RENEWAL_CODES - drift = status == 403 and error_code in {_EXECUTION_POLICY_CHANGED[0], _CAPABILITY_CHANGED[0]} - if renewal or drift: - return _REAUTHORIZATION_MESSAGES[error_code] - return f"peer rejected {method} {path} with HTTP {status}" - - class PeerRunsHTTPError(RuntimeError): """Controlled peer HTTP failure.""" def __init__( self, message: str, *, retryable: bool = False, ambiguous: bool = False, not_admitted: bool = False, status_code: int | None = None, error_code: str | None = None, - error_message: str | None = None) -> None: + ) -> None: super().__init__(message) - self.retryable = retryable - self.ambiguous = ambiguous - self.not_admitted = not_admitted - self.status_code = status_code - self.error_code = error_code - self.error_message = error_message + self.retryable, self.ambiguous, self.not_admitted = retryable, ambiguous, not_admitted + self.status_code, self.error_code = status_code, error_code self.needs_reauthorization = ( status_code in {401, 403} and error_code in _REAUTHORIZATION_CODES) - self.needs_capability_refresh = status_code == 403 and error_code == _CAPABILITY_CHANGED[0] - self.needs_execution_policy_refresh = ( - status_code == 403 and error_code == _EXECUTION_POLICY_CHANGED[0]) def digest_reauthorization_error( @@ -225,15 +212,13 @@ class PeerRunsHTTPClient: base_url, self.transport_security = validate_room_link_url(base_url) if api_key and len(api_key) < 16: raise ValueError("peer API key is missing or too short") - self.base_url = base_url - self.api_key = api_key + self.base_url, self.api_key, self.clock = base_url, api_key, clock self.timeout_seconds = float(timeout_seconds) self.receipt_db_path = Path(receipt_db_path) if receipt_db_path else None if poll_min_seconds <= 0 or poll_max_seconds < poll_min_seconds: raise ValueError("peer polling bounds are invalid") self.poll_min_seconds = float(poll_min_seconds) self.poll_max_seconds = float(poll_max_seconds) - self.clock = clock self._runs: dict[tuple[str, int], dict[str, Any]] = {} self._observation_key: tuple[str, int] | None = None self._status_cache: dict[str, dict[str, Any]] = {} @@ -253,11 +238,9 @@ class PeerRunsHTTPClient: authority_epoch: int, member_id: str, target_install_id: str, target_profile: str) -> None: """Fence every in-memory and durable receipt to one room authority.""" epoch = int(authority_epoch or 0) - names = [ - str(value or "") - for value in ( - room_id, home_install_id, authority_gateway_id, member_id, target_install_id, - target_profile)] + names = [str(v or "") for v in ( + room_id, home_install_id, authority_gateway_id, member_id, target_install_id, + target_profile)] if not all(names): raise PeerRunsHTTPError("peer room receipt scope is incomplete") if epoch < 1: @@ -265,15 +248,10 @@ class PeerRunsHTTPClient: scope = dict(zip(_RECEIPT_SCOPE_FIELDS, names[:3] + [epoch] + names[3:])) if self._room_scope == scope: return - self._room_scope = scope - self._runs.clear() - self._observation_key = None - self._status_cache.clear() - self._recovery_backoff.clear() - self._terminal_receipts.clear() - - def _bind_dispatch_scope(self, dispatch: HostedMemberDispatch) -> None: - self.bind_room_scope(**{field: getattr(dispatch, field) for field in _RECEIPT_SCOPE_FIELDS}) + self._room_scope, self._observation_key = scope, None + for table in ( + self._runs, self._status_cache, self._recovery_backoff, self._terminal_receipts): + table.clear() def _receipt(self, task_id: str, execution_generation: int) -> dict[str, Any] | None: """Return the in-memory receipt, falling back to the durable store.""" @@ -281,7 +259,6 @@ class PeerRunsHTTPClient: if record is not None or self.receipt_db_path is None or self._room_scope is None: return record from gateway import hosted_rooms - identity = {"task_id": task_id, "execution_generation": execution_generation} return hosted_rooms.remote_run_receipt( self.receipt_db_path, record={**self._room_scope, **identity}) @@ -291,47 +268,33 @@ class PeerRunsHTTPClient: key = (str(task_id or ""), int(execution_generation or 0)) if not key[0] or key[1] < 1: raise PeerRunsHTTPError("peer observation identity is invalid") - if self._observation_key != key: - for terminal_key in self._terminal_receipts - {key}: - self._runs.pop(terminal_key, None) - self._terminal_receipts.intersection_update({key}) - self._observation_key = key - self._status_cache.clear() - self._recovery_backoff.clear() + if self._observation_key == key: + return + for terminal_key in self._terminal_receipts - {key}: + self._runs.pop(terminal_key, None) + self._terminal_receipts.intersection_update({key}) + self._observation_key = key + self._status_cache.clear() + self._recovery_backoff.clear() def _request( self, path: str, *, method: str = "GET", body: Mapping[str, Any] | None = None, headers: Mapping[str, str] | None = None, room_grant: str | None = None) -> dict[str, Any]: from hermes_cli.urllib_security import open_credentialed_url - - deadline = time.monotonic() + self.timeout_seconds - ambiguous = method == "POST" + deadline, ambiguous = time.monotonic() + self.timeout_seconds, method == "POST" request = urllib.request.Request( - f"{self.base_url}{path}", - data=( - json.dumps(body, separators=(",", ":")).encode("utf-8") - if body is not None - else None), - method=method, + f"{self.base_url}{path}", method=method, + data=None if body is None else json.dumps(body, separators=(",", ":")).encode("utf-8"), headers={ "Authorization": ( f"HermesRoom {room_grant}" if room_grant else f"Bearer {self.api_key}"), - "Content-Type": "application/json", - "User-Agent": "Hermes-RoomLink/1.0", + "Content-Type": "application/json", "User-Agent": "Hermes-RoomLink/1.0", **(headers or {})}) try: with open_credentialed_url(request, timeout=self.timeout_seconds) as response: - raw = _read_bounded_response( - response, max_bytes=MAX_PEER_RESPONSE_BYTES, deadline=deadline - ).decode("utf-8", "replace") - except _PeerResponseTooLarge as exc: - raise PeerRunsHTTPError( - _BUDGET_MESSAGES["size"].format(kind=""), ambiguous=ambiguous, - ) from exc - except _PeerResponseDeadlineExceeded as exc: - raise PeerRunsHTTPError( - _BUDGET_MESSAGES["time"].format(kind=""), retryable=True, ambiguous=ambiguous - ) from exc + raw = _read_body( + response, max_bytes=MAX_PEER_RESPONSE_BYTES, deadline=deadline, kind="", + ambiguous=ambiguous) except urllib.error.HTTPError as exc: self._raise_http_error(exc, method=method, path=path, deadline=deadline) except (urllib.error.URLError, TimeoutError, OSError) as exc: @@ -354,28 +317,25 @@ class PeerRunsHTTPClient: """Raise the classified PeerRunsHTTPError for an HTTP error response.""" # A 4xx on admission proves the peer never admitted the run. flags = { - "ambiguous": method == "POST" and exc.code >= 500, - "not_admitted": method == "POST" and path == "/v1/runs" and 400 <= exc.code < 500, - "status_code": exc.code} + "ambiguous": method == "POST" and exc.code >= 500, "status_code": exc.code, + "not_admitted": method == "POST" and path == "/v1/runs" and 400 <= exc.code < 500} try: - detail = _read_bounded_response( - exc, max_bytes=MAX_PEER_ERROR_RESPONSE_BYTES, deadline=deadline - ).decode("utf-8", "replace")[:500] - except _PeerResponseTooLarge as body_exc: - raise PeerRunsHTTPError( - _BUDGET_MESSAGES["size"].format(kind=" error"), **flags, - ) from body_exc - except _PeerResponseDeadlineExceeded as body_exc: - raise PeerRunsHTTPError( - _BUDGET_MESSAGES["time"].format(kind=" error"), retryable=True, **flags - ) from body_exc + detail = _read_body( + exc, max_bytes=MAX_PEER_ERROR_RESPONSE_BYTES, deadline=deadline, kind=" error", + **flags)[:500] + except PeerRunsHTTPError: + raise except Exception: detail = "" error_code = _response_error_code(detail) logger.debug( "Peer RoomLink request returned HTTP %s (%s)", exc.code, error_code or "no-code") + renewal = exc.code in {401, 403} and error_code in _GRANT_RENEWAL_CODES + drift = exc.code == 403 and error_code in { + _EXECUTION_POLICY_CHANGED[0], _CAPABILITY_CHANGED[0]} raise PeerRunsHTTPError( - _http_error_message(method, path, exc.code, error_code), + _REAUTHORIZATION_MESSAGES[error_code] if renewal or drift + else f"peer rejected {method} {path} with HTTP {exc.code}", retryable=exc.code in {408, 425, 429} or exc.code >= 500, error_code=error_code, **flags, ) from exc @@ -386,8 +346,8 @@ class PeerRunsHTTPClient: if source != "bot_room": raise PeerRunsHTTPError("peer room source must be bot_room") self._require_room_grant(grant) - logical_session = ( - "roomlink_" + hashlib.sha256(f"{room_id}\0{profile}".encode("utf-8")).hexdigest()[:32]) + logical_session = "roomlink_" + hashlib.sha256( + f"{room_id}\0{profile}".encode("utf-8")).hexdigest()[:32] if expected_session_id and expected_session_id != logical_session: raise PeerRunsHTTPError("peer room session identity changed") return {"session_id": logical_session, "title": f"Group: {room_id}", "source": source} @@ -396,7 +356,7 @@ class PeerRunsHTTPClient: """Validate a dispatch, its grant, and pin scope + observation to it.""" checked = HostedMemberDispatch.from_mapping(dispatch) self._require_room_grant(grant) - self._bind_dispatch_scope(checked) + self.bind_room_scope(**{f: getattr(checked, f) for f in _RECEIPT_SCOPE_FIELDS}) self.bind_observation( task_id=checked.task_id, execution_generation=checked.execution_generation) return checked @@ -412,8 +372,8 @@ class PeerRunsHTTPClient: if any(existing[field] != getattr(checked, field) for field in _RECEIPT_SCOPE_FIELDS): raise PeerRunsHTTPError("peer run receipt conflicts with the recovered dispatch") return self._accepted( - checked, run_id=str(existing["run_id"]), - session_id=str(existing["session_id"]), replayed=True) + checked, run_id=str(existing["run_id"]), session_id=str(existing["session_id"]), + replayed=True) key, now = (checked.task_id, checked.execution_generation), self.clock() backoff = self._recovery_backoff.get(key) if backoff is not None and now < float(backoff["next_attempt_at"]): @@ -434,12 +394,9 @@ class PeerRunsHTTPClient: checked: HostedMemberDispatch, *, run_id: str, session_id: str, replayed: bool, ) -> dict[str, Any]: return { - "status": "accepted", - "task_id": checked.task_id, - "execution_generation": checked.execution_generation, - "run_id": run_id, - "session_id": session_id, - "replayed": replayed} + "status": "accepted", "task_id": checked.task_id, + "execution_generation": checked.execution_generation, "run_id": run_id, + "session_id": session_id, "replayed": replayed} def _admit_dispatch(self, checked: HostedMemberDispatch, *, grant: str) -> Mapping[str, Any]: session_id = self._session_id(checked, grant=grant) @@ -462,14 +419,11 @@ class PeerRunsHTTPClient: if not run_id: raise PeerRunsHTTPError("peer did not return a run id") receipt = { - "run_id": run_id, - "session_id": session_id, + "run_id": run_id, "session_id": session_id, **{field: getattr(checked, field) for field in _RECEIPT_SCOPE_FIELDS}, - "task_id": checked.task_id, - "execution_generation": checked.execution_generation} + "task_id": checked.task_id, "execution_generation": checked.execution_generation} if self.receipt_db_path is not None: from gateway import hosted_rooms - hosted_rooms.upsert_remote_run_receipt(self.receipt_db_path, record=receipt) self._runs[(checked.task_id, checked.execution_generation)] = receipt self._status_cache.pop(run_id, None) @@ -490,9 +444,7 @@ class PeerRunsHTTPClient: def _observation_receipt( self, *, room_id: str, profile: str, session_id: str) -> dict[str, Any] | None: - if self._observation_key is None: - return None - record = self._receipt(*self._observation_key) + record = None if self._observation_key is None else self._receipt(*self._observation_key) if record is None: return None scope = (record["room_id"], record["target_profile"], record["session_id"]) @@ -500,27 +452,16 @@ class PeerRunsHTTPClient: raise PeerRunsHTTPError("peer observation receipt changed scope") return record - @staticmethod - def _compact_run_status(status: Mapping[str, Any]) -> dict[str, Any]: - return { - key: status[key] - for key in ("run_id", "status", "output", "error", "approval", "last_event") - if key in status} - def _next_poll_delay(self, cached: Mapping[str, Any] | None) -> float: previous = float(cached["delay"]) if cached is not None else self.poll_min_seconds / 2 return min(self.poll_max_seconds, max(self.poll_min_seconds, previous * 2)) - @staticmethod - def _run_is_terminal(status: Mapping[str, Any]) -> bool: - return status.get("status") in _TERMINAL_RUN_STATES - def _poll_receipt(self, record: Mapping[str, Any], *, grant: str) -> dict[str, Any]: run_id, now = str(record["run_id"]), self.clock() cached = self._status_cache.get(run_id) if cached is not None: status = cached["status"] - if self._run_is_terminal(status): + if status.get("status") in _TERMINAL_RUN_STATES: return status if now < float(cached["next_poll_at"]): error = cached.get("error") @@ -530,8 +471,8 @@ class PeerRunsHTTPClient: delay = self._next_poll_delay(cached) entry = {"delay": delay, "next_poll_at": now + delay} try: - status = self._compact_run_status( - self._request(_run_path(record), room_grant=self._require_room_grant(grant))) + full = self._request(_run_path(record), room_grant=self._require_room_grant(grant)) + status = {key: full[key] for key in _RUN_STATUS_KEYS if key in full} if ( str(status.get("run_id") or "") != run_id or status.get("status") not in _KNOWN_RUN_STATES): @@ -541,7 +482,7 @@ class PeerRunsHTTPClient: self._status_cache = {run_id: {"status": previous, "error": exc, **entry}} raise self._status_cache = {run_id: {"status": status, **entry}} - if self._run_is_terminal(status): + if status.get("status") in _TERMINAL_RUN_STATES: self._terminal_receipts.add( (str(record["task_id"]), int(record["execution_generation"]))) return status @@ -557,8 +498,7 @@ class PeerRunsHTTPClient: if state not in {"completed", "failed", "interrupted"}: return [] return [{ - "role": "assistant", - "task_id": receipt["task_id"], + "role": "assistant", "task_id": receipt["task_id"], "execution_generation": receipt["execution_generation"], "status": "settled" if state == "completed" else "failed", "message_id": f"peer-run:{status.get('run_id')}", @@ -571,11 +511,9 @@ class PeerRunsHTTPClient: return {"active": False, "task_id": None} status = self._poll_receipt(receipt, grant=grant) return { - "active": status.get("status") in _ACTIVE_RUN_STATES, - "task_id": receipt["task_id"], + "active": status.get("status") in _ACTIVE_RUN_STATES, "task_id": receipt["task_id"], "execution_generation": receipt["execution_generation"], - "status": status.get("status"), - "run_id": status.get("run_id"), + "status": status.get("status"), "run_id": status.get("run_id"), "approval": status.get("approval")} def approve_receipt( @@ -602,7 +540,7 @@ class PeerRunsHTTPClient: def stop(self, *, dispatch: Mapping[str, Any], grant: str) -> Mapping[str, Any] | None: checked = HostedMemberDispatch.from_mapping(dispatch) - self._bind_dispatch_scope(checked) + self.bind_room_scope(**{f: getattr(checked, f) for f in _RECEIPT_SCOPE_FIELDS}) return self.stop_receipt( task_id=checked.task_id, execution_generation=checked.execution_generation, grant=grant) @@ -614,7 +552,7 @@ class PeerRunsHTTPClient: return None result = self._post_run_action( record, "stop", body={}, grant=self._require_room_grant(grant)) - if self._run_is_terminal(result): + if result.get("status") in _TERMINAL_RUN_STATES: self._terminal_receipts.add((str(task_id), int(execution_generation))) return result @@ -628,17 +566,11 @@ class PeerRunsHTTPClient: return self._request( "/v1/room-members/invitations", method="POST", body={ - "room_id": room_id, - "home_install_id": home_install_id, - "authority_gateway_id": authority_gateway_id, - "authority_epoch": authority_epoch, - "member_id": member_id, - "grant_id": grant_id, - "ttl_seconds": ttl_seconds, - **( - {"status_ttl_seconds": status_ttl_seconds} - if status_ttl_seconds is not None - else {})}) + "room_id": room_id, "home_install_id": home_install_id, + "authority_gateway_id": authority_gateway_id, "authority_epoch": authority_epoch, + "member_id": member_id, "grant_id": grant_id, "ttl_seconds": ttl_seconds, + **({} if status_ttl_seconds is None else { + "status_ttl_seconds": status_ttl_seconds})}) def refresh_grant( self, *, grant: str, ttl_seconds: float = 24 * 60 * 60, @@ -650,8 +582,7 @@ class PeerRunsHTTPClient: replacement = str(refreshed.get("grant") or "") if not replacement: raise PeerRunsHTTPError("peer returned no refreshed room grant") - # Persist only after the target proves the replacement can authorize - # the same scoped capability endpoint. + # Persist only after the target proves the replacement authorizes the scoped endpoint. probe = self.probe(grant=replacement) error = digest_reauthorization_error( GatewayRoomCatalog.from_mapping(probe.get("catalog")), diff --git a/tui_gateway/hosted_room_peer_transport.py b/tui_gateway/hosted_room_peer_transport.py index 170c98a28b..0bda289022 100644 --- a/tui_gateway/hosted_room_peer_transport.py +++ b/tui_gateway/hosted_room_peer_transport.py @@ -1,9 +1,6 @@ -"""Peer-backed session transport for one hosted-room member task. - -This adapter implements :class:`InternalSessionRPC` without using canonical -Bot Chat. The remote client must resolve a hidden ``Group: `` session -with ``source=bot_room`` and verify the scoped grant at admission. -""" +"""Peer-backed session transport for one hosted-room member task: implements +:class:`InternalSessionRPC` without canonical Bot Chat. The remote client must resolve a hidden +``Group: `` session with ``source=bot_room`` and verify the scoped grant at admission.""" from __future__ import annotations @@ -24,27 +21,16 @@ class HostedRoomPeerClient(Protocol): """Authenticated client for a target gateway's narrow room-member API.""" def bind_room_scope(self, **scope: Any) -> None: ... - - def prepare( - self, *, room_id: str, profile: str, source: str, grant: str, create: bool, - expected_session_id: str | None = None, - ) -> Mapping[str, Any] | None: ... - + def prepare(self, *, room_id: str, profile: str, source: str, grant: str, create: bool, + expected_session_id: str | None = None) -> Mapping[str, Any] | None: ... def dispatch(self, *, dispatch: Mapping[str, Any], grant: str) -> Mapping[str, Any]: ... - - def history( - self, *, room_id: str, profile: str, session_id: str, grant: str - ) -> Sequence[Mapping[str, Any]]: ... - - def status( - self, *, room_id: str, profile: str, session_id: str, grant: str - ) -> Mapping[str, Any]: ... - + def history(self, *, room_id: str, profile: str, session_id: str, grant: str + ) -> Sequence[Mapping[str, Any]]: ... + def status(self, *, room_id: str, profile: str, session_id: str, grant: str + ) -> Mapping[str, Any]: ... def stop(self, *, dispatch: Mapping[str, Any], grant: str) -> Mapping[str, Any] | None: ... - - def stop_receipt( - self, *, task_id: str, execution_generation: int, grant: str - ) -> Mapping[str, Any] | None: ... + def stop_receipt(self, *, task_id: str, execution_generation: int, grant: str + ) -> Mapping[str, Any] | None: ... @dataclass(frozen=True) @@ -69,10 +55,8 @@ class FailoverHostedRoomPeerClient: raise ValueError("RoomLink candidates must target one installation") if reprobe_interval_seconds <= 0: raise ValueError("reprobe_interval_seconds must be positive") - self.candidates = tuple(candidates) - self._active = 0 + self.candidates, self._active, self.clock = tuple(candidates), 0, clock self.reprobe_interval_seconds = float(reprobe_interval_seconds) - self.clock = clock self._last_primary_probe = 0.0 @property @@ -80,11 +64,9 @@ class FailoverHostedRoomPeerClient: return self.candidates[self._active] def _call(self, method: str, **kwargs): - """Try the active link (re-probing the primary after a cooldown), then the rest. - - Ambiguous or non-retryable failures propagate immediately: failing over - after an ambiguous dispatch could run the same task twice. - """ + """Try the active link (re-probing the primary after a cooldown), then the rest. Ambiguous + or non-retryable failures propagate: failing over after an ambiguous dispatch could run + the same task twice.""" now = self.clock() order = [self._active] if self._active != 0 and now - self._last_primary_probe >= self.reprobe_interval_seconds: @@ -102,29 +84,20 @@ class FailoverHostedRoomPeerClient: continue self._active = index return result - if last_error is not None: - raise last_error - raise RuntimeError("no RoomLink candidate was attempted") + raise last_error if last_error is not None else RuntimeError("no RoomLink candidate was attempted") - def prepare(self, **kwargs): - return self._call("prepare", **kwargs) + def _delegate(method: str): + def call(self, **kwargs): + return self._call(method, **kwargs) + call.__name__ = method + return call - def dispatch(self, **kwargs): - return self._call("dispatch", **kwargs) - - def history(self, **kwargs): - return self._call("history", **kwargs) - - def status(self, **kwargs): - return self._call("status", **kwargs) - - def stop(self, **kwargs): - return self._call("stop", **kwargs) + prepare, dispatch, history, status, stop = map(_delegate, ("prepare", "dispatch", "history", "status", "stop")) + del _delegate def bind_room_scope(self, **kwargs): for candidate in self.candidates: - bind = getattr(candidate.client, "bind_room_scope", None) - if callable(bind): + if callable(bind := getattr(candidate.client, "bind_room_scope", None)): bind(**kwargs) @@ -149,23 +122,15 @@ def build_member_dispatch( trace_id: str) -> HostedMemberDispatch: """Build the fully fenced member dispatch shared by submit and recovery.""" return HostedMemberDispatch.from_mapping({ - "protocol_version": PROTOCOL_VERSION, - "room_id": room_id, - "home_install_id": route.home_install_id, - "authority_gateway_id": binding.gateway_id, - "authority_epoch": binding.authority_epoch, - "member_id": route.member_id, - "target_install_id": route.target_install_id, - "target_profile": target_profile, - "task_id": task_id, - "execution_generation": execution_generation, - "source_event_seq": source_event_seq, - "cancellation_scope_id": route.cancellation_scope_id, - "prompt": prompt, - "prompt_digest": hashlib.sha256(prompt.encode("utf-8")).hexdigest(), + "protocol_version": PROTOCOL_VERSION, "room_id": room_id, + "home_install_id": route.home_install_id, "authority_gateway_id": binding.gateway_id, + "authority_epoch": binding.authority_epoch, "member_id": route.member_id, + "target_install_id": route.target_install_id, "target_profile": target_profile, + "task_id": task_id, "execution_generation": execution_generation, + "source_event_seq": source_event_seq, "cancellation_scope_id": route.cancellation_scope_id, + "prompt": prompt, "prompt_digest": hashlib.sha256(prompt.encode("utf-8")).hexdigest(), "capability_digest": route.capability_digest, - "execution_policy_digest": route.execution_policy_digest, - "trace_id": trace_id}) + "execution_policy_digest": route.execution_policy_digest, "trace_id": trace_id}) class PeerHostedRoomTransport(InternalSessionRPC): @@ -185,8 +150,7 @@ class PeerHostedRoomTransport(InternalSessionRPC): self.execution_generation = execution_generation self._session_id: str | None = None self._dispatch: HostedMemberDispatch | None = None - bind_scope = getattr(self.client, "bind_room_scope", None) - if callable(bind_scope): + if callable(bind_scope := getattr(self.client, "bind_room_scope", None)): bind_scope( room_id=binding.room_id, home_install_id=route.home_install_id, authority_gateway_id=binding.gateway_id, authority_epoch=binding.authority_epoch, @@ -205,26 +169,23 @@ class PeerHostedRoomTransport(InternalSessionRPC): """Room id + grant keyword arguments shared by every scoped client call.""" return {"room_id": self.binding.room_id, "grant": self.route.grant, **extra} - def _prepare(self, *, profile: str, source: str, create: bool, **extra): - return self.client.prepare( - **self._scoped(profile=profile, source=source, create=create, **extra)) + def _prepare(self, *, profile: str, source: str, create: bool, title: str | None = None, **extra): + """Validate coordinates, then the scoped ``prepare`` call.""" + self._validate_coordinates(profile=profile, source=source, title=title) + return self.client.prepare(**self._scoped(profile=profile, source=source, create=create, **extra)) def resolve_exact(self, *, profile: str, title: str, source: str) -> Mapping[str, Any] | None: - self._validate_coordinates(profile=profile, source=source, title=title) - return self._prepare(profile=profile, source=source, create=False) + return self._prepare(profile=profile, source=source, create=False, title=title) def create(self, *, profile: str, title: str, source: str) -> Mapping[str, Any]: - self._validate_coordinates(profile=profile, source=source, title=title) - session = self._prepare(profile=profile, source=source, create=True) + session = self._prepare(profile=profile, source=source, create=True, title=title) if session is None: raise RuntimeError("peer did not create the room session") self._session_id = str(session.get("session_id") or session.get("id") or "") return session def resume(self, *, profile: str, session_id: str, source: str) -> Mapping[str, Any]: - self._validate_coordinates(profile=profile, source=source) - session = self._prepare( - profile=profile, source=source, create=False, expected_session_id=session_id) + session = self._prepare(profile=profile, source=source, create=False, expected_session_id=session_id) if session is None: raise RuntimeError("peer room session is unavailable") self._session_id = session_id @@ -262,15 +223,13 @@ class PeerHostedRoomTransport(InternalSessionRPC): ) -> Mapping[str, Any] | None: self._validate_coordinates(profile=profile, source=source) dispatch = self._dispatch - if dispatch is None: - if ( - self.task_id != expected_task_id - or not self.execution_generation - or not hasattr(self.client, "stop_receipt")): + if dispatch is not None: + if dispatch.task_id != expected_task_id: return None - return self.client.stop_receipt( - task_id=expected_task_id, execution_generation=self.execution_generation, - grant=self.route.grant) - if dispatch.task_id != expected_task_id: + return self.client.stop(dispatch=dispatch.as_mapping(), grant=self.route.grant) + if (self.task_id != expected_task_id or not self.execution_generation + or not hasattr(self.client, "stop_receipt")): return None - return self.client.stop(dispatch=dispatch.as_mapping(), grant=self.route.grant) + return self.client.stop_receipt( + task_id=expected_task_id, execution_generation=self.execution_generation, + grant=self.route.grant) diff --git a/tui_gateway/hosted_room_server_rpc.py b/tui_gateway/hosted_room_server_rpc.py index edbe5b8eaa..e470c91a09 100644 --- a/tui_gateway/hosted_room_server_rpc.py +++ b/tui_gateway/hosted_room_server_rpc.py @@ -1,10 +1,6 @@ -"""In-process session adapter for the hosted room driver. - -The room worker must not depend on a Desktop/WebSocket transport, but it should -still use the same session handlers as every other TUI/Desktop turn. This -adapter calls the installed handler registry directly and keeps the extra -task proof as an in-process-only Python object that JSON clients cannot forge. -""" +"""In-process session adapter for the hosted room driver: the room worker uses the same +installed session handlers as every TUI/Desktop turn (no WebSocket transport), passing the +task proof as an in-process-only Python object that JSON clients cannot forge.""" from __future__ import annotations @@ -57,50 +53,32 @@ class HostedRoomServerRPC: if not isinstance(rows, list) or not rows or not isinstance(rows[0], dict): return None row = rows[0] - return { - "session_id": row.get("resolved_id") or row.get("id"), - "title": row.get("title") or title} + return {"session_id": row.get("resolved_id") or row.get("id"), + "title": row.get("title") or title} def create(self, *, profile: str, title: str, source: str) -> Mapping[str, Any]: - return self._call( - "session.create", - { - "profile": profile, - "title": title, - "source": source, - "hidden": True, - "room_plumbing": True, - "follow_profile_config": True, - "close_on_disconnect": False}) + return self._call("session.create", { + "profile": profile, "title": title, "source": source, "hidden": True, + "room_plumbing": True, "follow_profile_config": True, "close_on_disconnect": False}) def resume(self, *, profile: str, session_id: str, source: str) -> Mapping[str, Any]: - return self._call( - "session.resume", - {"profile": profile, "session_id": session_id, "omit_messages": True, "source": source}) + return self._call("session.resume", { + "profile": profile, "session_id": session_id, "omit_messages": True, "source": source}) def submit( self, *, profile: str, session_id: str, prompt: str, source: str, task: state.TaskIdentity, execution_generation: int, on_terminal: Callable[[Mapping[str, Any]], None], ) -> Mapping[str, Any]: try: - return self._call( - "prompt.submit", - { - "profile": profile, - "session_id": session_id, - "text": prompt, - "source": source, - "_hosted_task": { - "room_id": task.room_id, - "task_id": task.task_id, - "thread_id": task.thread_id, - "turn_id": task.turn_id, - "execution_generation": execution_generation}, - "_hosted_terminal_callback": on_terminal}) + return self._call("prompt.submit", { + "profile": profile, "session_id": session_id, "text": prompt, "source": source, + "_hosted_task": { + "room_id": task.room_id, "task_id": task.task_id, "thread_id": task.thread_id, + "turn_id": task.turn_id, "execution_generation": execution_generation}, + "_hosted_terminal_callback": on_terminal}) except HostedRoomSessionError as exc: - # In-process prompt.submit error envelopes are returned before the - # background turn is admitted. Preserve that proof so the driver - # can defer or requeue without waiting out an ambiguity lease. + # In-process prompt.submit error envelopes come back before the background turn is + # admitted; keep that proof so the driver can defer/requeue without an ambiguity lease. exc.not_admitted = True raise @@ -115,10 +93,8 @@ class HostedRoomServerRPC: record = self.server._sessions.get(session_id) if record is not None: return record - for candidate in self.server._sessions.values(): - if str(candidate.get("session_key") or "") == session_id: - return candidate - return None + return next((c for c in self.server._sessions.values() + if str(c.get("session_key") or "") == session_id), None) def info(self, *, profile: str, session_id: str, source: str) -> Mapping[str, Any]: del profile, source @@ -130,32 +106,23 @@ class HostedRoomServerRPC: return {"active": bool(record.get("running")), "task_id": None} with lock: task = record.get("_hosted_room_task") - result = { - "active": bool(record.get("running")), - "task_id": task.get("task_id") if isinstance(task, dict) else None} + result = {"active": bool(record.get("running")), + "task_id": task.get("task_id") if isinstance(task, dict) else None} pending_reader = getattr(self.server, "_pending_approval_request_payload", None) - pending = ( - pending_reader(str(record.get("session_key") or "")) - if callable(pending_reader) - else None) - if pending: + if callable(pending_reader) and (pending := pending_reader(str(record.get("session_key") or ""))): result["status"] = "waiting_for_approval" result["pending_approval"] = pending return result def approve(self, *, session_id: str, request_id: str, choice: str) -> Mapping[str, Any]: """Resolve one exact local room approval without broad policy changes.""" - return self._call( - "approval.respond", - {"session_id": session_id, "request_id": request_id, "choice": choice, "all": False}) + return self._call("approval.respond", { + "session_id": session_id, "request_id": request_id, "choice": choice, "all": False}) def interrupt( self, *, profile: str, session_id: str, source: str, expected_task_id: str ) -> Mapping[str, Any] | None: del source - return self._call( - "session.interrupt", - { - "profile": profile, - "session_id": session_id, - "expected_hosted_task_id": expected_task_id}) + return self._call("session.interrupt", { + "profile": profile, "session_id": session_id, + "expected_hosted_task_id": expected_task_id}) diff --git a/tui_gateway/hosted_room_service.py b/tui_gateway/hosted_room_service.py index 7303058c45..0246d75e89 100644 --- a/tui_gateway/hosted_room_service.py +++ b/tui_gateway/hosted_room_service.py @@ -22,11 +22,11 @@ from gateway.hosted_room_peer import ( GatewayRoomCatalog, HostedMemberDispatch, PROTOCOL_VERSION, room_grant_needs_dispatch_refresh) from tui_gateway.hosted_room_driver import HostedRoomBinding, HostedRoomRuntime from tui_gateway.hosted_room_server_rpc import HostedRoomServerRPC -from tui_gateway.hosted_room_peer_http import PeerRunsHTTPClient, PeerRunsHTTPError +from tui_gateway.hosted_room_peer_http import ( + PeerRunsHTTPClient, PeerRunsHTTPError, digest_reauthorization_error) from tui_gateway.hosted_room_peer_transport import ( HostedRoomPeerClient, PeerHostedRoomTransport, PeerMemberRoute, build_member_dispatch) - _HOSTED_ROOM_IDLE_FALLBACK_SECONDS = 5.0 _HOSTED_ROOM_ACTIVE_POLL_SECONDS = 0.25 _HOSTED_ROOM_TERMINAL_GRACE_SECONDS = 30.0 @@ -41,10 +41,8 @@ def _hosted_room_turn_timeout_seconds() -> float: try: agent_timeout = float(os.getenv("HERMES_AGENT_TIMEOUT", "1800")) except (TypeError, ValueError): - agent_timeout = 1800.0 - if agent_timeout <= 0: - agent_timeout = 1800.0 - return agent_timeout + _HOSTED_ROOM_TERMINAL_GRACE_SECONDS + agent_timeout = 0.0 + return (agent_timeout if agent_timeout > 0 else 1800.0) + _HOSTED_ROOM_TERMINAL_GRACE_SECONDS def _grant_revoke_is_terminal(exc: PeerRunsHTTPError) -> bool: @@ -70,8 +68,7 @@ class HostedRoomService: self, server: ModuleType, *, db_path: Path | str | None = None, peer_routes: Mapping[tuple[str, str], PeerMemberRoute] | None = None, peer_clients: Mapping[Any, HostedRoomPeerClient] | None = None) -> None: - self.server = server - self.db_path = Path(db_path or hosted_rooms.default_db_path()) + self.server, self.db_path = server, Path(db_path or hosted_rooms.default_db_path()) hosted_rooms.prune_disbanded_rooms(self.db_path) self._policy_lock = threading.RLock() self._pending_actions: dict[tuple[str, str], dict[str, Any]] = {} @@ -85,13 +82,10 @@ class HostedRoomService: self._load_stored_links() except Exception as exc: self._link_load_error = str(exc) - supplied_routes = dict(peer_routes or {}) supplied_clients = dict(peer_clients or {}) - self.peer_routes.update(supplied_routes) - for key, route in supplied_routes.items(): - client = supplied_clients.get(key) - if client is None: - client = supplied_clients.get(route.target_install_id) + for key, route in dict(peer_routes or {}).items(): + self.peer_routes[key] = route + client = supplied_clients.get(key, supplied_clients.get(route.target_install_id)) if client is not None: self.peer_clients[key] = client self.runtime = HostedRoomRuntime( @@ -108,16 +102,15 @@ class HostedRoomService: stored_links, load_errors = hosted_room_links.load_room_links_tolerant(self.db_path) errors = list(load_errors) for stored in stored_links: - key = (stored.room_id, stored.member_id) - if PROTOCOL_VERSION not in stored.catalog.protocol_versions: + key, catalog = (stored.room_id, stored.member_id), stored.catalog + if PROTOCOL_VERSION not in catalog.protocol_versions: errors.append(f"{stored.room_id}:{stored.member_id}:protocol-upgrade-required") continue self.peer_routes[key] = PeerMemberRoute( home_install_id=hosted_rooms.local_authority_gateway_id(), - member_id=stored.member_id, target_install_id=stored.catalog.installation_id, - target_profile=stored.target_profile, - capability_digest=stored.catalog.catalog_digest, - execution_policy_digest=stored.catalog.execution_policy.policy_digest, + member_id=stored.member_id, target_install_id=catalog.installation_id, + target_profile=stored.target_profile, capability_digest=catalog.catalog_digest, + execution_policy_digest=catalog.execution_policy.policy_digest, cancellation_scope_id=stored.cancellation_scope_id, trace_id=stored.trace_id, grant=stored.grant) self.peer_clients[key] = PeerRunsHTTPClient( @@ -131,8 +124,7 @@ class HostedRoomService: return self.db_path.parent def local_profiles(self) -> tuple[str, ...]: - profiles = {"default"} - profiles_dir = self.root / "profiles" + profiles, profiles_dir = {"default"}, self.root / "profiles" if profiles_dir.is_dir(): profiles.update(path.name for path in profiles_dir.iterdir() if path.is_dir()) return tuple(sorted(profiles)) @@ -140,27 +132,24 @@ class HostedRoomService: def bindings(self) -> tuple[HostedRoomBinding, ...]: local_gateway_id = hosted_rooms.local_authority_gateway_id() return tuple( - HostedRoomBinding( - room_id=str(room["room_id"]), gateway_id=str(room["authority_gateway_id"]), - authority_epoch=int(room["authority_epoch"])) + HostedRoomBinding(str(room["room_id"]), *_authority(room)) for room in hosted_rooms.list_rooms(self.db_path) if str(room["authority_gateway_id"]) == local_gateway_id) def _room(self, room_id: str) -> dict[str, Any]: return hosted_rooms.room_state(self.db_path, room_id=room_id) - def _owned_room(self, room_id: str) -> dict[str, Any]: - room = self._room(room_id) - if str(room["authority_gateway_id"]) != hosted_rooms.local_authority_gateway_id(): + def _owned_authority(self, room_id: str) -> tuple[str, int]: + """(gateway_id, epoch) of a room this gateway owns; conflict error otherwise.""" + gateway_id, epoch = _authority(self._room(room_id)) + if gateway_id != hosted_rooms.local_authority_gateway_id(): raise hosted_rooms.AuthorityConflictError( "This Group Chat is managed by another gateway.") - return room + return gateway_id, epoch - @contextlib.contextmanager - def _turn_lock(self, profile: str) -> Iterator[None]: + def _turn_lock(self, profile: str) -> contextlib.AbstractContextManager[Path]: from tools.bot_relay import acquire_turn_lock - with acquire_turn_lock(self.root, profile): - yield + return acquire_turn_lock(self.root, profile) def start(self) -> None: self.runtime.start() @@ -175,15 +164,9 @@ class HostedRoomService: for status in statuses: yield from driver.list_tasks(self.db_path, room_id=room_id, status=status) - def _save_link( - self, *, room_id: str, member_id: str, target_url: str, target_profile: str, grant: str, - catalog: GatewayRoomCatalog, cancellation_scope_id: str, trace_id: str) -> None: - hosted_room_links.save_room_link( - self.db_path, - hosted_room_links.make_stored_link( - room_id=room_id, member_id=member_id, target_url=target_url, - target_profile=target_profile, grant=grant, catalog=catalog, - cancellation_scope_id=cancellation_scope_id, trace_id=trace_id)) + def _save_link(self, **link: Any) -> None: + """Persist one stored link (``make_stored_link`` keyword fields).""" + hosted_room_links.save_room_link(self.db_path, hosted_room_links.make_stored_link(**link)) def register_peer_route( self, *, room_id: str, member_id: str, route: PeerMemberRoute, @@ -208,18 +191,19 @@ class HostedRoomService: cancellation_scope_id=route.cancellation_scope_id, trace_id=route.trace_id) # Persistence is the publication boundary: a failed disk write must never # leave a process-local route that disappears after restart. - with self._policy_lock: - self.peer_routes[(room_id, member_id)] = route - self.peer_clients[(room_id, member_id)] = client - self._peer_route_status[(room_id, member_id)] = "ready" + self._publish_route((room_id, member_id), route, client) self.runtime.wakeup() - def revoke_room_routes(self, room_id: str) -> int: - """Revoke and forget every scoped peer route for one room. + def _publish_route(self, key: tuple[str, str], route: PeerMemberRoute, client=None) -> None: + """Make a persisted route live as ``ready`` (and bind its client when given).""" + with self._policy_lock: + self.peer_routes[key], self._peer_route_status[key] = route, "ready" + if client is not None: + self.peer_clients[key] = client - Remote revocation is the boundary: an unreachable target leaves the room - intact for retry rather than reporting a false disband with a live grant. - """ + def revoke_room_routes(self, room_id: str) -> int: + """Revoke and forget every scoped peer route for one room; an unreachable target + leaves the room intact for retry rather than a false disband with a live grant.""" with self._policy_lock: routes = [(key, route) for key, route in self.peer_routes.items() if key[0] == room_id] for key, route in routes: @@ -234,40 +218,38 @@ class HostedRoomService: hosted_rooms.delete_room_link_records(self.db_path, room_id=room_id) with self._policy_lock: for key, _route in routes: - self.peer_routes.pop(key, None) - self._peer_route_status.pop(key, None) - self.peer_clients.pop(key, None) + for table in (self.peer_routes, self._peer_route_status, self.peer_clients): + table.pop(key, None) return len(routes) def _resolve_member_transport(self, binding: HostedRoomBinding, task: Mapping[str, Any]): payload = task.get("payload", {}) member_id = str(payload.get("target_member_id") or payload.get("target_profile") or "") - route = self.peer_routes.get((binding.room_id, member_id)) + key = (binding.room_id, member_id) + route = self.peer_routes.get(key) if route is None: if self._member_is_peer(binding.room_id, member_id): raise RuntimeError("peer room route is unavailable") return self.rpc - client = self.peer_clients.get((binding.room_id, member_id)) + client = self.peer_clients.get(key) if client is None: raise RuntimeError("peer room client is unavailable") identity = task.get("identity") execution_generation = int(task.get("execution_generation") or 0) bind_observation = _hook(client, "bind_observation") if ( - bind_observation is not None - and isinstance(identity, driver.TaskIdentity) + bind_observation is not None and isinstance(identity, driver.TaskIdentity) and execution_generation > 0): bind_observation(task_id=identity.task_id, execution_generation=execution_generation) def set_status(status: str): - return lambda: self._set_route_status(binding.room_id, member_id, status) + return lambda: self._set_route_status(*key, status) tracked_client = _RouteStatusPeerClient( - client, - on_ready=set_status("ready"), + client, on_ready=set_status("ready"), on_reauthorization=set_status("needs_reauthorization"), on_unavailable=set_status("unavailable"), on_refreshed=lambda grant, catalog=None: self._rotate_route_grant( - binding.room_id, member_id, grant, catalog)) + *key, grant, catalog)) self._recover_peer_admission(binding, task, route, tracked_client) return PeerHostedRoomTransport( binding=binding, route=route, client=tracked_client, @@ -279,14 +261,11 @@ class HostedRoomService: client: Any) -> None: """Rediscover an admitted peer run without advancing its generation.""" recover = _hook(client, "recover_dispatch") - identity = task.get("identity") - payload = task.get("payload") + identity, payload = task.get("identity"), task.get("payload") execution_generation = int(task.get("execution_generation") or 0) if ( - recover is None - or not isinstance(identity, driver.TaskIdentity) - or not isinstance(payload, Mapping) - or execution_generation < 1 + recover is None or not isinstance(identity, driver.TaskIdentity) + or not isinstance(payload, Mapping) or execution_generation < 1 or task.get("status") not in {"running", "indeterminate", "stopping"}): return prompt = payload.get("prompt") @@ -300,32 +279,28 @@ class HostedRoomService: recover(dispatch=dispatch.as_mapping(), grant=route.grant) def _member_is_peer(self, room_id: str, member_id: str) -> bool: - for member in self._room(room_id).get("members") or []: - if not isinstance(member, Mapping): - continue - if str(member.get("member_id") or member.get("profile") or "") != member_id: - continue - target = member.get("target") - return isinstance(target, Mapping) and target.get("kind") == "peer" + for m in self._room(room_id).get("members") or []: + if isinstance(m, Mapping) and str( + m.get("member_id") or m.get("profile") or "") == member_id: + target = m.get("target") + return isinstance(target, Mapping) and target.get("kind") == "peer" return False def _set_route_status(self, room_id: str, member_id: str, status: str) -> None: - key = (room_id, member_id) with self._policy_lock: - if self._peer_route_status.get(key) == status: + if self._peer_route_status.get((room_id, member_id)) == status: return - self._peer_route_status[key] = status + self._peer_route_status[(room_id, member_id)] = status hosted_room_links.mark_room_link_status( self.db_path, room_id=room_id, member_id=member_id, status=status) def _set_pending_action( self, room_id: str, member_id: str, action: Mapping[str, Any] | None) -> None: - key = (room_id, member_id) with self._policy_lock: if action is None: - self._pending_actions.pop(key, None) + self._pending_actions.pop((room_id, member_id), None) else: - self._pending_actions[key] = {**action, "member_id": member_id} + self._pending_actions[(room_id, member_id)] = {**action, "member_id": member_id} def _rotate_route_grant( self, room_id: str, member_id: str, grant: str, catalog: GatewayRoomCatalog | None = None @@ -335,11 +310,9 @@ class HostedRoomService: route = self.peer_routes.get(key) if route is None: raise RuntimeError("peer room route is unavailable") - stored = next( - ( - link for link in hosted_room_links.load_room_links(self.db_path) - if (link.room_id, link.member_id) == key), - None) + stored = next(( + l for l in hosted_room_links.load_room_links(self.db_path) + if (l.room_id, l.member_id) == key), None) if stored is None: raise RuntimeError("peer room route cannot be renewed before persistence") digests = {} @@ -348,8 +321,7 @@ class HostedRoomService: catalog.installation_id != route.target_install_id or catalog.execution_policy.target_profile != route.target_profile or PROTOCOL_VERSION not in catalog.protocol_versions - or "direct" not in catalog.link_modes - or not catalog.text + or "direct" not in catalog.link_modes or not catalog.text or catalog.execution_policy.policy_digest != route.execution_policy_digest): self._set_route_status(room_id, member_id, "needs_reauthorization") raise RuntimeError( @@ -357,22 +329,18 @@ class HostedRoomService: digests = { "capability_digest": catalog.catalog_digest, "execution_policy_digest": catalog.execution_policy.policy_digest} - rotated_route = replace(route, grant=grant, **digests) self._save_link( room_id=room_id, member_id=member_id, target_url=stored.target_url, target_profile=stored.target_profile, grant=grant, catalog=catalog or stored.catalog, cancellation_scope_id=stored.cancellation_scope_id, trace_id=stored.trace_id) - with self._policy_lock: - self.peer_routes[key] = rotated_route - self._peer_route_status[key] = "ready" + self._publish_route(key, replace(route, grant=grant, **digests)) def _route_statuses(self, room_id: str | None = None) -> list[dict[str, str]]: with self._policy_lock: - rows = [ - {"room_id": key[0], "member_id": key[1], "status": status} - for key, status in self._peer_route_status.items() - if room_id is None or key[0] == room_id] - return sorted(rows, key=lambda row: (row["room_id"], row["member_id"])) + rows = sorted(self._peer_route_status.items()) + return [ + {"room_id": key[0], "member_id": key[1], "status": status} + for key, status in rows if room_id is None or key[0] == room_id] def _events(self, room_id: str) -> list[dict[str, Any]]: events: list[dict[str, Any]] = [] @@ -390,35 +358,29 @@ class HostedRoomService: raise RuntimeError("hosted room replay cursor did not advance") cursor = next_cursor - def _append_plan(self, room_id: str, plan: discussion.PublicationPlan) -> None: - for event in plan.events: - hosted_rooms.append_event(self.db_path, **event.append_kwargs(room_id)) - def _policy_snapshot(self, room: Mapping[str, Any]) -> PolicySnapshot: return self.policy_checkpoint.snapshot( room_id=str(room["room_id"]), latest_seq=int(room["latest_seq"])) def _publish_terminal_tasks(self, room: Mapping[str, Any]) -> bool: - changed = False - room_id = str(room["room_id"]) - local_profiles = self.local_profiles() - for status in _TERMINAL_STATUSES: - for task in driver.list_tasks(self.db_path, room_id=room_id, status=status): - execution_generation = int(task["execution_generation"]) - if self.policy_checkpoint.publication_exists( - room_id=room_id, task_id=task["identity"].task_id, status=status, - execution_generation=execution_generation): - continue - task_events = self.policy_checkpoint.events_for_task( - room_id=room_id, source_event_seq=int(task["payload"]["source_event_seq"])) - plan = discussion.reconstruct_task_plan( - room, task_events, task, local_profiles=local_profiles) - publication = discussion.plan_publication( - room, task_events, plan, status=status, result=task.get("result"), - execution_generation=execution_generation if status == "deferred" else None, - local_profiles=local_profiles) - self._append_plan(room_id, publication) - changed = True + changed, room_id, local_profiles = False, str(room["room_id"]), self.local_profiles() + for task in self._list_tasks(room_id, _TERMINAL_STATUSES): + status, execution_generation = task["status"], int(task["execution_generation"]) + if self.policy_checkpoint.publication_exists( + room_id=room_id, task_id=task["identity"].task_id, status=status, + execution_generation=execution_generation): + continue + task_events = self.policy_checkpoint.events_for_task( + room_id=room_id, source_event_seq=int(task["payload"]["source_event_seq"])) + plan = discussion.reconstruct_task_plan( + room, task_events, task, local_profiles=local_profiles) + publication = discussion.plan_publication( + room, task_events, plan, status=status, result=task.get("result"), + execution_generation=execution_generation if status == "deferred" else None, + local_profiles=local_profiles) + for event in publication.events: + hosted_rooms.append_event(self.db_path, **event.append_kwargs(room_id)) + changed = True return changed def _append_room_status( @@ -427,29 +389,26 @@ class HostedRoomService: return gateway_id, epoch = _authority(room) hosted_rooms.append_event( - self.db_path, - room_id=str(room["room_id"]), + self.db_path, room_id=str(room["room_id"]), event_id=f"dactivity:{decision.discussion_event_id}:{decision.reason}", - kind="room.activity", - actor={"kind": "gateway", "id": gateway_id}, + kind="room.activity", actor={"kind": "gateway", "id": gateway_id}, payload={ "status": decision.status, "reason_code": decision.reason, "thread_id": decision.thread_id, "discussion_event_id": decision.discussion_event_id}, - authority_gateway_id=gateway_id, - authority_epoch=epoch) + authority_gateway_id=gateway_id, authority_epoch=epoch) def prepare_room(self, binding: HostedRoomBinding) -> None: with self._policy_lock: room = self._room(binding.room_id) - snapshot = self._policy_snapshot(room) + snapshot = self._policy_snapshot(room) # sync() side effect feeds the publish below if self._publish_terminal_tasks(room): room = self._room(binding.room_id) snapshot = self._policy_snapshot(room) self.policy_checkpoint.compact_completed(room_id=binding.room_id) driver.prune_published_terminal_tasks( self.db_path, room_id=binding.room_id, clock=self.runtime.clock) - if any(True for _ in self._list_tasks(binding.room_id, _LIVE_STATUSES)): + if next(iter(self._list_tasks(binding.room_id, _LIVE_STATUSES)), None) is not None: return decision = discussion.plan_next_task( room, list(snapshot.events), local_profiles=self.local_profiles(), @@ -458,17 +417,11 @@ class HostedRoomService: driver.admit_task( self.db_path, decision.task.identity, payload=decision.task.payload, clock=time.time) - # A stop can race the policy read from another process. Re-read after - # admission and cancel before the runtime can execute a task whose - # source event is now behind the room stop fence. - stopped_through_seq = self._policy_snapshot( - self._room(binding.room_id) - ).stopped_through_seq - if ( - decision.source_event_seq is not None - and decision.source_event_seq < stopped_through_seq): - self.runtime.cancel( - decision.task.identity, cancel_id=f"stop-fence:{stopped_through_seq}") + # A stop can race the policy read from another process: re-read after admission + # and cancel a task whose source event is now behind the room stop fence. + fence = self._policy_snapshot(self._room(binding.room_id)).stopped_through_seq + if decision.source_event_seq is not None and decision.source_event_seq < fence: + self.runtime.cancel(decision.task.identity, cancel_id=f"stop-fence:{fence}") elif decision.status in {"settled", "bounded"}: self._append_room_status(room, decision) @@ -479,9 +432,7 @@ class HostedRoomService: def create_room(self, *, room_id: str, name: str, members: Any) -> dict[str, Any]: normalized = discussion.validate_roster(members, local_profiles=self.local_profiles()) room = hosted_rooms.create_room( - self.db_path, - room_id=room_id, - name=name, + self.db_path, room_id=room_id, name=name, members=[ { "member_id": member.member_id, "profile": member.profile, @@ -494,7 +445,7 @@ class HostedRoomService: def send(self, *, room_id: str, event_id: str, payload: Any) -> dict[str, Any]: normalized = discussion.validate_user_payload(payload) - gateway_id, epoch = _authority(self._owned_room(room_id)) + gateway_id, epoch = self._owned_authority(room_id) event = hosted_rooms.append_event( self.db_path, room_id=room_id, event_id=event_id, kind="message.user", actor={"kind": "user", "id": "desktop"}, payload=normalized, @@ -508,36 +459,30 @@ class HostedRoomService: def stop_room( self, room_id: str, *, cancel_id: str, require_acknowledged: bool = False) -> int: - gateway_id, epoch = _authority(self._owned_room(room_id)) + gateway_id, epoch = self._owned_authority(room_id) hosted_rooms.request_room_stop( self.db_path, room_id=room_id, cancel_id=cancel_id, expected_gateway_id=gateway_id, expected_epoch=epoch) - cancelled = 0 pending = 0 with self._policy_lock: tasks = { (task["identity"].room_id, task["identity"].task_id): task for task in self._list_tasks(room_id, _STOPPABLE_STATUSES)} for task in tasks.values(): - task_cancel_id = ( - str(task.get("cancel_id") or "") if task.get("status") == "stopping" else "") - result = self.runtime.cancel( - task["identity"], cancel_id=task_cancel_id or cancel_id) - cancelled += 1 + own_cancel_id = ( + task.get("status") == "stopping" and str(task.get("cancel_id") or "")) + result = self.runtime.cancel(task["identity"], cancel_id=own_cancel_id or cancel_id) if result["status"] == "stopping": pending += 1 if require_acknowledged and pending: raise RuntimeError("room work is still stopping; retry deletion after Stop completes") self.runtime.wakeup() - return cancelled + return len(tasks) def retry_room_task(self, room_id: str, *, task_id: str) -> dict[str, Any]: """Retry one uncertain or deferred task only after explicit user action.""" - task = next( - ( - candidate for candidate in self._list_tasks(room_id, _RETRYABLE_STATUSES) - if candidate["identity"].task_id == task_id), - None) + candidates = self._list_tasks(room_id, _RETRYABLE_STATUSES) + task = next((c for c in candidates if c["identity"].task_id == task_id), None) if task is None: raise driver.InvalidTaskTransitionError("no retryable room task matches task_id") return self.runtime.retry_indeterminate(task["identity"]) @@ -547,18 +492,16 @@ class HostedRoomService: choice: str, request_id: str | None = None) -> Mapping[str, Any]: """Resolve one exact local or peer approval and wake room observation.""" key = (room_id, member_id) - route = self.peer_routes.get(key) - client = self.peer_clients.get(key) + route, client = self.peer_routes.get(key), self.peer_clients.get(key) with self._policy_lock: action = self._pending_actions.get(key) requested_approval_id = str(request_id or "") def matches(pending: Mapping[str, Any] | None) -> bool: - return ( - pending is not None - and str(pending.get("request_id") or "") == requested_approval_id - and pending.get("task_id") == task_id - and int(pending.get("execution_generation") or 0) == execution_generation) + return pending is not None and ( + str(pending.get("request_id") or ""), pending.get("task_id"), + int(pending.get("execution_generation") or 0), + ) == (requested_approval_id, task_id, execution_generation) if not requested_approval_id or not matches(action): raise RuntimeError("room approval is no longer pending") if choice not in {"once", "deny"}: @@ -592,21 +535,16 @@ class HostedRoomService: counts = Counter(str(task["status"]) for task in tasks) pending_actions = [ {"kind": "retry", "task_id": task["identity"].task_id} - for task in tasks - if task["status"] in _RETRYABLE_STATUSES] + for task in tasks if task["status"] in _RETRYABLE_STATUSES] with self._policy_lock: pending_actions.extend( - dict(action) - for (action_room_id, _member_id), action in self._pending_actions.items() - if action_room_id == room_id) + dict(action) for (action_room_id, _member_id), action + in self._pending_actions.items() if action_room_id == room_id) return { - "running": runtime["running"], - "working": bool( - counts.get("running") or counts.get("queued") or counts.get("stopping")), + "running": runtime["running"], "working": any(counts.get(s) for s in _LIVE_STATUSES), "blocked": room_id in runtime["blocked_rooms"] or bool(counts.get("indeterminate") or counts.get("stopping")), - "counts": dict(counts), - "pending_actions": pending_actions, + "counts": dict(counts), "pending_actions": pending_actions, "peer_routes": self._route_statuses(room_id)} @@ -615,20 +553,14 @@ class _RouteStatusPeerClient: def __init__( self, client, *, on_ready, on_reauthorization, on_unavailable, on_refreshed) -> None: - self._client = client - self._on_ready = on_ready - self._on_reauthorization = on_reauthorization - self._on_unavailable = on_unavailable - self._on_refreshed = on_refreshed + self._client, self._on_ready, self._on_refreshed = client, on_ready, on_refreshed + self._on_reauthorization, self._on_unavailable = on_reauthorization, on_unavailable def _refresh_grant(self, kwargs: dict) -> dict: - """Rotate an expiring grant before dispatch; return the kwargs to send. - - Refresh failures only escalate to reauthorization when the peer says so or the - grant is already past its hard expiry; otherwise the original grant is tried - as-is. A refreshed catalog whose digests drift from the dispatch is a policy - change and is refused before any dispatch. - """ + """Rotate an expiring grant before dispatch; return the kwargs to send. Refresh + failures escalate to reauthorization only when the peer says so or the grant is + past its hard expiry; otherwise the original grant is tried as-is. A refreshed + catalog whose digests drift from the dispatch is a policy change: refused.""" grant = kwargs["grant"] if not room_grant_needs_dispatch_refresh(grant): return kwargs @@ -652,19 +584,12 @@ class _RouteStatusPeerClient: refreshed_catalog = None if refreshed.get("catalog") is not None: refreshed_catalog = GatewayRoomCatalog.from_mapping(refreshed.get("catalog")) - drift = None - if refreshed_catalog.execution_policy.policy_digest != checked.execution_policy_digest: - drift = ( - "peer room execution policy needs reauthorization", - "room_execution_policy_changed") - elif refreshed_catalog.catalog_digest != checked.capability_digest: - drift = ( - "peer room capabilities need reauthorization", "room_capability_catalog_changed" - ) + drift = digest_reauthorization_error( + refreshed_catalog, capability_digest=checked.capability_digest, + execution_policy_digest=checked.execution_policy_digest) if drift is not None: self._on_reauthorization() - raise PeerRunsHTTPError( - drift[0], status_code=403, error_code=drift[1], not_admitted=True) + raise drift self._on_refreshed(replacement, refreshed_catalog) return {**kwargs, "grant": replacement} diff --git a/tui_gateway/mcp_oauth_sessions.py b/tui_gateway/mcp_oauth_sessions.py index 4a1f1c8a16..650575fd0d 100644 --- a/tui_gateway/mcp_oauth_sessions.py +++ b/tui_gateway/mcp_oauth_sessions.py @@ -1,12 +1,9 @@ -"""Session-backed MCP OAuth flows for the gateway (mcp.servers.oauth.*). - -``start`` kicks off a background worker and returns ``{session_id, auth_url, flow}``; -``poll`` reports ``{status: pending|approved|error}`` until tokens land on disk. No OAuth -logic is reimplemented: ``hermes mcp login``'s probe under ``force_interactive_oauth`` -plus ``DashboardOAuthFlow`` as the bridge; the only new piece is a loopback listener -feeding ``deliver_callback``. Remote backends: the client hosts the listener, passes -``client_redirect_uri`` and relays via ``deliver_callback_flow`` (state check stays here). -""" +"""Session-backed MCP OAuth flows for the gateway (mcp.servers.oauth.*): ``start`` spawns a +worker and returns ``{session_id, auth_url, flow}``; ``poll`` reports ``{status}`` until tokens +land. Reuses ``hermes mcp login``'s probe under ``force_interactive_oauth`` plus +``DashboardOAuthFlow``; the only new piece is a loopback listener feeding ``deliver_callback``. +Remote backends host the listener (``client_redirect_uri``) and relay via +``deliver_callback_flow``.""" from __future__ import annotations @@ -23,18 +20,8 @@ from urllib.parse import parse_qs, urlparse _sessions: Dict[str, Dict[str, Any]] = {} _sessions_lock = threading.Lock() -# How long a completed/abandoned session lingers before GC (seconds). -_SESSION_TTL_SECONDS = 900 -# Cap concurrent in-flight flows so a runaway client can't exhaust ports/threads. -_MAX_PENDING = 12 - - -def _gc_sessions() -> None: - """Drop expired sessions. Called opportunistically on start.""" - cutoff = time.time() - _SESSION_TTL_SECONDS - with _sessions_lock: - for sid in [sid for sid, rec in _sessions.items() if rec["created_at"] < cutoff]: - _shutdown_listener(_sessions.pop(sid)) +_SESSION_TTL_SECONDS = 900 # completed/abandoned session lingers this long before GC +_MAX_PENDING = 12 # cap in-flight flows so a runaway client can't exhaust ports/threads def _shutdown_listener(rec: Dict[str, Any]) -> None: @@ -48,24 +35,20 @@ def _shutdown_listener(rec: Dict[str, Any]) -> None: def _validate_client_redirect_uri(uri: str) -> str: - """Accept only plain-http loopback URLs (RFC 8252 native-app rules) so the - gateway can't pin an attacker-controlled redirect into a DCR registration.""" + """Accept only plain-http loopback URLs (RFC 8252) so the gateway can't pin an + attacker-controlled redirect into a DCR registration.""" parsed = urlparse(str(uri or "").strip()) host = (parsed.hostname or "").lower() if (parsed.scheme != "http" or host not in ("127.0.0.1", "localhost", "::1") or not parsed.port or parsed.username is not None or parsed.password is not None): raise ValueError( - "client_redirect_uri must be a loopback http URL like " - "http://127.0.0.1:/callback" - ) + "client_redirect_uri must be a loopback http URL like http://127.0.0.1:/callback") return f"http://{'[' + host + ']' if ':' in host else host}:{parsed.port}{parsed.path or '/callback'}" def _start_loopback_listener(flow) -> "http.server.HTTPServer": """Bind a loopback callback listener feeding ``flow.deliver_callback``; returns the - HTTPServer already serving on a daemon thread. The caller pins ``flow.redirect_uri`` - from ``server_address`` BEFORE the worker starts (fixed at authorization).""" - + HTTPServer already serving on a daemon thread (caller pins ``flow.redirect_uri`` from it).""" class _Handler(http.server.BaseHTTPRequestHandler): def do_GET(self): # noqa: N802 — stdlib naming parsed = urlparse(self.path) @@ -74,11 +57,11 @@ def _start_loopback_listener(flow) -> "http.server.HTTPServer": self.end_headers() return qs = parse_qs(parsed.query) - code, state, error = ((qs.get(k) or [None])[0] for k in ("code", "state", "error")) body = b"

Authorization received

You can close this tab and return to Hermes.

" status = 200 try: - flow.deliver_callback(code=code, state=state, error=error) + flow.deliver_callback( + **{k: (qs.get(k) or [None])[0] for k in ("code", "state", "error")}) except Exception: body = b"

OAuth callback rejected

The callback was invalid or already used.

" status = 400 @@ -129,9 +112,9 @@ def _probe_with_rollback( raise -def _worker(session_id: str, hermes_home: str, server_name: str, cfg: dict, reconnect_live: bool) -> None: - """Drive the interactive MCP OAuth probe under the shared dashboard bridge (same - HERMES_HOME + secret-scope + force_interactive_oauth + dashboard_oauth_flow wrapping +def _worker( + session_id: str, hermes_home: str, server_name: str, cfg: dict, reconnect_live: bool) -> None: + """Drive the interactive MCP OAuth probe under the shared dashboard bridge (same wrapping as ``web_server._run_dashboard_mcp_oauth``), keyed to our session record.""" from hermes_constants import reset_hermes_home_override, set_hermes_home_override rec = _sessions.get(session_id) @@ -166,23 +149,18 @@ def _worker(session_id: str, hermes_home: str, server_name: str, cfg: dict, reco def start_flow( - hermes_home: str, - server_name: str, - cfg: dict, - *, - reconnect_live: bool = False, - url_timeout: float = 30.0, - client_redirect_uri: Optional[str] = None) -> Dict[str, Any]: - """Begin an MCP OAuth flow and return ``{session_id, auth_url, flow}``; blocks up - to ``url_timeout`` for the authorization URL. With ``client_redirect_uri`` (remote - backend; invalid values raise ``ValueError``) no gateway-side listener is bound and - the client relays ``code``/``state`` via ``deliver_callback_flow``.""" + hermes_home: str, server_name: str, cfg: dict, *, reconnect_live: bool = False, + url_timeout: float = 30.0, client_redirect_uri: Optional[str] = None) -> Dict[str, Any]: + """Begin an MCP OAuth flow and return ``{session_id, auth_url, flow}``; blocks up to + ``url_timeout`` for the authorization URL. With ``client_redirect_uri`` (invalid values + raise ``ValueError``) no gateway-side listener is bound.""" from tools.mcp_dashboard_oauth import DashboardOAuthFlow if client_redirect_uri is not None: client_redirect_uri = _validate_client_redirect_uri(client_redirect_uri) - - _gc_sessions() - + cutoff = time.time() - _SESSION_TTL_SECONDS # opportunistic GC of expired sessions + with _sessions_lock: + for sid in [sid for sid, rec in _sessions.items() if rec["created_at"] < cutoff]: + _shutdown_listener(_sessions.pop(sid)) with _sessions_lock: active = [r for r in _sessions.values() if not r["flow"].worker_done] if len(active) >= _MAX_PENDING: @@ -199,18 +177,14 @@ def start_flow( httpd = None if client_redirect_uri else _start_loopback_listener(flow) flow.redirect_uri = ( client_redirect_uri or f"http://127.0.0.1:{httpd.server_address[1]}/callback") - rec = { "session_id": session_id, "server_name": server_name, "hermes_home": hermes_home, - "flow": flow, "httpd": httpd, "created_at": time.time(), - } + "flow": flow, "httpd": httpd, "created_at": time.time()} with _sessions_lock: _sessions[session_id] = rec - threading.Thread( target=_worker, args=(session_id, hermes_home, server_name, dict(cfg), reconnect_live), daemon=True, name=f"mcp-oauth-{server_name}").start() - try: auth_url = None # wait_for_authorization_url is async; run its wait synchronously. @@ -220,7 +194,8 @@ def start_flow( if auth_url := snap.get("authorization_url"): break if snap.get("status") == "error": - raise RuntimeError(snap.get("error") or "MCP OAuth flow failed before authorization") + raise RuntimeError( + snap.get("error") or "MCP OAuth flow failed before authorization") time.sleep(0.1) if not auth_url: raise TimeoutError("Timed out waiting for MCP authorization URL") @@ -228,9 +203,7 @@ def start_flow( flow.mark_error("Timed out waiting for MCP authorization URL") _shutdown_listener(rec) raise - - # ``flow`` mirrors the provider-OAuth discriminator: open a URL then poll - # (no user_code to type, unlike device_code). + # ``flow`` mirrors the provider-OAuth discriminator: open a URL then poll (no user_code). return {"session_id": session_id, "auth_url": auth_url, "flow": "pkce"} @@ -246,21 +219,19 @@ def _lookup(session_id: str, server_name: str) -> "tuple[Dict[str, Any] | None, def poll_flow(session_id: str, server_name: str) -> Dict[str, Any]: - """Poll a session → ``{status, error_message?, auth_url?, tools?}``; ``status`` - is ``pending`` | ``approved`` | ``error`` (the bridge's ``authorization_required`` - maps to ``pending`` — the client only needs to know whether to keep waiting).""" + """Poll a session → ``{status, error_message?, auth_url?, tools?}``; ``status`` is + ``pending`` | ``approved`` | ``error`` (the bridge's ``authorization_required`` maps to + ``pending``).""" rec, err = _lookup(session_id, server_name) if rec is None: return {"status": "error", "error_message": err} - flow = rec["flow"] snap = flow.snapshot() raw = snap.get("status") status = raw if raw in ("approved", "error") else "pending" out: Dict[str, Any] = { "session_id": session_id, "status": status, "error_message": snap.get("error"), - "auth_url": snap.get("authorization_url"), - } + "auth_url": snap.get("authorization_url")} if status == "approved": out["tools"] = list(getattr(flow, "tools", []) or []) return out @@ -268,12 +239,10 @@ def poll_flow(session_id: str, server_name: str) -> Dict[str, Any]: def deliver_callback_flow( session_id: str, server_name: str, *, code: Optional[str], state: Optional[str], - error: Optional[str] = None, -) -> Dict[str, Any]: - """Relay a client-captured OAuth redirect into a session's flow (remote-backend - companion to ``start_flow(client_redirect_uri=...)``). Security is unchanged: - ``DashboardOAuthFlow.deliver_callback`` verifies ``state`` (constant-time) and - rejects replays. Returns ``{ok: true}`` or ``{ok: false, error_message}``.""" + error: Optional[str] = None) -> Dict[str, Any]: + """Relay a client-captured OAuth redirect into a session's flow (remote-backend companion + to ``start_flow(client_redirect_uri=...)``); ``deliver_callback`` still verifies ``state`` + and rejects replays. Returns ``{ok: true}`` or ``{ok: false, error_message}``.""" rec, err = _lookup(session_id, server_name) if rec is None: return {"ok": False, "error_message": err} diff --git a/tui_gateway/method_ctx.py b/tui_gateway/method_ctx.py index a35e1c6103..445f02c274 100644 --- a/tui_gateway/method_ctx.py +++ b/tui_gateway/method_ctx.py @@ -1,12 +1,9 @@ -"""Seam for the server.py handler/helper split. - -server.py's JSON-RPC handlers and helpers close over its module globals (``_sessions``, -``_ok``, ``_err``, ...). Split modules define their code normally and server.py calls -:func:`bind_module` at the end of its own import, once every global exists: bodies are -re-created with ``types.FunctionType`` against server.py's namespace, so they stay -byte-identical and ``global X`` statements keep mutating server.py state. No import -cycle: split modules never import server at module level — server passes itself in. -""" +"""Seam for the server.py handler/helper split. server.py's JSON-RPC handlers and helpers close +over its module globals (``_sessions``, ``_ok``, ``_err``, ...). Split modules define their code +normally and server.py calls :func:`bind_module` at the end of its own import, once every global +exists: bodies are re-created with ``types.FunctionType`` against server.py's namespace, so they +stay byte-identical and ``global X`` keeps mutating server.py state. No import cycle: split +modules never import server at module level — server passes itself in.""" import contextlib import types @@ -27,18 +24,15 @@ def rebind(fn, g: dict, _seen=None): return contextlib.contextmanager(rebind(wrapped, g, _seen)) closure = fn.__closure__ if closure: - cells = [] - for cell in closure: + def _cell(cell): try: val = cell.cell_contents except ValueError: # empty cell - cells.append(cell) - continue + return cell if isinstance(val, types.FunctionType) and val.__module__ == fn.__module__: - cells.append(types.CellType(rebind(val, g, _seen))) - else: - cells.append(cell) - closure = tuple(cells) + return types.CellType(rebind(val, g, _seen)) + return cell + closure = tuple(_cell(c) for c in closure) real = types.FunctionType(fn.__code__, g, fn.__name__, fn.__defaults__, closure) real.__kwdefaults__ = fn.__kwdefaults__ real.__doc__ = fn.__doc__ @@ -55,11 +49,9 @@ class HandlerRegistry: def method(self, name: str): """Drop-in for server.py's ``@method`` decorator (defers registration).""" - def dec(fn): self._pending.append((name, fn)) return fn - return dec def profile_scoped(self, fn): @@ -82,16 +74,11 @@ _PLUMBING = {"HandlerRegistry", "method", "_profile_scoped", "register", "rebind def bind_module(module_globals: dict, server, *, skip=()) -> None: """Publish everything a split module defines onto ``server``, rebound to its globals. - - ``module_globals`` is the caller's ``globals()`` (not ``sys.modules[__name__]``: tests - that ``patch.dict(sys.modules)`` around the server import drop the submodule entries - while the package attribute survives, so a re-import would KeyError). Functions are - rebound; classes get their methods rebound in place; dispatch tables (dict/tuple/list - holding this module's functions) get their values rebound; other values (constants, - ``global``-mutated state seeds) are copied as-is. Imported modules/functions, dunders - and registry plumbing are skipped, so no hand-maintained export list is needed. - Finally the module's ``_registry`` (if any) installs its @method handlers. - """ + ``module_globals`` is the caller's ``globals()`` (not ``sys.modules[__name__]``: tests that + ``patch.dict(sys.modules)`` around the server import drop the submodule entries). Functions + are rebound; classes get their methods rebound in place; dispatch tables (dict/tuple/list of + this module's functions) get their values rebound; other values are copied as-is. Imported + modules/functions, dunders and registry plumbing are skipped; finally ``_registry`` installs.""" g = vars(server) mod_name = module_globals["__name__"] seen: dict = {} @@ -104,18 +91,15 @@ def bind_module(module_globals: dict, server, *, skip=()) -> None: return rebind(v, g, seen) if isinstance(v, dict): return {k: _rebind_in(x) for k, x in v.items()} - if isinstance(v, (tuple, list)): - return type(v)(_rebind_in(x) for x in v) - return v + return type(v)(_rebind_in(x) for x in v) if isinstance(v, (tuple, list)) else v def _has_own_fn(v): items = v.values() if isinstance(v, dict) else v if isinstance(v, (tuple, list)) else None return _own_fn(v) if items is None else any(_has_own_fn(x) for x in items) for name, obj in list(module_globals.items()): - if name.startswith("__") or name in _PLUMBING or name in skip: - continue - if isinstance(obj, (types.ModuleType, HandlerRegistry)): + if (name.startswith("__") or name in _PLUMBING or name in skip + or isinstance(obj, (types.ModuleType, HandlerRegistry))): continue if isinstance(obj, types.FunctionType): if obj.__module__ == mod_name: diff --git a/tui_gateway/methods_bot_relay.py b/tui_gateway/methods_bot_relay.py index 31b46a2264..59fae4ecb6 100644 --- a/tui_gateway/methods_bot_relay.py +++ b/tui_gateway/methods_bot_relay.py @@ -1,12 +1,10 @@ -"""Bot-relay JSON-RPC handlers — the gateway side of cross-connection A2A. - -Connections ARE the peer set: the Desktop owns every gateway socket and relays between them via -four doors on EACH gateway: ``roster.sync`` (push OTHER connections' agents so ``message_agent`` -resolves them), ``outbox.drain`` (collect envelopes queued here for other connections), ``deliver`` -(one-turn Bot Chat delivery on the TARGET gateway, returns the reply), ``reply`` (write the -reply/error back on the SENDER gateway for its waiter). Plumbing: ``tools/bot_relay.py``. -Handlers are rebound onto server.py's globals (method_ctx.py) and reference ``_ok``/``_err`` bare. -""" +"""Bot-relay JSON-RPC handlers — the gateway side of cross-connection A2A. Connections ARE the +peer set: the Desktop owns every gateway socket and relays between them via four doors on EACH +gateway: ``roster.sync`` (push OTHER connections' agents so ``message_agent`` resolves them), +``outbox.drain`` (collect envelopes queued here for other connections), ``deliver`` (one-turn Bot +Chat delivery on the TARGET gateway, returns the reply), ``reply`` (write the reply/error back on +the SENDER gateway for its waiter). Plumbing: ``tools/bot_relay.py``; handlers are rebound onto +server.py's globals (method_ctx.py) and reference ``_ok``/``_err`` bare.""" import os import subprocess @@ -26,7 +24,6 @@ def _relay_root() -> Path: def _run_delivery(profile: str, tmp: str) -> subprocess.CompletedProcess: from tools.bot_relay import local_delivery_command - return subprocess.run( local_delivery_command(profile, tmp), capture_output=True, text=True, encoding="utf-8", errors="replace", timeout=600) @@ -34,14 +31,10 @@ def _run_delivery(profile: str, tmp: str) -> subprocess.CompletedProcess: @method("bot_relay.roster.sync") def _(rid, params: dict, _root=_relay_root) -> dict: - """Replace this gateway's view of agents on OTHER connections → ``{count}`` accepted rows. - - ``agents``: rows ``{profile, handle, connection_id, connection_label?, title?, description?}``; - rows failing validation are dropped, not fatal. - """ + """Replace this gateway's view of agents on OTHER connections → ``{count}`` accepted rows + (``agents`` rows ``{profile, handle, connection_id, ...}``; invalid rows are dropped).""" try: from tools.bot_relay import write_remote_roster - return _ok(rid, {"count": write_remote_roster(_root(), params.get("agents"))}) except Exception as e: return _err(rid, 5090, str(e)) @@ -49,13 +42,10 @@ def _(rid, params: dict, _root=_relay_root) -> dict: @method("bot_relay.outbox.drain") def _(rid, params: dict, _root=_relay_root) -> dict: - """Claim every pending cross-connection envelope queued on this gateway → ``{envelopes}``. - - Claimed envelopes move to ``claimed/`` atomically, so concurrent drains can't double-deliver. - """ + """Claim every pending cross-connection envelope queued here → ``{envelopes}``; claimed + envelopes move to ``claimed/`` atomically so concurrent drains can't double-deliver.""" try: from tools.bot_relay import claim_pending_envelopes - return _ok(rid, {"envelopes": claim_pending_envelopes(_root())}) except Exception as e: return _err(rid, 5091, str(e)) @@ -66,10 +56,7 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: """Deliver a relayed DM (``profile``, attribution-prefixed ``message``) into a Bot Chat ON THIS GATEWAY via the one-turn ``hermes -p chat -c "Bot Chat"`` transport local DMs use → ``{reply}``. Blocking by design (Desktop relay worker; the RPC pool keeps it off the reader).""" - import os - import subprocess import tempfile - profile = str(params.get("profile") or "").strip() message = str(params.get("message") or "").strip() if not profile or not message: @@ -77,14 +64,12 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: try: from tools.bot_mode_dm import MESSAGE_MAX_CHARS from tools.bot_relay import acquire_turn_lock - if len(message) > MESSAGE_MAX_CHARS + 200: # + attribution headroom return _err(rid, 4091, "message too long") root = _root() known = {"default"} - profiles_dir = root / "profiles" - if profiles_dir.is_dir(): - known.update(c.name for c in profiles_dir.iterdir() if c.is_dir()) + if (root / "profiles").is_dir(): + known.update(c.name for c in (root / "profiles").iterdir() if c.is_dir()) resolved = "default" if profile.lower() == "hermes" else profile if resolved not in known: return _err(rid, 4092, f"no profile '{profile}' on this gateway") @@ -92,23 +77,15 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: # When THIS gateway already hosts the target's Bot Chat live, the subprocess transport is # fenced out by the single-owner lease and the payload dropped. Land the DM in the live # session via prompt.submit — the composer's choke point, so role alternation, persistence - # and streaming behave as a typed message would. (Nested: needs server globals via rebind.) - def _live_bot_chat_sid(profile_name: str) -> str: - from tools.bot_mode_probe import BOT_CHAT_TITLE - - live_home = _profile_home(profile_name) - want_home = str(live_home) if live_home is not None else None - for live_sid, record in list(_sessions.items()): - if not isinstance(record, dict): - continue - if (record.get("profile_home") or None) != want_home: - continue - key = _session_lookup_key(record, fallback=live_sid) - if _session_live_title(record, key) == BOT_CHAT_TITLE: - return live_sid - return "" - - live_sid = _live_bot_chat_sid(resolved) + # and streaming behave as a typed message would. + from tools.bot_mode_probe import BOT_CHAT_TITLE + live_home = _profile_home(resolved) + want_home = str(live_home) if live_home is not None else None + live_sid = next(( + live_sid for live_sid, record in list(_sessions.items()) + if isinstance(record, dict) and (record.get("profile_home") or None) == want_home + and _session_live_title( + record, _session_lookup_key(record, fallback=live_sid)) == BOT_CHAT_TITLE), "") if live_sid: # queued=True: a teammate's DM runs as the NEXT turn and never interrupts or steers a # turn in flight (the default busy mode does); arrivals queue in order. @@ -118,6 +95,9 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: reply = f"Delivered into @{resolved}'s open Bot Chat; the reply will appear there." return _ok(rid, {"reply": reply}) + def _detail(p) -> str: + return (p.stderr or p.stdout or "").strip()[-500:] + fd, tmp = tempfile.mkstemp(prefix="hermes-relay-dm-", suffix=".txt", text=True) try: with os.fdopen(fd, "w", encoding="utf-8") as f: @@ -133,28 +113,22 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: # transcript first (no fresh session is minted). Auth/quota/config never retry. from tools.bot_failure_reasons import ( RETRY_NONE, classify_agent_error, retry_action) - - first_detail = (proc.stderr or proc.stdout or "").strip()[-500:] - if retry_action(classify_agent_error(first_detail)) != RETRY_NONE: + if retry_action(classify_agent_error(_detail(proc))) != RETRY_NONE: proc = _run(resolved, tmp) finally: with contextlib.suppress(OSError): os.unlink(tmp) if proc.returncode != 0: from tools.bot_failure_reasons import classify_agent_error - - detail = (proc.stderr or proc.stdout or "").strip()[-500:] - return _err( - rid, 5092, f"delivery turn failed: {detail or proc.returncode}", - data={"reason": classify_agent_error(detail)}) + detail = _detail(proc) + return _err(rid, 5092, f"delivery turn failed: {detail or proc.returncode}", + data={"reason": classify_agent_error(detail)}) return _ok(rid, {"reply": (proc.stdout or "").strip()}) except subprocess.TimeoutExpired: return _err(rid, 5093, "delivery turn timed out") except Exception as e: # 'target_busy' extends the structured refusal enum. - if getattr(e, "reason", "") == "target_busy": - return _err(rid, 5096, str(e)) - return _err(rid, 5094, str(e)) + return _err(rid, 5096 if getattr(e, "reason", "") == "target_busy" else 5094, str(e)) @method("bot_relay.reply") @@ -166,10 +140,8 @@ def _(rid, params: dict, _root=_relay_root) -> dict: return _err(rid, 4093, "id required") try: from tools.bot_relay import write_reply - - write_reply( - _root(), envelope_id, reply=str(params.get("reply") or ""), - error=str(params.get("error") or ""), reason=str(params.get("reason") or "")) + write_reply(_root(), envelope_id, reply=str(params.get("reply") or ""), + error=str(params.get("error") or ""), reason=str(params.get("reason") or "")) return _ok(rid, {"ok": True}) except ValueError as e: return _err(rid, 4094, str(e)) @@ -180,7 +152,6 @@ def _(rid, params: dict, _root=_relay_root) -> dict: def register(server) -> None: _registry.install(server) from . import methods_groups - server._LONG_HANDLERS = server._LONG_HANDLERS | methods_groups.LONG_HANDLERS for name in ( "get_hosted_room_service", "_WORKER_UNAVAILABLE", "_profile_name", "_requested_profile", diff --git a/tui_gateway/methods_browser.py b/tui_gateway/methods_browser.py index ac9a303718..af7cd12d4c 100644 --- a/tui_gateway/methods_browser.py +++ b/tui_gateway/methods_browser.py @@ -1,8 +1,5 @@ -"""Browser connect/disconnect helpers for the browser.* RPCs (CDP probing, no network I/O on status). - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so they reference server.py globals bare. -""" +"""Browser connect/disconnect helpers for the browser.* RPCs (CDP probing, no network I/O on +status). Bodies are rebound onto server.py's globals at install time (method_ctx.bind_module).""" from __future__ import annotations @@ -14,25 +11,17 @@ _CDP_SCHEMES = {"http", "https", "ws", "wss"} def _resolve_browser_cdp_url() -> str: - """Configured browser CDP override without network I/O. - - ``/browser status`` must be fast: ``tools.browser_tool._get_cdp_override`` runs an HTTP - probe with a multi-second timeout for discovery-style URLs. Mirrors its precedence (env var, - then ``browser.cdp_url``) minus the WS-resolution step, so the answer reflects user intent - even when the host is unreachable; ``browser_navigate`` normalizes on the next tool call. - """ - env_url = os.environ.get("BROWSER_CDP_URL", "").strip() - if env_url: + """Configured browser CDP override without network I/O (``/browser status`` must be fast; + ``tools.browser_tool._get_cdp_override`` HTTP-probes discovery URLs). Same precedence (env, + then ``browser.cdp_url``) minus WS resolution; ``browser_navigate`` normalizes on the next call.""" + if env_url := os.environ.get("BROWSER_CDP_URL", "").strip(): return env_url - try: + with contextlib.suppress(Exception): from hermes_cli.config import read_raw_config - cfg = read_raw_config() browser_cfg = cfg.get("browser", {}) if isinstance(cfg, dict) else {} if isinstance(browser_cfg, dict): return str(browser_cfg.get("cdp_url", "") or "").strip() - except Exception: - pass return "" @@ -40,59 +29,29 @@ def _is_default_local_cdp(parsed) -> bool: """Match the discovery-style local default; never the concrete WS form — a ``ws://127.0.0.1:9222/devtools/browser/`` is connectable as-is and collapsing it to bare ``http://...:9222`` would break the connect.""" - try: - port = parsed.port or 80 - except ValueError: - return False - return (parsed.scheme in {"http", "ws"} and parsed.hostname in {"127.0.0.1", "localhost"} - and port == 9222 and parsed.path in {"", "/", "/json", "/json/version"}) + with contextlib.suppress(ValueError): + return (parsed.scheme in {"http", "ws"} and parsed.hostname in {"127.0.0.1", "localhost"} + and (parsed.port or 80) == 9222 and parsed.path in {"", "/", "/json", "/json/version"}) + return False def _cdp_http_reachable(parsed, timeout: float = 2.0) -> bool: """True when ``/json/version`` or ``/json`` on the CDP host answers 2xx.""" import urllib.request - scheme = {"ws": "http", "wss": "https"}.get(parsed.scheme, parsed.scheme) root = f"{scheme}://{parsed.netloc}".rstrip("/") for url in (f"{root}/json/version", f"{root}/json"): - try: - with urllib.request.urlopen(url, timeout=timeout) as resp: - if 200 <= getattr(resp, "status", 200) < 300: - return True - except Exception: - pass + with contextlib.suppress(Exception), urllib.request.urlopen(url, timeout=timeout) as resp: + if 200 <= getattr(resp, "status", 200) < 300: + return True return False -def _normalize_cdp_url(parsed) -> str: - # Concrete ``/devtools/browser/`` endpoints stay as-is; discovery-style inputs - # collapse to ``scheme://host:port`` so ``_resolve_cdp_override`` can append ``/json/version``. - if parsed.path.startswith("/devtools/browser/"): - return parsed.geturl() - return parsed._replace(path="", params="", query="", fragment="").geturl() - - -def _launch_failure_hints(port: int, system: str) -> list[str]: - from hermes_cli.browser_connect import manual_chrome_debug_command - - command = manual_chrome_debug_command(port, system) - hint = ( - ["Start a Chromium-family browser with remote debugging, then retry /browser connect:", command] - if command - else [ - "No supported Chromium-family browser executable was found in this environment.", - f"Install one or start a Chromium-family browser with --remote-debugging-port={port}, then retry /browser connect.", - ]) - return [ - *hint, - "Browser not connected — start a Chromium-family browser with remote debugging and retry /browser connect", - ] - - def _connect_local_default(port: int, system: str, announce) -> str | None: """Discover (or launch) the default local debug browser → CDP URL, or None after announcing.""" from hermes_cli.browser_connect import ( - discover_local_cdp_url, find_free_debug_port, launch_chrome_debug, local_port_in_use) + discover_local_cdp_url, find_free_debug_port, launch_chrome_debug, local_port_in_use, + manual_chrome_debug_command) # Dual-stack discovery: when another app squats the IPv4 loopback on the debug port, a # browser bound there comes up on [::1] only; an IPv4-only probe misses it AND hangs @@ -104,10 +63,9 @@ def _connect_local_default(port: int, system: str, announce) -> str | None: launch_port = port if local_port_in_use(port): launch_port = find_free_debug_port(port) - announce( - f"Port {port} is occupied by another application that isn't a CDP browser " - "(an IDE debugger or dev server may be using it) — launching a debug browser " - f"on port {launch_port} instead...") + announce(f"Port {port} is occupied by another application that isn't a CDP browser " + "(an IDE debugger or dev server may be using it) — launching a debug browser " + f"on port {launch_port} instead...") else: announce("Chromium-family browser isn't running with remote debugging — attempting to launch...") launch = launch_chrome_debug(launch_port, system) @@ -115,8 +73,7 @@ def _connect_local_default(port: int, system: str, announce) -> str | None: # Bounded wait: the whole connect must finish inside the client RPC timeout. deadline = time.monotonic() + 10.0 while time.monotonic() < deadline: - discovered = discover_local_cdp_url(launch_port, timeout=1.0) - if discovered: + if discovered := discover_local_cdp_url(launch_port, timeout=1.0): break time.sleep(0.5) if discovered: @@ -124,18 +81,23 @@ def _connect_local_default(port: int, system: str, announce) -> str | None: return discovered if launch.hint: announce(launch.hint, level="error") - for line in _launch_failure_hints(launch_port, system): + command = manual_chrome_debug_command(launch_port, system) + hints = ( + ["Start a Chromium-family browser with remote debugging, then retry /browser connect:", command] + if command else [ + "No supported Chromium-family browser executable was found in this environment.", + f"Install one or start a Chromium-family browser with --remote-debugging-port={launch_port}, then retry /browser connect."]) + hints.append("Browser not connected — start a Chromium-family browser with remote debugging and retry /browser connect") + for line in hints: announce(line, level="error") return None def _browser_connect(rid, params: dict) -> dict: import platform - from hermes_cli.browser_connect import DEFAULT_BROWSER_CDP_URL from tools.browser_tool import cleanup_all_browsers from urllib.parse import urlparse - raw_url = params.get("url") if raw_url is not None and not isinstance(raw_url, str): return _err(rid, 4015, f"browser url must be a string, got {type(raw_url).__name__}") @@ -144,10 +106,9 @@ def _browser_connect(rid, params: dict) -> dict: def announce(message: str, *, level: str = "info") -> None: messages.append(message) - # Without a session id the TUI prints `messages` from the response; an event would double-render. + # Without a session id the TUI prints `messages` from the response (an event would double-render). if sid: _emit("browser.progress", sid, {"message": message, "level": level}) - parsed = urlparse(url if "://" in url else f"http://{url}") if parsed.scheme not in _CDP_SCHEMES: return _err(rid, 4015, f"unsupported browser url: {url}") @@ -167,7 +128,6 @@ def _browser_connect(rid, params: dict) -> dict: # check TCP reachability only and let browser_navigate handshake. if parsed.scheme in {"ws", "wss"} and parsed.path.startswith("/devtools/browser/"): import socket - try: with socket.create_connection((parsed.hostname, port), timeout=2.0): pass @@ -182,7 +142,10 @@ def _browser_connect(rid, params: dict) -> dict: parsed = urlparse(url) elif not _cdp_http_reachable(parsed): return _err(rid, 5031, f"could not reach browser CDP at {url}") - normalized = _normalize_cdp_url(parsed) + # Concrete ``/devtools/browser/`` endpoints stay as-is; discovery-style inputs collapse + # to ``scheme://host:port`` so ``_resolve_cdp_override`` can append ``/json/version``. + normalized = (parsed.geturl() if parsed.path.startswith("/devtools/browser/") + else parsed._replace(path="", params="", query="", fragment="").geturl()) # Reap BEFORE publishing the new env (an in-flight tool call sees the old supervisor closed) # and AFTER (the default task's cached supervisor drains against the new URL). cleanup_all_browsers() @@ -190,10 +153,8 @@ def _browser_connect(rid, params: dict) -> dict: cleanup_all_browsers() except Exception as e: return _err(rid, 5031, str(e)) - payload: dict[str, object] = {"connected": True, "url": normalized} - if messages: - payload["messages"] = messages - return _ok(rid, payload) + return _ok(rid, {"connected": True, "url": normalized, + **({"messages": messages} if messages else {})}) def _browser_disconnect(rid) -> dict: @@ -201,7 +162,6 @@ def _browser_disconnect(rid) -> dict: def reap() -> None: with contextlib.suppress(Exception): from tools.browser_tool import cleanup_all_browsers - cleanup_all_browsers() reap() diff --git a/tui_gateway/methods_complete.py b/tui_gateway/methods_complete.py index 7f2869070c..328d3d9354 100644 --- a/tui_gateway/methods_complete.py +++ b/tui_gateway/methods_complete.py @@ -4,7 +4,6 @@ Rebound onto server.py's globals at install time (``method_ctx.bind_module``), s bodies reference server globals bare (``_ok``, ``_err``, ``_sessions``, ...). """ - from .method_ctx import HandlerRegistry, bind_module _registry = HandlerRegistry() @@ -13,15 +12,10 @@ _profile_scoped = _registry.profile_scoped _BUILTIN_AT_PREFIXES = frozenset({"file", "folder", "url", "git", "diff", "staged"}) _AT_DIRECTIVE_HINTS = [ - ("@diff", "git diff"), - ("@staged", "staged diff"), - ("@file:", "attach file"), - ("@folder:", "attach folder"), - ("@url:", "fetch url"), - ("@git:", "git log")] + ("@diff", "git diff"), ("@staged", "staged diff"), ("@file:", "attach file"), + ("@folder:", "attach folder"), ("@url:", "fetch url"), ("@git:", "git log")] _SLASH_EXTRAS = [ - ("/density", "Toggle compact display mode"), - ("/details", "Control agent detail visibility"), + ("/density", "Toggle compact display mode"), ("/details", "Control agent detail visibility"), ("/logs", "Show recent gateway log lines"), ("/mouse", "Set mouse tracking preset [on|off|toggle|wheel|buttons|all]")] @@ -30,6 +24,20 @@ def _item(text: str, meta: str, display: str | None = None) -> dict: return {"text": text, "display": display if display is not None else text, "meta": meta} +def _catch(fail_code: int): + """Handler body exceptions → ``_err(rid, fail_code, str(e))``.""" + + def deco(body): + def handler(rid, params: dict) -> dict: + try: + return body(rid, params) + except Exception as e: + return _err(rid, fail_code, str(e)) + handler.__doc__ = body.__doc__ + return handler + return deco + + @method("paste.collapse") def _(rid, params: dict) -> dict: global _paste_counter @@ -56,13 +64,11 @@ def _profile_mention_items(prefix: str) -> list[dict]: from hermes_cli.profiles import list_profiles seen: set[str] = set() for p in list_profiles(): - name = (p.name or "").strip() - if not name: + if not (name := (p.name or "").strip()): continue seen.add(name.lower()) - desc = (getattr(p, "description", "") or "").strip() if name.lower().startswith(prefix.lower()): - out.append(_item(f"@{name}", desc or "agent profile")) + out.append(_item(f"@{name}", (getattr(p, "description", "") or "").strip() or "agent profile")) if "hermes".startswith(prefix.lower()) and "hermes" not in seen: out.append(_item("@hermes", "agent profile (primary)")) except Exception: @@ -75,21 +81,18 @@ def _plugin_reference_items(pfx: str, qval: str) -> list[dict] | None: no provider owns ``pfx`` or it fails.""" try: from agent.context_references import get_context_reference_providers - prov = get_context_reference_providers().get(pfx) - if prov is None: - return None import asyncio + if (prov := get_context_reference_providers().get(pfx)) is None: + return None coro = prov.autocomplete(qval, limit=20) try: - loop = asyncio.get_running_loop() + asyncio.get_running_loop() except RuntimeError: - loop = None - if loop and loop.is_running(): + ac = asyncio.run(coro) + else: # already inside a running loop: run the coroutine on a side thread import concurrent.futures with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: ac = pool.submit(asyncio.run, coro).result() - else: - ac = asyncio.run(coro) return [{"text": f"@{pfx}:{it.text}", "display": it.display, "meta": it.meta} for it in ac] except Exception: return None @@ -105,19 +108,16 @@ def _fuzzy_basename_items(root: str, path_part: str, prefix_tag: str) -> list[di def _consider(rel: str, name: str, is_dir: bool) -> None: if rel in seen or (name.startswith(".") and not want_hidden): return - rank = _fuzzy_basename_rank(name, path_part) - if rank is not None: + if (rank := _fuzzy_basename_rank(name, path_part)) is not None: seen.add(rel) ranked.append((rank, rel, name, is_dir)) # Seed with root's immediate children: `_list_repo_files` is capped at _FUZZY_CACHE_MAX_FILES # and the non-git fallback walk can burn the whole budget on one deep subtree. - try: + with contextlib.suppress(OSError): for entry in os.listdir(root): if entry not in _FUZZY_FALLBACK_EXCLUDES: _consider(entry, entry, os.path.isdir(os.path.join(root, entry))) - except OSError: - pass for rel in _list_repo_files(root): _consider(rel, os.path.basename(rel), False) # Rank each ancestor dir too — a folder with no name-matching file inside is otherwise invisible. @@ -133,181 +133,141 @@ def _fuzzy_basename_items(root: str, path_part: str, prefix_tag: str) -> list[di return [ _item( f"@{'folder' if is_dir else tag}:{rel}{'/' if is_dir else ''}", - "dir" if is_dir else os.path.dirname(rel), - basename + ("/" if is_dir else "")) + "dir" if is_dir else os.path.dirname(rel), basename + ("/" if is_dir else "")) for _, rel, basename, is_dir in ranked[:30]] +def _at_root_items() -> list[dict]: + """Completions for a bare ``@``: directive hints, agent profiles, plugin ``@:`` providers.""" + items = [_item(t, m) for t, m in _AT_DIRECTIVE_HINTS] + _profile_mention_items("") + with contextlib.suppress(Exception): + from agent.context_references import get_context_reference_providers + for pfx, prov in sorted(get_context_reference_providers().items()): + items.append(_item(f"@{pfx}:", prov.description or f"plugin: {pfx}")) + return items + + +def _dir_listing_items(root: str, word: str, path_part: str, prefix_tag: str, is_context: bool) -> list[dict]: + """Prefix-match entries of the directory ``path_part`` points at (max 30).""" + expanded = _normalize_completion_path(path_part) if path_part else "." + if expanded == "." or not expanded or expanded.endswith("/"): + search_dir, match = (expanded or "."), "" + else: + search_dir, match = os.path.dirname(expanded) or ".", os.path.basename(expanded) + search_dir = search_dir if os.path.isabs(search_dir) else os.path.join(root, search_dir) + items: list[dict] = [] + if not os.path.isdir(search_dir): + return items + for entry in sorted(os.listdir(search_dir)): + if match and not entry.lower().startswith(match.lower()): + continue + if is_context and (entry in _FUZZY_FALLBACK_EXCLUDES or (not prefix_tag and entry.startswith("."))): + continue + full = os.path.join(search_dir, entry) + is_dir = os.path.isdir(full) + if prefix_tag and (prefix_tag == "folder") != is_dir: # explicit `@folder:`/`@file:` skip the other kind + continue + rel = os.path.relpath(full, root).replace(os.sep, "/") + suffix = "/" if is_dir else "" + if is_context: + text = f"@{prefix_tag or ('folder' if is_dir else 'file')}:{rel}{suffix}" + elif word.startswith("~"): + text = "~/" + os.path.relpath(full, os.path.expanduser("~")) + suffix + else: + text = ("./" if word.startswith("./") else "") + rel + suffix + items.append(_item(text, "dir" if is_dir else "", entry + suffix)) + if len(items) >= 30: + break + return items + + @method("complete.path") +@_catch(5021) def _(rid, params: dict) -> dict: word = params.get("word", "") if not word: return _ok(rid, {"items": []}) - items: list[dict] = [] - try: - root = _completion_cwd(params) - is_context = word.startswith("@") - query = word[1:] if is_context else word - if is_context and not query: - items = [_item(t, m) for t, m in _AT_DIRECTIVE_HINTS] - items.extend(_profile_mention_items("")) # `@` alone reveals agent profiles too - try: - from agent.context_references import get_context_reference_providers - for _pfx, _prov in sorted(get_context_reference_providers().items()): - items.append(_item(f"@{_pfx}:", _prov.description or f"plugin: {_pfx}")) - except Exception: - pass - return _ok(rid, {"items": items}) - - # Plugin `@:` runs before the built-in file/folder branching. - if is_context and ":" in query: - _pfx, _, _qval = query.partition(":") - if _pfx not in _BUILTIN_AT_PREFIXES: - plugin_items = _plugin_reference_items(_pfx, _qval) - if plugin_items is not None: - return _ok(rid, {"items": plugin_items}) - - # Bare `@folder` lists as soon as the keyword is typed (the static `@folder:` hint is not accepted). - if is_context and query in {"file", "folder"}: - prefix_tag, path_part = query, "" - elif is_context and query.startswith(("file:", "folder:")): - prefix_tag, _, path_part = query.partition(":") - else: - prefix_tag, path_part = "", query - - # `@/foo` usually means "foo, from here": absolute only when that prefix exists, - # else resolve relative to cwd (`@/Desktop` must not dead-end; `@/usr/local` still resolves). - if is_context and path_part.startswith("/") and not path_part.startswith("//"): - if not _abs_completion_prefix_exists(path_part): - path_part = path_part.lstrip("/") - if is_context and path_part and len(path_part.strip()) >= 2 and "/" not in path_part and prefix_tag != "folder": - items = _fuzzy_basename_items(root, path_part, prefix_tag) - if not prefix_tag: # bare `@name` may be an agent mention: profiles rank ABOVE file hits - items = _profile_mention_items(path_part) + items - return _ok(rid, {"items": items}) - expanded = _normalize_completion_path(path_part) if path_part else "." - if expanded == "." or not expanded: - search_dir, match = ".", "" - elif expanded.endswith("/"): - search_dir, match = expanded, "" - else: - search_dir = os.path.dirname(expanded) or "." - match = os.path.basename(expanded) - search_dir = search_dir if os.path.isabs(search_dir) else os.path.join(root, search_dir) - if not os.path.isdir(search_dir): - return _ok(rid, {"items": []}) - want_dir = prefix_tag == "folder" - match_lower = match.lower() - for entry in sorted(os.listdir(search_dir)): - if match and not entry.lower().startswith(match_lower): - continue - if is_context and entry in _FUZZY_FALLBACK_EXCLUDES: - continue - if is_context and not prefix_tag and entry.startswith("."): - continue - full = os.path.join(search_dir, entry) - is_dir = os.path.isdir(full) - # Explicit `@folder:` / `@file:` skip the opposite kind (never rewrite the tag). - if prefix_tag and want_dir != is_dir: - continue - rel = os.path.relpath(full, root).replace(os.sep, "/") - suffix = "/" if is_dir else "" - if is_context and prefix_tag: - text = f"@{prefix_tag}:{rel}{suffix}" - elif is_context: - text = f"@{'folder' if is_dir else 'file'}:{rel}{suffix}" - elif word.startswith("~"): - text = "~/" + os.path.relpath(full, os.path.expanduser("~")) + suffix - elif word.startswith("./"): - text = "./" + rel + suffix - else: - text = rel + suffix - items.append(_item(text, "dir" if is_dir else "", entry + suffix)) - if len(items) >= 30: - break - except Exception as e: - return _err(rid, 5021, str(e)) - - # Bare-word `@name` (incl. single chars, which skip the fuzzy branch): profiles rank above paths. - try: - if is_context and not prefix_tag and path_part and "/" not in path_part: + root = _completion_cwd(params) + is_context = word.startswith("@") + query = word[1:] if is_context else word + if is_context and not query: + return _ok(rid, {"items": _at_root_items()}) + # Plugin `@:` runs before the built-in file/folder branching. + if is_context and ":" in query: + pfx, _, qval = query.partition(":") + if pfx not in _BUILTIN_AT_PREFIXES and (plugin_items := _plugin_reference_items(pfx, qval)) is not None: + return _ok(rid, {"items": plugin_items}) + # Bare `@folder` lists as soon as the keyword is typed (the static `@folder:` hint is not accepted). + if is_context and (query in {"file", "folder"} or query.startswith(("file:", "folder:"))): + prefix_tag, _, path_part = query.partition(":") + else: + prefix_tag, path_part = "", query + # `@/foo` usually means "foo, from here": absolute only when that prefix exists, + # else resolve relative to cwd (`@/Desktop` must not dead-end; `@/usr/local` still resolves). + if ( + is_context and path_part.startswith("/") and not path_part.startswith("//") + and not _abs_completion_prefix_exists(path_part)): + path_part = path_part.lstrip("/") + bare_word = is_context and path_part and "/" not in path_part + if bare_word and len(path_part.strip()) >= 2 and prefix_tag != "folder": + items = _fuzzy_basename_items(root, path_part, prefix_tag) + else: + items = _dir_listing_items(root, word, path_part, prefix_tag, is_context) + # Bare-word `@name` may be an agent mention: profiles rank ABOVE file hits. + if bare_word and not prefix_tag: + with contextlib.suppress(Exception): items = _profile_mention_items(path_part) + items - except Exception: - pass return _ok(rid, {"items": items}) @method("complete.slash") +@_catch(5020) def _(rid, params: dict) -> dict: text = params.get("text", "") if not text.startswith("/"): return _ok(rid, {"items": []}) - try: - from hermes_cli.commands import SlashCommandCompleter - from prompt_toolkit.document import Document - from prompt_toolkit.formatted_text import to_plain_text - from agent.skill_commands import get_skill_commands - from agent.skill_bundles import get_skill_bundles - completer = SlashCommandCompleter( - skill_commands_provider=lambda: get_skill_commands(), skill_bundles_provider=lambda: get_skill_bundles() - ) - # `kind` reaches the TUI as data (from the providers, not sniffed from ⚡/▣ glyphs): - # skills/bundles are the only completions for an inline `/skill` typed mid-message. - skill_names = {key.lstrip("/").lower() for key in (*get_skill_commands(), *get_skill_bundles())} + from hermes_cli.commands import SlashCommandCompleter + from prompt_toolkit.document import Document + from prompt_toolkit.formatted_text import to_plain_text + from agent.skill_commands import get_skill_commands + from agent.skill_bundles import get_skill_bundles + completer = SlashCommandCompleter( + skill_commands_provider=lambda: get_skill_commands(), skill_bundles_provider=lambda: get_skill_bundles()) + # `kind` reaches the TUI as data (from the providers, not sniffed from ⚡/▣ glyphs): + # skills/bundles are the only completions for an inline `/skill` typed mid-message. + skill_names = {key.lstrip("/").lower() for key in (*get_skill_commands(), *get_skill_bundles())} - def to_items(doc: Document) -> list[dict]: - # display/display_meta are FormattedText; the TUI contract is a plain string - # (the raw list trips Ink's row layout into 1-char truncation). - return [ - { - "text": c.text, - "display": to_plain_text(c.display) if c.display else c.text, - "meta": to_plain_text(c.display_meta) if c.display_meta else "", - "kind": "skill" if c.text.strip().lstrip("/").lower() in skill_names else "command", - } - for c in completer.get_completions(doc, None)] - items = to_items(Document(text, len(text))) - - # Rank + bound while a `/token` is under the cursor (the one stage skills are - # offered at); an argument stage (`/personality `) keeps its command's order. - if text.rsplit(" ", 1)[-1].startswith("/"): - score_of = None - # Command-token stage: the completer only emits name-prefix matches, so merge in - # catalog entries whose name SUBSTRING or DESCRIPTION words match (name outranks description). - if " " not in text and len(text) > 1: - from tui_gateway.slash_fuzzy import fuzzy_rank_slash_items, normalize_slash_search_query - items, score_of = fuzzy_rank_slash_items( - items, to_items(Document("/", 1)), normalize_slash_search_query(text)) - usage, origin_of = _skill_usage_lookup() - items = _rank_slash_completions(items, usage, origin_of, browsing=text == "/", score_of=score_of) - else: - items = items[:_SLASH_COMPLETION_LIMIT] - text_lower = text.lower() - for extra_text, extra_meta in _SLASH_EXTRAS: - if extra_text.startswith(text_lower) and not any(item["text"] == extra_text for item in items): - items.append({"text": extra_text, "display": extra_text, "meta": extra_meta, "kind": "command"}) - details_items = _details_completions(text) - if details_items is not None: - return _ok(rid, {"items": details_items, "replace_from": text.rfind(" ") + 1 if " " in text else len(text)}) - return _ok(rid, {"items": items, "replace_from": text.rfind(" ") + 1 if " " in text else 1}) - except Exception as e: - return _err(rid, 5020, str(e)) - - -def _catch(fail_code: int): - """Handler body exceptions → ``_err(rid, fail_code, str(e))``.""" - - def deco(body): - def handler(rid, params: dict) -> dict: - try: - return body(rid, params) - except Exception as e: - return _err(rid, fail_code, str(e)) - - handler.__doc__ = body.__doc__ - return handler - - return deco + def to_items(doc: Document) -> list[dict]: + # display/display_meta are FormattedText; the TUI contract is a plain string + # (the raw list trips Ink's row layout into 1-char truncation). + return [ + { + "text": c.text, "display": to_plain_text(c.display) if c.display else c.text, + "meta": to_plain_text(c.display_meta) if c.display_meta else "", + "kind": "skill" if c.text.strip().lstrip("/").lower() in skill_names else "command"} + for c in completer.get_completions(doc, None)] + items = to_items(Document(text, len(text))) + # Rank + bound while a `/token` is under the cursor (the one stage skills are + # offered at); an argument stage (`/personality `) keeps its command's order. + if text.rsplit(" ", 1)[-1].startswith("/"): + score_of = None + # Command-token stage: the completer only emits name-prefix matches, so merge in + # catalog entries whose name SUBSTRING or DESCRIPTION words match (name outranks description). + if " " not in text and len(text) > 1: + from tui_gateway.slash_fuzzy import fuzzy_rank_slash_items, normalize_slash_search_query + items, score_of = fuzzy_rank_slash_items( + items, to_items(Document("/", 1)), normalize_slash_search_query(text)) + usage, origin_of = _skill_usage_lookup() + items = _rank_slash_completions(items, usage, origin_of, browsing=text == "/", score_of=score_of) + else: + items = items[:_SLASH_COMPLETION_LIMIT] + text_lower = text.lower() + for extra_text, extra_meta in _SLASH_EXTRAS: + if extra_text.startswith(text_lower) and not any(item["text"] == extra_text for item in items): + items.append({**_item(extra_text, extra_meta), "kind": "command"}) + if (details_items := _details_completions(text)) is not None: + return _ok(rid, {"items": details_items, "replace_from": text.rfind(" ") + 1 if " " in text else len(text)}) + return _ok(rid, {"items": items, "replace_from": text.rfind(" ") + 1 if " " in text else 1}) def _session_agent(params: dict): @@ -322,13 +282,9 @@ def _(rid, params: dict) -> dict: from hermes_cli.inventory import build_model_options_payload # A spawned agent owns the live provider/model/base_url; empty attributes must # NOT clobber disk config (with_overrides is truthy-only). - ctx = _model_picker_context(_session_agent(params)) - payload = build_model_options_payload( - ctx, - explicit_only=bool(params.get("explicit_only")), - include_unconfigured=bool(params.get("include_unconfigured")), - refresh=bool(params.get("refresh"))) - return _ok(rid, payload) + return _ok(rid, build_model_options_payload( + _model_picker_context(_session_agent(params)), explicit_only=bool(params.get("explicit_only")), + include_unconfigured=bool(params.get("include_unconfigured")), refresh=bool(params.get("refresh")))) @method("model.save_key") @@ -337,26 +293,23 @@ def _(rid, params: dict) -> dict: """Save an API key for ``slug``; return its refreshed provider row (model.options shape + ``authenticated``).""" from hermes_cli.auth import PROVIDER_REGISTRY from hermes_cli.config import is_managed - from hermes_cli.inventory import build_models_payload - slug = (params.get("slug") or "").strip() - api_key = (params.get("api_key") or "").strip() + slug, api_key = (params.get("slug") or "").strip(), (params.get("api_key") or "").strip() if not slug or not api_key: return _err(rid, 4001, "slug and api_key are required") if is_managed(): return _err(rid, 4006, "managed install — credentials are read-only") - pconfig = PROVIDER_REGISTRY.get(slug) - if not pconfig: + if not (pconfig := PROVIDER_REGISTRY.get(slug)): return _err(rid, 4002, f"unknown provider: {slug}") if pconfig.auth_type != "api_key": return _err(rid, 4003, f"{pconfig.name} uses {pconfig.auth_type} auth — run `hermes model` to configure") if not pconfig.api_key_env_vars: return _err(rid, 4004, f"no env var defined for {pconfig.name}") - # Unified lifecycle rotates stale config.yaml mirrors of the old key too. env_var = pconfig.api_key_env_vars[0] - from hermes_cli.credential_lifecycle import save_provider_env_credential + from hermes_cli.credential_lifecycle import save_provider_env_credential # also rotates stale config.yaml mirrors save_provider_env_credential(env_var, api_key) os.environ[env_var] = api_key # so the refreshed inventory sees it # Shared inventory builder (lock-step with model.options / dashboard); picker_hints carries `authenticated`. + from hermes_cli.inventory import build_models_payload payload = build_models_payload(_model_picker_context(_session_agent(params)), picker_hints=True, max_models=50) provider_data = next((p for p in payload["providers"] if p["slug"] == slug), None) if provider_data is None: # key saved but provider didn't appear — still success @@ -371,16 +324,13 @@ def _(rid, params: dict) -> dict: """Remove all credentials (env keys AND OAuth/pool state) for provider ``slug``.""" from hermes_cli.auth import PROVIDER_REGISTRY, clear_provider_auth from hermes_cli.credential_lifecycle import remove_provider_env_credential - slug = (params.get("slug") or "").strip() - if not slug: + if not (slug := (params.get("slug") or "").strip()): return _err(rid, 4001, "slug is required") pconfig = PROVIDER_REGISTRY.get(slug) - # Remove EVERY env var plus its mirrors (env-seeded pool entries, model cache rows, - # value-matched config.yaml copies) or the provider resurrects in the picker after restart. + # Remove EVERY env var plus its mirrors or the provider resurrects in the picker after restart. env_vars = (pconfig.api_key_env_vars if pconfig else None) or () cleared_env = any([remove_provider_env_credential(ev).get("found") for ev in env_vars]) - # Full disconnect: removing OAuth grants is intended here, unlike key-only deletes. - cleared_auth = clear_provider_auth(slug) + cleared_auth = clear_provider_auth(slug) # full disconnect: OAuth grants go too if not cleared_env and not cleared_auth: return _err(rid, 4005, f"no credentials found for {slug}") return _ok(rid, {"slug": slug, "name": pconfig.name if pconfig else slug, "disconnected": True}) diff --git a/tui_gateway/methods_complete_helpers.py b/tui_gateway/methods_complete_helpers.py index b6631ab069..b4ea31533a 100644 --- a/tui_gateway/methods_complete_helpers.py +++ b/tui_gateway/methods_complete_helpers.py @@ -12,9 +12,6 @@ from .method_ctx import HandlerRegistry, bind_module _registry = HandlerRegistry() - -# ── Methods: complete ───────────────────────────────────────────────── - _FUZZY_CACHE_TTL_S = 5.0 _FUZZY_CACHE_MAX_FILES = 20000 _FUZZY_FALLBACK_EXCLUDES = frozenset( @@ -24,54 +21,54 @@ _fuzzy_cache_lock = threading.Lock() _fuzzy_cache: dict[str, tuple[float, list[str]]] = {} +def _git_repo_files(root: str): + """Yield ``git ls-files`` paths (tracked + untracked) relative to ``root``; empty outside a + repo or on git failure/timeout. Entries above ``root`` are skipped (Cmd-P workspace scope).""" + from hermes_cli._subprocess_compat import windows_hide_flags + run_kw = dict(capture_output=True, timeout=2.0, check=False, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags()) + try: + top_result = subprocess.run(["git", "-C", root, "rev-parse", "--show-toplevel"], **run_kw) + if top_result.returncode != 0: + return + top = top_result.stdout.decode("utf-8", "replace").strip() + list_result = subprocess.run( + ["git", "-C", top, "ls-files", "-z", "--cached", "--others", "--exclude-standard"], **run_kw) + if list_result.returncode != 0: + return + except (OSError, subprocess.TimeoutExpired): + return + for p in list_result.stdout.decode("utf-8", "replace").split("\0"): + if p: + rel = os.path.relpath(os.path.join(top, p), root).replace(os.sep, "/") + if not rel.startswith("../"): + yield rel + + +def _walk_repo_files(root: str): + """Non-git fallback: ``os.walk`` skipping vendor/build dirs + dot-dirs; dotfiles survive + (the ranker decides based on whether the query starts with `.`).""" + try: + for dirpath, dirnames, filenames in os.walk(root, followlinks=False): + dirnames[:] = [d for d in dirnames if d not in _FUZZY_FALLBACK_EXCLUDES and not d.startswith(".")] + rel_dir = os.path.relpath(dirpath, root) + for f in filenames: + yield (f if rel_dir == "." else f"{rel_dir}/{f}").replace(os.sep, "/") + except OSError: + return + + def _list_repo_files(root: str) -> list[str]: - """File paths relative to ``root`` (tracked + untracked via ``git ls-files`` from the - repo top; files outside ``root`` excluded so the picker stays Cmd-P scoped). Falls - back to a bounded ``os.walk(root)`` outside a git repo. Cached per-root for + """File paths relative to ``root`` (git listing, else a bounded walk), cached per-root for ``_FUZZY_CACHE_TTL_S`` so rapid keystrokes don't respawn git.""" now = time.monotonic() with _fuzzy_cache_lock: cached = _fuzzy_cache.get(root) if cached and now - cached[0] < _FUZZY_CACHE_TTL_S: return cached[1] - files: list[str] = [] - from hermes_cli._subprocess_compat import windows_hide_flags - run_kw = dict(capture_output=True, timeout=2.0, check=False, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags()) - try: - top_result = subprocess.run(["git", "-C", root, "rev-parse", "--show-toplevel"], **run_kw) - if top_result.returncode == 0: - top = top_result.stdout.decode("utf-8", "replace").strip() - list_result = subprocess.run( - ["git", "-C", top, "ls-files", "-z", "--cached", "--others", "--exclude-standard"], **run_kw - ) - if list_result.returncode == 0: - for p in list_result.stdout.decode("utf-8", "replace").split("\0"): - if not p: - continue - rel = os.path.relpath(os.path.join(top, p), root).replace(os.sep, "/") - if rel.startswith("../"): # parents/siblings of cwd: keep Cmd-P workspace scope - continue - files.append(rel) - if len(files) >= _FUZZY_CACHE_MAX_FILES: - break - except (OSError, subprocess.TimeoutExpired): - pass + from itertools import islice + files = list(islice(_git_repo_files(root), _FUZZY_CACHE_MAX_FILES)) if not files: - # Fallback walk skips vendor/build dirs + dot-dirs; dotfiles survive (the ranker - # decides based on whether the query starts with `.`). - try: - for dirpath, dirnames, filenames in os.walk(root, followlinks=False): - dirnames[:] = [d for d in dirnames if d not in _FUZZY_FALLBACK_EXCLUDES and not d.startswith(".")] - rel_dir = os.path.relpath(dirpath, root) - for f in filenames: - rel = f if rel_dir == "." else f"{rel_dir}/{f}" - files.append(rel.replace(os.sep, "/")) - if len(files) >= _FUZZY_CACHE_MAX_FILES: - break - if len(files) >= _FUZZY_CACHE_MAX_FILES: - break - except OSError: - pass + files = list(islice(_walk_repo_files(root), _FUZZY_CACHE_MAX_FILES)) with _fuzzy_cache_lock: _fuzzy_cache[root] = (now, files) return files @@ -83,44 +80,32 @@ def _fuzzy_basename_rank(name: str, query: str) -> tuple[int, int] | None: · 3 substring · 4 subsequence (query chars appear in order).""" if not query: return (3, len(name)) - nl = name.lower() - ql = query.lower() + nl, ql = name.lower(), query.lower() if nl == ql: return (0, len(name)) if nl.startswith(ql): return (1, len(name)) - - # Split on -_. and camelCase (`appChrome` → ["app","Chrome"]); cheap approximation, - # falls through to substring/subsequence if it misses. + # Word boundaries: split on -_. and camelCase (`appChrome` → ["app","Chrome"]); cheap + # approximation, falls through to substring/subsequence if it misses. parts: list[str] = [] buf = "" for ch in name: if ch in "-_." or (ch.isupper() and buf and not buf[-1].isupper()): - if buf: - parts.append(buf) + parts += [buf] if buf else [] buf = ch if ch not in "-_." else "" else: buf += ch - if buf: - parts.append(buf) - for p in parts: - if p.lower().startswith(ql): - return (2, len(name)) + if any(p.lower().startswith(ql) for p in parts + ([buf] if buf else [])): + return (2, len(name)) if ql in nl: return (3, len(name)) - i = 0 - for ch in nl: - if ch == ql[i]: - i += 1 - if i == len(ql): - return (4, len(name)) - return None + it = iter(nl) + return (4, len(name)) if all(any(c == q for c in it) for q in ql) else None def _abs_completion_prefix_exists(path_part: str) -> bool: - """True when ``path_part`` reads sensibly as an absolute path: the parent dir exists - and a partially-typed final segment matches at least one entry. Decides whether - `@/foo` is the absolute `/foo` or shorthand for `foo` under the cwd.""" + """True when ``path_part`` reads sensibly as an absolute path (parent exists and a + partially-typed final segment matches an entry): decides `@/foo` = `/foo` vs cwd `foo`.""" expanded = _normalize_completion_path(path_part) parent = os.path.dirname(expanded.rstrip("/")) or "/" tail = os.path.basename(expanded.rstrip("/")) @@ -135,14 +120,6 @@ def _abs_completion_prefix_exists(path_part: str) -> bool: return False -def _details_completion_item(value: str, meta: str = "") -> dict: - return {"text": value, "display": value, "meta": meta} - - -def _details_root_completion_item(value: str, meta: str, needs_leading_space: bool) -> dict: - return _details_completion_item(f" {value}" if needs_leading_space else value, meta) - - _DETAILS_SECTIONS = ("thinking", "tools", "subagents", "activity") _DETAILS_MODES = ("hidden", "collapsed", "expanded") @@ -154,39 +131,34 @@ def _details_root_meta(candidate: str) -> str: def _details_completions(text: str) -> list[dict] | None: + """Argument completions for ``/details [section] [mode]``; None when ``text`` is not that command.""" if not text.lower().startswith("/details"): return None stripped = text.strip() if stripped and not "/details".startswith(stripped.lower().split()[0]): return None - body = text[len("/details") :] - if body.startswith(" "): - body = body[1:] + body = text[len("/details") :].removeprefix(" ") parts = body.split() - has_trailing_space = text.endswith(" ") - sections, modes = _DETAILS_SECTIONS, _DETAILS_MODES - root_candidates = (*modes, "cycle", *sections) - if not body or (len(parts) == 0 and has_trailing_space): - return [_details_root_completion_item(c, _details_root_meta(c), not has_trailing_space) for c in root_candidates] - if len(parts) == 1 and not has_trailing_space: + trailing = text.endswith(" ") + root_candidates = (*_DETAILS_MODES, "cycle", *_DETAILS_SECTIONS) + if not body or (not parts and trailing): + lead = "" if trailing else " " + return [_item(f"{lead}{c}", _details_root_meta(c)) for c in root_candidates] + if len(parts) == 1 and not trailing: prefix = parts[0].lower() - return [ - _details_completion_item(c, _details_root_meta(c)) - for c in root_candidates - if c.startswith(prefix) and c != prefix] + return [_item(c, _details_root_meta(c)) for c in root_candidates if c.startswith(prefix) and c != prefix] section = parts[0].lower() if parts else "" - if section not in sections: + if section not in _DETAILS_SECTIONS: return [] def section_meta(candidate: str) -> str: return f"clear {section} override" if candidate == "reset" else f"set {section}" - if len(parts) == 1 and has_trailing_space: - return [_details_completion_item(c, section_meta(c)) for c in (*modes, "reset")] - if len(parts) == 2 and not has_trailing_space: + mode_candidates = (*_DETAILS_MODES, "reset") + if len(parts) == 1: # trailing space after the section + return [_item(c, section_meta(c)) for c in mode_candidates] + if len(parts) == 2 and not trailing: prefix = parts[1].lower() - return [ - _details_completion_item(c, section_meta(c)) for c in (*modes, "reset") if c.startswith(prefix) and c != prefix - ] + return [_item(c, section_meta(c)) for c in mode_candidates if c.startswith(prefix) and c != prefix] return [] @@ -194,22 +166,16 @@ def _model_picker_context(agent): """Layer live session state onto config without losing custom identity.""" from hermes_cli.inventory import load_picker_context ctx = load_picker_context() - provider = getattr(agent, "provider", "") if agent else "" - base_url = getattr(agent, "base_url", "") if agent else "" - model = getattr(agent, "model", "") if agent else "" + provider, base_url, model = (getattr(agent, k, "") if agent else "" for k in ("provider", "base_url", "model")) if str(provider or "").strip().lower() == "custom": try: from hermes_cli.runtime_provider import canonical_custom_identity - provider = ( - canonical_custom_identity( - base_url=base_url or None, config_provider=ctx.current_provider, model=model or None - ) - or provider) + provider = canonical_custom_identity( + base_url=base_url or None, config_provider=ctx.current_provider, model=model or None) or provider except Exception: logger.debug("custom provider identity recovery failed (model picker)", exc_info=True) return ctx.with_overrides( - current_provider=provider, current_model=model or _resolve_model(), current_base_url=base_url - ) + current_provider=provider, current_model=model or _resolve_model(), current_base_url=base_url) def register(server) -> None: diff --git a/tui_gateway/methods_config.py b/tui_gateway/methods_config.py index 140c7ab837..f469215435 100644 --- a/tui_gateway/methods_config.py +++ b/tui_gateway/methods_config.py @@ -5,6 +5,7 @@ from .method_ctx import HandlerRegistry, bind_module from hermes_constants import DEFAULT_INDICATOR_STYLE, INDICATOR_STYLES +from hermes_constants import display_hermes_home as _display_hermes_home _registry = HandlerRegistry() method = _registry.method @@ -24,8 +25,8 @@ def _projects_handler(name: str): def _reconcile_repo_discovery(pdb, conn, policy, policy_key): - pdb.reconcile_discovered_repos_policy( - conn, policy_key, preserve_unversioned=_repo_discovery_policy_is_default(policy)) + pdb.reconcile_discovered_repos_policy(conn, policy_key, + preserve_unversioned=_repo_discovery_policy_is_default(policy)) @_projects_handler("projects.discover_repos") @@ -38,8 +39,8 @@ def _(rid, params: dict) -> dict: policy = _repo_discovery_policy() with pdb.connect_closing() as conn: _reconcile_repo_discovery(pdb, conn, policy, _repo_discovery_policy_key(policy)) - # `scan=true` (remote-gateway desktop): its native scan only sees its own - # filesystem, so the host scans the policy roots so zero-session repos surface. + # `scan=true` (remote-gateway desktop): its native scan only sees its own filesystem, + # so the host scans the policy roots so zero-session repos surface. if params.get("scan") and policy["enabled"]: _scan_discovered_repos_remote(conn, policy) repos = _discover_repos_payload(db, conn=conn, include_cached=policy["enabled"]) @@ -52,29 +53,23 @@ def _(rid, params: dict) -> dict: from hermes_cli import projects_db as pdb policy = _repo_discovery_policy() policy_key = _repo_discovery_policy_key(policy) - incoming_raw = params.get("discovery_policy") - incoming_policy = ( - _repo_discovery_policy(incoming_raw) if isinstance(incoming_raw, dict) else None) - incoming_matches = (incoming_policy is not None - and _repo_discovery_policy_key(incoming_policy) == policy_key) - accept_legacy_default = (incoming_policy is None - and _repo_discovery_policy_is_default(policy)) - pairs: list[tuple[str, str | None]] = [] - for item in params.get("repos") or []: - if isinstance(item, str): - pairs.append((item, None)) - elif isinstance(item, dict) and item.get("root"): - pairs.append((str(item["root"]), item.get("label"))) + incoming = params.get("discovery_policy") + if isinstance(incoming, dict): + accepted = _repo_discovery_policy_key(_repo_discovery_policy(incoming)) == policy_key + else: + accepted = _repo_discovery_policy_is_default(policy) # legacy client without a policy + accepted = bool(policy["enabled"] and accepted) + pairs = [(item, None) if isinstance(item, str) else (str(item["root"]), item.get("label")) + for item in params.get("repos") or [] + if isinstance(item, str) or (isinstance(item, dict) and item.get("root"))] with pdb.connect_closing() as conn: _reconcile_repo_discovery(pdb, conn, policy, policy_key) - accepted = bool(policy["enabled"] and (incoming_matches or accept_legacy_default)) if accepted: pdb.record_discovered_repos(conn, pairs, replace=True, policy_key=policy_key) elif not policy["enabled"]: pdb.clear_discovered_repos(conn, policy_key=policy_key) with _profile_db(params) as db: - repos = ([] if db is None - else _discover_repos_payload(db, include_cached=policy["enabled"])) + repos = [] if db is None else _discover_repos_payload(db, include_cached=policy["enabled"]) return _ok(rid, {"repos": repos, "accepted": accepted, "discovery_policy": policy}) @@ -88,18 +83,17 @@ def _stamped_project_tree(db, params, **kwargs): @_projects_handler("projects.tree") def _(rid, params: dict) -> dict: - """Project -> repo -> lane overview with counts + a few preview sessions per project, plus - the flat set of session ids claimed by any project (excluded from flat Recents). Lanes carry - no session rows here; drill-in uses ``projects.project_sessions``.""" + """Project -> repo -> lane overview with counts + a few preview sessions per project, plus the + flat set of session ids claimed by any project (excluded from flat Recents). Lanes carry no + session rows; drill-in uses ``projects.project_sessions``.""" with _profile_db(params) as db: if db is None: return _ok(rid, {"projects": [], "active_id": None, "scoped_session_ids": []}) tree, active_id = _stamped_project_tree( db, params, preview_limit=int(params.get("preview_limit") or 3), hydrate=False, session_limit=int(params.get("session_limit") or 2000), include_discovered=True) - return _ok(rid, { - "projects": tree["projects"], "active_id": active_id, - "scoped_session_ids": tree["scoped_session_ids"]}) + return _ok(rid, {"projects": tree["projects"], "active_id": active_id, + "scoped_session_ids": tree["scoped_session_ids"]}) @_projects_handler("projects.project_sessions") @@ -115,68 +109,50 @@ def _(rid, params: dict) -> dict: tree, _active = _stamped_project_tree( db, params, preview_limit=0, hydrate=True, session_limit=int(params.get("session_limit") or 5000), include_discovered=False) - proj = next((p for p in tree["projects"] if p["id"] == project_id), None) - return _ok(rid, {"project": proj}) + return _ok(rid, {"project": next((p for p in tree["projects"] if p["id"] == project_id), None)}) -# ── config.get — one getter per key; returns the result payload or a full ``_err`` response -# (dicts containing "error" pass through untouched). +# ── config.get — one getter per key returning the result payload. + +def _display_raw() -> dict: + return _load_cfg().get("display") or {} -def _display_mode(cfg: dict, key: str, allowed: frozenset, default: str) -> str: - raw = str((cfg.get("display") or {}).get(key, default) or default).strip().lower() +def _display_word(key: str, default: str, allowed) -> str: + """Normalised ``display.``; unknown/garbage values read back as ``default``.""" + raw = str(_display_raw().get(key, default) or "").strip().lower() return raw if raw in allowed else default _THINKING_MODES = frozenset({"collapsed", "truncated", "full"}) -def _cfg_get_provider(rid, params): - try: - from hermes_cli.models import list_available_providers, normalize_provider - model = _resolve_model() - parts = model.split("/", 1) - return { - "model": model, - "provider": normalize_provider(parts[0]) if len(parts) > 1 else "unknown", +def _cfg_get_provider(params): + from hermes_cli.models import list_available_providers, normalize_provider + model = _resolve_model() + parts = model.split("/", 1) + return {"model": model, "provider": normalize_provider(parts[0]) if len(parts) > 1 else "unknown", "providers": list_available_providers()} - except Exception as e: - return _err(rid, 5013, str(e)) -def _cfg_get_profile(rid, params): - from hermes_constants import display_hermes_home - return {"home": str(_hermes_home), "display": display_hermes_home()} - - -def _cfg_get_project(rid, params): - cfg_terminal = _load_cfg().get("terminal") or {} - raw = str(params.get("cwd", "") or cfg_terminal.get("cwd", "") or "").strip() +def _cfg_get_project(params): + raw = str(params.get("cwd", "") or (_load_cfg().get("terminal") or {}).get("cwd", "") or "").strip() cwd = _completion_cwd({"cwd": raw} if raw else {}) return {"cwd": cwd, "branch": _git_branch_for_cwd(cwd)} -def _cfg_get_indicator(rid, params): - # Normalize so a hand-edited config.yaml (stray casing / unknown value) reads back the SAME - # value the TUI rendered (frontend falls back to DEFAULT_INDICATOR_STYLE for the same inputs). - norm = str((_load_cfg().get("display") or {}).get("tui_status_indicator", "")).strip().lower() - return {"value": norm if norm in INDICATOR_STYLES else DEFAULT_INDICATOR_STYLE} - - -def _cfg_get_personality(rid, params): +def _cfg_get_personality(params): # EFFECTIVE personality via the single owner — a stale/unknown name must not show as active. from hermes_cli.personality import active_personality_name return {"value": active_personality_name(_load_cfg()) or "none"} -def _cfg_get_reasoning(rid, params): +def _cfg_get_reasoning(params): cfg = _load_cfg() - session = _sessions.get(params.get("session_id", "")) - reasoning_config = None - if session is not None: - reasoning_config = session.get("create_reasoning_override") - if not isinstance(reasoning_config, dict): - reasoning_config = getattr(session.get("agent"), "reasoning_config", None) + session = _sessions.get(params.get("session_id", "")) or {} + reasoning_config = session.get("create_reasoning_override") + if session and not isinstance(reasoning_config, dict): + reasoning_config = getattr(session.get("agent"), "reasoning_config", None) if isinstance(reasoning_config, dict): enabled = reasoning_config.get("enabled") is not False effort = str(reasoning_config.get("effort") or "medium") if enabled else "none" @@ -184,48 +160,30 @@ def _cfg_get_reasoning(rid, params): raw_effort = (cfg.get("agent") or {}).get("reasoning_effort", "") # YAML `reasoning_effort: false` means thinking disabled, not "unset". effort = "none" if raw_effort is False else str(raw_effort or "medium") - display = "show" if bool((cfg.get("display") or {}).get("show_reasoning", True)) else "hide" + display = "show" if (cfg.get("display") or {}).get("show_reasoning", True) else "hide" return {"value": effort, "display": display} -def _cfg_get_fast(rid, params): +def _cfg_get_fast(params): # `config.set fast` is session-scoped: prefer the session's live/pinned value over the # global key (a pre-build session keeps its pin in create_service_tier_override). - session = _sessions.get(params.get("session_id", "")) - tier = None - if session is not None: - agent = session.get("agent") - if agent is not None: - tier = getattr(agent, "service_tier", None) - elif session.get("create_service_tier_override") is not None: - tier = session["create_service_tier_override"] + session = _sessions.get(params.get("session_id", "")) or {} + agent = session.get("agent") + tier = (getattr(agent, "service_tier", None) if agent is not None + else session.get("create_service_tier_override")) if tier is None: tier = _load_service_tier() return {"value": "fast" if tier == "priority" else "normal"} -def _cfg_get_approval_mode(rid, params): - try: - return {"value": _load_approval_mode()} - except Exception as e: - return _err(rid, 5001, str(e)) +def _cfg_get_thinking_mode(params): + raw = _display_word("thinking_mode", "", _THINKING_MODES) + if not raw: # legacy details_mode fallback + raw = "full" if _display_word("details_mode", "collapsed", _DETAIL_MODES) == "expanded" else "collapsed" + return {"value": raw} -def _cfg_get_thinking_mode(rid, params): - cfg = _load_cfg() - raw = str((cfg.get("display") or {}).get("thinking_mode", "") or "").strip().lower() - if raw in _THINKING_MODES: - return {"value": raw} - dm = _display_mode(cfg, "details_mode", _DETAIL_MODES, "collapsed") - return {"value": "full" if dm == "expanded" else "collapsed"} - - -def _cfg_get_theme(rid, params): - raw = str(_display_cfg().get("tui_theme", "auto")).strip().lower() - return {"value": raw if raw in {"auto", "light", "dark"} else "auto"} - - -def _cfg_get_mtime(rid, params): +def _cfg_get_mtime(params): cfg_path = _hermes_home / "config.yaml" try: mtime = cfg_path.stat().st_mtime if cfg_path.exists() else 0 @@ -236,82 +194,72 @@ def _cfg_get_mtime(rid, params): return {"mtime": mtime, "mcp_rev": _compute_mcp_rev()} -def _config_getters() -> dict: - """key -> getter(rid, params). Built per call so, once rebound onto server.py, every entry - resolves to the rebound helper copies rather than this module's originals.""" - return { - "provider": _cfg_get_provider, - "profile": _cfg_get_profile, - "project": _cfg_get_project, - "full": lambda rid, params: {"config": _load_cfg()}, - "prompt": lambda rid, params: {"prompt": _load_cfg().get("custom_prompt", "")}, - "skin": lambda rid, params: {"value": (_load_cfg().get("display") or {}).get("skin", "default")}, - "indicator": _cfg_get_indicator, - "personality": _cfg_get_personality, - "reasoning": _cfg_get_reasoning, - "fast": _cfg_get_fast, - "busy": lambda rid, params: {"value": _load_busy_input_mode()}, - "approval_mode": _cfg_get_approval_mode, - "approvals.mode": _cfg_get_approval_mode, - "details_mode": lambda rid, params: { - "value": _display_mode(_load_cfg(), "details_mode", _DETAIL_MODES, "collapsed")}, - "thinking_mode": _cfg_get_thinking_mode, - "density": lambda rid, params: { - "value": "on" if bool((_load_cfg().get("display") or {}).get("tui_compact", False)) else "off" - }, - "theme": _cfg_get_theme, - "statusbar": lambda rid, params: { - "value": _coerce_statusbar(_display_cfg().get("tui_statusbar", "top"))}, - "focus": lambda rid, params: { - "value": "on" if bool(_display_cfg().get("focus_view", False)) else "off", - "tool_progress": _load_tool_progress_mode()}, - "mouse": lambda rid, params: {"value": _display_mouse_tracking(_load_cfg().get("display"))}, - "mtime": _cfg_get_mtime} +# key -> getter(params); bind_module rebinds the table's functions onto server.py's globals. +_CONFIG_GETTERS = { + "provider": _cfg_get_provider, + "profile": lambda params: {"home": str(_hermes_home), "display": _display_hermes_home()}, + "project": _cfg_get_project, + "full": lambda params: {"config": _load_cfg()}, + "prompt": lambda params: {"prompt": _load_cfg().get("custom_prompt", "")}, + "skin": lambda params: {"value": _display_raw().get("skin", "default")}, + # Normalised like the TUI renders it (frontend falls back to the default for the same inputs). + "indicator": lambda params: { + "value": _display_word("tui_status_indicator", DEFAULT_INDICATOR_STYLE, INDICATOR_STYLES)}, + "personality": _cfg_get_personality, + "reasoning": _cfg_get_reasoning, + "fast": _cfg_get_fast, + "busy": lambda params: {"value": _load_busy_input_mode()}, + "approval_mode": lambda params: {"value": _load_approval_mode()}, + "approvals.mode": lambda params: {"value": _load_approval_mode()}, + "details_mode": lambda params: {"value": _display_word("details_mode", "collapsed", _DETAIL_MODES)}, + "thinking_mode": _cfg_get_thinking_mode, + "density": lambda params: {"value": "on" if bool(_display_raw().get("tui_compact", False)) else "off"}, + "theme": lambda params: {"value": _display_word("tui_theme", "auto", {"auto", "light", "dark"})}, + "statusbar": lambda params: {"value": _coerce_statusbar(_display_cfg().get("tui_statusbar", "top"))}, + "focus": lambda params: {"value": "on" if bool(_display_cfg().get("focus_view", False)) else "off", + "tool_progress": _load_tool_progress_mode()}, + "mouse": lambda params: {"value": _display_mouse_tracking(_load_cfg().get("display"))}, + "mtime": _cfg_get_mtime} +# Getters whose failure is a JSON-RPC error of this code (others propagate to dispatch). +_CONFIG_GET_ERR = {"provider": 5013, "approval_mode": 5001, "approvals.mode": 5001} @method("config.get") @_profile_scoped def _(rid, params: dict) -> dict: key = params.get("key", "") - getter = _config_getters().get(key) + getter = _CONFIG_GETTERS.get(key) if getter is None: return _err(rid, 4002, f"unknown config key: {key}") - payload = getter(rid, params) - if "error" in payload: - return payload - return _ok(rid, payload) + try: + return _ok(rid, getter(params)) + except Exception as e: + if key not in _CONFIG_GET_ERR: + raise + return _err(rid, _CONFIG_GET_ERR[key], str(e)) # ── setup readiness - -def _readiness_profile_scope(params: dict): - """``(profile, scope)`` for the readiness RPCs' optional ``profile`` param: ``scope`` binds - that profile's HERMES_HOME + ``.env`` secret scope (ContextVars, so concurrent checks stay - isolated); no param yields ``("", nullcontext())``. An unknown profile raises - ``FileNotFoundError`` — never quietly answer for the launch profile instead.""" +def _readiness_check(rid, params, probe): + """Shared shell of setup.status / setup.runtime_check. ``probe(profile, scoped)`` runs inside the + optional ``profile`` param's HERMES_HOME + ``.env`` secret scope (ContextVars: concurrent checks + stay isolated); ``scoped`` is the ``{"profile": ...}`` payload stamp (``{}`` for the launch + profile). An unknown profile answers ``ok=False`` (never a JSON-RPC error, never a quiet answer + for the launch profile instead).""" import contextlib profile = str(params.get("profile") or "").strip() if isinstance(params, dict) else "" - if not profile: - return "", contextlib.nullcontext() - from hermes_cli import profiles as profiles_mod - if not profiles_mod.profile_exists(profile): - raise FileNotFoundError(f"Profile '{profile}' does not exist on this backend.") - home = _profile_home(profile) - if home is None: - return profile, contextlib.nullcontext() - return profile, _session_profile_runtime_scope({"profile_home": str(home)}) - - -def _readiness_check(rid, params, probe): - """Shared shell of setup.status / setup.runtime_check: ``probe(profile)`` runs inside the - profile scope; an unknown profile answers ``ok=False`` (never a JSON-RPC error).""" - try: - profile, scope = _readiness_profile_scope(params) - except FileNotFoundError as e: - return _ok(rid, {"ok": False, "profile": params.get("profile"), "error": str(e)}) + scope = contextlib.nullcontext() + if profile: + from hermes_cli import profiles as profiles_mod + if not profiles_mod.profile_exists(profile): + return _ok(rid, {"ok": False, "profile": params.get("profile"), + "error": f"Profile '{profile}' does not exist on this backend."}) + home = _profile_home(profile) + if home is not None: + scope = _session_profile_runtime_scope({"profile_home": str(home)}) with scope: - payload = probe(profile) + payload = probe(profile, {"profile": profile} if profile else {}) return _ok(rid, payload) @@ -320,21 +268,18 @@ def _(rid, params: dict) -> dict: """Loose provider check; ``profile`` (optional) scopes it to that profile's home.""" try: from hermes_cli.main import _has_any_provider_configured - - def probe(profile): - configured = bool(_has_any_provider_configured(strict_profile_scope=bool(profile))) - return {"provider_configured": configured, **({"profile": profile} if profile else {})} - return _readiness_check(rid, params, probe) + return _readiness_check(rid, params, lambda profile, scoped: { + "provider_configured": bool(_has_any_provider_configured(strict_profile_scope=bool(profile))), + **scoped}) except Exception as e: return _err(rid, 5016, str(e)) @method("setup.runtime_check") def _(rid, params: dict) -> dict: - """Strict provider check: does the configured/default model resolve to a usable runtime? - Unlike setup.status (True if ANY provider auth state is discoverable), this runs the same - resolve_runtime_provider() the agent uses on session creation and returns ok=False with the - auth error when the model can't be served, so UIs surface onboarding before a doomed prompt. + """Strict provider check via the same resolve_runtime_provider() the agent uses on session + creation (setup.status is True if ANY provider auth state is discoverable): ok=False + the auth + error when the model can't be served, so UIs surface onboarding before a doomed prompt. ``profile`` answers for THAT profile's pin and ``.env``; unknown -> ``ok=False``.""" try: from hermes_cli.runtime_provider import resolve_runtime_provider @@ -342,11 +287,9 @@ def _(rid, params: dict) -> dict: from hermes_cli.main import _has_any_provider_configured requested = str(params.get("provider") or "").strip() or None - def probe(profile): + def probe(profile, scoped): runtime = resolve_runtime_provider(requested=requested) - provider_configured = bool( - _has_any_provider_configured(strict_profile_scope=bool(profile))) - scoped = {"profile": profile} if profile else {} + provider_configured = bool(_has_any_provider_configured(strict_profile_scope=bool(profile))) provider = runtime.get("provider") or "provider" source = str(runtime.get("source") or "") @@ -358,10 +301,8 @@ def _(rid, params: dict) -> dict: return fail("No Hermes provider is configured.", source) api_key = runtime.get("api_key") api_key_text = "" if callable(api_key) else str(api_key or "").strip() - credential_ok = ( - callable(api_key) or api_key_text in {"aws-sdk", "no-key-required"} - or has_usable_secret(api_key_text) or bool(runtime.get("command"))) - if not credential_ok: + if not (callable(api_key) or api_key_text in {"aws-sdk", "no-key-required"} + or has_usable_secret(api_key_text) or bool(runtime.get("command"))): return fail(f"No usable credentials found for {provider}.", runtime.get("source")) return {"ok": True, "provider": runtime.get("provider"), "model": runtime.get("model"), "source": runtime.get("source"), **scoped} @@ -370,14 +311,21 @@ def _(rid, params: dict) -> dict: return _ok(rid, {"ok": False, "error": str(e)}) +def _safe_client_label(label: str) -> str: + """Alnum/._- () only, ≤64 chars, dot-runs and leading dots collapsed (no traversal shapes).""" + safe = "".join(ch for ch in label if ch.isalnum() or ch in "._- ()").strip()[:64] + while ".." in safe: + safe = safe.replace("..", ".") + return safe.lstrip(".").strip() + + @method("diagnostics.share_nous") def _(rid, params: dict) -> dict: """Upload a redacted debug bundle to Nous-internal diagnostics storage — same collection + - force-redaction pipeline as ``hermes debug share --nous``; redaction is NOT - client-controllable and consent lives with the CALLER (privacy notice first). Structured - ``ok``/``error`` envelope so the client renders upload failures inline. Optional params: - ``error_context`` (redacted, attached as ``error-context.txt``), ``extra_files`` ({label -> - text}, force-redacted, labels sanitized and size-capped), ``log_lines`` (default 200).""" + force-redaction pipeline as ``hermes debug share --nous``; redaction is NOT client-controllable + and consent lives with the CALLER (privacy notice first). Structured ``ok``/``error`` envelope so + upload failures render inline. Optional: ``error_context`` (-> ``error-context.txt``), + ``extra_files`` ({label -> text}), ``log_lines`` (default 200); all force-redacted.""" try: from hermes_cli.debug import _redact_log_text, build_nous_bundle, collect_share_bundle from hermes_cli.diagnostics_upload import share_to_nous @@ -392,28 +340,17 @@ def _(rid, params: dict) -> dict: bundle["error-context.txt"] = _redact_log_text(error_context.strip()[:8_000]) # Bounded: at most 4 files, 512KB each, sanitized labels — not an arbitrary upload surface. extra_files = params.get("extra_files") - if isinstance(extra_files, dict): - for label, text in list(extra_files.items())[:4]: - if not isinstance(label, str) or not isinstance(text, str): - continue - safe_label = "".join(ch for ch in label if ch.isalnum() or ch in "._- ()").strip()[:64] - # Collapse dot-runs / leading dots so traversal-shaped labels can't survive. - while ".." in safe_label: - safe_label = safe_label.replace("..", ".") - safe_label = safe_label.lstrip(".").strip() - if not safe_label or not text.strip(): - continue + for label, text in list(extra_files.items())[:4] if isinstance(extra_files, dict) else (): + safe_label = _safe_client_label(label) if isinstance(label, str) else "" + if safe_label and isinstance(text, str) and text.strip(): bundle[f"client/{safe_label}"] = _redact_log_text(text[:524_288]) res = share_to_nous(build_nous_bundle(bundle, redact=True)) view_url = res.get("viewUrl") or res.get("view_url") upload_id = res.get("id") - if not view_url and not upload_id: - # An upload the user can't reference is useless to support. - return _ok(rid, {"ok": False, - "error": "upload succeeded but returned no view URL or id"}) - return _ok(rid, { - "ok": True, "view_url": view_url, "upload_id": upload_id, - "expires_at": res.get("expiresAt") or res.get("expires_at")}) + if not view_url and not upload_id: # an upload the user can't reference is useless to support + return _ok(rid, {"ok": False, "error": "upload succeeded but returned no view URL or id"}) + return _ok(rid, {"ok": True, "view_url": view_url, "upload_id": upload_id, + "expires_at": res.get("expiresAt") or res.get("expires_at")}) except Exception as e: return _ok(rid, {"ok": False, "error": str(e)}) diff --git a/tui_gateway/methods_config_set.py b/tui_gateway/methods_config_set.py index e9ff07fc87..a18d5864de 100644 --- a/tui_gateway/methods_config_set.py +++ b/tui_gateway/methods_config_set.py @@ -1,5 +1,5 @@ -"""``config.set`` — one JSON-RPC method, dispatched on ``key`` through a table. Bodies are -rebound onto server.py's globals (method_ctx.bind_module) and reference them bare. Each +"""``config.set`` — one JSON-RPC method, dispatched on ``key`` through ``_CONFIG_SETTERS``. Bodies +are rebound onto server.py's globals (method_ctx.bind_module) and reference them bare. Each ``_set_*`` takes ``(rid, params, key, value, session)`` and returns the JSON-RPC envelope. Keys match exactly except ``details_mode.
`` (prefix) and ``_DISPLAY_TOGGLE_KEYS``. """ @@ -20,10 +20,8 @@ _profile_scoped = _registry.profile_scoped def _write_display_sections(*, sections=None, drop_sections=(), **display_fields) -> None: """Persist ``display.`` + ``display.sections`` edits via the raw (uncached) write-back.""" cfg = _load_cfg_raw() - display = cfg.get("display") - display = display if isinstance(display, dict) else {} - cur = display.get("sections") - cur = cur if isinstance(cur, dict) else {} + display = cfg.get("display") if isinstance(cfg.get("display"), dict) else {} + cur = display.get("sections") if isinstance(display.get("sections"), dict) else {} display.update(display_fields) cur.update(sections or {}) for name in drop_sections: @@ -44,24 +42,19 @@ def _emit_all_session_info() -> None: _emit_session_info(sid, sess) -def _toggle_display_bool(rid, key, value, *, cfg_key, on_words, off_words): - """Shared body of the on/off/toggle display booleans (``density``, ``battery``).""" - raw = _word(value) - cur_b = bool(_display_cfg().get(cfg_key, False)) - if raw in {"", "toggle"}: - nv_b = not cur_b - elif raw in on_words or raw in off_words: - nv_b = raw in on_words - else: - return _err(rid, 4002, f"unknown {key} value: {value}") - _write_config_key(f"display.{cfg_key}", nv_b) - return _ok(rid, {"key": key, "value": "on" if nv_b else "off"}) - - def _word(value) -> str: return str(value or "").strip().lower() +def _raw_word(value) -> str: + """Like ``_word`` but only None is blank: falsy non-strings (0, False, []) keep their text.""" + return ("" if value is None else str(value)).strip().lower() + + +def _kv(rid, key, value, **extra): + return _ok(rid, {"key": key, "value": value, **extra}) + + def _cfgset_await_agent(session, rid): """Wait for an in-progress agent build; the error envelope if it failed, else None.""" init_err = _wait_agent(session, rid) @@ -70,82 +63,87 @@ def _cfgset_await_agent(session, rid): return _err(rid, 5032, "agent initialization failed") if session.get("agent") is None else None -def _cfgset_model_ok(rid, key, value, warning, confirm_required, confirm_message, scope, **extra): - return _ok(rid, {"key": key, "value": value, "warning": warning, "confirm_required": confirm_required, - "confirm_message": confirm_message, "scope": scope, **extra}) +def _cfgset_model_ok(rid, key, value, warning="", confirm_message="", scope="session", **extra): + """Model-switch envelope; ``confirm_required`` follows ``confirm_message`` (canonical; ``warning`` + is its legacy alias on the deferred path).""" + return _kv(rid, key, value, warning=warning, confirm_required=bool(confirm_message), + confirm_message=confirm_message, scope=scope, **extra) + + +def _stash_pending_model_switch(rid, key, value, session, confirmed, parsed): + """No live swap while a turn streams (agent.switch_model() mutates fields the worker thread + reads every iteration): stash the pick for the NEXT turn start. Selection guards run HERE (the + only moment a confirm round-trip is possible; an unconfirmed stashed pick is dropped at turn + start) — on a warning nothing is stashed.""" + try: + pending_model = parsed.model_input + except Exception: + pending_model = str(value) + pending_provider = (getattr(parsed, "explicit_provider", "") or "").strip() + if not confirmed: + pending_warning = _pending_switch_selection_warning(pending_model, pending_provider) + if pending_warning is not None: + return _cfgset_model_ok(rid, key, pending_model, pending_warning, pending_warning, deferred=False) + # display_*: _session_info shows the user's pick while pending, not the live old model. + session["pending_model_switch"] = { + "raw": value, "confirm_expensive_model": confirmed, + "display_model": pending_model, "display_provider": pending_provider} + return _cfgset_model_ok(rid, key, pending_model, deferred=True) + + +def _cfgset_guarded(fn): + """Setter whose uncaught exception becomes ``_err(rid, 5001, str(e))``.""" + def setter(rid, params, key, value, session): + try: + return fn(rid, params, key, value, session) + except Exception as e: + return _err(rid, 5001, str(e)) + return setter # ── per-key handlers +@_cfgset_guarded def _set_model(rid, params, key, value, session): """Live/deferred model switch; see _apply_model_switch and _apply_pending_model_switch.""" - try: - if not value: - return _err(rid, 4002, "model value required") - confirmed = bool(params.get("confirm_expensive_model", False)) - if session: - from hermes_cli.model_switch import parse_model_switch_args - sid = params.get("session_id", "") - # No live swap while a turn streams (agent.switch_model() mutates fields the worker - # thread reads every iteration): stash the pick for the NEXT turn start. - if session.get("running"): - parsed = parse_model_switch_args(value) - try: - pending_model = parsed.model_input - except Exception: - pending_model = str(value) - pending_provider = (getattr(parsed, "explicit_provider", "") or "").strip() - # Selection guards run HERE (the only moment a confirm round-trip is possible); - # otherwise an unconfirmed stashed pick is dropped at turn start. - if not confirmed: - pending_warning = _pending_switch_selection_warning(pending_model, pending_provider) - if pending_warning is not None: - # Nothing stashed; the client re-sends with confirm_expensive_model. - # `confirm_message` is canonical, `warning` its legacy alias. - return _cfgset_model_ok( - rid, key, pending_model, pending_warning, True, pending_warning, "session", deferred=False - ) - session["pending_model_switch"] = { - "raw": value, - "confirm_expensive_model": confirmed, - # _session_info reports these while pending so the end-of-turn settle keeps - # showing the user's pick, not the still-live old model. - "display_model": pending_model, - "display_provider": pending_provider} - return _cfgset_model_ok(rid, key, pending_model, "", False, "", "session", deferred=True) - parsed_flags = parse_model_switch_args(value) - explicit_provider = parsed_flags.explicit_provider - failed_agent_init = session.get("agent") is None and session.get("agent_error") is not None - failed_ready = session.get("agent_ready") if failed_agent_init else None - if failed_agent_init: - if failed_ready is None: - return _err(rid, 5032, session.get("agent_error") or "agent initialization failed") - if not failed_ready.wait(timeout=30.0): - return _err(rid, 5032, "agent initialization timed out") - failed_agent_init = ( - failed_agent_init and session.get("agent") is None and session.get("agent_error") is not None - and session.get("agent_ready") is failed_ready and failed_ready.is_set()) - if session.get("agent") is None and not explicit_provider.strip() and not failed_agent_init: - _start_agent_build(sid, session) - if init_err := _cfgset_await_agent(session, rid): - return init_err + if not value: + return _err(rid, 4002, "model value required") + confirmed = bool(params.get("confirm_expensive_model", False)) + if session: + from hermes_cli.model_switch import parse_model_switch_args + sid = params.get("session_id", "") + parsed_flags = parse_model_switch_args(value) + if session.get("running"): + return _stash_pending_model_switch(rid, key, value, session, confirmed, parsed_flags) + explicit_provider = parsed_flags.explicit_provider + failed_agent_init = session.get("agent") is None and session.get("agent_error") is not None + failed_ready = session.get("agent_ready") if failed_agent_init else None + if failed_agent_init: + if failed_ready is None: + return _err(rid, 5032, session.get("agent_error") or "agent initialization failed") + if not failed_ready.wait(timeout=30.0): + return _err(rid, 5032, "agent initialization timed out") + failed_agent_init = ( + failed_agent_init and session.get("agent") is None and session.get("agent_error") is not None + and session.get("agent_ready") is failed_ready and failed_ready.is_set()) + if session.get("agent") is None and not explicit_provider.strip() and not failed_agent_init: + _start_agent_build(sid, session) + if init_err := _cfgset_await_agent(session, rid): + return init_err + with _session_profile_runtime_scope(session): + result = _apply_model_switch(sid, session, value, confirm_expensive_model=confirmed, + parsed_flags=parsed_flags) + if failed_agent_init and not result.get("confirm_required"): + _restart_completed_failed_agent_build(sid, session, failed_ready) + if init_err := _cfgset_await_agent(session, rid): + return init_err with _session_profile_runtime_scope(session): - result = _apply_model_switch( - sid, session, value, confirm_expensive_model=confirmed, parsed_flags=parsed_flags - ) - if failed_agent_init and not result.get("confirm_required"): - _restart_completed_failed_agent_build(sid, session, failed_ready) - if init_err := _cfgset_await_agent(session, rid): - return init_err - with _session_profile_runtime_scope(session): - _persist_live_session_runtime(session) - else: - result = _apply_model_switch("", {"agent": None}, value, confirm_expensive_model=confirmed) - return _cfgset_model_ok( - rid, key, result["value"], result["warning"], result.get("confirm_required", False), - result.get("confirm_message", ""), result.get("scope", "session")) - except Exception as e: - return _err(rid, 5001, str(e)) + _persist_live_session_runtime(session) + else: + result = _apply_model_switch("", {"agent": None}, value, confirm_expensive_model=confirmed) + return _kv(rid, key, result["value"], warning=result["warning"], + confirm_required=result.get("confirm_required", False), + confirm_message=result.get("confirm_message", ""), scope=result.get("scope", "session")) _FAST_WORDS = {"fast": "fast", "on": "fast", "normal": "normal", "off": "normal", @@ -158,15 +156,12 @@ def _set_fast(rid, params, key, value, session): if agent is not None: current_tier = getattr(agent, "service_tier", None) elif session is not None and session.get("create_service_tier_override") is not None: - # Pre-build session with a pinned tier: report/toggle from the pin, not the global. - current_tier = session["create_service_tier_override"] or None + current_tier = session["create_service_tier_override"] or None # pre-build pin beats global else: current_tier = _load_service_tier() - current_fast = current_tier == "priority" if raw == "status": - return _ok(rid, {"key": key, "value": {"priority": "fast", None: "normal"}.get(current_tier, current_tier)}) - toggled = ("normal" if current_fast else "fast") if raw in {"", "toggle"} else None - nv = _FAST_WORDS.get(raw, toggled) + return _kv(rid, key, {"priority": "fast", None: "normal"}.get(current_tier, current_tier)) + nv = _FAST_WORDS.get(raw, ("normal" if current_tier == "priority" else "fast") if raw in {"", "toggle"} else None) if nv is None: return _err(rid, 4002, f"unknown fast mode: {value}") overrides = None @@ -174,8 +169,7 @@ def _set_fast(rid, params, key, value, session): from hermes_cli.models import resolve_fast_mode_overrides if agent is not None: target_model = getattr(agent, "model", None) - else: - # A pre-build session may carry a picked model (desktop draft) — validate against THAT. + else: # a pre-build session may carry a picked model (desktop draft): validate against THAT session_override = (session or {}).get("model_override") or {} target_model = (isinstance(session_override, dict) and session_override.get("model")) or _resolve_model() if not target_model: @@ -185,56 +179,46 @@ def _set_fast(rid, params, key, value, session): if overrides is None: return _err(rid, 4002, "fast mode is not available for this model") if session is not None: - # Session-scoped like `reasoning` (global persistence is `--global` / Settings → Model): - # writing config.yaml here flipped fast mode for every other surface. The create - # override keeps the choice across lazy builds and rebuilds; "" pins normal. + # Session-scoped like `reasoning` (global = `--global` / Settings → Model): writing config.yaml + # here flipped fast mode for every surface. The create override survives rebuilds; "" pins normal. session["create_service_tier_override"] = {"fast": "priority", "normal": ""}.get(nv, nv) else: _write_config_key("agent.service_tier", nv) if agent is not None: agent.service_tier = {"fast": "priority", "normal": None}.get(nv, nv) - current_overrides = dict(getattr(agent, "request_overrides", {}) or {}) - current_overrides.pop("service_tier", None) - current_overrides.pop("speed", None) - if nv == "fast": - current_overrides.update(overrides) - agent.request_overrides = current_overrides + current_overrides = {k: v for k, v in (getattr(agent, "request_overrides", {}) or {}).items() + if k not in ("service_tier", "speed")} + agent.request_overrides = {**current_overrides, **(overrides or {})} _persist_live_session_runtime(session) _emit_session_info(params.get("session_id", ""), session) - return _ok(rid, {"key": key, "value": nv}) + return _kv(rid, key, nv) def _set_busy(rid, params, key, value, session): - raw = _word(value) - if raw in {"", "status"}: - return _ok(rid, {"key": key, "value": _load_busy_input_mode()}) - if raw not in {"queue", "steer", "interrupt"}: - return _err(rid, 4002, f"unknown busy mode: {value}") - _write_config_key("display.busy_input_mode", raw) - return _ok(rid, {"key": key, "value": raw}) + if _word(value) in {"", "status"}: + return _kv(rid, key, _load_busy_input_mode()) + return _set_word(rid, params, key, value, session) def _set_verbose(rid, params, key, value, session): cycle = ["off", "new", "all", "verbose"] - cur = session.get("tool_progress_mode", _load_tool_progress_mode()) if session else _load_tool_progress_mode() if value and value != "cycle": nv = str(value).strip().lower() if nv not in cycle: return _err(rid, 4002, f"unknown verbose mode: {value}") else: - idx = cycle.index(cur) if cur in cycle else 2 - nv = cycle[(idx + 1) % len(cycle)] + cur = session.get("tool_progress_mode", _load_tool_progress_mode()) if session else _load_tool_progress_mode() + nv = cycle[((cycle.index(cur) if cur in cycle else 2) + 1) % len(cycle)] _write_config_key("display.tool_progress", nv) if session: session["tool_progress_mode"] = nv if session.get("agent") is not None: session["agent"].verbose_logging = nv == "verbose" - return _ok(rid, {"key": key, "value": nv}) + return _kv(rid, key, nv) def _set_focus(rid, params, key, value, session): - # Focus view (/focus): enabling stashes the configured tool_progress mode and pins it - # "off"; disabling restores the stash. + # /focus: enabling stashes the configured tool_progress mode and pins it "off"; disabling restores. from hermes_cli.focus_view import FOCUS_TOOL_PROGRESS_MODE, normalize_tool_progress_mode, resolve_focus_arg d_f = _display_cfg() cur_focus = bool(d_f.get("focus_view", False)) @@ -242,15 +226,14 @@ def _set_focus(rid, params, key, value, session): if action == "usage": return _err(rid, 4002, f"unknown focus value: {value} (use on|off|status)") if action == "status" or target is None: - return _ok(rid, {"key": key, "value": "on" if cur_focus else "off", "tool_progress": _load_tool_progress_mode()}) + return _kv(rid, key, "on" if cur_focus else "off", tool_progress=_load_tool_progress_mode()) if target: saved = (cur_focus and d_f.get("focus_saved_tool_progress")) or _load_tool_progress_mode() _write_config_key("display.focus_saved_tool_progress", normalize_tool_progress_mode(saved)) - _write_config_key("display.tool_progress", FOCUS_TOOL_PROGRESS_MODE) effective = FOCUS_TOOL_PROGRESS_MODE else: effective = normalize_tool_progress_mode(d_f.get("focus_saved_tool_progress") or "all") - _write_config_key("display.tool_progress", effective) + _write_config_key("display.tool_progress", effective) _write_config_key("display.focus_view", bool(target)) if session: session["focus_view"] = bool(target) @@ -258,51 +241,39 @@ def _set_focus(rid, params, key, value, session): if session.get("agent") is not None: with contextlib.suppress(Exception): session["agent"].tool_progress_mode = effective - return _ok(rid, {"key": key, "value": "on" if target else "off", "tool_progress": effective}) + return _kv(rid, key, "on" if target else "off", tool_progress=effective) def _set_approval_mode(rid, params, key, value, session): - raw = _word(value) - if raw not in _APPROVAL_MODES: - return _err(rid, 4002, f"unknown approval mode: {value}; pick one of manual|smart|off") - _write_config_key("approvals.mode", raw) - _emit_all_session_info() - return _ok(rid, {"key": "approvals.mode", "value": raw}) + return _set_word(rid, params, "approvals.mode", value, session) # legacy alias reports the real key +@_cfgset_guarded def _set_yolo(rid, params, key, value, session): - # scope="session" (default; Shift+Tab) toggles ONLY this session's flag. scope="global" + # scope="session" (default; Shift+Tab) toggles ONLY this session's flag; scope="global" # (Shift+click the zap) flips persistent approvals.mode between "off" and "manual". scope = _word(params.get("scope") or "session") - try: - from tools.approval import disable_session_yolo, enable_session_yolo, is_session_yolo_enabled - raw = _word(value) - - def _resolve_toggle(current: bool) -> bool: - return _BOOL_WORDS.get(raw, not current) - if scope == "global": - from tools.approval import _normalize_approval_mode - appr = _load_cfg().get("approvals") - appr = appr if isinstance(appr, dict) else {} - enable = _resolve_toggle(_normalize_approval_mode(appr.get("mode", "manual")) == "off") - # Binary affordance: no restore of a prior "smart"/custom mode (those live in config.yaml). - _write_config_key("approvals.mode", "off" if enable else "manual") - _emit_all_session_info() # reflect the flip in every live indicator - return _ok(rid, {"key": key, "value": "1" if enable else "0", "scope": "global"}) - if session: - skey = session["session_key"] - enable = _resolve_toggle(is_session_yolo_enabled(skey)) - (enable_session_yolo if enable else disable_session_yolo)(skey) - _emit_session_info(params.get("session_id", ""), session) + from tools.approval import disable_session_yolo, enable_session_yolo, is_session_yolo_enabled + raw = _word(value) + if scope == "global": + from tools.approval import _normalize_approval_mode + appr = _load_cfg().get("approvals") + appr = appr if isinstance(appr, dict) else {} + enable = _BOOL_WORDS.get(raw, _normalize_approval_mode(appr.get("mode", "manual")) != "off") + _write_config_key("approvals.mode", "off" if enable else "manual") # binary: no "smart" restore + _emit_all_session_info() # reflect the flip in every live indicator + elif session: + skey = session["session_key"] + enable = _BOOL_WORDS.get(raw, not is_session_yolo_enabled(skey)) + (enable_session_yolo if enable else disable_session_yolo)(skey) + _emit_session_info(params.get("session_id", ""), session) + else: + enable = _BOOL_WORDS.get(raw, not is_truthy_value(os.environ.get("HERMES_YOLO_MODE"))) + if enable: + os.environ["HERMES_YOLO_MODE"] = "1" else: - enable = _resolve_toggle(is_truthy_value(os.environ.get("HERMES_YOLO_MODE"))) - if enable: - os.environ["HERMES_YOLO_MODE"] = "1" - else: - os.environ.pop("HERMES_YOLO_MODE", None) - return _ok(rid, {"key": key, "value": "1" if enable else "0", "scope": "session"}) - except Exception as e: - return _err(rid, 5001, str(e)) + os.environ.pop("HERMES_YOLO_MODE", None) + return _kv(rid, key, "1" if enable else "0", scope=scope if scope == "global" else "session") # /reasoning display words: (accepted inputs, reported value, display field, sections.thinking, @@ -314,125 +285,106 @@ _REASONING_DISPLAY_WORDS = ( ({"clamp", "collapse", "short"}, "clamp", {"reasoning_full": False}, "collapsed", None)) +@_cfgset_guarded def _set_reasoning(rid, params, key, value, session): - try: - from hermes_constants import parse_reasoning_effort - arg = _word(value) - scope = _word(params.get("scope")) - for words, reported, fields, thinking, show in _REASONING_DISPLAY_WORDS: - if arg in words: - _write_display_sections(sections={"thinking": thinking}, **fields) - if show is not None and session: - session["show_reasoning"] = show - return _ok(rid, {"key": key, "value": reported}) - parsed = parse_reasoning_effort(arg) - if parsed is None: - return _err(rid, 4002, f"unknown reasoning value: {value}") - if scope == "global" or session is None: - _write_config_key("agent.reasoning_effort", arg) - if session is not None: - session.pop("create_reasoning_override", None) - else: - # Session-scoped like the gateway's `/reasoning `; otherwise every desktop - # model-menu pick rewrote the global default. - session["create_reasoning_override"] = parsed - if session and session.get("agent") is not None: - session["agent"].reasoning_config = parsed - _persist_live_session_runtime(session) - _emit_session_info(params.get("session_id", ""), session) - return _ok(rid, {"key": key, "value": arg}) - except Exception as e: - return _err(rid, 5001, str(e)) + from hermes_constants import parse_reasoning_effort + arg = _word(value) + scope = _word(params.get("scope")) + for words, reported, fields, thinking, show in _REASONING_DISPLAY_WORDS: + if arg in words: + _write_display_sections(sections={"thinking": thinking}, **fields) + if show is not None and session: + session["show_reasoning"] = show + return _kv(rid, key, reported) + parsed = parse_reasoning_effort(arg) + if parsed is None: + return _err(rid, 4002, f"unknown reasoning value: {value}") + if scope == "global" or session is None: + _write_config_key("agent.reasoning_effort", arg) + if session is not None: + session.pop("create_reasoning_override", None) + else: # session-scoped like the gateway's `/reasoning `; a menu pick must not rewrite the global + session["create_reasoning_override"] = parsed + if session and session.get("agent") is not None: + session["agent"].reasoning_config = parsed + _persist_live_session_runtime(session) + _emit_session_info(params.get("session_id", ""), session) + return _kv(rid, key, arg) -def _set_details_mode(rid, params, key, value, session): - nv = _word(value) - if nv not in _DETAIL_MODES: - return _err(rid, 4002, f"unknown details_mode: {value}") - _write_display_sections(sections={section: nv for section in _DETAIL_SECTION_NAMES}, details_mode=nv) - return _ok(rid, {"key": key, "value": nv}) +def _word_setters() -> dict: + """key -> (normaliser, accepted words, error template, apply(word)); the reported value is the + accepted word. Built per call: the specs reference server.py globals (rebound at install).""" + return { + "busy": (_word, {"queue", "steer", "interrupt"}, "unknown busy mode: {value}", + lambda w: _write_config_key("display.busy_input_mode", w)), + "approvals.mode": (_word, _APPROVAL_MODES, "unknown approval mode: {value}; pick one of manual|smart|off", + lambda w: (_write_config_key("approvals.mode", w), _emit_all_session_info())), + "details_mode": (_word, _DETAIL_MODES, "unknown details_mode: {value}", lambda w: _write_display_sections( + sections={section: w for section in _DETAIL_SECTION_NAMES}, details_mode=w)), + # thinking_mode also keeps details_mode aligned (compat bridge). + "thinking_mode": (_word, {"collapsed", "truncated", "full"}, "unknown thinking_mode: {value}", lambda w: ( + _write_config_key("display.thinking_mode", w), + _write_config_key("display.details_mode", "expanded" if w == "full" else "collapsed"))), + # 'light'/'dark' pin beats background auto-detection (xterm.js hosts misreport OSC 11). + "theme": (_word, {"auto", "light", "dark"}, "unknown theme value: {value} (use auto|light|dark)", + lambda w: _write_config_key("display.tui_theme", w)), + # _raw_word: 0/False/[] keep their text so the error names what was sent. + "indicator": (_raw_word, INDICATOR_STYLES, "unknown indicator: {raw!r}; pick one of " + "|".join(INDICATOR_STYLES), + lambda w: _write_config_key("display.tui_status_indicator", w))} + + +def _set_word(rid, params, key, value, session): + norm, allowed, err, apply = _word_setters()[key] + raw = norm(value) + if raw not in allowed: + return _err(rid, 4002, err.format(value=value, raw=raw)) + apply(raw) + return _kv(rid, key, raw) def _set_details_section(rid, params, key, value, session): - # `details_mode.
` -> `display.sections.
`; empty clears the override so the - # frontend applies built-in section defaults before the global details_mode. + # `details_mode.
` -> `display.sections.
`; empty clears the override (frontend + # then applies built-in section defaults before the global details_mode). section = key.split(".", 1)[1] if section not in _DETAIL_SECTION_NAMES: return _err(rid, 4002, f"unknown section: {section}") nv = _word(value) - if not nv: - _write_display_sections(drop_sections=(section,)) - elif nv not in _DETAIL_MODES: + if nv and nv not in _DETAIL_MODES: return _err(rid, 4002, f"unknown details_mode: {value}") - else: - _write_display_sections(sections={section: nv}) - return _ok(rid, {"key": key, "value": nv}) + _write_display_sections(sections={section: nv} if nv else None, drop_sections=() if nv else (section,)) + return _kv(rid, key, nv) -def _set_thinking_mode(rid, params, key, value, session): - nv = _word(value) - if nv not in {"collapsed", "truncated", "full"}: - return _err(rid, 4002, f"unknown thinking_mode: {value}") - _write_config_key("display.thinking_mode", nv) - # Backward compatibility bridge: keep details_mode aligned. - _write_config_key("display.details_mode", "expanded" if nv == "full" else "collapsed") - return _ok(rid, {"key": key, "value": nv}) +def _toggle_setters() -> dict: + """key -> (normaliser, cfg key, alias word -> value, flipped(current), report). ``""``/``toggle`` + flips the current value; an alias word maps directly; anything else is 4002. Built per call: + the specs reference server.py globals (rebound at install).""" + def on_off(v): + return "on" if v else "off" + return { + # density/battery are on/off/toggle booleans on display.. + "density": (_word, "display.tui_compact", {"on": True, "off": False}, + lambda: not bool(_display_cfg().get("tui_compact", False)), on_off), + "battery": (_word, "display.battery", + {"on": True, "true": True, "yes": True, "off": False, "false": False, "no": False}, + lambda: not bool(_display_cfg().get("battery", False)), on_off), + "statusbar": (_word, "display.tui_statusbar", {"on": "top", **{m: m for m in _STATUSBAR_MODES}}, + lambda: "top" if _coerce_statusbar(_display_cfg().get("tui_statusbar", "top")) == "off" else "off", + lambda v: v), + # _raw_word: falsy non-strings (0, False) reach the alias map as themselves (-> 'off'), not toggle. + "mouse": (_raw_word, "display.mouse_tracking", _MOUSE_TRACKING_ALIASES, + lambda: "all" if _display_mouse_tracking(_display_cfg()) == "off" else "off", lambda v: v)} -def _set_density(rid, params, key, value, session): - return _toggle_display_bool(rid, key, value, cfg_key="tui_compact", on_words={"on"}, off_words={"off"}) - - -def _set_battery(rid, params, key, value, session): - return _toggle_display_bool( - rid, key, value, cfg_key="battery", on_words={"on", "true", "yes"}, off_words={"off", "false", "no"} - ) - - -def _set_theme(rid, params, key, value, session): - # 'light'/'dark' pin beats background auto-detection (xterm.js hosts misreport OSC 11). - raw = _word(value) - if raw not in {"auto", "light", "dark"}: - return _err(rid, 4002, f"unknown theme value: {value} (use auto|light|dark)") - _write_config_key("display.tui_theme", raw) - return _ok(rid, {"key": key, "value": raw}) - - -def _set_statusbar(rid, params, key, value, session): - raw = _word(value) - current = _coerce_statusbar(_display_cfg().get("tui_statusbar", "top")) - if raw in {"", "toggle"}: - nv = "top" if current == "off" else "off" - elif raw == "on" or raw in _STATUSBAR_MODES: - nv = "top" if raw == "on" else raw - else: - return _err(rid, 4002, f"unknown statusbar value: {value}") - _write_config_key("display.tui_statusbar", nv) - return _ok(rid, {"key": key, "value": nv}) - - -def _set_mouse(rid, params, key, value, session): - # Explicit None check so falsy non-string inputs (0, False) reach the alias map as - # themselves (-> 'off') instead of toggling. - raw = ("" if value is None else str(value)).strip().lower() - current = _display_mouse_tracking(_display_cfg()) - if raw in {"", "toggle"}: - nv = "all" if current == "off" else "off" - elif raw in _MOUSE_TRACKING_ALIASES: - nv = _MOUSE_TRACKING_ALIASES[raw] - else: - return _err(rid, 4002, f"unknown mouse value: {value}") - _write_config_key("display.mouse_tracking", nv) - return _ok(rid, {"key": key, "value": nv}) - - -def _set_indicator(rid, params, key, value, session): - # Explicit None check so falsy non-string inputs (0, False, []) surface in the error message. - raw = ("" if value is None else str(value)).strip().lower() - if raw not in INDICATOR_STYLES: - return _err(rid, 4002, f"unknown indicator: {raw!r}; pick one of {'|'.join(INDICATOR_STYLES)}") - _write_config_key("display.tui_status_indicator", raw) - return _ok(rid, {"key": key, "value": raw}) +def _set_toggle(rid, params, key, value, session): + norm, cfg_key, aliases, flipped, report = _toggle_setters()[key] + raw = norm(value) + nv = flipped() if raw in {"", "toggle"} else aliases.get(raw) + if nv is None: + return _err(rid, 4002, f"unknown {key} value: {value}") + _write_config_key(cfg_key, nv) + return _kv(rid, key, report(nv)) def _set_cwd(rid, params, key, value, session): @@ -444,41 +396,38 @@ def _set_cwd(rid, params, key, value, session): return _err(rid, 4002, f"working directory does not exist: {raw}") _write_config_key("terminal.cwd", cwd) os.environ["TERMINAL_CWD"] = cwd - return _ok(rid, {"key": "terminal.cwd", "value": cwd, "cwd": cwd, "branch": _git_branch_for_cwd(cwd)}) + return _kv(rid, "terminal.cwd", cwd, cwd=cwd, branch=_git_branch_for_cwd(cwd)) -def _set_prompt_like(rid, params, key, value, session): - try: - cfg = _load_cfg_raw() # write-back round-trip ("prompt" saves cfg) - resp = {"key": key, "value": value} - if key == "prompt": - if value == "clear": - cfg.pop("custom_prompt", None) - resp["value"] = "" - else: - cfg["custom_prompt"] = value - _save_cfg(cfg) - elif key == "personality": - pname, new_prompt = _validate_personality(str(value or ""), cfg) - # Personality persists through hermes_cli.personality (single owner), never the - # user-owned global system prompt. - from hermes_cli.personality import persist_personality - persist_personality(pname) - resp["value"] = str(value or "none") - history_reset, info = _apply_personality_to_session(params.get("session_id", ""), session, new_prompt, pname) - resp["history_reset"] = history_reset - if info is not None: - resp["info"] = info - else: - _write_config_key(f"display.{key}", value) - if key == "skin": - # Every surface repaints; sync the watcher baseline so the poll loop doesn't - # re-broadcast the skin this RPC just applied. - _broadcast_global_event("skin.changed", resolve_skin()) - _note_skin_broadcast() - return _ok(rid, resp) - except Exception as e: - return _err(rid, 5001, str(e)) +@_cfgset_guarded +def _set_prompt(rid, params, key, value, session): + cfg = _load_cfg_raw() # write-back round-trip + if value == "clear": + cfg.pop("custom_prompt", None) + else: + cfg["custom_prompt"] = value + _save_cfg(cfg) + return _kv(rid, key, "" if value == "clear" else value) + + +@_cfgset_guarded +def _set_personality(rid, params, key, value, session): + pname, new_prompt = _validate_personality(str(value or ""), _load_cfg_raw()) + # Persists via hermes_cli.personality (single owner), never the user-owned system prompt. + from hermes_cli.personality import persist_personality + persist_personality(pname) + history_reset, info = _apply_personality_to_session(params.get("session_id", ""), session, new_prompt, pname) + return _kv(rid, key, str(value or "none"), history_reset=history_reset, + **({"info": info} if info is not None else {})) + + +@_cfgset_guarded +def _set_skin(rid, params, key, value, session): + _write_config_key("display.skin", value) + # Every surface repaints; sync the watcher baseline so the poll loop doesn't re-broadcast. + _broadcast_global_event("skin.changed", resolve_skin()) + _note_skin_broadcast() + return _kv(rid, key, value) def _set_display_toggle(rid, params, key, value, session): @@ -486,19 +435,19 @@ def _set_display_toggle(rid, params, key, value, session): if on is None: return _err(rid, 4002, f"{key} takes true or false") _write_config_key(key, on) - return _ok(rid, {"key": key, "value": on}) + return _kv(rid, key, on) # ── dispatch _CONFIG_SETTERS = { "model": _set_model, "fast": _set_fast, "busy": _set_busy, "verbose": _set_verbose, "focus": _set_focus, - "approval_mode": _set_approval_mode, "approvals.mode": _set_approval_mode, "yolo": _set_yolo, - "reasoning": _set_reasoning, "details_mode": _set_details_mode, "thinking_mode": _set_thinking_mode, - "density": _set_density, "battery": _set_battery, "theme": _set_theme, "statusbar": _set_statusbar, - "mouse": _set_mouse, "indicator": _set_indicator, + "approval_mode": _set_approval_mode, "approvals.mode": _set_word, "yolo": _set_yolo, + "reasoning": _set_reasoning, "details_mode": _set_word, "thinking_mode": _set_word, + "density": _set_toggle, "battery": _set_toggle, "theme": _set_word, + "statusbar": _set_toggle, "mouse": _set_toggle, "indicator": _set_word, "cwd": _set_cwd, "terminal.cwd": _set_cwd, "workdir": _set_cwd, - "prompt": _set_prompt_like, "personality": _set_prompt_like, "skin": _set_prompt_like} + "prompt": _set_prompt, "personality": _set_personality, "skin": _set_skin} @method("config.set") diff --git a/tui_gateway/methods_groups.py b/tui_gateway/methods_groups.py index e286266174..44a5c9b12b 100644 --- a/tui_gateway/methods_groups.py +++ b/tui_gateway/methods_groups.py @@ -1,17 +1,13 @@ -"""Hosted-room JSON-RPC contract. +"""Hosted-room JSON-RPC contract: durable room identity, replay, and the process-owned +same-gateway Discussion driver; ``groups.capabilities`` keeps that boundary machine-readable. -These methods expose durable room identity, replay, and the process-owned -same-gateway Discussion driver. ``groups.capabilities`` keeps that boundary -machine-readable so older clients stay on the renderer-owned room path. - -Handlers are rebound onto server.py's globals at install (see method_ctx.py), so -bodies see only server globals plus the names methods_bot_relay.register publishes; -module-private helpers reach them through keyword defaults. ``_room_method`` wraps -each handler with the shared service-lookup / error-code envelope. -""" +Handlers are rebound onto server.py's globals at install (method_ctx.py); module-private +helpers reach them through keyword defaults. ``_room_method`` is the shared envelope.""" from .method_ctx import HandlerRegistry +import contextlib +import importlib import os import threading @@ -93,6 +89,14 @@ def _current_profile() -> str: return str(_bound_server._current_profile_name() or "").strip() +def _foreign_profile_home(profile: str): + """Home of a routed profile other than the process's own, or ``ValueError``.""" + home = _bound_server._profile_home(profile) + if home is None: + raise ValueError(f"profile '{profile}' is unavailable") + return home + + def _requested_profile(params: dict) -> str: requested = str(params.get("profile") or "").strip() if not requested: @@ -101,28 +105,24 @@ def _requested_profile(params: dict) -> str: raise ValueError("profile routing is unavailable") if requested == _current_profile(): return requested - if _bound_server._profile_home(requested) is None: - raise ValueError(f"profile '{requested}' is unavailable") + _foreign_profile_home(requested) return str(_bound_server._response_profile_name(requested) or requested) def _api_server_key(profile: str | None = None) -> str: + # Published onto the server by methods_bot_relay.register (an explicit routed profile is + # authoritative: never borrow the process profile's key on a multiplexed gateway). if profile and _bound_server is not None and profile != _current_profile(): from agent.secret_scope import build_profile_secret_scope home = _bound_server._profile_home(profile) if home is None: return "" - # An explicit routed profile is authoritative. Never borrow the - # process/default profile's API key on a multiplexed gateway. return str(build_profile_secret_scope(home).get("API_SERVER_KEY") or "").strip() - try: + scoped = "" + with contextlib.suppress(Exception): from agent.secret_scope import get_secret scoped = (get_secret("API_SERVER_KEY", "") or "").strip() - if scoped: - return scoped - except Exception: - pass - return (os.getenv("API_SERVER_KEY") or "").strip() + return scoped or (os.getenv("API_SERVER_KEY") or "").strip() def _profile_execution_policy(profile: str) -> dict: @@ -131,10 +131,7 @@ def _profile_execution_policy(profile: str) -> dict: from hermes_constants import reset_hermes_home_override, set_hermes_home_override token = None if _bound_server is not None and profile not in {_current_profile(), _profile_name()}: - home = _bound_server._profile_home(profile) - if home is None: - raise ValueError(f"profile '{profile}' is unavailable") - token = set_hermes_home_override(str(home)) + token = set_hermes_home_override(str(_foreign_profile_home(profile))) try: return execution_policy_mapping(target_profile=profile) finally: @@ -145,20 +142,17 @@ def _profile_execution_policy(profile: str) -> dict: def _room_link_run_storage_durable() -> bool: """Return whether peer-run replay survives this gateway process.""" if _bound_server is None: - # Direct method-contract tests and embedded callers without a bound API - # server do not expose peer-run transport; production always binds first. + # Embedded callers without a bound server expose no peer-run transport. return True store = getattr(_bound_server, "_run_idempotency_store", None) if store is None: - # The dashboard/TUI process owns groups.* but does not construct the API adapter - # that owns this store. Open the same shared SQLite-backed store lazily so - # capability negotiation reflects the real /v1/runs replay boundary. + # This process does not construct the API adapter that owns the store; open the + # same shared SQLite store lazily so negotiation reflects the real replay boundary. from gateway.platforms.api_server import RunIdempotencyStore with _run_store_lock: store = getattr(_bound_server, "_run_idempotency_store", None) if store is None: - store = RunIdempotencyStore() - _bound_server._run_idempotency_store = store + store = _bound_server._run_idempotency_store = RunIdempotencyStore() return bool(getattr(store, "durable", False)) @@ -175,18 +169,28 @@ def _grant_expiry(claims: dict) -> float: return float(claims.get("status_expires_at", claims["expires_at"])) +def _include_disbanded(params: dict) -> bool: + return params.get("include_disbanded") is True + + +def _room_error_class(replica_only: bool) -> type: + if replica_only: + from gateway.hosted_room_replicas import ReplicaError + return ReplicaError + from gateway.hosted_rooms import HostedRoomError + return HostedRoomError + + def _room_method( name: str, *, code: int, room_code: int | None = None, replica_only: bool = False, with_reason: bool = True, service_code: int | None = None, service_message: str = _DRIVER_UNAVAILABLE, db: bool = False): """Register ``fn`` under ``name`` with the shared hosted-room error envelope. - - ``service_code`` set: the live service is required and passed as a third argument; - when absent the handler fails with that code. ``db``: the default room db path is - passed as the next argument. ``room_code`` maps ``HostedRoomError`` (or only - ``ReplicaError`` when ``replica_only``) to a 4xxx client error, attaching - ``{"reason"}`` data when ``with_reason``; any other exception maps to ``code``. - """ + ``service_code``: the live service is required (else that error) and passed as a third + argument; ``db``: the default room db path follows. ``room_code`` maps ``HostedRoomError`` + (only ``ReplicaError`` when ``replica_only``) to a client error with ``{"reason"}`` data + when ``with_reason``; anything else maps to ``code``.""" + error_class = _room_error_class # closure cell: handlers run under server.py globals def dec(fn): def handler(rid, params: dict) -> dict: @@ -202,16 +206,9 @@ def _room_method( try: return fn(*args) except Exception as exc: - if room_code is not None: - from gateway.hosted_rooms import HostedRoomError - klass = HostedRoomError - if replica_only: - from gateway.hosted_room_replicas import ReplicaError - klass = ReplicaError - if isinstance(exc, klass): - reason = getattr(exc, "reason", None) if with_reason else None - data = {"reason": reason} if reason else None - return _err(rid, room_code, str(exc), data) + if room_code is not None and isinstance(exc, error_class(replica_only)): + reason = getattr(exc, "reason", None) if with_reason else None + return _err(rid, room_code, str(exc), {"reason": reason} if reason else None) return _err(rid, code, str(exc)) handler.__doc__ = fn.__doc__ return method(name)(handler) @@ -233,26 +230,21 @@ def _(rid, params: dict, _catalog=_local_catalog, _methods=_METHODS) -> dict: policy = _profile_execution_policy(profile) catalog = _catalog(local_authority_gateway_id(), profile, policy) room_link = { - "enabled": True, "profile": profile, "catalog": catalog, "endpoint": catalog["endpoint"] - } + "enabled": True, "profile": profile, "catalog": catalog, + "endpoint": catalog["endpoint"]} except Exception: - room_link = { - "enabled": False, - "reason": ( - "durable_run_storage_required" if not _room_link_run_storage_durable() - else "gateway_roomlink_secret_unavailable")} + room_link = {"enabled": False, "reason": ( + "durable_run_storage_required" if not _room_link_run_storage_durable() + else "gateway_roomlink_secret_unavailable")} return _ok(rid, { - "protocol_version": PROTOCOL_VERSION, - "driver": driver_ready, + "protocol_version": PROTOCOL_VERSION, "driver": driver_ready, "persistent_process": bool(room_link.get("catalog", {}).get("persistent_process", False)), - "authority_gateway_id": local_authority_gateway_id(), - "room_link": room_link, + "authority_gateway_id": local_authority_gateway_id(), "room_link": room_link, "features": [ "authority_epoch", "coordinator_fencing", "room_identity", "monotonic_log", "idempotent_send", "replayable_disband", "typed_events", "actor_identity", "log_replication", "authority_takeover"], - "methods": list(_methods), - "max_log_limit": MAX_LOG_LIMIT}) + "methods": list(_methods), "max_log_limit": MAX_LOG_LIMIT}) @_room_method("groups.peer.invite", code=4120, db=True) @@ -295,9 +287,8 @@ def _(rid, params: dict, db_path, _expiry=_grant_expiry) -> dict: profile = _requested_profile(params) claims = decode_room_grant( gateway_room_grant_secret(), str(params.get("grant") or ""), permission="status") - if ( - claims["target_profile"] != profile - or claims["target_install_id"] != local_authority_gateway_id()): + if (claims["target_profile"] != profile + or claims["target_install_id"] != local_authority_gateway_id()): raise ValueError("room grant target does not match this profile") revoke_room_grant_scope(db_path, claims=claims, expires_at=_expiry(claims)) return _ok(rid, {"revoked": True}) @@ -321,36 +312,29 @@ def _(rid, params: dict, service) -> dict: grant = str(params.get("grant") or "") client = PeerRunsHTTPClient(base_url=target_url, api_key="", receipt_db_path=service.db_path) probe = client.probe(grant=grant) - live_catalog = GatewayRoomCatalog.from_mapping(probe.get("catalog")) - if live_catalog != catalog: + # Frozen dataclass equality: an equal live catalog already passed the checks above. + if GatewayRoomCatalog.from_mapping(probe.get("catalog")) != catalog: raise ValueError("target capability catalog changed during setup") - if ( - ROOM_LINK_PROTOCOL_VERSION not in live_catalog.protocol_versions - or "direct" not in live_catalog.link_modes): - raise ValueError("target RoomLink capability is incompatible") room_id = str(params.get("room_id") or "") member_id = str(params.get("member_id") or "") home_install_id = local_authority_gateway_id() home_room = room_state(service.db_path, room_id=room_id) - if ( - probe.get("room_id") != room_id - or probe.get("home_install_id") != home_install_id - or probe.get("authority_gateway_id") != home_room.get("authority_gateway_id") - or int(probe.get("authority_epoch") or 0) != int(home_room.get("authority_epoch") or 0) - or probe.get("member_id") != member_id - or probe.get("target_profile") != target_profile): + expected_scope = { + "room_id": room_id, "home_install_id": home_install_id, + "authority_gateway_id": home_room.get("authority_gateway_id"), + "member_id": member_id, "target_profile": target_profile} + if (any(probe.get(k) != v for k, v in expected_scope.items()) + or int(probe.get("authority_epoch") or 0) + != int(home_room.get("authority_epoch") or 0)): raise ValueError("room grant scope does not match this route") route = PeerMemberRoute( - home_install_id=home_install_id, - member_id=member_id, - target_install_id=catalog.installation_id, - target_profile=target_profile, + home_install_id=home_install_id, member_id=member_id, + target_install_id=catalog.installation_id, target_profile=target_profile, capability_digest=catalog.catalog_digest, execution_policy_digest=catalog.execution_policy.policy_digest, cancellation_scope_id=str( params.get("cancellation_scope_id") or f"cancel-{params.get('room_id') or ''}"), - trace_id=str(params.get("trace_id") or f"trace-{os.urandom(16).hex()}"), - grant=grant) + trace_id=str(params.get("trace_id") or f"trace-{os.urandom(16).hex()}"), grant=grant) service.register_peer_route( room_id=room_id, member_id=member_id, route=route, client=client, target_url=target_url, catalog=catalog) @@ -376,11 +360,7 @@ def _(rid, params: dict, db_path) -> dict: "groups.create", code=5111, room_code=4110, service_code=4123, service_message=_WORKER_UNAVAILABLE) def _(rid, params: dict, service) -> dict: - """Create a hosted room idempotently. - - Required params: ``room_id``, ``name``, and ``members``. Authority is - derived from this gateway's stable install identity, never from the client. - """ + """Create a hosted room idempotently; authority is this gateway's stable install identity.""" room = service.create_room( room_id=params.get("room_id"), name=params.get("name"), members=params.get("members")) return _ok(rid, {"room": room}) @@ -401,33 +381,18 @@ def _(rid, params: dict, db_path) -> dict: @_room_method( - "groups.send", code=5112, room_code=4111, service_code=4123, service_message=_WORKER_UNAVAILABLE -) + "groups.send", code=5112, room_code=4111, service_code=4123, + service_message=_WORKER_UNAVAILABLE) def _(rid, params: dict, service) -> dict: - """Append one typed event to a hosted room idempotently. - - Required params: ``room_id``, ``event_id``, and object ``payload``. Only - inert ``message.user`` events are accepted through this client-facing - method; the actor is server-owned rather than trusted from params. - """ + """Append one typed event idempotently (inert ``message.user`` only; actor is server-owned).""" from gateway.hosted_rooms import user_event_id client_event_id = params.get("event_id") event = service.send( room_id=params.get("room_id"), event_id=user_event_id(client_event_id), payload=params.get("payload")) return _ok(rid, { - "event": event, "client_event_id": client_event_id, "accepted": True, "driver_started": True - }) - - -@_room_method("groups.rename", code=5117, room_code=4117, db=True) -def _(rid, params: dict, db_path) -> dict: - """Rename one hosted room atomically with its replay event.""" - from gateway.hosted_rooms import rename_room - renamed = rename_room( - db_path, room_id=params.get("room_id"), event_id=params.get("event_id"), - name=params.get("name")) - return _ok(rid, {"room": renamed}) + "event": event, "client_event_id": client_event_id, "accepted": True, + "driver_started": True}) @_room_method( @@ -491,79 +456,72 @@ def _(rid, params: dict, service) -> dict: task = {} identity = task.get("identity") receipt = { - **{ - field: str(getattr(identity, field, "") or "") - for field in ("room_id", "task_id", "thread_id", "turn_id")}, + **{f: str(getattr(identity, f, "") or "") + for f in ("room_id", "task_id", "thread_id", "turn_id")}, "status": str(task.get("status") or ""), "execution_generation": int(task.get("execution_generation") or 0), "cancel_generation": int(task.get("cancel_generation") or 0)} return _ok(rid, {"retried": True, "task": receipt}) -@_room_method("groups.log", code=5113, room_code=4112, db=True) -def _(rid, params: dict, db_path) -> dict: - """Return a monotonic room-log delta after ``since_seq``.""" - from gateway.hosted_rooms import read_events - delta = read_events( - db_path, room_id=params.get("room_id"), since_seq=params.get("since_seq", 0), - limit=params.get("limit", 100), include_disbanded=params.get("include_disbanded") is True) - return _ok(rid, delta) +def _passthrough( + name: str, module: str, fn_name: str, doc: str, *, code: int, room_code: int, + params: tuple, replica_only: bool = False, wrap: str | None = None) -> None: + """Register a method whose result is ``module.fn(db_path, **params)`` verbatim (or under key + ``wrap``). ``params`` items are ``key`` (-> ``params.get(key)``) or ``(key, extractor)``.""" + @_room_method( + name, code=code, room_code=room_code, replica_only=replica_only, + with_reason=not replica_only, db=True) + def handler(rid, params_in: dict, db_path, _import=importlib.import_module) -> dict: + kwargs = { + (spec if isinstance(spec, str) else spec[0]): + (params_in.get(spec) if isinstance(spec, str) else spec[1](params_in)) + for spec in params} + result = getattr(_import(module), fn_name)(db_path, **kwargs) + return _ok(rid, {wrap: result} if wrap else result) + handler.__doc__ = doc -@_room_method( - "groups.replicate", code=5116, room_code=4116, replica_only=True, with_reason=False, db=True) -def _(rid, params: dict, db_path) -> dict: - """Persist one authority-stamped replay page into the local replica store. - - ``page`` is the verbatim ``groups.log`` result read from the room's - authority gateway; ingest is idempotent and refuses sequence gaps and - authority-epoch regressions. - """ - from gateway.hosted_room_replicas import ingest_page - result = ingest_page( - db_path, room_id=params.get("room_id"), room_name=params.get("room_name"), - members=params.get("members"), page=params.get("page")) - return _ok(rid, result) - - -@_room_method( - "groups.replica_state", code=5117, room_code=4117, replica_only=True, with_reason=False, db=True -) -def _(rid, params: dict, db_path) -> dict: - """Report the local replica's coverage and authority lineage.""" - from gateway.hosted_room_replicas import replica_state - return _ok(rid, replica_state(db_path, room_id=params.get("room_id"))) +_passthrough( + "groups.rename", "gateway.hosted_rooms", "rename_room", + """Rename one hosted room atomically with its replay event.""", + code=5117, room_code=4117, params=("room_id", "event_id", "name"), wrap="room") +_passthrough( + "groups.log", "gateway.hosted_rooms", "read_events", + """Return a monotonic room-log delta after ``since_seq``.""", + code=5113, room_code=4112, + params=( + "room_id", ("since_seq", lambda p: p.get("since_seq", 0)), + ("limit", lambda p: p.get("limit", 100)), ("include_disbanded", _include_disbanded))) +_passthrough( + "groups.replicate", "gateway.hosted_room_replicas", "ingest_page", + """Persist one authority-stamped replay page (a verbatim ``groups.log`` result) into + the local replica store; idempotent, refuses sequence gaps and epoch regressions.""", + code=5116, room_code=4116, params=("room_id", "room_name", "members", "page"), + replica_only=True) +_passthrough( + "groups.replica_state", "gateway.hosted_room_replicas", "replica_state", + """Report the local replica's coverage and authority lineage.""", + code=5117, room_code=4117, params=("room_id",), replica_only=True) @_room_method("groups.promote", code=5118, room_code=4118, with_reason=False, db=True) def _(rid, params: dict, db_path) -> dict: - """Continue a replicated room on THIS gateway at ``epoch + 1``. - - Requires ``confirm: true`` — the caller asserts the previous authority can - no longer commit (explicit user action; a lease/quorum driver later). - """ + """Continue a replicated room on THIS gateway at ``epoch + 1``. Requires ``confirm: + true`` — the caller asserts the previous authority can no longer commit.""" from gateway.hosted_room_replicas import promote_replica if params.get("confirm") is not True: - return _err( - rid, 4118, - "promotion requires confirm=true acknowledging the previous " - "authority can no longer commit") - result = promote_replica( - db_path, room_id=params.get("room_id"), reason=params.get("reason", "authority-unreachable") - ) - return _ok(rid, result) + return _err(rid, 4118, "promotion requires confirm=true acknowledging the previous " + "authority can no longer commit") + reason = params.get("reason", "authority-unreachable") + return _ok(rid, promote_replica(db_path, room_id=params.get("room_id"), reason=reason)) -@_room_method( - "groups.demote", code=5119, room_code=4119, replica_only=True, with_reason=False, db=True) -def _(rid, params: dict, db_path) -> dict: - """Fence this gateway's stale room authority against a proven newer epoch.""" - from gateway.hosted_room_replicas import demote_room - result = demote_room( - db_path, room_id=params.get("room_id"), - observed_gateway_id=params.get("observed_gateway_id"), - observed_epoch=params.get("observed_epoch")) - return _ok(rid, result) +_passthrough( + "groups.demote", "gateway.hosted_room_replicas", "demote_room", + """Fence this gateway's stale room authority against a proven newer epoch.""", + code=5119, room_code=4119, params=("room_id", "observed_gateway_id", "observed_epoch"), + replica_only=True) def register(server) -> None: diff --git a/tui_gateway/methods_images.py b/tui_gateway/methods_images.py index 7eec464413..595f089d8f 100644 --- a/tui_gateway/methods_images.py +++ b/tui_gateway/methods_images.py @@ -1,8 +1,7 @@ """Image-generation JSON-RPC handler (ws twin of the image_generate tool) for UI surfaces (avatar pickers, artifact panes). The result is a data URL: a remote desktop can't read a -gateway file path and hosted URLs are often CORS-opaque to a renderer canvas. - -Bodies are rebound onto server.py's globals (method_ctx.bind_module) and reference them bare. +gateway file path and hosted URLs are often CORS-opaque to a renderer canvas. Bodies are +rebound onto server.py's globals (method_ctx.bind_module) and reference them bare. """ from .method_ctx import HandlerRegistry, bind_module @@ -11,14 +10,6 @@ _registry = HandlerRegistry() method = _registry.method -def _image_gen_available() -> bool: - try: - from tools.image_generation_tool import check_image_generation_requirements - return bool(check_image_generation_requirements()) - except Exception: - return False - - def _image_to_data_url(ref: str, cap: int): """Fetch a URL or read a local path into a data URL; None when missing, over *cap*, or failing.""" import base64 @@ -43,8 +34,7 @@ def _image_to_data_url(ref: str, cap: int): return None if len(data) > cap: return None - if not mime.startswith("image/"): - mime = "image/png" + mime = mime if mime.startswith("image/") else "image/png" return f"data:{mime};base64,{base64.b64encode(data).decode('ascii')}" except Exception: return None @@ -57,14 +47,17 @@ def _(rid, params: dict) -> dict: on the data URL, default 8MB, max 16MB). Result: ``{available, success, image, image_data, error}`` — ``image_data`` is omitted when the download fails, so callers fall back to ``image`` (the backend's URL/path).""" - available = _image_gen_available() + try: + from tools.image_generation_tool import check_image_generation_requirements + available = bool(check_image_generation_requirements()) + except Exception: + available = False if is_truthy_value(params.get("probe", False)): return _ok(rid, {"available": available}) if not available: return _ok(rid, { "available": False, "success": False, - "error": "No image generation backend configured (run `hermes tools` to enable one).", - }) + "error": "No image generation backend configured (run `hermes tools` to enable one)."}) prompt = str(params.get("prompt") or "").strip() if not prompt: return _err(rid, 4071, "prompt required") @@ -74,9 +67,9 @@ def _(rid, params: dict) -> dict: except (TypeError, ValueError): cap = 8_000_000 try: - from tools.image_generation_tool import _handle_image_generate # Full provider dispatcher — same path as the model tool (source-image confinement, # plugin providers, managed routing, FAL fallback); the FAL leaf bypassed providers. + from tools.image_generation_tool import _handle_image_generate result = json.loads(_handle_image_generate({"prompt": prompt, "aspect_ratio": aspect})) except Exception as e: return _err(rid, 5071, str(e)) @@ -84,11 +77,9 @@ def _(rid, params: dict) -> dict: return _ok(rid, {"available": True, "success": False, "error": str(result.get("error") or "generation failed")}) image_ref = str(result.get("image") or "") - payload = {"available": True, "success": True, "image": image_ref} data_url = _image_to_data_url(image_ref, cap) if image_ref else None - if data_url: - payload["image_data"] = data_url - return _ok(rid, payload) + return _ok(rid, {"available": True, "success": True, "image": image_ref, + **({"image_data": data_url} if data_url else {})}) def register(server) -> None: diff --git a/tui_gateway/methods_profiles.py b/tui_gateway/methods_profiles.py index 9792e3c5eb..65165ce0b4 100644 --- a/tui_gateway/methods_profiles.py +++ b/tui_gateway/methods_profiles.py @@ -1,8 +1,7 @@ """Profile JSON-RPC handlers — the ws twin of the dashboard's /api/profiles (desktop plugins -only have the ws door), on the same `hermes_cli.profiles` primitives. - -Bodies are rebound onto server.py's globals (method_ctx.bind_module) and use them bare; -module-level names are published onto server.py, so they must not collide with its globals. +only have the ws door), on the same `hermes_cli.profiles` primitives. Bodies are rebound onto +server.py's globals (method_ctx.bind_module) and use them bare; module-level names are published +onto server.py, so they must not collide with its globals. """ import contextlib @@ -14,11 +13,13 @@ method = _registry.method # ext -> mime; iteration order is the on-disk lookup order for assets. _ASSET_EXTS = {"png": "image/png", "jpg": "image/jpeg", "webp": "image/webp"} +# ext -> [(start, end, magic bytes)]; format is sniffed, the declared mime is never trusted. +_ASSET_MAGIC = {"png": [(0, 8, b"\x89PNG\r\n\x1a\n")], "jpg": [(0, 3, b"\xff\xd8\xff")], + "webp": [(0, 4, b"RIFF"), (8, 12, b"WEBP")]} def _profile_handler(name: str, code: int): """``@method(name)`` whose body's uncaught exception becomes ``_err(rid, code, str(e))``.""" - def deco(fn): def handler(rid, params: dict) -> dict: try: @@ -39,13 +40,12 @@ def _pin_profile_model(profile_dir, provider, model) -> None: _lazy("hermes_cli.web_routers.profiles", "_write_profile_model")(profile_dir, provider, model) -def _launch_mcp_catalog() -> dict: - mcp = (_lazy("hermes_cli.config", "load_config_readonly")() or {}).get("mcp_servers") - return mcp if isinstance(mcp, dict) else {} +def _model_provider_params(params) -> tuple: + return str(params.get("model") or "").strip(), str(params.get("provider") or "").strip() def _try(fn, default): - """``fn()`` or ``default`` on any exception — best-effort sections must never fail each other.""" + """``fn()`` or ``default`` on any exception (best-effort sections must never fail each other).""" try: return fn() except Exception: @@ -53,7 +53,6 @@ def _try(fn, default): def _best_effort(fn) -> bool: - """Run ``fn``; True on success, False on any exception.""" return _try(lambda: (fn(), True)[1], False) @@ -68,7 +67,7 @@ def _hermes_home_scope(path): def _resolve_profile(rid, params): - """``(name, profile_dir, err)`` — err is the 4063 (name required) / 4064 (not found) response.""" + """``(name, profile_dir, err)``; err = 4063 (name required) / 4064 (not found) response.""" name = str(params.get("name") or "").strip() if not name: return name, None, _err(rid, 4063, "name required") @@ -80,10 +79,12 @@ def _resolve_profile(rid, params): def _read_profile_yaml(profile_dir) -> dict: - """profile.yaml as a mapping; ``{}`` when missing, unparseable, or not a mapping.""" - import yaml - meta_path = profile_dir / "profile.yaml" - loaded = (yaml.safe_load(meta_path.read_text(encoding="utf-8")) or {}) if meta_path.is_file() else {} + """profile.yaml as a mapping; ``{}`` when missing, unreadable, unparseable, or not a mapping.""" + def load(): + import yaml + meta_path = profile_dir / "profile.yaml" + return (yaml.safe_load(meta_path.read_text(encoding="utf-8")) or {}) if meta_path.is_file() else {} + loaded = _try(load, {}) return loaded if isinstance(loaded, dict) else {} @@ -94,7 +95,7 @@ def _clean_revisions(raw: dict) -> dict: def _latest_message_preview(db, session_id): """≤80-char excerpt of the NEWEST active user/assistant message, or "" (roster semantics). - Same query shape as ``SessionDB.latest_message_row_id`` — keep them in step.""" + Same query shape as ``SessionDB.latest_message_row_id``; keep them in step.""" try: with db._lock: row = db._conn.execute( @@ -105,24 +106,13 @@ def _latest_message_preview(db, session_id): (session_id,)).fetchone() except Exception: return "" - if not row: - return "" - text = " ".join(str(row[0] or "").split()).strip() + text = " ".join(str(row[0] or "").split()).strip() if row else "" return text[:80] + "..." if len(text) > 80 else text -def _open_profile_session_db_readonly(profile_path): - """Read-only attach for roster previews, or None (a writable ``SessionDB()`` waits up to 20s - for the write lock + runs DDL and stalled the 5s roster poll).""" - db_path = Path(profile_path) / "state.db" - if not _try(db_path.exists, False): - return None - return _try(lambda: _lazy("hermes_state", "SessionDB")(db_path=db_path, read_only=True), None) - - def _resurrect_recoverable_canonical(db, profile_path, session_id): """Un-archive an accidentally archived canonical row (judged read-only, written via a - short-lived writable handle), or False.""" + short-lived writable handle); False otherwise.""" try: row = db.get_session(session_id) if not row or not row.get("archived"): @@ -145,13 +135,9 @@ def _canonical_session_row(db, profile_path): """Summary of the profile's canonical "Bot Chat" row (identity is the NAME), or None. Lineages via ``get_compression_tip`` (NOT the resume walker's unmarked-child fallback); worker sources count as absent. ``id`` is the registry row, ``resolved_id`` the live tip.""" - if db is None: - return None try: row = db.get_session_by_title("Bot Chat") - if not row: - return None - session_id = str(row.get("id") or "").strip() + session_id = str((row or {}).get("id") or "").strip() if not session_id or _denied_source(row): return None # Archived = retired (absent), except accidental reaper archives: resurrect those. @@ -171,29 +157,22 @@ def _canonical_session_row(db, profile_path): def _latest_profile_session_rows(db): - """(newest human-facing session, newest worker session). The worker row lets rosters show - a profile as working (workers heartbeat ``last_activity_at`` every ≤60s).""" - if db is None: - return None, None + """(newest human-facing session, newest worker session); the worker row lets rosters show a + profile as working (workers heartbeat ``last_activity_at`` every ≤60s).""" try: human = worker = None for s in db.list_sessions_rich(source=None, limit=20, order_by_last_active=True, compact_rows=True): - title = s.get("title") or "" - last_active = s.get("last_active") or s.get("started_at") or 0 + title, last_active = s.get("title") or "", s.get("last_active") or s.get("started_at") or 0 if _denied_source(s): if worker is None: - src = (s.get("source") or "").strip().lower() - worker = {"id": s["id"], "source": src, "title": title, "last_active": last_active} - continue - if human is not None: - continue - # Rosters want "where the conversation IS": prefer the newest text. - human = { - "id": s["id"], "title": title, - "preview": _latest_message_preview(db, s["id"]) or s.get("preview") or "", - "started_at": s.get("started_at") or 0, "last_active": last_active, - "message_count": s.get("message_count") or 0} - if worker is not None: + worker = {"id": s["id"], "source": (s.get("source") or "").strip().lower(), + "title": title, "last_active": last_active} + elif human is None: # rosters want "where the conversation IS": prefer the newest text + human = {"id": s["id"], "title": title, + "preview": _latest_message_preview(db, s["id"]) or s.get("preview") or "", + "started_at": s.get("started_at") or 0, "last_active": last_active, + "message_count": s.get("message_count") or 0} + if human is not None and worker is not None: break return human, worker except Exception: @@ -201,8 +180,13 @@ def _latest_profile_session_rows(db): def _profile_session_fields(row, profile_path): - """Attach last_session / worker_session / canonical_session to a roster row.""" - db = _open_profile_session_db_readonly(profile_path) + """Attach last_session / worker_session / canonical_session to a roster row. The DB is a + read-only attach (a writable ``SessionDB()`` waits up to 20s for the write lock + runs DDL + and stalled the 5s roster poll); no/unreadable DB -> every field None (the readers swallow).""" + db_path = Path(profile_path) / "state.db" + db = None + if _try(db_path.exists, False): + db = _try(lambda: _lazy("hermes_state", "SessionDB")(db_path=db_path, read_only=True), None) try: row["last_session"], row["worker_session"] = _latest_profile_session_rows(db) # Resolved server-side on every listing so no client carries a session pointer. @@ -214,16 +198,13 @@ def _profile_session_fields(row, profile_path): def _profile_ui_meta_fields(row: dict, profile_dir) -> None: """Attach ``ui_meta`` / ``ui_meta_revisions`` / ``has_avatar`` from profile.yaml + assets. - ``ui_meta_revisions`` is always present: it feature-detects gateway-owned CAS even for a - brand-new profile.""" - row["ui_meta_revisions"] = {} - raw_meta = _try(lambda: _read_profile_yaml(profile_dir), {}) - ui_meta = raw_meta.get("ui_meta") + ``ui_meta_revisions`` is always present: it feature-detects gateway-owned CAS for a new profile.""" + raw_meta = _read_profile_yaml(profile_dir) + ui_meta, revisions = raw_meta.get("ui_meta"), raw_meta.get("_ui_meta_revisions") + # Key order is wire-visible: ui_meta_revisions precedes ui_meta. + row["ui_meta_revisions"] = _try(lambda: _clean_revisions(revisions), {}) if isinstance(revisions, dict) else {} if isinstance(ui_meta, dict) and ui_meta: row["ui_meta"] = ui_meta - revisions = raw_meta.get("_ui_meta_revisions") - if isinstance(revisions, dict) and revisions: - row["ui_meta_revisions"] = _try(lambda: _clean_revisions(revisions), {}) # Cheap existence flag so rosters skip a get_asset probe per paint. row["has_avatar"] = _try(lambda: any((profile_dir / "assets" / f"avatar.{e}").is_file() for e in _ASSET_EXTS), False) @@ -236,29 +217,22 @@ def _(rid, params: dict) -> dict: include_sessions = is_truthy_value(params.get("include_sessions", True)) out = [] for p in list_profiles(): - row = { - "name": p.name, "path": str(p.path), "is_default": bool(p.is_default), - "model": p.model, "provider": p.provider, - "description": p.description or "", "display_name": p.display_name or "", - "skill_count": p.skill_count or 0} + row = {"name": p.name, "path": str(p.path), "is_default": bool(p.is_default), "model": p.model, + "provider": p.provider, "description": p.description or "", + "display_name": p.display_name or "", "skill_count": p.skill_count or 0} if include_sessions: _profile_session_fields(row, p.path) _profile_ui_meta_fields(row, Path(str(p.path))) out.append(row) - # Capability flag: this backend injects the Bot Mode teammate-messaging - # protocol into every session, so clients must not append it to SOUL.md. + # bot_mode_protocol: this backend injects the Bot Mode teammate-messaging protocol into every + # session, so clients must not append it to SOUL.md. return _ok(rid, {"profiles": out, "bot_mode_protocol": True}) -def _has_real_env_content(env_path) -> bool: - """True when .env has any non-comment, non-blank line.""" - lines = env_path.read_text(encoding="utf-8", errors="replace").splitlines() - return any(s and not s.startswith("#") for s in map(str.strip, lines)) - - -def _copy_secret_file(src, dst, wanted: bool) -> bool: - """Copy ``src`` -> ``dst`` (0600) when ``src`` exists and ``wanted``; True if copied.""" - if not (src.is_file() and wanted): +def _mirror_secret(path, launch_home, name: str, wanted) -> bool: + """Copy the launch ``name`` file into the profile (0600) when it exists and ``wanted(src, dst)``.""" + src, dst = launch_home / name, path / name + if not (src.is_file() and wanted(src, dst)): return False import shutil shutil.copy2(src, dst) @@ -267,28 +241,14 @@ def _copy_secret_file(src, dst, wanted: bool) -> bool: return True -def _mirror_env(path, launch_home) -> bool: - """Copy the launch .env only over the seeded comment-only stub (never a clone's secrets).""" - src, dst = launch_home / ".env", path / ".env" - return _copy_secret_file( - src, dst, _has_real_env_content(src) and not _try(lambda: _has_real_env_content(dst), False)) - - -def _mirror_auth(path, launch_home) -> bool: - """Copy the launch auth.json when absent (skipped under ``share_auth``: a copy forks token - state and the first refresh in either store strands the other).""" - src, dst = launch_home / "auth.json", path / "auth.json" - if not _copy_secret_file(src, dst, not dst.exists()): - return False - # Drop single-use OAuth grants (first refresh strands every sibling); they read from - # the root grant via the pool fallback. API keys stay. - _best_effort(lambda: _lazy("hermes_cli.auth", "strip_cloned_single_use_oauth_grants")(path)) - return True +def _env_has_content(env_path) -> bool: + lines = env_path.read_text(encoding="utf-8", errors="replace").splitlines() + return any(s and not s.startswith("#") for s in map(str.strip, lines)) def _mirror_voice_sections(path) -> bool: """Copy stt/tts/voice sections from the launch profile (a fresh profile has only ``model``, - so voice fell back to defaults); True if written. Canonical loaders under the home override.""" + so voice fell back to defaults); True if written.""" try: from hermes_cli.config import load_config_readonly, read_user_config_raw, save_config src_cfg = load_config_readonly() or {} @@ -300,8 +260,7 @@ def _mirror_voice_sections(path) -> bool: dst_cfg = read_user_config_raw() or {} missing = {k: v for k, v in sections.items() if k not in dst_cfg} if missing: - dst_cfg.update(missing) - save_config(dst_cfg) + save_config({**dst_cfg, **missing}) return bool(missing) except Exception: return False @@ -316,10 +275,9 @@ def _inherit_launch_model(path) -> bool: if dst_model.get("provider") and dst_model.get("default"): return False model_cfg = (load_config_readonly() or {}).get("model") or {} - provider, model = str(model_cfg.get("provider") or ""), str(model_cfg.get("default") or "") - if not (provider and model): + if not (model_cfg.get("provider") and model_cfg.get("default")): return False - _pin_profile_model(path, provider, model) + _pin_profile_model(path, str(model_cfg["provider"]), str(model_cfg["default"])) return True @@ -327,16 +285,22 @@ def _mirror_launch_credentials(path, params: dict) -> dict: """Copy launch .env / auth.json / voice sections into a new profile (best-effort per item). ``share_auth`` reports ``auth: "shared"`` and skips the auth copy; ``mirror_credentials`` false skips everything. ``model_inherited`` is filled in by the caller.""" - mirrored = {"env": False, "auth": False, "model_inherited": False, "voice": False} share_auth = is_truthy_value(params.get("share_auth", False)) - if share_auth: - mirrored["auth"] = "shared" + mirrored = {"env": False, "auth": "shared" if share_auth else False, "model_inherited": False, + "voice": False} if not is_truthy_value(params.get("mirror_credentials", True)): return mirrored launch_home = get_hermes_home() - mirrored["env"] = _try(lambda: _mirror_env(path, launch_home), False) - if not share_auth: - mirrored["auth"] = _try(lambda: _mirror_auth(path, launch_home), False) + # .env: only over the seeded comment-only stub (never a clone's secrets). + mirrored["env"] = _try(lambda: _mirror_secret(path, launch_home, ".env", lambda src, dst: ( + _env_has_content(src) and not _try(lambda: _env_has_content(dst), False))), False) + if not share_auth: # a copy forks token state: the first refresh in either store strands the other + mirrored["auth"] = _try(lambda: _mirror_secret(path, launch_home, "auth.json", + lambda src, dst: not dst.exists()), False) + if mirrored["auth"]: + # Drop single-use OAuth grants (first refresh strands every sibling); they read from the + # root grant via the pool fallback. API keys stay. + _best_effort(lambda: _lazy("hermes_cli.auth", "strip_cloned_single_use_oauth_grants")(path)) mirrored["voice"] = _mirror_voice_sections(path) return mirrored @@ -345,9 +309,8 @@ def _mirror_launch_credentials(path, params: dict) -> dict: def _(rid, params: dict) -> dict: """Create a profile (ws twin of POST /api/profiles). Params: ``name``, ``description``, ``clone_from`` (omitted = fresh + bundled skills), ``clone_all``, ``no_skills``, ``soul``, - ``model`` + ``provider``, ``share_auth``, ``mirror_credentials`` (default true — a - ``create_profile()`` seeds a comment-only .env and no auth.json, so a headless profile had - NO provider).""" + ``model`` + ``provider``, ``share_auth``, ``mirror_credentials`` (default true: a bare + ``create_profile()`` seeds a comment-only .env and no auth.json = NO provider headless).""" name = str(params.get("name") or "").strip() if not name: return _err(rid, 4061, "name required") @@ -369,12 +332,10 @@ def _(rid, params: dict) -> dict: _best_effort(lambda: profiles_mod.seed_profile_skills(path, quiet=True)) _best_effort(lambda: profiles_mod.check_alias_collision(name) or profiles_mod.create_wrapper_script(name)) soul = params.get("soul") - soul_written = False - if isinstance(soul, str) and soul.strip(): - soul_written = _best_effort(lambda: (path / "SOUL.md").write_text(soul, encoding="utf-8")) + soul_written = isinstance(soul, str) and bool(soul.strip()) and _best_effort( + lambda: (path / "SOUL.md").write_text(soul, encoding="utf-8")) mirrored = _mirror_launch_credentials(path, params) - model = str(params.get("model") or "").strip() - provider = str(params.get("provider") or "").strip() + model, provider = _model_provider_params(params) model_set = False if model and provider: model_set = _best_effort(lambda: _pin_profile_model(path, provider, model)) @@ -386,7 +347,7 @@ def _(rid, params: dict) -> dict: def _describe_toolsets(cfg): """``(toolsets, pinned_set)`` as the `hermes tools` checklist presents them (the raw registry - leaks platform composites and reports everything "enabled" without a pin).""" + leaks platform composites and reports everything enabled without a pin).""" from hermes_cli.tools_config import ( _get_effective_configurable_toolsets, _get_platform_tools, _toolset_allowed_for_platform) from toolsets import resolve_toolset @@ -396,34 +357,21 @@ def _describe_toolsets(cfg): default_off = _try(lambda: _lazy("hermes_cli.tools_config", "_DEFAULT_OFF_TOOLSETS"), set()) toolsets_out = [] for ts_name, ts_label, ts_desc in _get_effective_configurable_toolsets(): - if not _toolset_allowed_for_platform(ts_name, "cli"): - continue - enabled = ts_name in pinned_set if pinned_set is not None else ts_name in platform_enabled + enabled = ts_name in (pinned_set if pinned_set is not None else platform_enabled) # Default-off integrations (+ opt-in yuanbao) are noise unless already enabled. - if (ts_name in default_off or ts_name == "yuanbao") and not enabled: + if not _toolset_allowed_for_platform(ts_name, "cli") or ( + (ts_name in default_off or ts_name == "yuanbao") and not enabled): continue - tool_count = _try(lambda: len(set(resolve_toolset(ts_name))), 0) toolsets_out.append({"name": ts_name, "label": ts_label, "description": ts_desc or "", - "tool_count": tool_count, "enabled": enabled}) + "tool_count": _try(lambda: len(set(resolve_toolset(ts_name))), 0), + "enabled": enabled}) return toolsets_out, pinned_set -def _describe_mcp_servers(cfg): - """``[{name, enabled, transport}]`` for the profile's ``mcp_servers`` (best-effort).""" - mcp_cfg = cfg.get("mcp_servers") - if not isinstance(mcp_cfg, dict): - return [] - return _try(lambda: [ - {"name": str(srv_name), "enabled": not is_truthy_value(entry.get("disabled", False)), - "transport": str(entry.get("transport") or "http") if entry.get("url") else "stdio"} - for srv_name in sorted(mcp_cfg.keys()) for entry in (mcp_cfg[srv_name],) - if isinstance(entry, dict) - ], []) - - @_profile_handler("profiles.describe", 5063) def _(rid, params: dict) -> dict: - """Editor snapshot; installed skills are enabled unless in ``skills.disabled``.""" + """Editor snapshot; installed skills are enabled unless in ``skills.disabled``; ``mcp_servers`` + is ``[{name, enabled, transport}]`` (best-effort).""" name, profile_dir, err = _resolve_profile(rid, params) if err is not None: return err @@ -439,7 +387,13 @@ def _(rid, params: dict) -> dict: toolsets_out, pinned_set = _describe_toolsets(cfg) soul_path = profile_dir / "SOUL.md" soul = _try(lambda: soul_path.read_text(encoding="utf-8", errors="replace") if soul_path.is_file() else "", "") - mcp_out = _describe_mcp_servers(cfg) + mcp_cfg = cfg.get("mcp_servers") + mcp_out = _try(lambda: [ + {"name": str(srv_name), "enabled": not is_truthy_value(entry.get("disabled", False)), + "transport": str(entry.get("transport") or "http") if entry.get("url") else "stdio"} + for srv_name in sorted(mcp_cfg.keys()) for entry in (mcp_cfg[srv_name],) + if isinstance(entry, dict) + ], []) if isinstance(mcp_cfg, dict) else [] model_cfg = cfg.get("model") if isinstance(cfg.get("model"), dict) else {} meta = _try(lambda: _lazy("hermes_cli.profiles", "read_profile_meta")(profile_dir), {}) return _ok(rid, { @@ -454,16 +408,16 @@ def _configure_ui_meta(profile_dir, params, applied) -> None: """Merge ``params["ui_meta"]`` key-wise into profile.yaml (None deletes). 64KB cap (rides every roster paint). ``ui_meta_expected_revisions``: per-key CAS, any mismatch rejects the whole write; revisions survive deletion so a stale client cannot recreate a removed key.""" + applied["ui_meta"] = False try: incoming = params["ui_meta"] if len(json.dumps(incoming)) > 65536: - applied["ui_meta"] = False return expected = params.get("ui_meta_expected_revisions") if expected is not None and not isinstance(expected, dict): raise ValueError("ui_meta_expected_revisions must be an object") with _profile_ui_meta_lock: - existing = _try(lambda: _read_profile_yaml(profile_dir), {}) + existing = _read_profile_yaml(profile_dir) raw_revisions = existing.get("_ui_meta_revisions") revisions = _clean_revisions(raw_revisions if isinstance(raw_revisions, dict) else {}) conflicts = {} @@ -472,7 +426,6 @@ def _configure_ui_meta(profile_dir, params, applied) -> None: if not isinstance(wanted, int) or isinstance(wanted, bool) or wanted < 0 or wanted != actual: conflicts[key] = {"expected": wanted, "actual": actual} if conflicts: - applied["ui_meta"] = False applied["ui_meta_conflicts"] = conflicts applied["ui_meta_revisions"] = {key: revisions.get(key, 0) for key in incoming} return @@ -499,12 +452,11 @@ def _configure_ui_meta(profile_dir, params, applied) -> None: def _configure_model(profile_dir, params, applied): """Apply a ``model`` + ``provider`` pin, or return a confirm message and write NOTHING (client - resends with ``confirm_expensive_model``). A failing guard = "no warning" (as _apply_model_switch).""" - model = str(params.get("model") or "").strip() - provider = str(params.get("provider") or "").strip() - confirm_message = None + resends with ``confirm_expensive_model``). A failing guard = no warning (as _apply_model_switch).""" + model, provider = _model_provider_params(params) if not (model and provider): return None + confirm_message = None if not is_truthy_value(params.get("confirm_expensive_model", False)): warn = _lazy("hermes_cli.model_selection_guards", "combined_selection_warning") confirm_message = _try(lambda: getattr(warn(model, provider=provider or None), "message", None), None) @@ -513,31 +465,6 @@ def _configure_model(profile_dir, params, applied): return confirm_message -def _configure_cfg_sections(profile_dir, params, applied) -> None: - """Apply ``disabled_skills`` / ``enabled_toolsets`` / ``enabled_mcp_servers`` (replace - semantics; empty toolsets clears the pin). An undefined MCP server is copied from the LAUNCH - catalog (unknown names skipped); credentials stay in .env/auth.""" - want_mcp = isinstance(params.get("enabled_mcp_servers"), list) - # Launch catalog read BEFORE the home override flips config resolution. - launch_mcp = _try(_launch_mcp_catalog, {}) if want_mcp else {} - with _hermes_home_scope(profile_dir): - from hermes_cli.config import load_config, save_config - cfg = load_config() or {} - if isinstance(params.get("disabled_skills"), list): - try: - from hermes_cli.skills_config import save_disabled_skills - save_disabled_skills(cfg, _clean_names(params["disabled_skills"])) - applied["skills"] = True - cfg = load_config() or {} - except Exception: - applied["skills"] = False - if isinstance(params.get("enabled_toolsets"), list): - applied["toolsets"] = _best_effort(lambda: _save_toolset_pin(cfg, params["enabled_toolsets"], save_config)) - if want_mcp: - applied["mcp_servers"] = _best_effort(lambda: _save_mcp_toggles( - load_config() or {}, params["enabled_mcp_servers"], launch_mcp, save_config)) - - def _clean_names(values) -> set: return {str(v).strip() for v in values if str(v).strip()} @@ -557,10 +484,9 @@ def _save_mcp_toggles(cfg, enabled, launch_mcp, save_config) -> None: wanted = _clean_names(enabled) mcp_cfg = cfg.get("mcp_servers") if isinstance(cfg.get("mcp_servers"), dict) else {} for srv in wanted: - if srv in mcp_cfg and isinstance(mcp_cfg[srv], dict): - mcp_cfg[srv].pop("disabled", None) - elif srv in launch_mcp and isinstance(launch_mcp[srv], dict): + if not isinstance(mcp_cfg.get(srv), dict) and isinstance(launch_mcp.get(srv), dict): mcp_cfg[srv] = dict(launch_mcp[srv]) + if isinstance(mcp_cfg.get(srv), dict): mcp_cfg[srv].pop("disabled", None) for srv, entry in mcp_cfg.items(): if srv not in wanted and isinstance(entry, dict): @@ -570,11 +496,39 @@ def _save_mcp_toggles(cfg, enabled, launch_mcp, save_config) -> None: save_config(cfg) +def _configure_cfg_sections(profile_dir, params, applied) -> None: + """Apply ``disabled_skills`` / ``enabled_toolsets`` / ``enabled_mcp_servers`` (replace + semantics; empty toolsets clears the pin). An undefined MCP server is copied from the LAUNCH + catalog (unknown names skipped); credentials stay in .env/auth.""" + want_mcp = isinstance(params.get("enabled_mcp_servers"), list) + launch_mcp = {} + if want_mcp: # launch catalog read BEFORE the home override flips config resolution + load_launch = _lazy("hermes_cli.config", "load_config_readonly") + launch_mcp = _try(lambda: (load_launch() or {}).get("mcp_servers"), {}) + launch_mcp = launch_mcp if isinstance(launch_mcp, dict) else {} + with _hermes_home_scope(profile_dir): + from hermes_cli.config import load_config, save_config + cfg = load_config() or {} + if isinstance(params.get("disabled_skills"), list): + try: + from hermes_cli.skills_config import save_disabled_skills + save_disabled_skills(cfg, _clean_names(params["disabled_skills"])) + applied["skills"] = True + cfg = load_config() or {} + except Exception: + applied["skills"] = False + if isinstance(params.get("enabled_toolsets"), list): + applied["toolsets"] = _best_effort(lambda: _save_toolset_pin(cfg, params["enabled_toolsets"], save_config)) + if want_mcp: + applied["mcp_servers"] = _best_effort(lambda: _save_mcp_toggles( + load_config() or {}, params["enabled_mcp_servers"], launch_mcp, save_config)) + + @_profile_handler("profiles.configure", 5064) def _(rid, params: dict) -> dict: """Editor Save: ``name`` plus any of ``ui_meta`` (+ ``ui_meta_expected_revisions``), ``soul``, ``description``, ``model`` + ``provider`` (+ ``confirm_expensive_model``), ``disabled_skills``, - ``enabled_toolsets``, ``enabled_mcp_servers``. Sections are independent; ``applied`` reports each.""" + ``enabled_toolsets``, ``enabled_mcp_servers``; sections are independent, ``applied`` reports each.""" _name, profile_dir, err = _resolve_profile(rid, params) if err is not None: return err @@ -590,35 +544,22 @@ def _(rid, params: dict) -> dict: confirm_message = _configure_model(profile_dir, params, applied) if any(isinstance(params.get(k), list) for k in ("disabled_skills", "enabled_toolsets", "enabled_mcp_servers")): _configure_cfg_sections(profile_dir, params, applied) - result = {"ok": all(applied.values()) if applied else True, "applied": applied} - if confirm_message is not None: - # Same shape config.set returns, so clients reuse one confirm handler. - result["confirm_required"] = True - result["confirm_message"] = confirm_message - return _ok(rid, result) - - -def _sniff_asset_ext(blob): - """Extension for a PNG/JPEG/WebP blob by magic bytes (never trust declared mime), or None.""" - if blob[:8] == b"\x89PNG\r\n\x1a\n": - return "png" - if blob[:3] == b"\xff\xd8\xff": - return "jpg" - return "webp" if blob[:4] == b"RIFF" and blob[8:12] == b"WEBP" else None + # confirm_* is the shape config.set returns, so clients reuse one confirm handler. + return _ok(rid, {"ok": all(applied.values()) if applied else True, "applied": applied, + **({"confirm_required": True, "confirm_message": confirm_message} + if confirm_message is not None else {})}) def _unlink_asset_files(assets_dir, asset) -> int: """Delete every ``.`` in ``assets_dir``; returns how many existed.""" present = [t for t in (assets_dir / f"{asset}.{ext}" for ext in _ASSET_EXTS) if t.is_file()] - for target in present: - target.unlink() - return len(present) + return len([t.unlink() for t in present]) @_profile_handler("profiles.set_asset", 5065) def _(rid, params: dict) -> dict: """Store ``assets/.`` atomically. Params: ``name``, ``asset`` (``"avatar"`` only), - ``data`` (data URL or base64; PNG/JPEG/WebP ≤2MB) or ``clear: true``.""" + ``data`` (data URL or base64; PNG/JPEG/WebP ≤2MB, format sniffed) or ``clear: true``.""" asset = str(params.get("asset") or "avatar").strip().lower() if not str(params.get("name") or "").strip(): return _err(rid, 4063, "name required") @@ -631,8 +572,7 @@ def _(rid, params: dict) -> dict: return err assets_dir = profile_dir / "assets" if is_truthy_value(params.get("clear", False)): - removed = _unlink_asset_files(assets_dir, asset) - return _ok(rid, {"ok": True, "asset": asset, "size": 0, "removed": removed}) + return _ok(rid, {"ok": True, "asset": asset, "size": 0, "removed": _unlink_asset_files(assets_dir, asset)}) data = str(params.get("data") or "") if not data: return _err(rid, 4067, "data required (data URL or base64)") @@ -643,15 +583,14 @@ def _(rid, params: dict) -> dict: return _err(rid, 4068, "data is not valid base64") if len(blob) > 2_000_000: return _err(rid, 4069, f"asset too large ({len(blob)} bytes; max 2MB)") - ext = _sniff_asset_ext(blob) + ext = next((e for e, magic in _ASSET_MAGIC.items() if all(blob[a:b] == m for a, b, m in magic)), None) if ext is None: return _err(rid, 4070, "unsupported image format (PNG/JPEG/WebP only)") assets_dir.mkdir(parents=True, exist_ok=True) _unlink_asset_files(assets_dir, asset) # one canonical file per asset - target = assets_dir / f"{asset}.{ext}" - tmp = target.with_suffix(target.suffix + ".tmp") + tmp = assets_dir / f"{asset}.{ext}.tmp" tmp.write_bytes(blob) - tmp.replace(target) + tmp.replace(assets_dir / f"{asset}.{ext}") return _ok(rid, {"ok": True, "asset": asset, "size": len(blob)}) @@ -667,8 +606,8 @@ def _(rid, params: dict) -> dict: target = profile_dir / "assets" / f"{asset}.{ext}" if target.is_file(): blob = target.read_bytes() - data = f"data:{mime};base64,{base64.b64encode(blob).decode('ascii')}" - return _ok(rid, {"found": True, "mime": mime, "size": len(blob), "data": data}) + return _ok(rid, {"found": True, "mime": mime, "size": len(blob), + "data": f"data:{mime};base64,{base64.b64encode(blob).decode('ascii')}"}) return _ok(rid, {"found": False}) diff --git a/tui_gateway/methods_projects.py b/tui_gateway/methods_projects.py index 6a6d7cac0d..0023eb4d5c 100644 --- a/tui_gateway/methods_projects.py +++ b/tui_gateway/methods_projects.py @@ -1,8 +1,5 @@ """Projects RPC surface: per-profile multi-folder workspaces, repo discovery, sidebar tree. - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so they reference server.py globals bare. -""" +Bodies are rebound onto server.py's globals at install (method_ctx.bind_module).""" from __future__ import annotations @@ -12,10 +9,8 @@ _registry = HandlerRegistry() method = _registry.method -# JSON-RPC error codes for the projects surface. -_E_PROJECTS = 5061 # generic failure -_E_NO_PROJECT = 5062 # id resolved to nothing -_E_PROJECT_ARG = 5063 # invalid argument (e.g. bad name/slug) +# JSON-RPC error codes: generic failure / id resolved to nothing / invalid argument. +_E_PROJECTS, _E_NO_PROJECT, _E_PROJECT_ARG = 5061, 5062, 5063 class _NoProject(Exception): @@ -30,13 +25,8 @@ def _projects_payload(conn) -> dict: def _projects_method(name: str): - """Register a projects RPC, injecting (pdb, conn) and unifying error mapping. - - Binds ``params['profile']`` (via ``@_profile_scoped``) so app-global remote - mode reads that profile's ``projects.db``. Missing id maps to 5062, bad args - to 5063, everything else to 5061. - """ - + """Register a projects RPC, injecting (pdb, conn) and unifying error mapping; profile-scoped + so app-global remote mode reads that profile's ``projects.db``.""" def decorator(fn): @method(name) @_registry.profile_scoped @@ -67,20 +57,9 @@ def _pick(params: dict, *keys: str) -> dict: return {k: params.get(k) for k in keys} -# Per-project mutators: (rpc suffix, pdb function, takes params['path'], extra kwargs). -# Each resolves ``params['id']`` (5062 when missing), mutates, and answers with the -# refreshed project. -_PROJECT_MUTATORS = ( - ("update", "update_project", False, - lambda p: _pick(p, "name", "description", "icon", "color", "board_slug")), - ("add_folder", "add_folder", True, - lambda p: {"label": p.get("label"), "is_primary": bool(p.get("is_primary"))}), - ("remove_folder", "remove_folder", True, lambda p: {}), - ("set_primary", "set_primary", True, lambda p: {}), -) - - def _register_project_mutator(suffix: str, fn_name: str, takes_path: bool, kwargs_of) -> None: + """``projects.``: resolve ``params['id']`` (5062 when missing), call + ``pdb.(conn, id[, path], **kwargs_of(params))``, answer with the refreshed project.""" @_projects_method(f"projects.{suffix}") def _(rid, params, pdb, conn) -> dict: proj = _require_project(pdb, conn, params) @@ -89,9 +68,14 @@ def _register_project_mutator(suffix: str, fn_name: str, takes_path: bool, kwarg return _ok(rid, {"project": pdb.get_project(conn, proj.id).to_dict()}) -for _spec in _PROJECT_MUTATORS: - _register_project_mutator(*_spec) -del _spec +_register_project_mutator( + "update", "update_project", False, + lambda p: _pick(p, "name", "description", "icon", "color", "board_slug")) +_register_project_mutator( + "add_folder", "add_folder", True, + lambda p: {"label": p.get("label"), "is_primary": bool(p.get("is_primary"))}) +_register_project_mutator("remove_folder", "remove_folder", True, lambda p: {}) +_register_project_mutator("set_primary", "set_primary", True, lambda p: {}) @_projects_method("projects.list") @@ -136,25 +120,26 @@ def _(rid, params, pdb, conn) -> dict: @_projects_method("projects.for_cwd") def _(rid, params, pdb, conn) -> dict: - cwd = _completion_cwd({"cwd": str(params.get("cwd") or "").strip()} if params.get("cwd") else {}) + cwd = _completion_cwd( + {"cwd": str(params.get("cwd") or "").strip()} if params.get("cwd") else {}) proj = pdb.project_for_path(conn, cwd) - return _ok(rid, {"project": proj.to_dict() if proj else None, "cwd": cwd, "branch": _git_branch_for_cwd(cwd)}) + return _ok(rid, { + "project": proj.to_dict() if proj else None, "cwd": cwd, + "branch": _git_branch_for_cwd(cwd)}) def _non_workspace_dirs() -> set[str]: - """Never-a-workspace dirs: ``/``, the user's home, and the dir homes live in. Both - POSIX spellings are excluded on every host (macOS ships an empty ``/home`` autofs - stub; containers/remote shells hand back Linux paths) — promoting one mints a - catch-all project and ``/home`` renders as a second "home" row beside Home.""" + """Never-a-workspace dirs: ``/``, the user's home, the dir homes live in, plus both POSIX + spellings on every host (remote shells hand back Linux paths; promoting one mints a + catch-all project).""" home = os.path.realpath(os.path.expanduser("~")) candidates = (os.sep, home, os.path.dirname(home), "/home", "/Users") return {os.path.normcase(os.path.realpath(path)) for path in candidates if path} def _is_repo_junk(root: str) -> bool: - """A git root never auto-surfaced as a project: a non-workspace dir or anything - under HERMES_HOME (config/sessions/skills). User-created projects pointing - there are still honored.""" + """A git root never auto-surfaced as a project: a non-workspace dir or anything under + HERMES_HOME. User-created projects pointing there are still honored.""" if not root: return True from hermes_constants import get_hermes_home @@ -167,9 +152,8 @@ def _is_repo_junk(root: str) -> bool: def _is_session_cwd_junk(cwd: str) -> bool: - """A non-git cwd that stays in flat Recents rather than auto-grouping. Unlike git - roots, a selected DESCENDANT of HERMES_HOME may be an intentional prose/data - workspace, so only HERMES_HOME itself and ``_non_workspace_dirs`` are excluded.""" + """A non-git cwd that stays in flat Recents. A DESCENDANT of HERMES_HOME may be an + intentional prose/data workspace, so only HERMES_HOME itself is excluded here.""" if not cwd: return True from hermes_constants import get_hermes_home @@ -194,25 +178,19 @@ def _repo_discovery_policy(raw: dict | None = None) -> dict: if not isinstance(values, list): return list(defaults[long]) return [v.strip() for v in values if isinstance(v, str) and v.strip()] - enabled = _get("enabled", "repo_scan_enabled") return { "enabled": enabled if isinstance(enabled, bool) else defaults["repo_scan_enabled"], "roots": _paths("roots", "repo_scan_roots"), - "exclude_paths": _paths("exclude_paths", "repo_scan_exclude_paths"), - } + "exclude_paths": _paths("exclude_paths", "repo_scan_exclude_paths")} def _repo_discovery_policy_key(policy: dict) -> str: def _paths(values: list[str]) -> list[str]: - normalized = set() home = os.path.expanduser("~") - for value in values: - expanded = os.path.expanduser(value) - if not os.path.isabs(expanded): - expanded = os.path.join(home, expanded) - normalized.add(os.path.normcase(os.path.abspath(expanded))) - return sorted(normalized) + return sorted({ + os.path.normcase(os.path.abspath(os.path.join(home, os.path.expanduser(v)))) + for v in values}) canonical = { "enabled": bool(policy["enabled"]), "roots": _paths(policy["roots"]), "exclude_paths": _paths(policy["exclude_paths"])} @@ -226,11 +204,10 @@ def _repo_discovery_policy_is_default(policy: dict) -> bool: def _scan_discovered_repos_remote(conn, policy: dict) -> bool: - """Backend-side disk scan of the policy roots into the discovery cache (the desktop's - native scan only sees the local filesystem). Best-effort: failures log and leave - the cache untouched. Returns True only when the scan is authoritative (every root - walked to completion, cap not hit) — only then is the cache write ``replace=True``; - a partial/errored scan must MERGE, never wipe, or a failed refresh blanks the sidebar.""" + """Backend-side disk scan of the policy roots into the discovery cache. Best-effort: + failures log and leave the cache untouched. True only when the scan is authoritative + (every root walked to completion, cap not hit) — only then is the cache write + ``replace=True``; a partial/errored scan must MERGE, or a failed refresh blanks the sidebar.""" from hermes_cli import projects_db as pdb roots = policy.get("roots") or [] excludes = policy.get("exclude_paths") or [] @@ -239,11 +216,11 @@ def _scan_discovered_repos_remote(conn, policy: dict) -> bool: authoritative = True def _is_excluded(path: str) -> bool: - return any(path == ex or path.startswith(ex.rstrip("/\\") + os.sep) for ex in excludes if ex) + return any( + path == ex or path.startswith(ex.rstrip("/\\") + os.sep) for ex in excludes if ex) for root in roots: if not os.path.isdir(root): - # `os.walk` on a missing root yields nothing instead of raising; an unmounted - # volume would look like an empty scan and let the replace wipe its cache. + # `os.walk` on a missing root yields nothing; an unmounted volume must not wipe. authoritative = False logger.debug("discover_repos scan root missing, skipping: %s", root) continue @@ -264,8 +241,7 @@ def _scan_discovered_repos_remote(conn, policy: dict) -> bool: except Exception: authoritative = False logger.debug("discover_repos scan failed for root %s", root, exc_info=True) - if len(pairs) >= 500: - # Cap hit: the walk didn't cover the full roots -> not authoritative. + if len(pairs) >= 500: # cap hit: the walk didn't cover the full roots authoritative = False break if pairs: @@ -280,18 +256,16 @@ def _scan_discovered_repos_remote(conn, policy: dict) -> bool: def _discover_repos_payload( db, *, conn=None, backfill: bool = True, include_cached: bool = True) -> list[dict]: - """Merge filesystem-scanned repos (cached; may have zero sessions) with - session-derived roots, junk-filtered, with session totals. ``conn`` reuses an open - projects.db connection; ``backfill`` persists resolved roots onto session rows — - kept OFF the per-turn tree path and done only on explicit refresh.""" + """Merge cached filesystem-scanned repos with session-derived roots, junk-filtered, with + session totals. ``backfill`` persists resolved roots onto session rows — kept OFF the + per-turn tree path and done only on explicit refresh.""" repos: dict[str, dict] = {} def _agg(root: str) -> dict: - return repos.setdefault(root, {"root": root, "label": "", "sessions": 0, "last_active": 0.0}) - + return repos.setdefault( + root, {"root": root, "label": "", "sessions": 0, "last_active": 0.0}) cwd_rows = list(db.distinct_session_cwds()) - # Warm the per-cwd git probes in parallel so a cold first paint doesn't - # serialize one subprocess per distinct cwd before this loop reads the cache. + # Parallel-warm the per-cwd git probes so a cold first paint doesn't serialize them. git_probe.warm_roots(str(r.get("cwd") or "") for r in cwd_rows) cwd_to_root: dict[str, str] = {} for row in cwd_rows: @@ -311,8 +285,7 @@ def _discover_repos_payload( except Exception: logger.debug("failed to backfill repo roots", exc_info=True) if include_cached: - # `last_seen` is scan time, not user activity — never fold it into - # `last_active` (made every scanned repo "just now"). + # `last_seen` is scan time, not user activity — never fold it into `last_active`. try: from hermes_cli import projects_db as pdb with (contextlib.nullcontext(conn) if conn is not None else pdb.connect_closing()) as c: @@ -330,27 +303,23 @@ def _discover_repos_payload( return out -# Not user conversations (cron has its own section; kanban runs are read on -# the board). Subagent/compression children are dropped by include_children=False. +# Not user conversations; subagent/compression children are dropped by include_children=False. _PROJECT_TREE_EXCLUDED_SOURCES = ["cron", "kanban"] def _project_tree_row(r: dict) -> dict: - """Project a SessionDB row to the minimal shape the sidebar renders: the - grouping fields (cwd/git_branch/git_repo_root) + everything ``SidebarSessionRow`` - reads (parent_session_id for the └─ connector, cost for Show → cost), minus the - heavy columns.""" + """Project a SessionDB row to the minimal shape the sidebar renders (grouping fields + + what ``SidebarSessionRow`` reads), minus the heavy columns.""" row = {k: r.get(k) for k in ( "id", "_lineage_root_id", "_lineage_ids", "parent_session_id", "title", "preview")} row.update( started_at=r.get("started_at") or 0, ended_at=r.get("ended_at"), last_active=r.get("last_active") or r.get("started_at") or 0, - source=r.get("source"), archived=bool(r.get("archived"))) - row.update({k: r.get(k) or 0 for k in ( - "message_count", "tool_call_count", "input_tokens", "output_tokens")}) - row.update({k: r.get(k) for k in ("actual_cost_usd", "estimated_cost_usd", "model")}) - row["is_active"] = False - row.update({k: r.get(k) for k in ("cwd", "git_branch", "git_repo_root")}) + source=r.get("source"), archived=bool(r.get("archived")), + **{k: r.get(k) or 0 for k in ( + "message_count", "tool_call_count", "input_tokens", "output_tokens")}, + **{k: r.get(k) for k in ("actual_cost_usd", "estimated_cost_usd", "model")}, + is_active=False, **{k: r.get(k) for k in ("cwd", "git_branch", "git_repo_root")}) return row @@ -358,17 +327,15 @@ def _project_tree_inputs( db, session_limit: int, *, include_discovered: bool ) -> tuple[list[dict], list[dict], list[dict], str | None]: """Gather (sessions, projects, discovered_repos, active_id) for build_tree. - ``include_discovered`` is the zero-session-repo overview tier; drill-in skips it, - avoiding the distinct-cwd scan + git probes on that per-turn path.""" - # compact_rows: `_project_tree_row` drops the system-prompt blob; selecting it - # only to discard it costs tens of MB of B-tree reads per build on a big DB. + ``include_discovered`` is the zero-session-repo overview tier; drill-in skips it (and + the distinct-cwd scan + git probes) on that per-turn path.""" + # compact_rows: selecting the system-prompt blob only to drop it costs tens of MB of reads. rows = db.list_sessions_rich( limit=session_limit, offset=0, order_by_last_active=True, min_message_count=1, include_children=False, exclude_sources=_PROJECT_TREE_EXCLUDED_SOURCES, include_archived=False, compact_rows=True) sessions = [_project_tree_row(r) for r in rows] - # Parallel-warm the git cache so build_tree's resolver reads it instead of - # cold-probing each cwd in sequence (matters on the drill-in path). + # Parallel-warm the git cache so build_tree's resolver doesn't cold-probe each cwd in turn. git_probe.warm_roots(s["cwd"] for s in sessions if s.get("cwd")) from hermes_cli import projects_db as pdb policy = _repo_discovery_policy() @@ -387,19 +354,15 @@ def _project_tree_inputs( return sessions, projects, discovered, active_id -# Per-build memo for `_dir_exists_cached`; cleared at the top of every -# `_build_project_tree` so a dir created/deleted between refreshes is seen. +# Per-build memo for `_dir_exists_cached`; cleared by every `_build_project_tree`. _DIR_EXISTS_CACHE: dict[str, bool] = {} def _dir_exists_cached(path: str) -> bool: - """``os.path.isdir`` memoized per build — ``build_tree`` asks per SESSION, not - per distinct path, so hundreds of sessions in a few dirs would otherwise fire - hundreds of redundant stats per sidebar open.""" + """``os.path.isdir`` memoized per build — ``build_tree`` asks per SESSION, not per path.""" hit = _DIR_EXISTS_CACHE.get(path) if hit is None: - hit = os.path.isdir(path) - _DIR_EXISTS_CACHE[path] = hit + hit = _DIR_EXISTS_CACHE[path] = os.path.isdir(path) return hit @@ -411,8 +374,7 @@ def _build_project_tree( _DIR_EXISTS_CACHE.clear() sessions, projects, discovered, active_id = _project_tree_inputs( db, session_limit, include_discovered=include_discovered) - # build_tree also resolves every declared project folder and discovered repo - # root — not session cwds, so warm them too or they probe git one at a time. + # build_tree also resolves declared project folders and discovered roots — warm them too. git_probe.warm_roots( [str(f.get("path") or "") for p in projects for f in (p.get("folders") or [])] + [str(r.get("root") or "") for r in discovered]) diff --git a/tui_gateway/methods_prompt.py b/tui_gateway/methods_prompt.py index dd38dd6b19..99208d1f3c 100644 --- a/tui_gateway/methods_prompt.py +++ b/tui_gateway/methods_prompt.py @@ -24,24 +24,19 @@ def _history_user_indices(history: list) -> list: def _message_row_id(msg: dict): - """Parse durable SQLite row id from a history entry, or None.""" + """Durable SQLite row id from a history entry (``_row_id`` else ``row_id``), or None.""" raw = msg.get("_row_id") if raw is None: raw = msg.get("row_id") - try: + with contextlib.suppress(TypeError, ValueError): return None if raw is None else int(raw) - except (TypeError, ValueError): - return None + return None def _mem_db_pair_agrees(mem, db_msg) -> bool: - """True when a live-memory entry plausibly corresponds to a durable row. - - Positional trust needs more than equal lengths: roles and display-marker - status must match (a marker on one side shifts every later position), and an - addressable user turn must show the same text. Multimodal content can't be - compared cheaply — role/marker agreement suffices. - """ + """True when a live-memory entry plausibly corresponds to a durable row: roles and + display-marker status must match (a marker shifts every later position) and an + addressable user turn must show the same text (multimodal: role/marker suffice).""" if not isinstance(mem, dict) or not isinstance(db_msg, dict): return False if mem.get("role") != db_msg.get("role"): @@ -64,11 +59,10 @@ def _mem_db_pair_agrees(mem, db_msg) -> bool: def _find_user_turn_by_row_id(history: list, target_row_id: int): - """Return ``(user_ordinal, history_index)`` for ``target_row_id``, or None.""" - for u_ord, h_idx in enumerate(_history_user_indices(history)): - if _message_row_id(history[h_idx]) == target_row_id: - return u_ord, h_idx - return None + """``(user_ordinal, history_index)`` for ``target_row_id``, or None.""" + return next( + ((u_ord, h_idx) for u_ord, h_idx in enumerate(_history_user_indices(history)) + if _message_row_id(history[h_idx]) == target_row_id), None) def _load_durable_truncation_history( @@ -93,91 +87,48 @@ def _load_durable_truncation_history( def _resolve_truncate_row_id(session: dict, history: list, target_row_id: int): - """Resolve ``truncate_before_row_id`` to ``(user_ordinal, history_index)``. - - Prefer in-memory ``_row_id``/``row_id`` stamps; when a live turn rewrote - ``session["history"]`` without them, load the durable transcript and map the - matched user-turn ordinal onto the live list. Never falls back to a - client-supplied ordinal — unknown row ids refuse. - """ - hit = _find_user_turn_by_row_id(history, target_row_id) - if hit is not None: + """Resolve ``truncate_before_row_id`` to ``(user_ordinal, history_index)``: in-memory + stamps first, else the durable transcript mapped onto the live list by user ordinal. + Never falls back to a client-supplied ordinal — unknown row ids refuse.""" + if (hit := _find_user_turn_by_row_id(history, target_row_id)) is not None: return hit db_history = _load_durable_truncation_history(session) if db_history is None: return None - # Heal missing stamps only when EVERY pair agrees (all-or-nothing): the - # durable copy is alternation-repaired (may merge/drop rows) while the live - # list can carry optimistic/marker rows; a stamp on a misaligned pair is - # sticky and re-aims every later rewind at the wrong durable row. + # Heal missing stamps only when EVERY pair agrees: the durable copy is alternation- + # repaired while the live list can carry optimistic/marker rows, and a stamp on a + # misaligned pair is sticky (re-aims every later rewind at the wrong durable row). if len(db_history) == len(history) and all( - _mem_db_pair_agrees(mem, db_msg) for mem, db_msg in zip(history, db_history)): + _mem_db_pair_agrees(mem, db_msg) for mem, db_msg in zip(history, db_history)): for mem, db_msg in zip(history, db_history): - db_rid = _message_row_id(db_msg) if isinstance(db_msg, dict) else None - if db_rid is not None and _message_row_id(mem) is None: + if (db_rid := _message_row_id(db_msg)) is not None and _message_row_id(mem) is None: mem["_row_id"] = db_rid - hit = _find_user_turn_by_row_id(history, target_row_id) - if hit is not None: + if (hit := _find_user_turn_by_row_id(history, target_row_id)) is not None: return hit - db_hit = _find_user_turn_by_row_id(db_history, target_row_id) - if db_hit is None: + if (db_hit := _find_user_turn_by_row_id(db_history, target_row_id)) is None: return None db_ord, db_idx = db_hit mem_user_indices = _history_user_indices(history) - if db_ord < 0 or db_ord >= len(mem_user_indices): + # Same-ordinal mapping across lists that can diverge (repair may merge a user;user + # pair): trust it only when the mapped live turn shows the durable target's content. + if db_ord >= len(mem_user_indices) or not _mem_db_pair_agrees( + history[mem_user_indices[db_ord]], db_history[db_idx]): return None - mem_idx = mem_user_indices[db_ord] - # Same-ordinal mapping across lists that can diverge (repair may have merged - # a user;user pair): trust it only when the mapped live turn shows the same - # content as the durable target — else refuse (caller fails closed, 4018). - if not _mem_db_pair_agrees(history[mem_idx], db_history[db_idx]): - return None - return db_ord, mem_idx + return db_ord, mem_user_indices[db_ord] def _coerce_truncate_int(rid, value, param_name="truncate_before_user_ordinal"): - """``(int_value, error_response)`` for a client integer param. bool is refused - like any non-integer: JSON ``true`` would int() to 1 and aim at the wrong turn.""" + """``(int_value, error_response)`` for a client integer param. bool is refused like + any non-integer: JSON ``true`` would int() to 1 and aim at the wrong turn.""" if not isinstance(value, bool): with contextlib.suppress(TypeError, ValueError): return int(value), None return None, _err(rid, 4004, f"{param_name} must be an integer") -def _reconcile_client_ordinal( - rid, sid, client_ordinal, msg_ordinal, param_name, target_repr, prefix_user_count=0): - """Cross-check a client ordinal against a resolved durable target. - - Returns ``(ordinal, error_response)``: the target's tip-relative ordinal when - the client sent none or agreed, else the 4004/4030 refusal — a stale ordinal - beside a *resolved* durable id is drift; never guess which the user meant. - Client ordinals count the full displayed lineage, so after compression - ``msg_ordinal + prefix_user_count`` is the SAME turn. The cut is always aimed - by the durable target, so this can never re-aim a truncation. - """ - if client_ordinal is None: - return msg_ordinal, None - ordinal, err = _coerce_truncate_int(rid, client_ordinal) - if err is not None: - return None, err - if ordinal == msg_ordinal or ( - prefix_user_count > 0 and ordinal == msg_ordinal + prefix_user_count): - return msg_ordinal, None - logger.warning( - "prompt.submit: REFUSED truncation due to ordinal mismatch for session %s " - "(ordinal=%d, %s_ordinal=%d, %s=%s, prefix_user_count=%d). " - "Stale truncate_before_user_ordinal detected.", - sid, ordinal, param_name, msg_ordinal, param_name, target_repr, prefix_user_count) - return None, _err( - rid, 4030, - f"truncate_before_user_ordinal ({ordinal}) does not match " - f"{param_name} target turn ({msg_ordinal})") - - def _pending_reaction_notes(session: dict) -> str: - """Note block for reactions added since the last turn, or "". Applied to the - MODEL INPUT only, never the persisted prompt; each reaction is announced once - (rows are stamped ``seen`` on read). Feature-gated (display.message_reactions).""" + """Note block for reactions since the last turn (model input only, announced once — + rows are stamped ``seen`` on read), or "". Gated on display.message_reactions.""" session_key = str(session.get("session_key") or "") if not session_key: return "" @@ -189,9 +140,7 @@ def _pending_reaction_notes(session: dict) -> str: return "" try: with _session_db(session) as db: - if db is None: - return "" - pending = db.take_unseen_reactions(session_key, author="user") + pending = None if db is None else db.take_unseen_reactions(session_key, author="user") except Exception: logger.debug("Failed to read pending reactions", exc_info=True) return "" @@ -212,20 +161,16 @@ def _pending_reaction_notes(session: dict) -> str: # ── prompt.submit pieces ──────────────────────────────────────────────────── - def _typed_stop_phrase_response(rid, text): - """End the voice chat when a bare stop phrase is TYPED while backend voice mode - is on (typed twin of the spoken stop phrase). Returns the RPC reply, or None - for a normal message. The desktop's renderer-owned voice chat never flips the - backend flag and handles its own typed stop.""" + """RPC reply ending the voice chat when a bare stop phrase is TYPED while backend voice + mode is on (typed twin of the spoken stop phrase), or None for a normal message.""" if not (isinstance(text, str) and _voice_mode_enabled()): return None try: from tools.voice_mode import is_voice_stop_phrase - typed_stop = is_voice_stop_phrase(text) + if not is_voice_stop_phrase(text): + return None except Exception: - typed_stop = False - if not typed_stop: return None _end_voice_chat(stop_loop=True, stop_tts=True) _voice_emit("voice.transcript", {"stop_phrase": True, "typed": True}) @@ -240,25 +185,21 @@ def _hosted_submit_error(rid, session, hosted_task, hosted_terminal_callback): """Validate the hosted-room turn proof carried by an internal submit.""" if session.get("source") != "bot_room": return _err(rid, 4120, "hosted room turns require a bot_room session") - if ( - not isinstance(hosted_task, dict) or not callable(hosted_terminal_callback) - or set(hosted_task) != _HOSTED_TASK_FIELDS - or not all( - isinstance(hosted_task.get(field), str) and hosted_task[field] - for field in _HOSTED_TASK_FIELDS - {"execution_generation"}) - or not isinstance(hosted_task.get("execution_generation"), int)): - return _err(rid, 4120, "invalid hosted room turn proof") - return None + valid = ( + isinstance(hosted_task, dict) and callable(hosted_terminal_callback) + and set(hosted_task) == _HOSTED_TASK_FIELDS + and all(isinstance(hosted_task.get(f), str) and hosted_task[f] + for f in _HOSTED_TASK_FIELDS - {"execution_generation"}) + and isinstance(hosted_task.get("execution_generation"), int)) + return None if valid else _err(rid, 4120, "invalid hosted room turn proof") def _legacy_group_fence_error(rid, session, params): - """Older Desktop builds know the ``Group: `` title but not the hosted - authority marker; once a gateway owns that room a direct prompt would start a - second renderer driver. Fence server-side instead of trusting the client.""" + """Fence direct prompts into a hosted room from older Desktop builds (they know the + ``Group: `` title but not the authority marker; a direct prompt would start a + second renderer driver).""" title = str(session.get("title") or "") - if not title.startswith("Group: "): - return None - room_id = title.removeprefix("Group: ").strip() + room_id = title.removeprefix("Group: ").strip() if title.startswith("Group: ") else "" if not room_id: return None try: @@ -270,17 +211,15 @@ def _legacy_group_fence_error(rid, session, params): if not hosted: from hermes_constants import named_profile_home session_profile_home = named_profile_home(str(session.get("profile_home") or "")) - requested_profile = ( - (session_profile_home.name if session_profile_home is not None else "") - or str(params.get("profile") or "").strip() - or str(_current_profile_name() or "default").strip()) peer = probe_peer_room_reservation( - default_db_path(), room_id=room_id, target_profile=requested_profile) + default_db_path(), room_id=room_id, target_profile=( + (session_profile_home.name if session_profile_home is not None else "") + or str(params.get("profile") or "").strip() + or str(_current_profile_name() or "default").strip())) except RoomProbeUnavailableError: return _err(rid, 5122, _GROUP_PROBE_FAILED_MSG) except HostedRoomError: - # Legacy Desktop sessions used the display name after "Group: "; those - # names are not hosted room ids. + # Legacy Desktop sessions used the display name after "Group: " — not a room id. return None except Exception: return _err(rid, 5122, _GROUP_PROBE_FAILED_MSG) @@ -293,32 +232,25 @@ def _legacy_group_fence_error(rid, session, params): def _parse_truncation_params(rid, sid, session, params, history): """Coerce + admit the truncation params; ``(target_row_id, client_ordinal, err)``. - - Precedence: malformed params (4004) -> unconfirmed (4029, checked BEFORE - target resolution so a leaked-state request never pays the durable read or - heal-stamps live dicts). An ordinal/id alone is not consent: a leftover - ordinal on an ORDINARY submit is indistinguishable from a real rewind, and - the cut is a destructive replace_messages(). - """ - truncate_user_ordinal = params.get("truncate_before_user_ordinal") - truncate_row_id = params.get("truncate_before_row_id") + Malformed (4004) -> unconfirmed (4029; checked BEFORE target resolution so a + leaked-state request never pays the durable read or heal-stamps live dicts). An + ordinal/id alone is not consent: a leftover ordinal on an ORDINARY submit is + indistinguishable from a real rewind, and the cut is a destructive replace.""" target_row_id = client_ordinal = None - if truncate_row_id is not None: + if (truncate_row_id := params.get("truncate_before_row_id")) is not None: target_row_id, err = _coerce_truncate_int(rid, truncate_row_id, "truncate_before_row_id") if err is not None: return None, None, err - if truncate_user_ordinal is not None: + if (truncate_user_ordinal := params.get("truncate_before_user_ordinal")) is not None: client_ordinal, err = _coerce_truncate_int(rid, truncate_user_ordinal) if err is not None: return None, None, err if is_truthy_value(params.get("confirm_truncate")): return target_row_id, client_ordinal, None logger.warning( - "prompt.submit: REFUSED unconfirmed truncation of session %s " - "(%d messages held; ordinal=%s, row_id=%s, message_id=%s). " - "The client attached truncation parameters without " - "confirm_truncate — likely stale truncation parameters on " - "an ordinary submit.", + "prompt.submit: REFUSED unconfirmed truncation of session %s (%d messages held; " + "ordinal=%s, row_id=%s, message_id=%s). The client attached truncation parameters without " + "confirm_truncate — likely stale truncation parameters on an ordinary submit.", sid, len(history), client_ordinal, target_row_id, params.get("truncate_before_message_id")) return None, None, _err( rid, 4029, @@ -327,48 +259,22 @@ def _parse_truncation_params(rid, sid, session, params, history): "(update your Hermes client if a rewind was intended)") -def _ordinal_only_truncation_error(rid, sid, session, history, user_indices, client_ordinal): - """4004 refusal when an ordinal-only cut targets a durable session, else None. - - Durability is a state.db property, not an annotation on the live copy (resume - paths historically omitted _row_id stamps). An unreadable durable state - fails closed too: absence of proof is not proof of an ephemeral conversation. - """ - has_stamped_user = any(_message_row_id(history[h_idx]) is not None for h_idx in user_indices) - durable_history = [] if has_stamped_user else _load_durable_truncation_history(session, sid) - if not (has_stamped_user or durable_history is None or durable_history): - return None - logger.warning( - "prompt.submit: REFUSED ordinal-only truncation of durable " - "session %s (ordinal=%d); truncate_before_row_id required", - sid, client_ordinal) - return _err( - rid, 4004, - "ordinal-only truncation is unsafe for durable session history; " - "include truncate_before_row_id") - - def _resolve_truncation_ordinal(rid, sid, session, params, history): - """Resolve the truncation target to ``(ordinal, cut_index, err)``. - - After ``_parse_truncation_params``: unresolvable target (4018, fail closed — - never degrade a missing row_id/message_id into an ordinal cut) -> ordinal - drift (4030) -> ordinal-only on a durable session (4004). - """ + """Resolve the truncation target to ``(ordinal, cut_index, err)``: unresolvable target + (4018, fail closed — never degrade a missing row_id/message_id into an ordinal cut) -> + ordinal drift (4030) -> ordinal-only on a durable session (4004).""" target_row_id, client_ordinal, err = _parse_truncation_params( rid, sid, session, params, history) if err is not None: return None, None, err truncate_message_id = params.get("truncate_before_message_id") - # Client ordinals count the full displayed lineage; after compression the tip - # is session["history"] and ancestors live in display_history_prefix. Count - # ancestor user turns once so client and tip-relative ordinals translate. + # Client ordinals count the full displayed lineage; after compression ancestors live in + # display_history_prefix, so count their user turns once to translate ordinals. prefix_user_count = len(_history_user_indices(session.get("display_history_prefix") or [])) user_indices = _history_user_indices(history) def _stale(resolved_ordinal=None): - # Structured recovery fields: Desktop resyncs + retries on a stale target - # and shows "compressed away" when segment_ordinal < 0 (ancestor-only). + # Recovery fields: Desktop resyncs + retries, "compressed away" when segment < 0. segment = ( client_ordinal - prefix_user_count if client_ordinal is not None else resolved_ordinal) return None, None, _err(rid, 4018, _STALE_TARGET_MSG, data={ @@ -384,32 +290,52 @@ def _resolve_truncation_ordinal(rid, sid, session, params, history): param_name = "truncate_before_message_id" target_repr = msg_id_str = str(truncate_message_id) found_match = next( - ( - (u_ord, h_idx) for u_ord, h_idx in enumerate(user_indices) - if history[h_idx].get("id") == msg_id_str - or history[h_idx].get("message_id") == msg_id_str), - None) + ((u_ord, h_idx) for u_ord, h_idx in enumerate(user_indices) + if history[h_idx].get("id") == msg_id_str + or history[h_idx].get("message_id") == msg_id_str), None) not_found = "target message_id %s not found in history for session %s" if found_match is None: logger.warning( "prompt.submit: " + not_found + "; refusing truncation without fallback", target_repr, sid) return _stale() - ordinal, err = _reconcile_client_ordinal( - rid, sid, client_ordinal, found_match[0], param_name, target_repr, - prefix_user_count=prefix_user_count) - if err is not None: - return None, None, err + ordinal = found_match[0] + # A stale client ordinal beside a *resolved* durable id is drift — never guess + # which the user meant. Client ordinals count the full displayed lineage, so + # after compression ``ordinal + prefix_user_count`` is the SAME turn. The cut is + # always aimed by the durable target, so this can never re-aim a truncation. + if client_ordinal is not None and client_ordinal != ordinal and not ( + prefix_user_count > 0 and client_ordinal == ordinal + prefix_user_count): + logger.warning( + "prompt.submit: REFUSED truncation due to ordinal mismatch for session %s " + "(ordinal=%d, %s_ordinal=%d, %s=%s, prefix_user_count=%d). " + "Stale truncate_before_user_ordinal detected.", + sid, client_ordinal, param_name, ordinal, param_name, target_repr, + prefix_user_count) + return None, None, _err( + rid, 4030, + f"truncate_before_user_ordinal ({client_ordinal}) does not match " + f"{param_name} target turn ({ordinal})") else: ordinal = client_ordinal - prefix_user_count if ordinal < 0 or ordinal >= len(user_indices): return _stale() - err = _ordinal_only_truncation_error( - rid, sid, session, history, user_indices, client_ordinal) - if err is not None: - return None, None, err - # Reject out-of-range on BOTH ends: a negative ordinal would hit Python's - # negative indexing (user_indices[-1] -> the LAST user turn) and persist the loss. + # Ordinal-only cut on a durable session: durability is a state.db property, not a + # live-copy annotation (resume paths omitted _row_id stamps); an unreadable + # durable state fails closed too. + has_stamped_user = any( + _message_row_id(history[h_idx]) is not None for h_idx in user_indices) + durable = [] if has_stamped_user else _load_durable_truncation_history(session, sid) + if has_stamped_user or durable is None or durable: + logger.warning( + "prompt.submit: REFUSED ordinal-only truncation of durable " + "session %s (ordinal=%d); truncate_before_row_id required", + sid, client_ordinal) + return None, None, _err( + rid, 4004, + "ordinal-only truncation is unsafe for durable session history; " + "include truncate_before_row_id") + # BOTH ends: a negative ordinal would index user_indices[-1] and persist the loss. if ordinal < 0 or ordinal >= len(user_indices): return _stale(resolved_ordinal=ordinal) return ordinal, user_indices[ordinal], None @@ -419,82 +345,16 @@ def _row_ids_of(messages) -> set: return {row_id for message in messages if isinstance((row_id := _message_row_id(message)), int)} -def _persist_truncation(rid, sid, session, history, truncated, ordinal, requested_rebind_ids): - """Write the truncated transcript BEFORE touching memory (fail closed). - - If replace_messages failed after session["history"] was rewritten, the turn - would run against the short list while state.db kept the old tail, and the - append-only flush would stack the new exchange on the "undone" turns — zombie - history on resume. Writes through ``_session_db`` (owner of this session's - row), never ``_get_db()``: a profile session's transcript lives in its own - profile's state.db. Returns ``(err, survivor_user_row_ids, survivor_row_id_map)``. - """ - survivor_user_row_ids = survivor_row_id_map = None - with _session_db(session) as db: - if db is not None: - try: - # session_key can be NULL for old CLI-origin sessions; fall back to - # sid or replace_messages(None) trips an FK violation. - truncation_key = session.get("session_key") or sid - old_active_row_ids = _row_ids_of(history) - if requested_rebind_ids is not None: - # Row-id fallback can resolve a target the live list is too - # misaligned to stamp, and repair can merge a user;user pair: - # read the un-repaired pre-write active-id set so a rewritten - # row is never mistaken for an untouched archived/ancestor row. - durable_rebind_history = _load_durable_truncation_history( - session, truncation_key, repair_alternation=False) - if durable_rebind_history is None: - raise RuntimeError("could not load durable row identities for truncation") - old_active_row_ids.update(_row_ids_of(durable_rebind_history)) - old_survivor_row_ids = [_message_row_id(message) for message in truncated] - # active_only=True: compaction keeps the pre-compaction transcript - # as active=0 rows under this key; a bare replace would DELETE that - # archive on every edit. archive_dropped=True: soft-archive the - # dropped turns (active=0, still in FTS) so a mis-aimed cut is - # recoverable. - db.replace_messages( - truncation_key, truncated, active_only=True, archive_dropped=True, - reject_active_turn_lease=True) - except Exception as exc: - logger.error( - "prompt.submit: replace_messages failed for session %s " - "(ordinal=%d); refusing turn so memory and DB stay " - "aligned: %s", - sid, ordinal, exc, exc_info=True) - return _err(rid, 5008, f"failed to persist history truncation: {exc}"), None, None - # replace_messages re-inserted the survivors as NEW rows with fresh - # _row_id stamps. Surface the surviving user-turn ids (visible-user- - # ordinal order) so the client rebinds its cached rowIds — else a - # second rewind sends the pre-rewind id and the resolver refuses with - # 4018. None entries: the client must drop its cached id for that turn. - survivor_user_row_ids = [ - _message_row_id(truncated[i]) for i in _history_user_indices(truncated)] - if requested_rebind_ids is not None: - survivor_row_id_map = { - str(old_row_id): new_row_id - for old_row_id, new_row_id in zip( - old_survivor_row_ids, (_message_row_id(message) for message in truncated)) - if isinstance(old_row_id, int) and isinstance(new_row_id, int) - and old_row_id in requested_rebind_ids} - for dropped_row_id in requested_rebind_ids.intersection(old_active_row_ids): - survivor_row_id_map.setdefault(str(dropped_row_id), None) - return None, survivor_user_row_ids, survivor_row_id_map - - def _truncate_history_for_submit(rid, sid, session, params, requested_rebind_ids): - """Rewind/regenerate cut, under ``history_lock``. Returns - ``(err, survivor_user_row_ids, survivor_row_id_map)``; on success - ``session["history"]`` is replaced and ``history_version`` bumped.""" + """Rewind/regenerate cut under ``history_lock``: ``(err, survivor_fields)``; the fields + are the client rowId-rebind payload.""" history = _history_without_ephemeral_scaffolding(session.get("history", [])) ordinal, cut_index, err = _resolve_truncation_ordinal(rid, sid, session, params, history) if err is not None: - return err, None, None + return err, {} from agent.context_compressor import history_before_user_originated_turn truncated, _live_view = history_before_user_originated_turn(history, cut_index) - # Second gate on top of confirm_truncate: ordinal 0 -> history[:0] == [] and - # replace_messages() DELETEs every durable row. Wiping the whole transcript - # needs its own opt-in (legitimate restore/regenerate of the first turn). + # Second gate: ordinal 0 would DELETE every durable row; wiping needs its own opt-in. if not truncated and history and not is_truthy_value(params.get("confirm_empty_truncate")): logger.warning( "prompt.submit: REFUSED empty truncation of session %s " @@ -503,36 +363,64 @@ def _truncate_history_for_submit(rid, sid, session, params, requested_rebind_ids return _err( rid, 4028, "truncation would erase the entire session transcript; " - "resubmit with confirm_empty_truncate=true if this is intended", - ), None, None + "resubmit with confirm_empty_truncate=true if this is intended"), {} log_fn = logger.warning if not truncated else logger.info log_fn( "prompt.submit: truncating session %s history %d -> %d messages (ordinal=%d)", sid, len(history), len(truncated), ordinal) - err, survivor_user_row_ids, survivor_row_id_map = _persist_truncation( - rid, sid, session, history, truncated, ordinal, requested_rebind_ids) - if err is not None: - return err, None, None + # Write the truncated transcript BEFORE touching memory (fail closed: a failed write + # after the in-memory rewrite would stack the new exchange on the "undone" turns). + # Writes through _session_db (profile sessions own their state.db). + fields = {} + with _session_db(session) as db: + if db is not None: + try: + # NULL session_key (old CLI-origin sessions) would trip an FK violation. + truncation_key = session.get("session_key") or sid + old_active_row_ids = _row_ids_of(history) + if requested_rebind_ids is not None: + # Un-repaired pre-write active-id set: a rewritten row must never be + # mistaken for an untouched archived/ancestor row. + durable_rebind_history = _load_durable_truncation_history( + session, truncation_key, repair_alternation=False) + if durable_rebind_history is None: + raise RuntimeError("could not load durable row identities for truncation") + old_active_row_ids.update(_row_ids_of(durable_rebind_history)) + old_survivor_row_ids = [_message_row_id(message) for message in truncated] + # active_only: a bare replace would DELETE the compaction archive (active=0 + # rows) on every edit. archive_dropped: a mis-aimed cut stays recoverable. + db.replace_messages( + truncation_key, truncated, active_only=True, archive_dropped=True, + reject_active_turn_lease=True) + except Exception as exc: + logger.error( + "prompt.submit: replace_messages failed for session %s (ordinal=%d); refusing " + "turn so memory and DB stay aligned: %s", + sid, ordinal, exc, exc_info=True) + return _err(rid, 5008, f"failed to persist history truncation: {exc}"), {} + # Survivors were re-inserted as NEW rows: surface the fresh ids so the client + # rebinds its cached rowIds (else a second rewind refuses with 4018). None + # entries: the client must drop its cached id for that turn. + if requested_rebind_ids is None: + fields["survivor_user_row_ids"] = [ + _message_row_id(truncated[i]) for i in _history_user_indices(truncated)] + else: + fields["survivor_row_id_map"] = row_id_map = { + str(old_row_id): new_row_id + for old_row_id, new_row_id in zip( + old_survivor_row_ids, (_message_row_id(message) for message in truncated)) + if isinstance(old_row_id, int) and isinstance(new_row_id, int) + and old_row_id in requested_rebind_ids} + for dropped_row_id in requested_rebind_ids.intersection(old_active_row_ids): + row_id_map.setdefault(str(dropped_row_id), None) session["history"] = truncated session["history_version"] = int(session.get("history_version", 0)) + 1 - return None, survivor_user_row_ids, survivor_row_id_map - - -def _survivor_fields(survivor_user_row_ids, survivor_row_id_map, requested_rebind_ids) -> dict: - """Client rowId-rebind payload for a submit that truncated a durable session.""" - fields = {} - if survivor_user_row_ids is not None and requested_rebind_ids is None: - fields["survivor_user_row_ids"] = survivor_user_row_ids - if survivor_row_id_map is not None: - fields["survivor_row_id_map"] = survivor_row_id_map - return fields + return None, fields def _persist_session_row_for_submit(rid, session): - """Lazily persist the DB row now that the user actually sent a message; a - branch becomes real here (parent transcript copied as its seed). Returns an - error reply (the only user-visible signal; desktop maps it to a toast) or - None. On failure the in-flight turn is released.""" + """Lazily persist the DB row now that the user sent a message (a branch becomes real + here); the error reply is the only user-visible signal (desktop maps it to a toast).""" try: if _ensure_session_db_row(session) is False: return _err( @@ -550,23 +438,21 @@ def _persist_session_row_for_submit(rid, session): if is_disk_full_error(exc): return _err( rid, 5070, - "disk full: session storage could not be written — free some disk space and try again", - ) + "disk full: session storage could not be written — free some disk space and try again") logger.warning("prompt.submit: session persist failed: %s", exc, exc_info=True) return _err(rid, 5071, f"session storage could not be written: {exc}") return None def _run_after_agent_ready(rid, sid, session, text, display_kind, hosted_terminal_callback): - """Turn thread body: patient wait for a deferred build (the message is already - the accepted in-flight turn, so a slow build must not eat it), then run.""" + """Turn thread body: patient wait for a deferred build (a slow build must not eat the + accepted in-flight message), then run.""" err = _wait_agent_for_prompt(session, rid, sid) if err: - # Terminal frame + retained snapshot (not a bare "error" event): if the - # client is disconnected, the snapshot is the only way resume shows this. + # Terminal frame + retained snapshot (not a bare "error" event): the snapshot is + # the only way resume shows this to a disconnected client. _emit_terminal_turn_error( sid, session, (err.get("error") or {}).get("message", "agent initialization failed"), - # Construction never reached the provider: local-runtime failure. error_surface={"layer": "runtime", "code": "agent_init_failed", "retryable": True}) with session["history_lock"]: session["running"] = False @@ -577,12 +463,11 @@ def _run_after_agent_ready(rid, sid, session, text, display_kind, hosted_termina if session.get("_turn_cancel_requested") or not session.get("running"): session["running"] = False _clear_inflight_turn(session) - # Without this emit the turn vanishes silently: the client saw - # {"status": "streaming"} but never gets message.start or error. - _emit("error", sid, { - "message": "Turn cancelled before the agent was ready" + # Without this emit the turn vanishes silently after {"status": "streaming"}. + _emit("error", sid, {"message": ( + "Turn cancelled before the agent was ready" if session.get("_turn_cancel_requested") - else "Session no longer running before the agent was ready"}) + else "Session no longer running before the agent was ready")}) return _run_prompt_submit( rid, sid, session, text, display_kind=display_kind, @@ -593,58 +478,33 @@ _TRUNCATION_PARAMS = ( "truncate_before_user_ordinal", "truncate_before_row_id", "truncate_before_message_id") -def _claim_submit_slot(rid, sid, session, text, params, transport, internal_hosted_submit): - """Claim the turn against a possibly-running session; returns an early RPC - reply (busy/queued) or None once ``running`` is observed False. - - A mid-turn prompt is queued (by default interrupting the live turn) instead of - rejected. The provider interrupt happens after ``history_lock`` is released: a - non-interruptible tool may hold it. If the old turn finished between the two - lock acquisitions, retry the claim rather than strand this prompt in a queue - whose drain already ran. - """ - while True: - with session["history_lock"]: - if not session.get("running"): - return None - if internal_hosted_submit: - return _err(rid, 4091, "hosted room member session is busy") - busy_transport = transport or session.get("transport") - busy_response = _handle_busy_submit( - rid, sid, session, text, busy_transport, queued=bool(params.get("queued"))) - if busy_response is not None: - return busy_response - - def _lock_in_submit_turn( rid, sid, session, text, params, has_truncation, requested_rebind_ids, hosted_task): - """Under ``history_lock``: refuse watch-child races and malformed truncation, - apply the cut, then mark the turn running + in flight. - Returns ``(err, survivor_user_row_ids, survivor_row_id_map)``.""" - survivor_user_row_ids = survivor_row_id_map = None + """Under ``history_lock``: refuse watch-child races / malformed truncation, apply the + cut, mark the turn running + in flight. Returns ``(err, survivor_fields)``.""" + fields = {} with session["history_lock"]: - # A watch session's run lives in the PARENT turn, so its own running flag - # is False; typing mid-run would build a second agent racing the child - # on the same stored session. After the run completes, submitting is fine. + # A watch session's run lives in the PARENT turn (own running flag False); typing + # mid-run would build a second agent racing the child on the same stored session. if session.get("lazy") and _child_run_active(str(session.get("session_key") or "")): - return _err(rid, 4009, "subagent still running — wait for it to finish"), None, None + return _err(rid, 4009, "subagent still running — wait for it to finish"), fields if is_truthy_value(params.get("confirm_truncate")) and not has_truncation: return _err( rid, 4004, "confirm_truncate requires truncate_before_user_ordinal, truncate_before_message_id, or truncate_before_row_id", - ), None, None + ), fields if has_truncation: - err, survivor_user_row_ids, survivor_row_id_map = _truncate_history_for_submit( + err, fields = _truncate_history_for_submit( rid, sid, session, params, requested_rebind_ids) if err is not None: - return err, None, None + return err, {} session["running"] = True session["_turn_cancel_requested"] = False session["last_active"] = time.time() if hosted_task is not None: session["_hosted_room_task"] = dict(hosted_task) _start_inflight_turn(session, text) - return None, survivor_user_row_ids, survivor_row_id_map + return None, fields @method("prompt.submit") @@ -653,14 +513,13 @@ def _(rid, params: dict) -> dict: sid = params.get("session_id", "") raw_text = params.get("text", "") text = sanitize_user_prompt_text(raw_text) if isinstance(raw_text, str) else raw_text - # Off-screen sends (widget intents) type the persisted row so no client - # renders a bubble. Whitelisted to "hidden": this RPC must not mint kinds. + # Off-screen sends (widget intents) type the row so no client renders a bubble; + # whitelisted to "hidden" — this RPC must not mint kinds. display_kind = "hidden" if params.get("display_kind") == "hidden" else None if (stopped := _typed_stop_phrase_response(rid, text)) is not None: return stopped if params.get("interrupted"): - # Client-side barge-in (desktop VAD / typing over playback): latch it so - # this turn's model message carries the interruption note. + # Client-side barge-in: latch so this turn's model message carries the note. from tools.tts_streaming import mark_speech_interrupted mark_speech_interrupted() session, err = _sess_nowait(params, rid) @@ -669,49 +528,54 @@ def _(rid, params: dict) -> dict: hosted_task = params.get("_hosted_task") hosted_terminal_callback = params.get("_hosted_terminal_callback") internal_hosted_submit = hosted_task is not None or hosted_terminal_callback is not None - if internal_hosted_submit: - err = _hosted_submit_error(rid, session, hosted_task, hosted_terminal_callback) - else: - err = _legacy_group_fence_error(rid, session, params) + err = ( + _hosted_submit_error(rid, session, hosted_task, hosted_terminal_callback) + if internal_hosted_submit else _legacy_group_fence_error(rid, session, params)) if err is not None: return err if (limit_message := _ensure_active_session_slot(sid, session)) is not None: - # Refused HERE — before the busy queue, _ensure_session_db_row and - # _start_agent_build — so a refusal leaves the session exactly as it was. - # The reason travels as machine-readable data ("at capacity, retry" vs - # "live owner, your write would interleave"), never as matched prose. + # Refused HERE — before the busy queue, db row and agent build — so a refusal + # leaves the session untouched. The reason travels as machine-readable data. reason = getattr(limit_message, "reason", None) return _err(rid, 4090, str(limit_message), {"reason": reason} if reason else None) - # Rewritten on every submit: one session can be driven from the app window - # and the HUD in turn, and a stale "hud" misinforms the model. + # Rewritten every submit: a session alternates app window / HUD; stale "hud" misinforms. session["client_surface"] = "hud" if params.get("surface") == "hud" else "" has_truncation = any(params.get(k) is not None for k in _TRUNCATION_PARAMS) if has_truncation and isinstance(text, str): - # A rewind replays what the transcript shows; a skill turn shows its - # invocation, so re-expand it or `/work fix it` sends nine literal chars. + # A rewind replays what the transcript shows: re-expand a skill invocation or + # `/work fix it` sends nine literal chars. text = _expand_skill_invocation_for_replay(text, str(session.get("session_key") or "")) turn_isolation = _session_uses_compute_host(session, _load_dashboard_process_isolation_config()) if internal_hosted_submit and turn_isolation: return _err(rid, 4121, "hosted room turns do not support isolated compute workers yet") - # Re-bind to the current client transport so streaming stays on the active - # websocket even if a disconnect/fallback moved the session to stdio. + # Re-bind to the current transport: streaming must stay on the active websocket even + # if a disconnect/fallback moved the session to stdio. if (t := current_transport()) is not None: session["transport"] = t - busy = _claim_submit_slot(rid, sid, session, text, params, t, internal_hosted_submit) - if busy is not None: - return busy + # Claim the turn against a possibly-running session (busy/queued reply, else fall + # through once ``running`` is observed False). The provider interrupt happens after + # history_lock is released (a non-interruptible tool may hold it); if the old turn + # finished between the two acquisitions, retry the claim rather than strand this + # prompt in a queue whose drain already ran. + while True: + with session["history_lock"]: + if not session.get("running"): + break + if internal_hosted_submit: + return _err(rid, 4091, "hosted room member session is busy") + busy_transport = t or session.get("transport") + busy_response = _handle_busy_submit( + rid, sid, session, text, busy_transport, queued=bool(params.get("queued"))) + if busy_response is not None: + return busy_response raw_rebind_ids = params.get("rebind_survivor_row_ids") requested_rebind_ids = ( - { - row_id for row_id in raw_rebind_ids - if isinstance(row_id, int) and not isinstance(row_id, bool)} + {r for r in raw_rebind_ids if isinstance(r, int) and not isinstance(r, bool)} if isinstance(raw_rebind_ids, list) else None) - err, survivor_user_row_ids, survivor_row_id_map = _lock_in_submit_turn( + err, survivor_fields = _lock_in_submit_turn( rid, sid, session, text, params, has_truncation, requested_rebind_ids, hosted_task) if err is not None: return err - survivor_fields = _survivor_fields( - survivor_user_row_ids, survivor_row_id_map, requested_rebind_ids) if turn_isolation: isolated_response = _submit_prompt_to_compute_host( rid, sid, session, text, display_kind=display_kind) @@ -724,8 +588,7 @@ def _(rid, params: dict) -> dict: isolated_response["error"].get("message", "unknown error")) if (err := _persist_session_row_for_submit(rid, session)) is not None: return err - # A completed FAILED build must not wedge the session: rebuild with fresh - # provider resolution instead of replaying the cached failure forever. + # A completed FAILED build must not wedge the session: rebuild, don't replay it. if not _restart_completed_failed_agent_build(sid, session, session.get("agent_ready")): _start_agent_build(sid, session) run_thread = threading.Thread( @@ -740,7 +603,6 @@ def _(rid, params: dict) -> dict: # ── attachments ───────────────────────────────────────────────────────────── - def _attached_image_result(session, image_path, **extra) -> dict: """Common ``{attached, path, count, ...meta}`` reply after queuing an image.""" return { @@ -765,10 +627,9 @@ def _(rid, params: dict) -> dict: # Save-first (CLI keybinding parity): more robust than a has_image() precheck. if not save_clipboard_image(img_path): session["image_counter"] = max(0, session["image_counter"] - 1) - msg = ( + return _ok(rid, {"attached": False, "message": ( "Clipboard has image but extraction failed" if has_clipboard_image() - else "No image found in clipboard") - return _ok(rid, {"attached": False, "message": msg}) + else "No image found in clipboard")}) session.setdefault("attached_images", []).append(str(img_path)) return _ok(rid, _attached_image_result(session, img_path)) @@ -784,10 +645,8 @@ def _(rid, params: dict) -> dict: try: from cli import ( _IMAGE_EXTENSIONS, _detect_file_drop, _resolve_attachment_path, _split_path_input) - dropped = _detect_file_drop(raw) - if dropped: - image_path = dropped["path"] - remainder = dropped["remainder"] + if dropped := _detect_file_drop(raw): + image_path, remainder = dropped["path"], dropped["remainder"] else: path_token, remainder = _split_path_input(raw) image_path = _resolve_attachment_path(path_token) @@ -805,10 +664,8 @@ def _(rid, params: dict) -> dict: @method("image.attach_bytes") def _(rid, params: dict) -> dict: - """Attach an image from base64 bytes (remote client: its file isn't on our disk). - Reply shape mirrors ``image.attach``. ``content_base64``/``data`` accept a - ``data:image/...;base64,`` prefix; ``filename``/``ext`` hint the extension, else - magic bytes decide (PNG/JPEG/GIF/WebP/BMP, fallback ``.png``).""" + """Attach an image from base64 bytes (remote client); reply mirrors ``image.attach``. + ``filename``/``ext`` hint the extension, else magic bytes decide (fallback ``.png``).""" session, err = _sess_building(params, rid) if err: return err @@ -854,22 +711,21 @@ def _pdf_attach_source(rid, params, td_path, raw_path, raw_b64): resolved = _resolve_attachment_path(raw_path) except Exception: resolved = None - if resolved is None or not Path(resolved).is_file(): + if resolved is None or not (pdf := Path(resolved)).is_file(): return None, None, _err(rid, 4016, f"PDF not found: {raw_path}") - if Path(resolved).suffix.lower() != ".pdf": - return None, None, _err(rid, 4016, f"not a PDF: {Path(resolved).name}") - if Path(resolved).stat().st_size > _PDF_ATTACH_MAX_BYTES: + if pdf.suffix.lower() != ".pdf": + return None, None, _err(rid, 4016, f"not a PDF: {pdf.name}") + if pdf.stat().st_size > _PDF_ATTACH_MAX_BYTES: mb = _PDF_ATTACH_MAX_BYTES // (1024 * 1024) return None, None, _err(rid, 4018, f"PDF too large; cap is {mb} MB") - return Path(resolved), Path(resolved).name, None + return pdf, pdf.name, None def _pdf_page_range(rid, params): """Validate first/last page against the per-call cap: ``(first, last, err)``.""" try: first_page = int(params.get("first_page") or 1) - last_page_param = params.get("last_page") - last_page = int(last_page_param) if last_page_param is not None else None + last_page = None if params.get("last_page") is None else int(params.get("last_page")) except (TypeError, ValueError): return None, None, _err(rid, 4015, "first_page/last_page must be integers") if first_page < 1: @@ -886,9 +742,8 @@ def _pdf_page_range(rid, params): @method("pdf.attach") def _(rid, params: dict) -> dict: - """Attach a PDF by rendering each page to PNG (``pdftoppm`` @150 DPI, poppler-utils; - 5028 if missing) and queuing the pages as images. Accepts a host ``path`` or - base64 ``content_base64``. Caps: 50 MB / 25 pages per call.""" + """Attach a PDF by rendering each page to PNG (``pdftoppm``; 5028 if missing) and + queuing the pages as images. Host ``path`` or base64 ``content_base64``.""" import shutil import subprocess import tempfile @@ -914,12 +769,11 @@ def _(rid, params: dict) -> dict: str(pdf_path), str(td_path / "page")] from hermes_cli._subprocess_compat import windows_hide_flags try: + # UTF-8 + lossy decode: non-UTF-8 child output must not crash the gateway + # thread on locale-mismatched Windows. res = subprocess.run( argv, capture_output=True, text=True, timeout=120, stdin=subprocess.DEVNULL, - # UTF-8 + lossy decode: non-UTF-8 child output must not crash the - # gateway thread on locale-mismatched Windows. - encoding="utf-8", errors="replace", - creationflags=windows_hide_flags()) + encoding="utf-8", errors="replace", creationflags=windows_hide_flags()) except subprocess.TimeoutExpired: return _err(rid, 5028, "pdftoppm timed out (>120s)") if res.returncode != 0: @@ -946,16 +800,13 @@ def _(rid, params: dict) -> dict: @method("file.attach") def _(rid, params: dict) -> dict: - """Stage a non-image file into the session workspace and return a - workspace-relative ``@file:`` ref the agent's file tools can read. ``path`` is - the client/host path (naming + local resolution); ``data_url`` carries the bytes - when the path isn't visible to the gateway; ``name`` overrides the filename.""" + """Stage a non-image file into the session workspace; returns a workspace-relative + ``@file:`` ref. ``data_url`` carries the bytes when ``path`` isn't gateway-visible.""" session, err = _sess_building(params, rid) if err: return err - raw = str(params.get("path", "") or "").strip() - data_url = str(params.get("data_url", "") or "").strip() - name = str(params.get("name", "") or "").strip() + raw, data_url, name = ( + str(params.get(k, "") or "").strip() for k in ("path", "data_url", "name")) if not raw and not data_url: return _err(rid, 4015, "path or data_url required") try: @@ -978,8 +829,7 @@ def _(rid, params: dict) -> dict: raw = str(params.get("path", "") or "").strip() if not raw: return _err(rid, 4015, "path required") - images = session.setdefault("attached_images", []) - before = len(images) + before = len(images := session.setdefault("attached_images", [])) session["attached_images"] = [path for path in images if path != raw] return _ok(rid, { "detached": len(session["attached_images"]) != before, @@ -993,12 +843,10 @@ def _(rid, params: dict) -> dict: return err try: from cli import _detect_file_drop - raw = str(params.get("text", "") or "") - dropped = _detect_file_drop(raw) + dropped = _detect_file_drop(str(params.get("text", "") or "")) if not dropped: return _ok(rid, {"matched": False}) - drop_path = dropped["path"] - remainder = dropped["remainder"] + drop_path, remainder = dropped["path"], dropped["remainder"] if dropped["is_image"]: session.setdefault("attached_images", []).append(str(drop_path)) return _ok(rid, { @@ -1016,38 +864,27 @@ def _(rid, params: dict) -> dict: # ── side agents (background / btw / preview.restart) ──────────────────────── - -@contextlib.contextmanager -def _session_profile_home_scope(session): - """Bind the session's HERMES_HOME override for an ephemeral agent thread: the - ContextVar set on the session-create thread doesn't propagate, so a turn under - a non-default profile would otherwise run against the wrong home.""" - profile_home = session.get("profile_home") - home_token = set_hermes_home_override(profile_home) if profile_home else None - try: - yield - finally: - if home_token is not None: - reset_hermes_home_override(home_token) - - def _final_response_text(result) -> str: return (result.get("final_response", str(result)) if isinstance(result, dict) else str(result)) def _spawn_side_agent( rid, session, task_id, parent, event, body, *, cwd="", extra=None, cleanup=None): - """Run ``body()`` (an ephemeral agent call) on a daemon thread under the - session's profile home and cwd; its text — or ``error: `` — lands on - ``parent`` as ``event`` with ``task_id`` (+ ``extra``). ``cleanup`` runs in - the finally before the session context is cleared. Replies ``{task_id}``.""" + """Run ``body()`` on a daemon thread under the session's profile home (the ContextVar + doesn't propagate across threads) and cwd; its text — or ``error: `` — lands on + ``parent`` as ``event`` with ``task_id`` (+ ``extra``). Replies ``{task_id}``.""" extra = extra or {} def run(): session_tokens = _set_session_context(task_id, cwd=(cwd or _session_cwd(session))) + profile_home = session.get("profile_home") + home_token = set_hermes_home_override(profile_home) if profile_home else None try: - with _session_profile_home_scope(session): + try: text = body() + finally: + if home_token is not None: + reset_hermes_home_override(home_token) _emit(event, parent, {"task_id": task_id, **extra, "text": text}) except Exception as e: _emit(event, parent, {"task_id": task_id, **extra, "text": f"error: {e}"}) @@ -1060,15 +897,22 @@ def _spawn_side_agent( return _ok(rid, {"task_id": task_id}) -@method("prompt.background") -def _(rid, params: dict) -> dict: +def _side_agent_args(rid, params, prefix): + """Shared admission for the side-agent RPCs: ``(session, text, parent, task_id, err)``.""" session, err = _sess(params, rid) if err: - return err + return None, None, None, None, err text, parent = params.get("text", ""), params.get("session_id", "") if not text: - return _err(rid, 4012, "text required") - task_id = f"bg_{uuid.uuid4().hex[:6]}" + return None, None, None, None, _err(rid, 4012, "text required") + return session, text, parent, f"{prefix}_{uuid.uuid4().hex[:6]}", None + + +@method("prompt.background") +def _(rid, params: dict) -> dict: + session, text, parent, task_id, err = _side_agent_args(rid, params, "bg") + if err: + return err def body(): from run_agent import AIAgent @@ -1081,28 +925,21 @@ def _(rid, params: dict) -> dict: @method("prompt.btw") def _(rid, params: dict) -> dict: - """Answer a side question without touching session history: snapshot the live - conversation (in-flight ``_session_messages`` else ``session["history"]``) and - run a one-shot auxiliary call (``agent/side_question.py``). History, role - alternation and prompt cache stay untouched; answer arrives as ``btw.complete``.""" - session, err = _sess(params, rid) + """Side question over a snapshot of the live conversation (``agent/side_question.py``); + history, alternation and prompt cache stay untouched. Answer: ``btw.complete``.""" + session, text, parent, task_id, err = _side_agent_args(rid, params, "btw") if err: return err - text, parent = params.get("text", ""), params.get("session_id", "") - if not text: - return _err(rid, 4012, "text required") - task_id = f"btw_{uuid.uuid4().hex[:6]}" agent = session.get("agent") snapshot = list(getattr(agent, "_session_messages", None) or session.get("history") or []) main_runtime = { - k: getattr(agent, k, None) for k in ("model", "provider", "base_url", "api_key", "api_mode") - } + k: getattr(agent, k, None) + for k in ("model", "provider", "base_url", "api_key", "api_mode")} def body(): from agent.side_question import answer_side_question return answer_side_question( - text, snapshot, parent_agent=agent, main_runtime=main_runtime, - ) or "" + text, snapshot, parent_agent=agent, main_runtime=main_runtime) or "" return _spawn_side_agent( rid, session, task_id, parent, "btw.complete", body, extra={"question": text}) @@ -1135,9 +972,7 @@ def _(rid, params: dict) -> dict: session, err = _sess(params, rid) if err: return err - url = str(params.get("url") or "").strip() - cwd = str(params.get("cwd") or "").strip() - context = str(params.get("context") or "").strip() + url, cwd, context = (str(params.get(k) or "").strip() for k in ("url", "cwd", "context")) if not url: return _err(rid, 4012, "url required") task_id = f"preview_{uuid.uuid4().hex[:6]}" @@ -1153,8 +988,7 @@ def _(rid, params: dict) -> dict: _PREVIEW_RESTART_HISTORY_NOTE if parent_history else None, *_PREVIEW_RESTART_RULES] if line) - # A malformed client path (embedded NUL, etc.) must not blow up the restart: - # treat it as "no validated cwd". + # A malformed client path (embedded NUL, etc.) is "no validated cwd". try: preview_cwd = os.path.abspath(os.path.expanduser(cwd)) if cwd else "" if preview_cwd and not os.path.isdir(preview_cwd): @@ -1173,9 +1007,8 @@ def _(rid, params: dict) -> dict: _emit( "preview.restart.progress", parent, {"task_id": task_id, "text": f"Starting hidden restart agent{history_note}"}) - # Deliberately NOT closed through task-wide process cleanup: the whole - # point is to leave a background server running under this task_id, - # and AIAgent.close() would kill every process for it. + # Deliberately NOT closed via AIAgent.close(): it would kill the background + # server this task exists to leave running. result = AIAgent( **_ephemeral_preview_agent_kwargs(session["agent"], task_id), **_preview_restart_callbacks(parent, task_id), @@ -1188,20 +1021,15 @@ def _(rid, params: dict) -> dict: from tools.terminal_tool import clear_task_env_overrides clear_task_env_overrides(task_id) - # Pin the validated preview cwd, else the parent workspace — never an - # invalid client path (which would silently fall back to the launch dir). + # Pin the validated preview cwd, else the parent workspace — never an invalid path. return _spawn_side_agent( rid, session, task_id, parent, "preview.restart.complete", body, cwd=preview_cwd, cleanup=cleanup) # ── late-answer RPCs for tool-driven UI cards ─────────────────────────────── - - -# All use allow_expired=True: each tool's bounded wait (read_terminal 30s, -# setup_mcp 10min, clarify ...) can expire — its _pending entry popped — while the -# card is still visible (e.g. a WS reconnect dropped tool.complete). A late answer -# must resolve gracefully instead of the raw 4009 "no pending answer request". +# allow_expired=True everywhere: a tool's bounded wait can expire (its _pending entry +# popped) while the card is still visible; a late answer must not surface the raw 4009. @method("clarify.respond") @@ -1215,22 +1043,13 @@ _LATE_RESPOND_KEYS = { "terminal.read.respond": "text", "preview.read.respond": "text", "preview.act.respond": "text", "window.read.respond": "text", "tour.respond": "text", "mcp.setup.respond": "result", "sudo.respond": "password", "secret.respond": "value"} - - -def _late_respond(key: str): - def handler(rid, params: dict) -> dict: - return _respond(rid, params, key, allow_expired=True) - return handler - - for _name, _key in _LATE_RESPOND_KEYS.items(): - method(_name)(_late_respond(_key)) + method(_name)(lambda rid, params, _k=_key: _respond(rid, params, _k, allow_expired=True)) del _name, _key # ── approvals ─────────────────────────────────────────────────────────────── - def _approval_reply(rid, result_key, call): """``_ok({result_key: call(tools.approval)})``, 5004 on any failure.""" try: @@ -1254,19 +1073,16 @@ def _(rid, params: dict) -> dict: session, err = _sess(params, rid) if err: return err - request_id = params.get("request_id") - if not isinstance(request_id, str) or not request_id: + if not isinstance(request_id := params.get("request_id"), str) or not request_id: return _err(rid, 4006, "request_id required") return _approval_reply( rid, "acknowledged", lambda a: a.ack_gateway_approval(session["session_key"], request_id)) def _approval_respond_session_fallback(params: dict): - """Durable-identity fallback for ``approval.respond``: the desktop can answer - with a stale live sid (runtime re-minted after a reconnect while the prompt - stayed on screen). Try (1) the approval ``request_id`` (unique across sessions) - against every live session's pending approvals, then (2) ``session_id`` as a - STORED id mapped to its live record. Returns the live session or None.""" + """Durable-identity fallback for a stale live sid (re-minted after a reconnect while + the prompt stayed on screen): (1) the ``request_id`` against every live session's + pending approvals, then (2) ``session_id`` as a STORED id. Live session or None.""" request_id = str(params.get("request_id") or "") if request_id: try: @@ -1281,11 +1097,9 @@ def _approval_respond_session_fallback(params: dict): return session except Exception: logger.debug("approval.respond request_id fallback failed", exc_info=True) - target = str(params.get("session_id") or "") - if target: + if target := str(params.get("session_id") or ""): try: - live = _find_live_session_by_key(target) - if live is not None: + if (live := _find_live_session_by_key(target)) is not None: return live[1] except Exception: logger.debug("approval.respond stored-id fallback failed", exc_info=True) diff --git a/tui_gateway/methods_session.py b/tui_gateway/methods_session.py index 79a8647199..d0cd8c723b 100644 --- a/tui_gateway/methods_session.py +++ b/tui_gateway/methods_session.py @@ -1,13 +1,10 @@ """Session / delegation / spawn-tree / billing / pet JSON-RPC handlers. -Bodies are rebound onto server.py's globals at install time (method_ctx.py), so -they use server helpers (``_sessions``, ``_ok``, ``_err``, ...) bare; module-level -helpers are published onto server.py the same way (tests monkeypatching ``server.X`` -still intercept). -""" +Bodies are rebound onto server.py's globals at install time (method_ctx.py), so they use server +helpers (``_sessions``, ``_ok``, ``_err``, ...) bare; module-level helpers are published onto +server.py the same way (tests monkeypatching ``server.X`` still intercept).""" import contextlib -from dataclasses import dataclass from .method_ctx import HandlerRegistry, bind_module @@ -17,18 +14,12 @@ _profile_scoped = _registry.profile_scoped # ── shared handler plumbing ────────────────────────────────────────── - - def _session_arg(resolve): - """Resolve ``params.session_id`` with ``resolve`` and pass the record as a 3rd arg. ``resolve`` is - a lambda over the server helper: decoration runs before bind_module publishes ``_sess*``.""" - + """Resolve ``params.session_id`` via ``resolve`` (a lambda — decoration precedes bind_module) → 3rd arg.""" def deco(fn): def handler(rid, params: dict) -> dict: session, err = resolve(params, rid) - if err: - return err - return fn(rid, params, session) + return err or fn(rid, params, session) return handler return deco @@ -37,25 +28,30 @@ _with_session = _session_arg(lambda params, rid: _sess_nowait(params, rid)) # n _with_live_session = _session_arg(lambda params, rid: _sess(params, rid)) # waits for the agent build -def _with_session_db(code: int): - """:func:`_with_session` plus the session's db as a 4th arg (``_db_unavailable_error(code)`` when None).""" +def _session_method(name: str, *, live: bool = False): + """``@method(name)`` over ``_with_live_session`` (waits for the agent build) or ``_with_session``.""" + return lambda fn: method(name)((_with_live_session if live else _with_session)(fn)) + +def _with_db(code: int, *, session_scoped: bool): + """Append a db arg — the session's db (after ``_with_session``) or ``_profile_db(params)``; ``code`` when None.""" def deco(fn): - def handler(rid, params: dict) -> dict: - session, err = _sess_nowait(params, rid) - if err: - return err - with _session_db(session) as db: + def handler(rid, params: dict, *session) -> dict: + with (_session_db(session[0]) if session_scoped else _profile_db(params)) as db: if db is None: return _db_unavailable_error(rid, code=code) - return fn(rid, params, session, db) - return handler + return fn(rid, params, *session, db) + return _with_session(handler) if session_scoped else handler return deco -def _new_runtime_ids(params: dict) -> tuple[str, str]: - """Fresh runtime sid + resolved DB ``source`` for a session minted from ``params``.""" - return (uuid.uuid4().hex[:8], _resolve_session_source(str(params.get("source") or "").strip() or None)) +def _str_param(params: dict, key: str, default: str = "") -> str: + """``str(params[key]).strip()`` with ``default`` for missing / falsy values.""" + return str(params.get(key) or "").strip() or default + + +def _flag(params: dict, name: str) -> bool: + return is_truthy_value(params.get(name, False)) def _int_param(params: dict, key: str, default: int) -> int: @@ -66,10 +62,14 @@ def _int_param(params: dict, key: str, default: int) -> int: return default +def _new_runtime_ids(params: dict) -> tuple[str, str]: + """Fresh runtime sid + resolved DB ``source`` for a session minted from ``params``.""" + return uuid.uuid4().hex[:8], _resolve_session_source(_str_param(params, "source") or None) + + @contextlib.contextmanager def _profile_build_scope(profile_home): - """Bind HERMES_HOME + the profile's secret scope while building/initializing an agent. The home - override alone only moves config/skills/memory; unscoped get_secret() reads the LAUNCH .env.""" + """Bind HERMES_HOME + secret scope for an agent build (home alone leaves get_secret() on the LAUNCH .env).""" if not profile_home: yield return @@ -91,6 +91,20 @@ def _make_agent_in_context(sid: str, key: str, **kwargs): _clear_session_context(tokens) +def _profile_session_db(profile_home): + """``(db, owns)``: a DEDICATED handle on ``profile_home``'s state.db, else the shared launch db.""" + if profile_home: + from hermes_state import get_shared_session_db + return get_shared_session_db(Path(profile_home) / "state.db"), True + return _get_db(), False + + +def _release_db(db) -> None: + with contextlib.suppress(Exception): + from hermes_state import release_or_close + release_or_close(db) + + def _branch_title(db, parent_key: str) -> str: """Next title in the parent's lineage (mirrors the TUI /branch naming).""" current = db.get_session_title(parent_key) or "branch" @@ -101,26 +115,22 @@ def _branch_title(db, parent_key: str) -> str: def _cwd_info(session: dict, cwd: str, branch=None) -> dict: """session.info after a cwd change: the full agent view, or the lazy shape.""" - agent = session.get("agent") - if agent is not None: + if (agent := session.get("agent")) is not None: return _session_info(agent, session) - return { - "cwd": cwd, "branch": _git_branch_for_cwd(cwd) if branch is None else branch, - "project": _project_info_for_cwd(cwd), "lazy": True} + return {"cwd": cwd, "branch": _git_branch_for_cwd(cwd) if branch is None else branch, + "project": _project_info_for_cwd(cwd), "lazy": True} def _session_row_summary(row: dict, *, tip_row: dict | None = None, resolved_id=None) -> dict: """Compact session.list row; ``tip_row``/``resolved_id`` come from the compression tip.""" tip_row = tip_row or row - return { - "id": row["id"], **({} if resolved_id is None else {"resolved_id": resolved_id}), - "title": row.get("title") or "", "preview": tip_row.get("preview") or "", - "started_at": row.get("started_at") or 0, "message_count": tip_row.get("message_count") or 0, - "source": row.get("source") or ""} + return {"id": row["id"], **({} if resolved_id is None else {"resolved_id": resolved_id}), + "title": row.get("title") or "", "preview": tip_row.get("preview") or "", + "started_at": row.get("started_at") or 0, "message_count": tip_row.get("message_count") or 0, + "source": row.get("source") or ""} -# Hidden from human-facing listings (sub-agent runs, kanban workers). A deny-list so -# new platforms / custom HERMES_SESSION_SOURCE values surface automatically. +# Hidden from human listings (sub-agent runs, kanban workers); a deny-list so new platforms surface automatically. _LISTING_DENY_SOURCES = frozenset({"kanban", "tool"}) @@ -130,8 +140,7 @@ def _denied_source(row: dict) -> bool: def _listing_rows(db, limit: int, **kwargs) -> list: """Human-facing ``list_sessions_rich`` rows (most recent first), deny-list applied.""" - rows = db.list_sessions_rich( - source=None, limit=limit, order_by_last_active=True, compact_rows=True, **kwargs) + rows = db.list_sessions_rich(source=None, limit=limit, order_by_last_active=True, compact_rows=True, **kwargs) return [row for row in rows if not _denied_source(row)] @@ -155,23 +164,6 @@ def _pet_display_cfg() -> dict: return {} -def _pet_guard(name: str, *, fail_open=None): - """Pet handlers never break the surface: exceptions log at debug and yield ``fail_open`` - (payload or ``params -> payload`` callable) or, without it, ``_err(5031, " failed: ...")``.""" - - def deco(fn): - def handler(rid, params: dict) -> dict: - try: - return fn(rid, params) - except Exception as exc: # noqa: BLE001 - cosmetic surface - logger.debug("%s failed: %s", name, exc) - if fail_open is not None: - return _ok(rid, fail_open(params) if callable(fail_open) else dict(fail_open)) - return _err(rid, 5031, f"{name} failed: {exc}") - return handler - return deco - - def _pet_emit(event: str, payload: dict, what: str) -> None: """Best-effort progress emit: a transport hiccup must never abort generation.""" try: @@ -186,22 +178,20 @@ def _pet_gen_abort(rid, token: str, code: int, message: str) -> dict: return _err(rid, code, message) -def _with_slug(fn): - """Require ``params.slug`` (4004 "missing slug") and pass it as a 3rd arg.""" - - def handler(rid, params: dict) -> dict: - slug = str(params.get("slug") or "").strip() - if not slug: - return _err(rid, 4004, "missing slug") - return fn(rid, params, slug) - return handler - - def _pet_method(name: str, *, fail_open=None, slug: bool = False, scoped: bool = True): - """``@method(name)`` + ``@_profile_scoped`` (unless ``scoped=False``) + ``_pet_guard`` (+ ``_with_slug``).""" - + """``@method`` (+ ``@_profile_scoped`` unless ``scoped=False``) whose exceptions never break the surface: logged + at debug, then ``fail_open`` (payload or ``params -> payload``) or ``_err(5031)``. ``slug``: 3rd arg (4004).""" def deco(fn): - handler = _pet_guard(name, fail_open=fail_open)(_with_slug(fn) if slug else fn) + def handler(rid, params: dict) -> dict: + try: + if slug and not (value := _str_param(params, "slug")): + return _err(rid, 4004, "missing slug") + return fn(rid, params, value) if slug else fn(rid, params) + except Exception as exc: # noqa: BLE001 - cosmetic surface + logger.debug("%s failed: %s", name, exc) + if fail_open is not None: + return _ok(rid, fail_open(params) if callable(fail_open) else dict(fail_open)) + return _err(rid, 5031, f"{name} failed: {exc}") return method(name)(_profile_scoped(handler) if scoped else handler) return deco @@ -209,14 +199,11 @@ def _pet_method(name: str, *, fail_open=None, slug: bool = False, scoped: bool = def _active_pet(): """``(pet, scale)`` when the pet display is enabled and the pet exists, else None.""" enabled, pet, scale = _pet_active_selection() - if not enabled or pet is None or not pet.exists: - return None - return pet, scale + return None if not enabled or pet is None or not pet.exists else (pet, scale) def _billing_call(rid, fn, extra: dict | None = None) -> dict: - """Run a portal call; BillingError → serialized envelope, anything else → generic. - ``extra`` rides both ERROR envelopes (e.g. the idempotency key the TUI reuses on retry).""" + """Portal call → ok; BillingError → serialized envelope, else generic; ``extra`` rides both ERROR envelopes.""" from hermes_cli.nous_billing import BillingError try: return _ok(rid, fn()) @@ -240,70 +227,56 @@ def _billing_pending_change(result: dict) -> dict: # ── session.create / list / most_recent / facts ────────────────────── - - -def _create_branch_row(db, new_key: str, parent_key: str, *, source, cwd, profile_name) -> None: - """Create a branch child row. ``_branched_from`` keeps it visible in list_sessions_rich() (the parent - stays live, so the legacy end_reason='branched' heuristic never matches); ``profile_name`` is stamped - explicitly — NULL rows drop out of profile-keyed sidebar matching / deep links.""" - db.create_session( - new_key, source=source, model=_resolve_model(), model_config={"_branched_from": parent_key}, - parent_session_id=parent_key, cwd=cwd, profile_name=profile_name) - - -def _copy_branch_transcript(db, new_key: str, title: str, history: list, copy_fields=()) -> None: - """Copy the parent transcript in bounded-chunk transactions, then title the child.""" - db.append_messages_batch( - new_key, - [{"role": msg.get("role", "user"), "content": msg.get("content"), - **{field: msg.get(field) for field in copy_fields}} for msg in history], - chunk_rows=500) - db.set_session_title(new_key, title) - - -def _seed_branch_row(sid: str, key: str, parent_session_id: str, history: list, source: str, profile_home) -> None: - """Persist a seeded desktop branch child up front (the one session.create exception to lazy rows): - the renderer's post-create resume re-fetches the child via REST/defer_history, so an unpersisted - child 404s and the fail-latch spins forever. Best-effort: on failure the lazy first-prompt path - is the fallback, as for plain drafts.""" +def _persist_branch(db, new_key: str, parent_key: str, title: str, history: list, *, source, cwd, profile_name, + copy_fields=(), compensate: bool = False) -> None: + """Branch child row + parent transcript (bounded-chunk transactions) + title. ``_branched_from`` keeps the + row visible in list_sessions_rich() (the live parent never matches the legacy end_reason='branched' + heuristic); NULL ``profile_name`` rows drop out of profile-keyed sidebar matching / deep links. ``compensate`` + deletes a committed row whose transcript/title failed (a durable-but-empty row would defeat the INSERT OR + IGNORE first-prompt seed) — except on disk-full, where the delete cannot land.""" + db.create_session(new_key, source=source, model=_resolve_model(), model_config={"_branched_from": parent_key}, + parent_session_id=parent_key, cwd=cwd, profile_name=profile_name) try: - with _session_db(_sessions[sid]) as db: + db.append_messages_batch( + new_key, [{"role": msg.get("role", "user"), "content": msg.get("content"), + **{field: msg.get(field) for field in copy_fields}} for msg in history], chunk_rows=500) + db.set_session_title(new_key, title) + except Exception as exc: + from hermes_state import is_disk_full_error + if compensate and not is_disk_full_error(exc): + try: + db.delete_session(new_key) + except Exception: + logger.debug("branch seed compensation delete failed for %s", new_key, exc_info=True) + raise + + +def _seed_branch_row(record: dict, key: str, parent_session_id: str, history: list, source: str, profile_home): + """Persist a seeded desktop branch child NOW (the one session.create exception to lazy rows): the + renderer's post-create resume re-fetches it via REST/defer_history, so an unpersisted child 404s and + the fail-latch spins forever. Best-effort — on failure the lazy first-prompt path is the fallback.""" + try: + with _session_db(record) as db: if db is None: return - branch_title = _branch_title(db, parent_session_id) - _create_branch_row( - db, key, parent_session_id, source=source, cwd=_sessions[sid]["cwd"], - profile_name=(Path(profile_home).name if profile_home else None)) - try: - _copy_branch_transcript(db, key, branch_title, history) - except Exception as exc: - # Row committed but transcript/title failed: a durable-but-empty row would defeat the - # INSERT OR IGNORE first-prompt seed — roll back this child. - from hermes_state import is_disk_full_error - if is_disk_full_error(exc): - raise - try: - db.delete_session(key) - except Exception: - logger.debug("branch seed compensation delete failed for %s", key, exc_info=True) - raise - _sessions[sid]["pending_title"] = None + _persist_branch(db, key, parent_session_id, _branch_title(db, parent_session_id), history, + source=source, cwd=record["cwd"], + profile_name=(Path(profile_home).name if profile_home else None), compensate=True) + record["pending_title"] = None except Exception: - logger.warning( - "seeded-branch persistence failed for %s; falling back to lazy row creation", key, exc_info=True, - ) + logger.warning("seeded-branch persistence failed for %s; falling back to lazy row creation", key, + exc_info=True) def _create_overrides(params: dict) -> tuple: - """(model_override, reasoning_override, service_tier_override) from the composer's UI state. - PER-SESSION only — never a global config write. ``fast`` presence is the contract: omitted - inherits, true pins priority, false pins normal ("" — _make_agent uses None for inheritance).""" - create_model = str(params.get("model") or "").strip() - model_override = ( - {"model": create_model, "provider": str(params.get("provider") or "").strip() or None} - if create_model else None) + """PER-SESSION (model, reasoning, service_tier) overrides from the composer — never a global config + write. ``fast`` presence is the contract: omitted inherits, true pins priority, false pins normal ("").""" + create_model = _str_param(params, "model") + model_override = None + if create_model: + model_override = {"model": create_model, "provider": _str_param(params, "provider") or None} reasoning_override = None - if effort := str(params.get("reasoning_effort") or "").strip(): + if effort := _str_param(params, "reasoning_effort"): with contextlib.suppress(Exception): from hermes_constants import parse_reasoning_effort reasoning_override = parse_reasoning_effort(effort) @@ -315,52 +288,43 @@ def _create_overrides(params: dict) -> tuple: @method("session.create") def _(rid, params: dict) -> dict: - sid = uuid.uuid4().hex[:8] - key = _new_session_key() - cols = int(params.get("cols", 80)) + (sid, source), key = _new_runtime_ids(params), _new_session_key() history = _coerce_seed_history(params.get("messages")) - title = str(params.get("title") or "").strip() # Branch: links back so list_sessions_rich keeps it visible and the sidebar nests it. - parent_session_id = str(params.get("parent_session_id") or "").strip() or None - # Only an explicitly chosen existing workspace persists as cwd; the launch-dir fallback lands - # in "No workspace". - raw_cwd = str(params.get("cwd") or "").strip() + parent_session_id = _str_param(params, "parent_session_id") or None + # Only an explicitly chosen existing workspace persists as cwd; the launch-dir fallback is "No workspace". explicit_cwd = False with contextlib.suppress(Exception): + raw_cwd = _str_param(params, "cwd") explicit_cwd = bool(raw_cwd) and os.path.isdir(os.path.abspath(os.path.expanduser(raw_cwd))) - resolved_cwd = _completion_cwd(params) - source = _resolve_session_source(str(params.get("source") or "").strip() or None) _enable_gateway_prompts() - # ``profile`` (app-global remote mode): stored on the session so the build and every turn - # re-bind HERMES_HOME. - profile = (params.get("profile") or "").strip() or None - profile_home = _profile_home(profile) + # ``profile`` (app-global remote mode): stored so the build and every turn re-bind HERMES_HOME. + profile_home = _profile_home(profile := (params.get("profile") or "").strip() or None) session_model_override, create_reasoning_override, create_service_tier_override = _create_overrides(params) now = time.time() with _sessions_lock: _sessions[sid] = { "agent": None, "agent_error": None, "agent_ready": threading.Event(), "attached_images": [], - "close_on_disconnect": is_truthy_value(params.get("close_on_disconnect", False)), + "close_on_disconnect": _flag(params, "close_on_disconnect"), "active_session_lease": None, # claimed lazily on the first turn (_ensure_active_session_slot) - "cols": cols, "created_at": now, "edit_snapshots": {}, "explicit_cwd": explicit_cwd, + "cols": int(params.get("cols", 80)), "created_at": now, "edit_snapshots": {}, + "explicit_cwd": explicit_cwd, "history": history, "history_lock": threading.Lock(), "history_version": 0, "image_counter": 0, - "cwd": resolved_cwd, "inflight_turn": None, "last_active": now, + "cwd": _completion_cwd(params), "inflight_turn": None, "last_active": now, "model_override": session_model_override, "create_reasoning_override": create_reasoning_override, "create_service_tier_override": create_service_tier_override, - "parent_session_id": parent_session_id, "pending_title": title or None, - "pending_hidden": is_truthy_value(params.get("hidden", False)), - "room_plumbing": is_truthy_value(params.get("room_plumbing", False)), - "follow_profile_config": is_truthy_value(params.get("follow_profile_config", False)), + "parent_session_id": parent_session_id, "pending_title": _str_param(params, "title") or None, + "pending_hidden": _flag(params, "hidden"), "room_plumbing": _flag(params, "room_plumbing"), + "follow_profile_config": _flag(params, "follow_profile_config"), "profile_home": str(profile_home) if profile_home is not None else None, "running": False, "session_key": key, "show_reasoning": _load_show_reasoning(), "source": source, "slash_worker": None, "tool_progress_mode": _load_tool_progress_mode(), "tool_started_at": {}, "transport": current_transport() or _stdio_transport} _register_session_cwd(_sessions[sid]) - # No DB row here (drafts left "Untitled" litter): created on the first prompt — except seeded - # branch children, which must exist now. + # No DB row here (drafts left "Untitled" litter): created on the first prompt — except seeded branch children. if parent_session_id and history: - _seed_branch_row(sid, key, parent_session_id, history, source, profile_home) + _seed_branch_row(_sessions[sid], key, parent_session_id, history, source, profile_home) # Return immediately so Ink can paint; the AIAgent builds right after the flush. _schedule_agent_build(sid) _schedule_session_cap_enforcement() # trim detached idle sessions over the cap @@ -369,81 +333,67 @@ def _(rid, params: dict) -> dict: return _ok(rid, { "session_id": sid, "stored_session_id": key, "message_count": len(history), "messages": _history_to_messages(history), - "info": { - # Reflect the override now so the client doesn't clobber its sticky pick. - "model": override.get("model") if override else _resolve_model(), - **({"provider": override["provider"]} if override.get("provider") else {}), - "tools": {}, "skills": {}, "cwd": cwd, "branch": _git_branch_for_cwd(cwd), - "project": _project_info_for_cwd(cwd), "lazy": True, "desktop_contract": DESKTOP_BACKEND_CONTRACT, - "profile_name": _response_profile_name(profile)}}) + # Reflect the override now so the client doesn't clobber its sticky pick. + "info": {"model": override.get("model") if override else _resolve_model(), + **({"provider": override["provider"]} if override.get("provider") else {}), + "tools": {}, "skills": {}, "cwd": cwd, "branch": _git_branch_for_cwd(cwd), + "project": _project_info_for_cwd(cwd), "lazy": True, "desktop_contract": DESKTOP_BACKEND_CONTRACT, + "profile_name": _response_profile_name(profile)}}) def _session_list_by_title(rid, db, title_lookup: str) -> dict: - """EXACT-title lookup for callers that treat a title as identity; window-free on purpose (a busy - profile's windowed listing can push the row out). Hidden rows resolve (canonical chats are born - hidden); archived / deny-listed do not; lineages resolve to the live tip (``resolved_id``).""" + """EXACT-title lookup (title as identity), window-free on purpose (a busy profile's windowed listing can + push the row out). Hidden rows resolve (canonical chats are born hidden); archived / deny-listed do not; + lineages resolve to the live tip (``resolved_id``).""" row = db.get_session_by_title(title_lookup) if row and row.get("archived"): from tools.bot_mode_probe import BOT_CHAT_TITLE - # A Bot Chat archived by the ws-orphan reaper / agent_close is an accident (the desktop - # would mint replacements forever): resurrect recoverable reasons only. Re-fetch by ID — - # title is not UNIQUE. + # A Bot Chat archived by the ws-orphan reaper / agent_close is an accident (the desktop would mint + # replacements forever): resurrect recoverable reasons only. Re-fetch by ID — title is not UNIQUE. if title_lookup == BOT_CHAT_TITLE and db.unarchive_recoverable_session(row["id"]): row = db.get_session(row["id"]) if not row or row.get("archived") or _denied_source(row): return _ok(rid, {"sessions": []}) - try: - # Real compression continuation only: the resume resolver's unmarked-child - # fallback could redirect the canonical Bot Chat to an unrelated child. + tip = row["id"] + with contextlib.suppress(Exception): + # Real compression continuation only: the resolver's unmarked-child fallback could redirect Bot Chat. tip = db.get_compression_tip(row["id"]) or row["id"] - except Exception: - tip = row["id"] tip_row = (db.get_session(tip) or row) if tip != row["id"] else row return _ok(rid, {"sessions": [_session_row_summary(row, tip_row=tip_row, resolved_id=tip)]}) @method("session.list") -def _(rid, params: dict) -> dict: - with _profile_db(params) as db: - if db is None: - return _db_unavailable_error(rid, code=5006) - try: - if title_lookup := str(params.get("title") or "").strip(): - return _session_list_by_title(rid, db, title_lookup) - limit = int(params.get("limit", 200) or 200) - # Over-fetch: per-source filtering + tip merging must not leave us short. - # ``include_hidden`` is for surfaces that OWN hidden sessions (Bots pane, pickers). - rows = _listing_rows( - db, max(limit * 2, 200), include_hidden=is_truthy_value(params.get("include_hidden", False)), - )[:limit] - return _ok(rid, {"sessions": [_session_row_summary(s) for s in rows]}) - except Exception as e: - return _err(rid, 5006, str(e)) +@_with_db(5006, session_scoped=False) +def _(rid, params: dict, db) -> dict: + try: + if title_lookup := _str_param(params, "title"): + return _session_list_by_title(rid, db, title_lookup) + limit = int(params.get("limit", 200) or 200) + # Over-fetch: per-source filtering + tip merging must not leave us short. ``include_hidden`` is for + # surfaces that OWN hidden sessions (Bots pane, pickers). + rows = _listing_rows(db, max(limit * 2, 200), include_hidden=_flag(params, "include_hidden"))[:limit] + return _ok(rid, {"sessions": [_session_row_summary(s) for s in rows]}) + except Exception as e: + return _err(rid, 5006, str(e)) @method("session.most_recent") def _(rid, params: dict) -> dict: - """Most recent human-facing session id (same deny-list as session.list), honoring ``params.profile``. - Errors fold into ``{"session_id": null}`` (and log) so callers never special-case envelopes.""" + """Most recent human-facing session (session.list deny-list); errors fold into ``session_id: null``.""" with _profile_db(params) as db: - if db is None: - return _ok(rid, {"session_id": None}) try: # Generous over-fetch: many ``tool`` rows must not yield a false "none". - for row in _listing_rows(db, 200)[:1]: - return _ok(rid, { - "session_id": row.get("id"), "title": row.get("title") or "", - "started_at": row.get("started_at") or 0, "source": row.get("source") or ""}) - return _ok(rid, {"session_id": None}) + for row in _listing_rows(db, 200)[:1] if db is not None else (): + return _ok(rid, {"session_id": row.get("id"), "title": row.get("title") or "", + "started_at": row.get("started_at") or 0, "source": row.get("source") or ""}) except Exception: logger.exception("session.most_recent failed") - return _ok(rid, {"session_id": None}) + return _ok(rid, {"session_id": None}) @method("project.facts") def _(rid, params: dict) -> dict: - """Project facts for a cwd — the coding-context detection the system prompt uses, so UIs - don't re-sniff. ``{"facts": null}`` = not a code workspace.""" + """The system prompt's coding-context detection for a cwd (UIs don't re-sniff); null = not code.""" try: from agent.coding_context import project_facts_for return _ok(rid, {"facts": project_facts_for(params.get("cwd"))}) @@ -459,50 +409,44 @@ def _(rid, params: dict) -> dict: never upgrades targeted evidence into a repository-wide guarantee.""" try: from agent.verification_evidence import verification_status - return _ok( - rid, - { - "verification": verification_status( - session_id=params.get("session_id") or params.get("session_key"), cwd=params.get("cwd"), - )}) + return _ok(rid, {"verification": verification_status( + session_id=params.get("session_id") or params.get("session_key"), cwd=params.get("cwd"))}) except Exception: logger.exception("verification.status failed") return _ok(rid, {"verification": {"status": "unknown", "evidence": None}}) # ── session.resume ─────────────────────────────────────────────────── - - -# repr/eq off: dataclass-generated methods read their own module globals, which -# bind_module cannot rebind. -@dataclass(repr=False, eq=False) class _Resume: """Per-call ``session.resume`` state. ``owns_db``: the DEDICATED profile handle is ours to close (handler ``finally``) until handed to the hydration worker or the agent.""" - rid: object - params: dict - target: str - cols: int - profile: str | None - profile_home: object - lazy: bool - defer_history: bool - omit_messages: bool - eager_build: bool - db: object = None - owns_db: bool = False - found: dict | None = None - profile_resume_cwd: str = "" + def __init__(self, rid, params: dict, target: str) -> None: + self.rid, self.params, self.target = rid, params, target + self.db, self.owns_db, self.found, self.profile_resume_cwd = None, False, None, "" + self.cols = _int_param(params, "cols", 80) + # ``profile`` (app-global remote mode): resume from another local profile's state.db. + self.profile = (params.get("profile") or "").strip() or None + self.profile_home = _profile_home(self.profile) + self.lazy, self.defer_history = _flag(params, "lazy"), _flag(params, "defer_history") + # Desktop hydrates over REST; suppress the duplicate WS copy only when asked. + self.omit_messages, self.eager_build = _flag(params, "omit_messages"), _flag(params, "eager_build") - def cwd(self) -> str: - return self.profile_resume_cwd or _default_session_cwd() + def mint(self, prompts: bool = True) -> tuple: + """``(runtime sid, source, cwd)`` for the live record this resume registers (+ gateway prompts on).""" + ids = _new_runtime_ids(self.params) + if prompts: + _enable_gateway_prompts() + return *ids, self.profile_resume_cwd or _default_session_cwd() - def record(self, source: str, cwd: str, history: list, **extra) -> dict: - """``_deferred_session_record`` with this resume's common fields (lease claimed lazily on turn 1).""" + def record(self, source: str, cwd: str, history: list, overrides: dict | None = None, **extra) -> dict: + """``_deferred_session_record`` with this resume's common fields (lease claimed lazily on turn 1); + ``overrides`` restores the stored model/provider/reasoning/tier so the deferred build matches eager.""" + if overrides is not None: + extra.update(model_override=overrides.get("model_override"), resume_runtime_overrides=overrides or None) return _deferred_session_record( self.target, cols=self.cols, cwd=cwd, history=history, lease=None, source=source, - close_on_disconnect=is_truthy_value(self.params.get("close_on_disconnect", False)), + close_on_disconnect=_flag(self.params, "close_on_disconnect"), profile_home=self.profile_home, explicit_cwd=bool(self.profile_resume_cwd), **extra) def claim(self, sid: str, record: dict) -> dict | None: @@ -510,19 +454,33 @@ class _Resume: live = _claim_or_reuse_live(sid, self.target, record, None) return None if live is None else _resume_reuse_live(self, *live) - def resume_failed(self, exc) -> dict: - return _err(self.rid, 5000, f"resume failed: {exc}") + def restore(self): + """``(sanitized model history, display history, raw history)`` for a cold/eager resume.""" + raw, display = self.read_history() + return sanitize_replay_history(raw), display, raw def info(self, cwd: str, overrides: dict) -> dict: - model_override = overrides.get("model_override") or {} - return _lazy_resume_info( - cwd, model=model_override.get("model") or "", provider=overrides.get("provider_override") or "", - profile=self.profile) + return _lazy_resume_info(cwd, model=(overrides.get("model_override") or {}).get("model") or "", + provider=overrides.get("provider_override") or "", profile=self.profile) def child_history(self, repair: bool) -> list: """The child's OWN conversation (no ancestors), row ids included.""" - return self.db.get_messages_as_conversation( - self.target, repair_alternation=repair, include_row_ids=True) + return self.db.get_messages_as_conversation(self.target, repair_alternation=repair, include_row_ids=True) + + def messages(self, display: list) -> list: + return [] if self.omit_messages else _history_to_messages(display) + + def read_history(self) -> tuple: + """One lineage SELECT, two projections: model-fed copy alternation-repaired (healed once + here instead of every turn's pre-request repair), display copy verbatim.""" + self.db.reopen_session(self.target) + if self.omit_messages: + return self.child_history(repair=True), [] + return self.db.get_resume_conversations(self.target) + + def display_prefix(self) -> list: + """Ancestor display rows (model-fed history drops a dangling tool-call tail — display keeps it).""" + return [] if self.omit_messages else self.db.get_ancestor_display_prefix(self.target) def _find_live_unpersisted(needle: str, home) -> str: @@ -531,21 +489,17 @@ def _find_live_unpersisted(needle: str, home) -> str: return next(( live_sid for live_sid, record in list(_sessions.items()) if isinstance(record, dict) and (record.get("profile_home") or None) == want_home - and (str(record.get("session_key") or "") == needle or (record.get("pending_title") or "") == needle) - ), "") + and (str(record.get("session_key") or "") == needle or (record.get("pending_title") or "") == needle)), "") def _resume_live_unpersisted(ctx: _Resume, live_sid: str, live: dict) -> dict: - """Reattach a LIVE lazy session with no state.db row yet (every fresh Bot Chat; a hard 404 here - killed messaging for bots that had never spoken). A WS drop may have sentinel-parked the record: - rebind the transport and cancel the armed orphan-reap Timer or it fires against this client.""" + """Reattach a LIVE lazy session with no state.db row yet (every fresh Bot Chat; a 404 here killed messaging + for never-spoken bots). Rebind the transport and cancel the armed orphan-reap Timer (a WS drop may have + sentinel-parked the record) or it fires against this client.""" if ctx.owns_db: - with contextlib.suppress(Exception): - from hermes_state import release_or_close - release_or_close(ctx.db) + _release_db(ctx.db) live["last_active"] = time.time() - transport = current_transport() - if transport is not None: + if (transport := current_transport()) is not None: with live.setdefault("history_lock", threading.Lock()): live["transport"] = transport live.setdefault("viewers", {})[transport] = time.time() @@ -553,15 +507,13 @@ def _resume_live_unpersisted(ctx: _Resume, live_sid: str, live: dict) -> dict: history = live.get("history") or [] return _ok(ctx.rid, _attach_todo_state({ "session_id": live_sid, "stored_session_id": str(live.get("session_key") or ""), - "message_count": len(history), "messages": [] if ctx.omit_messages else _history_to_messages(history), - "info": {"model": _resolve_model(), "lazy": True, "profile_name": ctx.profile or ""}, - }, live)) + "message_count": len(history), "messages": ctx.messages(history), + "info": {"model": _resolve_model(), "lazy": True, "profile_name": ctx.profile or ""}}, live)) def _resume_adopt_stranded(ctx: _Resume) -> None: - """Adopt a lineage stranded in the DEFAULT store into this profile's db (older builds ran a profile - bot's turns on the focused tile's backend; without adoption that chat 4001s forever). Exact-id match - ONLY — bot titles collide by design. Never re-adopt a retired donor (two "canonical" clones).""" + """Adopt a lineage stranded in the DEFAULT store (older builds ran a profile bot's turns on the focused + tile's backend; unadopted it 4001s forever). Exact-id ONLY — bot titles collide; never a retired donor.""" try: default_db = _get_db() donor_row = default_db.get_session(ctx.target) if default_db is not None else None @@ -569,11 +521,10 @@ def _resume_adopt_stranded(ctx: _Resume) -> None: return adoption = ctx.db.adopt_session_lineage_from(default_db, donor_row["id"]) if adoption.get("adopted"): - logger.info( - "adopted stranded session %s (lineage of %s segment(s)) from default store into profile %s", - donor_row["id"], - len(adoption.get("imported_ids") or []) + len(adoption.get("skipped_ids") or []), - ctx.profile or "?") + logger.info("adopted stranded session %s (lineage of %s segment(s)) from default store into profile %s", + donor_row["id"], + len(adoption.get("imported_ids") or []) + len(adoption.get("skipped_ids") or []), + ctx.profile or "?") ctx.found = ctx.db.get_session(donor_row["id"]) if ctx.found: ctx.target = ctx.found["id"] @@ -583,46 +534,39 @@ def _resume_adopt_stranded(ctx: _Resume) -> None: def _resume_locate(ctx: _Resume) -> dict | None: """Resolve ``ctx.target`` to a stored row (``ctx.found``); a dict is an early response.""" - db = ctx.db - ctx.found = db.get_session(ctx.target) + ctx.found = ctx.db.get_session(ctx.target) if ctx.found: return None - ctx.found = db.get_session_by_title(ctx.target) + ctx.found = ctx.db.get_session_by_title(ctx.target) if ctx.found: ctx.target = ctx.found["id"] return None if ctx.lazy and _child_run_active(ctx.target): - # Fresh subagent watch window: `subagent.start` relays BEFORE the child's first DB flush. - # Proceed lazily with empty history — the live mirror streams the turn and the row exists - # by upgrade time. + # Fresh subagent watch window: `subagent.start` relays BEFORE the child's first DB flush. Proceed lazily + # with empty history — the live mirror streams the turn and the row exists by upgrade time. ctx.found = {} return None live_sid = _find_live_unpersisted(ctx.target, ctx.profile_home) - live = _sessions.get(live_sid) if live_sid else None - if live is not None: + if (live := _sessions.get(live_sid) if live_sid else None) is not None: return _resume_live_unpersisted(ctx, live_sid, live) if ctx.owns_db: _resume_adopt_stranded(ctx) - if not ctx.found: - return _err(ctx.rid, 4007, "session not found") - return None + return None if ctx.found else _err(ctx.rid, 4007, "session not found") def _resume_follow_tip(ctx: _Resume) -> None: - """Rebind a rotated-out parent id to its compression-continuation tip (resuming the original would - reload the parent transcript and lose the post-compression reply; the live fast path also reuses the - rotated key). Skipped for lazy watch windows (exact child). Bot Chat follows proven compression - edges only; others keep the unmarked-child walker.""" + """Rebind a rotated-out parent id to its compression tip (resuming the original reloads the parent + transcript and loses the post-compression reply). Skipped for lazy watch windows (exact child); Bot Chat + follows proven compression edges only.""" if not ctx.found or ctx.lazy: return - try: + tip = ctx.target + with contextlib.suppress(Exception): from tools.bot_mode_probe import BOT_CHAT_TITLE if (ctx.found.get("title") or "").strip() == BOT_CHAT_TITLE: tip = ctx.db.get_compression_tip(ctx.target) or ctx.target else: tip = ctx.db.resolve_resume_session_id(ctx.target) - except Exception: - tip = ctx.target if tip and tip != ctx.target: ctx.target = tip ctx.found = ctx.db.get_session(tip) or ctx.found @@ -630,20 +574,15 @@ def _resume_follow_tip(ctx: _Resume) -> None: def _resume_guard(ctx: _Resume) -> dict | None: """Refuse a runaway transcript before any history read (sessions.max_resume_messages). Deferred / - omit_messages / lazy paths load the TIP segment only, so they are guarded tip-only (a full-lineage - count rejected exactly the well-compressed conversations). Metadata fallback keeps lightweight - adaptor DBs compatible. Fails OPEN on guard errors.""" + omit_messages / lazy paths load the TIP segment only and are guarded tip-only (a lineage count rejected + exactly the well-compressed chats). Metadata fallback for lightweight adaptor DBs; fails OPEN on errors.""" from hermes_state import SessionResumeTooLargeError, resolved_max_resume_messages - guard_tip_only = ctx.lazy or ctx.omit_messages or (ctx.defer_history and not ctx.eager_build) - safety_check = getattr(ctx.db, "assert_resume_safe", None) + tip_only = ctx.lazy or ctx.omit_messages or (ctx.defer_history and not ctx.eager_build) try: - if callable(safety_check): - safety_check(ctx.target, **({"tip_only": True} if guard_tip_only else {})) - else: - resume_limit = resolved_max_resume_messages() - stored_message_count = int(ctx.found.get("message_count") or 0) - if resume_limit and stored_message_count > resume_limit: - raise SessionResumeTooLargeError(stored_message_count, resume_limit) + if callable(safety_check := getattr(ctx.db, "assert_resume_safe", None)): + safety_check(ctx.target, **({"tip_only": True} if tip_only else {})) + elif (limit := resolved_max_resume_messages()) and (n := int(ctx.found.get("message_count") or 0)) > limit: + raise SessionResumeTooLargeError(n, limit) except SessionResumeTooLargeError as exc: return _err(ctx.rid, 4130, str(exc)) except Exception as exc: @@ -652,27 +591,23 @@ def _resume_guard(ctx: _Resume) -> dict | None: def _resume_reuse_live(ctx: _Resume, sid: str, session: dict) -> dict: - """Reattach an already-live session under the resume lock: holding it across the - client-gone check, transport rebind and reap cancel makes grace expiry atomic.""" + """Reattach an already-live session under the resume lock (held across the client-gone check, + transport rebind and reap cancel so grace expiry is atomic).""" with _session_resume_lock: if _sessions.get(sid) is not session: return _err(ctx.rid, 4007, "session no longer live; retry resume") if session.get("_client_gone_interrupt_requested"): return _err(ctx.rid, 4009, "session disconnect interrupt settling") - # Cancel unconditionally so the fast path can never race the reap Timer. - _cancel_ws_orphan_reap(sid) - payload = _live_session_payload( - sid, session, cols=ctx.cols, touch=True, transport=current_transport() or _stdio_transport, - omit_messages=ctx.omit_messages) + _cancel_ws_orphan_reap(sid) # unconditionally: the fast path must never race the reap Timer + payload = _live_session_payload(sid, session, cols=ctx.cols, touch=True, omit_messages=ctx.omit_messages, + transport=current_transport() or _stdio_transport) payload["resumed"] = ctx.target if ctx.defer_history: - payload["messages"] = [] - payload["message_count"] = int(session.get("resume_message_count") or payload["message_count"]) - payload["hydrating"] = bool(session.get("resume_hydrating")) + payload.update(messages=[], hydrating=bool(session.get("resume_hydrating")), + message_count=int(session.get("resume_message_count") or payload["message_count"])) # A lazy watch session never owns a run loop — overlay the child-run registry. if session.get("agent") is None and _child_run_active(ctx.target): - payload["running"] = True - payload["status"] = "streaming" + payload.update(running=True, status="streaming") return _ok(ctx.rid, payload) @@ -681,139 +616,99 @@ def _resume_response( messages: list | None = None, message_count: int | None = None, running: bool = False, status: str = "idle", hydrating: bool | None = None, started_at=None, auto_continue=None, ) -> dict: - """Common resume payload. With omit_messages the count falls back to ``count_source`` - so the client still learns the stored size. ``hydrating`` replaces ``messages_omitted``.""" + """Common resume payload; omit_messages counts ``count_source`` (client still learns the stored size).""" if messages is None: - messages = [] if ctx.omit_messages else _history_to_messages(display) + messages = ctx.messages(display) if message_count is None: message_count = len(count_source) if ctx.omit_messages else len(messages) - payload = {"session_id": sid, "resumed": ctx.target, "message_count": message_count, "messages": messages} - if hydrating is None: - payload["messages_omitted"] = ctx.omit_messages - else: - payload["hydrating"] = hydrating - payload.update({ - "info": info, "inflight": None, "running": running, "session_key": ctx.target, - "started_at": record["created_at"] if started_at is None else started_at, "status": status}) + payload = {"session_id": sid, "resumed": ctx.target, "message_count": message_count, "messages": messages, + **({"messages_omitted": ctx.omit_messages} if hydrating is None else {"hydrating": hydrating}), + "info": info, "inflight": None, "running": running, "session_key": ctx.target, + "started_at": record["created_at"] if started_at is None else started_at, "status": status} if auto_continue is not None: payload["auto_continue"] = auto_continue return _ok(ctx.rid, _attach_todo_state(payload, record)) -def _resume_read_history(ctx: _Resume): - """One lineage SELECT, two projections: model-fed copy alternation-repaired (healed once - here instead of every turn's pre-request repair), display copy verbatim.""" - ctx.db.reopen_session(ctx.target) - if ctx.omit_messages: - return ctx.child_history(repair=True), [] - return ctx.db.get_resume_conversations(ctx.target) - - def _resume_lazy(ctx: _Resume) -> dict: - """Lazy/watch resume (desktop subagent windows): register the live session WITHOUT an - agent — the child runs inside the parent's turn, so the window needs stored history - plus a transport. A later prompt.submit upgrades it via _start_agent_build.""" - sid, source = _new_runtime_ids(ctx.params) + """Lazy/watch resume (desktop subagent windows): a live session WITHOUT an agent — the child runs + inside the parent's turn, so the window needs stored history + a transport; prompt.submit upgrades it.""" + sid, source, cwd = ctx.mint(prompts=False) try: ctx.db.reopen_session(ctx.target) # repair_alternation heals a durable ``user;user`` once here. history = ctx.child_history(repair=True) except Exception as e: - return ctx.resume_failed(e) - cwd = ctx.cwd() + return _err(ctx.rid, 5000, f"resume failed: {e}") record = ctx.record(source, cwd, history, lazy=True, todo_state=_todo_state_from_history(history)) if (reused := ctx.claim(sid, record)) is not None: return reused # A child mid-run emits no session events — liveness comes from the relay registry. - child_running = _child_run_active(ctx.target) - # Display uses the VERBATIM child-only projection so model-invisible rows survive; - # the repaired ``history`` still feeds live replay. + running = _child_run_active(ctx.target) + # Display uses the VERBATIM child-only projection so model-invisible rows survive; repaired ``history`` + # still feeds live replay. + display = history try: - display_history = ctx.child_history(repair=False) + display = ctx.child_history(repair=False) except Exception: logger.debug("child-watch display projection read failed", exc_info=True) - display_history = history - return _resume_response( - ctx, sid, record, info=_lazy_resume_info(cwd, profile=ctx.profile), display=display_history, - count_source=display_history, running=child_running, status="streaming" if child_running else "idle", - ) + return _resume_response(ctx, sid, record, info=_lazy_resume_info(cwd, profile=ctx.profile), display=display, + count_source=display, running=running, status="streaming" if running else "idle") def _resume_deferred(ctx: _Resume) -> dict: - """Bounded ack; the transcript hydrates in the background and pages over REST. - defer_history SUPERSEDES omit_messages: the ONE history read happens in the worker.""" - sid, source = _new_runtime_ids(ctx.params) - _enable_gateway_prompts() - overrides = _stored_session_runtime_overrides(ctx.found) or {} - cwd = ctx.cwd() - record = ctx.record( - source, cwd, [], model_override=overrides.get("model_override"), - resume_runtime_overrides=overrides or None) - record["resume_history_ready"] = threading.Event() - record["resume_hydrating"] = True - record["resume_message_count"] = int(ctx.found.get("message_count") or 0) + """Bounded ack; the transcript hydrates in the background (the ONE history read) and pages over REST.""" + sid, source, cwd = ctx.mint() + overrides = _stored_session_runtime_overrides(ctx.found) + record = ctx.record(source, cwd, [], overrides) + record.update(resume_history_ready=threading.Event(), resume_hydrating=True, + resume_message_count=int(ctx.found.get("message_count") or 0)) if (reused := ctx.claim(sid, record)) is not None: return reused _schedule_resume_hydration(sid, ctx.target, ctx.db, close_db=ctx.owns_db) - # The hydration worker now owns (and closes) the profile-scoped handle. - ctx.owns_db = False + ctx.owns_db = False # the hydration worker now owns (and closes) the profile-scoped handle _schedule_session_cap_enforcement() - return _resume_response( - ctx, sid, record, info=ctx.info(cwd, overrides), messages=[], - message_count=record["resume_message_count"], status="resuming", hydrating=True) + return _resume_response(ctx, sid, record, info=ctx.info(cwd, overrides), messages=[], + message_count=record["resume_message_count"], status="resuming", hydrating=True) def _resume_cold(ctx: _Resume) -> dict: - """Default cold resume: read the transcript, build the agent OFF the response path - (_make_agent can block for seconds; callers await this RPC before painting). Pre-warms - on a timer; _sess() builds on demand if the first prompt beats it. Unlike lazy, restores - full ancestor history + persisted runtime identity.""" - sid, source = _new_runtime_ids(ctx.params) - _enable_gateway_prompts() + """Default cold resume: transcript now, agent OFF the response path (_make_agent can block for seconds; + callers await this RPC before painting) — pre-warmed on a timer, _sess() builds on demand if the first + prompt beats it. Unlike lazy, restores full ancestor history + persisted runtime identity.""" + sid, source, cwd = ctx.mint() try: - raw_history, display_history = _resume_read_history(ctx) + history, display_history, raw_history = ctx.restore() except Exception as e: - return ctx.resume_failed(e) - # Model-fed history drops a dangling tool-call tail (killed mid-loop) — display keeps it. - prefix = [] if ctx.omit_messages else ctx.db.get_ancestor_display_prefix(ctx.target) - history = sanitize_replay_history(raw_history) - # Restore model/provider/reasoning/tier so the deferred build matches eager. - overrides = _stored_session_runtime_overrides(ctx.found) or {} - cwd = ctx.cwd() - record = ctx.record( - source, cwd, history, display_history_prefix=prefix, model_override=overrides.get("model_override"), - resume_runtime_overrides=overrides or None, todo_state=_todo_state_from_history(history)) + return _err(ctx.rid, 5000, f"resume failed: {e}") + overrides = _stored_session_runtime_overrides(ctx.found) + record = ctx.record(source, cwd, history, overrides, display_history_prefix=ctx.display_prefix(), + todo_state=_todo_state_from_history(history)) if (reused := ctx.claim(sid, record)) is not None: return reused _schedule_agent_build(sid) _schedule_session_cap_enforcement() # trim detached idle sessions over the cap - auto_continue = _maybe_schedule_auto_continue(sid, record, ctx.target) - return _resume_response( - ctx, sid, record, info=ctx.info(cwd, overrides), display=display_history, - count_source=raw_history, auto_continue=auto_continue) + return _resume_response(ctx, sid, record, info=ctx.info(cwd, overrides), display=display_history, + count_source=raw_history, + auto_continue=_maybe_schedule_auto_continue(sid, record, ctx.target)) def _resume_eager(ctx: _Resume) -> dict: - """Synchronous build (``eager_build: true``). Built OUTSIDE _session_resume_lock (would - stall session.close), then double-checked: a concurrent winner's agent is reused.""" - sid, source = _new_runtime_ids(ctx.params) - _enable_gateway_prompts() + """Synchronous build OUTSIDE _session_resume_lock (it would stall session.close), then double-checked.""" + sid, source, _cwd = ctx.mint() with _profile_build_scope(ctx.profile_home): try: - raw_history, display_history = _resume_read_history(ctx) - display_history_prefix = [] if ctx.omit_messages else ctx.db.get_ancestor_display_prefix(ctx.target) - history = sanitize_replay_history(raw_history) - messages = [] if ctx.omit_messages else _history_to_messages(display_history) - # Profile db so turns persist to the right state.db; runtime identity from the stored row so - # switching chats does not inherit another chat's global model. + history, display_history, raw_history = ctx.restore() + display_history_prefix = ctx.display_prefix() + # Profile db so turns persist to the right state.db; stored runtime identity so switching chats does + # not inherit another chat's global model. stored_runtime_overrides = _stored_session_runtime_overrides(ctx.found) agent = _make_agent_in_context( sid, ctx.target, session_db=ctx.db, platform_override=source, - context_cwd_is_launch_artifact=( - source in _LAUNCH_CWD_NOT_A_WORKSPACE and not ctx.profile_resume_cwd), + context_cwd_is_launch_artifact=(source in _LAUNCH_CWD_NOT_A_WORKSPACE and not ctx.profile_resume_cwd), **stored_runtime_overrides) except Exception as e: - return ctx.resume_failed(e) + return _err(ctx.rid, 5000, f"resume failed: {e}") with _session_resume_lock: live = _find_live_session_by_key(ctx.target, ctx.profile_home) if live is not None: @@ -822,62 +717,43 @@ def _resume_eager(ctx: _Resume) -> dict: return _resume_reuse_live(ctx, *live) try: with _profile_build_scope(ctx.profile_home): - _init_session( - sid, ctx.target, agent, history, cols=ctx.cols, cwd=ctx.profile_resume_cwd, - session_db=ctx.db, source=source, explicit_cwd=bool(ctx.profile_resume_cwd)) - # Ownership TRANSFER: the agent holds the handle for life (AIAgent.close() releases - # it). The owns_db drop is UNCONDITIONAL — the session is registered against the - # handle, so the finally must not close it even if the transfer was refused (a leak - # beats "closed database" every turn). Gated on owns_db: the SHARED launch handle - # must never move onto one session. + _init_session(sid, ctx.target, agent, history, cols=ctx.cols, cwd=ctx.profile_resume_cwd, + session_db=ctx.db, source=source, explicit_cwd=bool(ctx.profile_resume_cwd)) + # Ownership TRANSFER: the agent holds the handle for life (AIAgent.close() releases it). The + # owns_db drop is UNCONDITIONAL — the session is registered against the handle, so the finally + # must not close it even if the transfer was refused (a leak beats "closed database" every + # turn). Gated on owns_db: the SHARED launch handle must never move onto one session. if ctx.owns_db: _transfer_db_to_agent(agent, ctx.db) ctx.owns_db = False if (session := _sessions.get(sid)) is not None: if stored_runtime_overrides.get("model_override") is not None: session["model_override"] = stored_runtime_overrides["model_override"] - session["display_history_prefix"] = display_history_prefix - # Each turn re-binds HERMES_HOME (mid-turn memory/skills reads). + # Each turn re-binds HERMES_HOME (mid-turn memory/skills reads); lease claimed lazily on turn 1. if ctx.profile_home is not None: session["profile_home"] = str(ctx.profile_home) - session["active_session_lease"] = None # claimed lazily on the first turn + session.update(display_history_prefix=display_history_prefix, active_session_lease=None) except Exception as e: - # _init_session registers _sessions[sid] BEFORE its first db read; left in place the - # fast path would serve that dead session forever. + # _init_session registers _sessions[sid] BEFORE its first db read; left in place the fast path + # would serve that dead session forever. if ctx.owns_db: with _sessions_lock: _sessions.pop(sid, None) - return ctx.resume_failed(e) + return _err(ctx.rid, 5000, f"resume failed: {e}") session = _sessions.get(sid) or {} - auto_continue = _maybe_schedule_auto_continue(sid, session, ctx.target) if session else None return _resume_response( - ctx, sid, session, info=_session_info(agent, session), messages=messages, count_source=raw_history, - started_at=float(session.get("created_at") or time.time()), auto_continue=auto_continue) + ctx, sid, session, info=_session_info(agent, session), display=display_history, count_source=raw_history, + started_at=float(session.get("created_at") or time.time()), + auto_continue=_maybe_schedule_auto_continue(sid, session, ctx.target) if session else None) @method("session.resume") def _(rid, params: dict) -> dict: - target = params.get("session_id", "") - if not target: + if not (target := params.get("session_id", "")): return _err(rid, 4006, "session_id required") - # ``profile`` (app-global remote mode): resume from another local profile's state.db. - profile = (params.get("profile") or "").strip() or None - - def flag(name: str) -> bool: - return is_truthy_value(params.get(name, False)) - ctx = _Resume( - rid=rid, params=params, target=target, cols=_int_param(params, "cols", 80), profile=profile, - profile_home=_profile_home(profile), - lazy=flag("lazy"), defer_history=flag("defer_history"), - # Desktop hydrates over REST; suppress the duplicate WS copy only when asked. - omit_messages=flag("omit_messages"), eager_build=flag("eager_build")) + ctx = _Resume(rid, params, target) # Profile scope: a DEDICATED handle we own until the agent takes it; else the shared launch db. - if ctx.profile_home is not None: - from hermes_state import get_shared_session_db - ctx.db = get_shared_session_db(ctx.profile_home / "state.db") - ctx.owns_db = True - else: - ctx.db = _get_db() + ctx.db, ctx.owns_db = _profile_session_db(ctx.profile_home) try: if ctx.db is None: return _db_unavailable_error(rid, code=5000) @@ -886,8 +762,7 @@ def _(rid, params: dict) -> dict: _resume_follow_tip(ctx) if (resp := _resume_guard(ctx)) is not None: return resp - ctx.profile_resume_cwd = ( - str(ctx.found.get("cwd") or "").strip() or _profile_configured_cwd(ctx.profile_home)) + ctx.profile_resume_cwd = _str_param(ctx.found, "cwd") or _profile_configured_cwd(ctx.profile_home) # Fast path: reuse a session live IN THIS PROFILE (never another profile's runtime). with _session_resume_lock: live = _find_live_session_by_key(ctx.target, ctx.profile_home) @@ -895,27 +770,23 @@ def _(rid, params: dict) -> dict: return _resume_reuse_live(ctx, *live) if ctx.lazy: return _resume_lazy(ctx) - if ctx.defer_history and not ctx.eager_build: - return _resume_deferred(ctx) - return _resume_eager(ctx) if ctx.eager_build else _resume_cold(ctx) + if ctx.eager_build: + return _resume_eager(ctx) + return _resume_deferred(ctx) if ctx.defer_history else _resume_cold(ctx) finally: - # Refcounting alone does not release the sqlite fds: SessionDB pins ITSELF once its background - # token writer starts (atexit.register); only close() unregisters. + # Refcounting alone does not release the sqlite fds: SessionDB pins ITSELF (atexit.register) once its + # background token writer starts; only close() unregisters. if ctx.owns_db and ctx.db is not None: with contextlib.suppress(Exception): ctx.db.close() # ── cwd / workspace / live-session bookkeeping ─────────────────────── - - -@method("session.cwd.set") -@_with_session +@_session_method("session.cwd.set") def _(rid, params: dict, session: dict) -> dict: if session.get("running"): return _err(rid, 4009, "session busy") - raw = str(params.get("cwd", "") or "").strip() - if not raw: + if not (raw := _str_param(params, "cwd")): return _err(rid, 4016, "cwd required") try: cwd = _set_session_cwd(session, raw) @@ -928,15 +799,12 @@ def _(rid, params: dict, session: dict) -> dict: @method("session.workspace.move") def _(rid, params: dict) -> dict: - """Re-home a STORED session's workspace (by ``session_key``); no live agent required. git branch/root - columns are REPLACED (a stale ``git_repo_root`` would keep the session under the project it left). A - live agent follows too, even mid-turn (refusing made the UI claim success while state.db kept the - old cwd); the NEXT tool call moves.""" - target = str(params.get("session_key") or "").strip() - if not target: + """Re-home a STORED session's workspace (by ``session_key``; no live agent required). git branch/root are + REPLACED (a stale ``git_repo_root`` kept the session under the project it left); a live agent follows even + mid-turn (refusing made the UI claim success while state.db kept the old cwd).""" + if not (target := _str_param(params, "session_key")): return _err(rid, 4007, "session_key required") - raw = str(params.get("cwd", "") or "").strip() - if not raw: + if not (raw := _str_param(params, "cwd")): return _err(rid, 4016, "cwd required") from hermes_constants import translate_cwd_for_wsl_backend resolved = os.path.abspath(os.path.expanduser(translate_cwd_for_wsl_backend(raw))) @@ -945,18 +813,16 @@ def _(rid, params: dict) -> dict: # Snapshot under the lock — concurrent RPCs mutate _sessions. with _sessions_lock: live_sid, live = next( - ((sid, sess) for sid, sess in list(_sessions.items()) if sess.get("session_key") == target), - ("", None)) - branch = _git_branch_for_cwd(resolved) - root = _git_common_repo_root_for_cwd(resolved) + ((sid, sess) for sid, sess in list(_sessions.items()) if sess.get("session_key") == target), ("", None)) + branch, root = _git_branch_for_cwd(resolved), _git_common_repo_root_for_cwd(resolved) with _profile_db(params) as db: if db is None: return _db_unavailable_error(rid, code=5007) # A draft has no row yet; the live re-home still applies (row inherits cwd on write). - row_exists = bool(db.get_session(target)) - if not row_exists and live is None: - return _err(rid, 4007, "session not found") - if row_exists: + if not db.get_session(target): + if live is None: + return _err(rid, 4007, "session not found") + else: try: db.update_session_cwd(target, resolved, branch, root, replace_git_meta=True) except Exception as e: @@ -973,168 +839,134 @@ def _(rid, params: dict) -> dict: @method("session.active_list") def _(rid, params: dict) -> dict: """Live TUI sessions in this process (not a DB browser).""" - current = str(params.get("current_session_id") or "") snapshot, err = _snapshot_sessions(rid) if err: return err - # ``_finalized`` sessions linger until the reaper pops them (they inflated the footer). Do NOT - # filter on the WS-detached sentinel: detached is still attachable until grace-reap, and - # ``hermes --tui`` rides stdio. Keep insertion order (focused must not jump). - rows = [ - _session_live_item(sid, session, current) for sid, session in snapshot if not session.get("_finalized") - ] + current = str(params.get("current_session_id") or "") + # ``_finalized`` sessions linger until the reaper pops them (they inflated the footer). Do NOT filter on + # the WS-detached sentinel: detached is attachable until grace-reap, and ``hermes --tui`` rides stdio. + # Keep insertion order (focused must not jump). + rows = [_session_live_item(sid, session, current) for sid, session in snapshot if not session.get("_finalized")] return _ok(rid, {"sessions": rows}) -@method("session.activate") -def _(rid, params: dict) -> dict: +@_session_method("session.activate") +def _(rid, params: dict, session: dict) -> dict: """Attach the frontend to a live TUI session without closing the previously focused one.""" - sid = str(params.get("session_id") or "") - session, err = _sess_nowait({"session_id": sid}, rid) - if err: - return err return _ok(rid, _live_session_payload( - sid, session, touch=True, transport=current_transport() or _stdio_transport, + str(params.get("session_id") or ""), session, touch=True, transport=current_transport() or _stdio_transport, omit_messages=is_truthy_value(params.get("omit_messages", False)))) @method("session.delete") def _(rid, params: dict) -> dict: - """Delete a stored session + transcript files (honors ``params.profile``). Refuses sessions - live in this process — deleting under a live agent trips FK constraints on the next flush.""" - target = params.get("session_id", "") - if not target: + """Delete a stored session + transcripts; refused while live here (FK trips on the agent's next flush).""" + if not (target := params.get("session_id", "")): return _err(rid, 4006, "session_id required") snapshot, err = _snapshot_sessions(rid) if err: return err - active = {s.get("session_key") for _sid, s in snapshot if s.get("session_key")} - if target in active: + if any(s.get("session_key") == target for _sid, s in snapshot): return _err(rid, 4023, "cannot delete an active session") profile_home = _profile_home((params.get("profile") or "").strip() or None) with _profile_db(params) as db: if db is None: return _db_unavailable_error(rid, code=5036) - sessions_dir = (Path(profile_home) if profile_home is not None else get_hermes_home()) / "sessions" try: - deleted = db.delete_session(target, sessions_dir=sessions_dir) + home = Path(profile_home) if profile_home is not None else get_hermes_home() + deleted = db.delete_session(target, sessions_dir=home / "sessions") except Exception as e: return _err(rid, 5036, f"delete failed: {e}") - if not deleted: - return _err(rid, 4007, "session not found") - return _ok(rid, {"deleted": target}) + return _ok(rid, {"deleted": target}) if deleted else _err(rid, 4007, "session not found") -def _title_read(rid, params: dict, session: dict, db) -> dict: +def _title_read(session: dict, db, key: str) -> str: """``session.title`` without ``title``: read it, applying a queued pending_title if possible.""" - key = session["session_key"] fallback = session.get("pending_title") or "" try: resolved_title = db.get_session_title(key) or "" if not fallback: if resolved_title: session["pending_title"] = None - elif db.set_session_title(key, fallback): + elif (db.set_session_title(key, fallback) + or ((db.get_session(key) or {}).get("title") or "").strip() == fallback): session["pending_title"] = None resolved_title = fallback - else: - existing_title = ((db.get_session(key) or {}).get("title") or "").strip() - if existing_title == fallback: - session["pending_title"] = None - resolved_title = fallback - elif not resolved_title: - resolved_title = fallback + elif not resolved_title: + resolved_title = fallback except Exception: resolved_title = fallback - _emit_session_info_for_session(params.get("session_id", ""), session) - return _ok(rid, {"title": resolved_title, "session_key": key}) + return resolved_title @method("session.title") -@_with_session_db(5007) +@_with_db(5007, session_scoped=True) def _(rid, params: dict, session: dict, db) -> dict: - if "title" not in params: - return _title_read(rid, params, session, db) key = session["session_key"] - title = (params.get("title", "") or "").strip() - if not title: + if "title" not in params: + result = {"title": _title_read(session, db, key), "session_key": key} + elif not (title := (params.get("title", "") or "").strip()): return _err(rid, 4021, "title required") - - def _done(pending: bool, value: str) -> dict: + else: + try: + if db.set_session_title(key, title): + pending, value = False, title + # rowcount == 0 can mean "same value" as well as "missing row". + elif existing_row := db.get_session(key): + pending, value = False, existing_row.get("title") or title + else: + # No row yet: an explicit /title is clear intent, so persist the row NOW (as the gateway's + # _handle_title_command); the min-messages sidebar filter hides a titled 0-message row. If + # row creation didn't take, queue so the post-turn apply block can recover. + _ensure_session_db_row(session) + with _session_db(session) as scoped_db: + pending, value = not (scoped_db is not None and scoped_db.set_session_title(key, title)), title + except ValueError as e: + return _err(rid, 4022, str(e)) + except Exception as e: + return _err(rid, 5007, str(e)) session["pending_title"] = value if pending else None - _emit_session_info_for_session(params.get("session_id", ""), session) - return _ok(rid, {"pending": pending, "title": value}) - try: - if db.set_session_title(key, title): - return _done(False, title) - # rowcount == 0 can mean "same value" as well as "missing row". - if existing_row := db.get_session(key): - return _done(False, existing_row.get("title") or title) - # No row yet: an explicit /title is clear intent, so persist the row NOW (as the gateway's - # _handle_title_command). The min-messages sidebar filter hides a titled 0-message row. - _ensure_session_db_row(session) - with _session_db(session) as scoped_db: - if scoped_db is not None and scoped_db.set_session_title(key, title): - return _done(False, title) - # Row creation didn't take — queue so the post-turn apply block can recover. - return _done(True, title) - except ValueError as e: - return _err(rid, 4022, str(e)) - except Exception as e: - return _err(rid, 5007, str(e)) + result = {"pending": pending, "title": value} + _emit_session_info_for_session(params.get("session_id", ""), session) + return _ok(rid, result) @method("session.set_hidden") def _(rid, params: dict) -> dict: - """Set/clear ``hidden`` on a session (and its compression lineage); hidden sessions leave - the default list but stay resumable by their owner. Resolution: LIVE runtime id first - (covers unpersisted drafts via ``pending_hidden``), then a stored id/key in the profile db.""" + """Set/clear ``hidden`` (leaves the default list, stays resumable by its owner) on a session + lineage: + LIVE runtime id first (unpersisted drafts via ``pending_hidden``), then a stored id/key in the profile db.""" hidden = is_truthy_value(params.get("hidden", True)) session, err = _sess_nowait(params, rid) - if session is not None: - with _session_db(session) as db: - if db is None: - return _db_unavailable_error(rid, code=5007) - key = session["session_key"] - try: - if not db.set_session_hidden(key, hidden): - # No row yet: _ensure_session_db_row is born hidden (as pending_title). - session["pending_hidden"] = hidden - return _ok(rid, {"hidden": hidden, "session_key": key}) - except Exception as e: - return _err(rid, 5007, str(e)) - # ``resolve_session_id`` follows key/title aliases like the REST pin/archive path. - target = str(params.get("session_id") or "").strip() - with _profile_db(params) as db: + with (_profile_db(params) if session is None else _session_db(session)) as db: if db is None: return _db_unavailable_error(rid, code=5007) try: - resolved = db.resolve_session_id(target) if hasattr(db, "resolve_session_id") else target - if not resolved: - return err - db.set_session_hidden(resolved, hidden) - return _ok(rid, {"hidden": hidden, "session_key": resolved}) + if session is not None: + key = session["session_key"] + if not db.set_session_hidden(key, hidden): + session["pending_hidden"] = hidden # no row yet: _ensure_session_db_row is born hidden + else: + # ``resolve_session_id`` follows key/title aliases like the REST pin/archive path. + target = _str_param(params, "session_id") + if not (key := db.resolve_session_id(target) if hasattr(db, "resolve_session_id") else target): + return err + db.set_session_hidden(key, hidden) + return _ok(rid, {"hidden": hidden, "session_key": key}) except Exception as e: return _err(rid, 5007, str(e)) -@method("message.react") -@_with_session +@_session_method("message.react") def _(rid, params: dict, session: dict) -> dict: - """Set/clear one author's emoji reaction (Tapback semantics in the DB layer: one per - author, same emoji retracts, ``emoji: null`` clears). ``row_id`` is ``messages.id``; a - live message not yet round-tripped can name ``newest_role`` instead.""" - newest_role = str(params.get("newest_role") or "").strip() + """Set/clear one author's emoji reaction (Tapback semantics: one per author, same emoji retracts, null + clears). ``row_id`` is ``messages.id``; a not-yet-persisted live message names ``newest_role`` instead.""" + newest_role = _str_param(params, "newest_role") row_id = params.get("row_id") if row_id is None and newest_role not in {"user", "assistant"}: return _err(rid, 4023, "row_id or newest_role required") - emoji = params.get("emoji") - if emoji is not None: - emoji = str(emoji).strip() - if not emoji: - return _err(rid, 4024, "emoji must be a non-empty string or null") - author = str(params.get("author") or "user").strip() - if author not in {"user", "agent"}: + if (emoji := params.get("emoji")) is not None and not (emoji := str(emoji).strip()): + return _err(rid, 4024, "emoji must be a non-empty string or null") + if (author := str(params.get("author") or "user").strip()) not in {"user", "agent"}: return _err(rid, 4025, "author must be 'user' or 'agent'") with _session_db(session) as db: if db is None: @@ -1154,57 +986,42 @@ def _(rid, params: dict, session: dict) -> dict: @method("llm.oneshot") def _(rid, params: dict) -> dict: - """Stateless one-shot LLM request (``template``+``variables`` or ``instructions``/``input``). - A live ``session_id`` lends its model, else the auxiliary ``task`` backend. Never touches history.""" + """Stateless one-shot LLM request; a live ``session_id`` lends its model, else the ``task`` backend.""" template = (params.get("template") or "").strip() or None instructions = params.get("instructions") or "" user_input = params.get("input") or "" variables = params.get("variables") if isinstance(params.get("variables"), dict) else {} - task = (params.get("task") or "title_generation").strip() or "title_generation" - max_tokens = _int_param(params, "max_tokens", 1024) or 1024 try: - temperature = float(params["temperature"]) if params.get("temperature") is not None else None + temperature = float(params["temperature"]) if params.get("temperature") is not None else 0.3 except (TypeError, ValueError): - temperature = None + temperature = 0.3 if not template and not str(instructions).strip() and not str(user_input).strip(): return _err(rid, 4030, "llm.oneshot requires a template or instructions/input") session = _sessions.get(params.get("session_id") or "") - main_runtime = _main_runtime_from_agent(session.get("agent")) if session else None try: from agent.oneshot import run_oneshot - text = run_oneshot( + return _ok(rid, {"text": run_oneshot( instructions=instructions, user_input=user_input, template=template, variables=variables, - task=task, max_tokens=max_tokens, temperature=temperature if temperature is not None else 0.3, - main_runtime=main_runtime) - except KeyError as e: - return _err(rid, 4031, str(e)) - except ValueError as e: - return _err(rid, 4032, str(e)) + task=(params.get("task") or "title_generation").strip() or "title_generation", + max_tokens=_int_param(params, "max_tokens", 1024) or 1024, temperature=temperature, + main_runtime=_main_runtime_from_agent(session.get("agent")) if session else None)}) + except (KeyError, ValueError) as e: + return _err(rid, 4031 if isinstance(e, KeyError) else 4032, str(e)) except Exception as e: logger.warning("llm.oneshot failed: %s", e) return _err(rid, 5030, f"one-shot generation failed: {e}") - return _ok(rid, {"text": text}) # ── handoff ────────────────────────────────────────────────────────── - - -@method("handoff.request") -@_with_session +@_session_method("handoff.request") def _(rid, params: dict, session: dict) -> dict: - """Queue a handoff to a messaging platform (desktop /handoff). Only writes - ``handoff_state='pending'``; the gateway's ``_handoff_watcher`` claims it and re-binds - the session to the home channel. The desktop polls ``handoff.state``.""" + """Queue a handoff (desktop /handoff): only writes ``pending``; the gateway watcher claims and re-binds.""" if session.get("running"): return _err(rid, 4009, "session busy — wait for the current turn to finish, then retry the handoff") - platform_name = (params.get("platform", "") or "").strip().lower() - if not platform_name: + if not (platform_name := (params.get("platform", "") or "").strip().lower()): return _err(rid, 4023, "platform required") # Validate up front: an unconfigured platform / missing home channel pends forever. - try: - from gateway.config import Platform, load_gateway_config - except Exception as e: # pragma: no cover — gateway pkg always ships - return _err(rid, 5021, f"could not load gateway config: {e}") + from gateway.config import Platform, load_gateway_config try: platform = Platform(platform_name) except (ValueError, KeyError): @@ -1214,34 +1031,29 @@ def _(rid, params: dict, session: dict) -> dict: gw_config = load_gateway_config() except Exception as e: return _err(rid, 5021, f"could not load gateway config: {e}") - pcfg = gw_config.platforms.get(platform) - if not pcfg or not pcfg.enabled: + if not getattr(gw_config.platforms.get(platform), "enabled", False): return _err(rid, 4025, f"platform '{platform_name}' is not configured/enabled in the gateway") - home = gw_config.get_home_channel(platform) - if not home or not home.chat_id: - return _err( - rid, 4026, - f"no home channel configured for {platform_name} — set one with " - "/sethome on the destination chat first") + if not (home := gw_config.get_home_channel(platform)) or not home.chat_id: + return _err(rid, 4026, f"no home channel configured for {platform_name} — set one with " + "/sethome on the destination chat first") # The watcher transfers a persisted row, so make sure one exists for an empty chat. _ensure_session_db_row(session) + key = session["session_key"] with _session_db(session) as db: if db is None: return _db_unavailable_error(rid, code=5007) - key = session["session_key"] try: if not db.get_session(key): db.set_session_title(key, f"handoff-{key[:8]}") - ok = db.request_handoff(key, platform_name) + if not db.request_handoff(key, platform_name): + return _err(rid, 4027, "session is already in flight for handoff — wait for it to settle, then retry") except Exception as e: return _err(rid, 5007, str(e)) - if not ok: - return _err(rid, 4027, "session is already in flight for handoff — wait for it to settle, then retry") return _ok(rid, {"queued": True, "session_key": key, "platform": platform_name, "home_name": home.name}) @method("handoff.state") -@_with_session_db(5007) +@_with_db(5007, session_scoped=True) def _(rid, params: dict, session: dict, db) -> dict: """Poll ``{state, platform, error}``; ``state`` is pending|running|completed|failed or empty.""" record = db.get_handoff_state(session["session_key"]) or {} @@ -1250,9 +1062,7 @@ def _(rid, params: dict, session: dict, db) -> dict: @method("handoff.fail") def _(rid, params: dict) -> dict: - """Mark a not-yet-claimed handoff failed (desktop poll timeout). Only PENDING rows change - (CAS in ``fail_handoff``): a claimed ``running`` row is the watcher's to finish and yields - ``{"failed": False, "state": "running"}``.""" + """Fail a not-yet-claimed handoff (poll timeout); a claimed ``running`` row is the watcher's (CAS).""" # Undecorated on purpose: tests rebind this handler's __code__ directly. session, err = _sess_nowait(params, rid) if err: @@ -1266,25 +1076,17 @@ def _(rid, params: dict) -> dict: failed = db.fail_handoff(key, reason, only_states=("pending",)) except TypeError: # Older SessionDB without only_states: fail only when still pending. - record = db.get_handoff_state(key) or {} - failed = (record.get("state") or "") == "pending" - if failed: + if failed := ((db.get_handoff_state(key) or {}).get("state") or "") == "pending": db.fail_handoff(key, reason) - if failed: - return _ok(rid, {"failed": True, "state": "failed"}) - record = db.get_handoff_state(key) or {} - return _ok(rid, {"failed": False, "state": record.get("state") or ""}) + state = "failed" if failed else (db.get_handoff_state(key) or {}).get("state") or "" + return _ok(rid, {"failed": bool(failed), "state": state}) # ── usage ──────────────────────────────────────────────────────────── - - -@method("session.usage") -@_with_session +@_session_method("session.usage") def _(rid, params: dict, session: dict) -> dict: - agent = session.get("agent") usage: dict = _session_usage_snapshot(session) - if agent is None and not usage: + if session.get("agent") is None and not usage: usage = {"calls": 0, "input": 0, "output": 0, "total": 0} # Nous credits are agent-independent (portal fetch); fail-open when absent. with contextlib.suppress(Exception): @@ -1294,11 +1096,9 @@ def _(rid, params: dict, session: dict) -> dict: return _ok(rid, usage) -@method("session.context_breakdown") -@_with_session +@_session_method("session.context_breakdown") def _(rid, params: dict, session: dict) -> dict: - agent = session.get("agent") - if agent is None: + if (agent := session.get("agent")) is None: usage = _session_usage_snapshot(session) or _get_usage(None) return _ok(rid, { "categories": [], "context_max": usage.get("context_max", 0) or 0, @@ -1310,28 +1110,24 @@ def _(rid, params: dict, session: dict) -> dict: history = list(session.get("history", [])) try: from agent.context_breakdown import compute_session_context_breakdown - payload = compute_session_context_breakdown(agent, history) + return _ok(rid, compute_session_context_breakdown(agent, history)) except Exception as exc: return _err(rid, 5000, f"Could not compute context breakdown: {exc}") - return _ok(rid, payload) # ── pet ────────────────────────────────────────────────────────────── - _PET_OFF = {"enabled": False} @_pet_method("pet.info", fail_open=_PET_OFF) def _(rid, params: dict) -> dict: - """Active pet for sprite-rendering surfaces: spritesheet (base64) + frame geometry + - state-row taxonomy so the renderer is a thin consumer.""" + """Active pet for sprite renderers: spritesheet (base64) + frame geometry + state-row taxonomy.""" if (active := _active_pet()) is None: return _ok(rid, {"enabled": False}) pet, scale = active payload = {"enabled": True, **_pet_sprite_payload(pet, scale=scale)} # Send-once for the multi-MB sheet: same revision → metadata only. - known_revision = str(params.get("knownRevision", "") or "") - if known_revision and known_revision == payload.get("spritesheetRevision"): + if (known := str(params.get("knownRevision", "") or "")) and known == payload.get("spritesheetRevision"): payload.pop("spritesheetBase64", None) payload["spritesheetUnchanged"] = True return _ok(rid, payload) @@ -1343,42 +1139,37 @@ def _(rid, params: dict) -> dict: if (active := _active_pet()) is None: return _ok(rid, {"enabled": False}) pet, scale = active - return _ok(rid, { - "enabled": True, "slug": pet.slug, "displayName": pet.display_name, "scale": scale, - "spritesheetRevision": _pet_sheet_revision(pet.spritesheet)}) + return _ok(rid, {"enabled": True, "slug": pet.slug, "displayName": pet.display_name, "scale": scale, + "spritesheetRevision": _pet_sheet_revision(pet.spritesheet)}) def _pet_kitty_cells(pet, pet_cfg: dict, state: str, scale: float) -> dict | None: - """kitty graphics payload for a TTY that speaks it (env shared with the Ink process; the - dashboard PTY falls through). Only kitty is grid-safe in Ink — iTerm/sixel stay on half-blocks.""" + """kitty payload for a TTY that speaks it (dashboard PTY falls through); only kitty is grid-safe in Ink.""" from agent.pet import constants, render from agent.pet.render import PetRenderer configured = str(pet_cfg.get("render_mode", "auto") or "auto").lower() - gmode = render.detect_terminal_graphics() if configured in ("", "auto") else configured - if gmode != "kitty": + if (render.detect_terminal_graphics() if configured in ("", "auto") else configured) != "kitty": return None image_id = render.kitty_image_id(pet.slug) # kitty sizes from scaled pixels, so unicode_cols is moot here. payload = PetRenderer(str(pet.spritesheet), mode="kitty", scale=scale).kitty_payload(state, image_id=image_id) if not payload: return None - return { - "graphics": "kitty", "imageId": image_id, "color": render.kitty_color_hex(image_id), - "cols": payload["cols"], "rows": payload["rows"], "placeholder": payload["placeholder"], - "frames": payload["frames"], "frameMs": constants.LOOP_MS / max(1, len(payload["frames"]) or 1), - "scale": scale} + return {"graphics": "kitty", "imageId": image_id, "color": render.kitty_color_hex(image_id), + "cols": payload["cols"], "rows": payload["rows"], "placeholder": payload["placeholder"], + "frames": payload["frames"], "frameMs": constants.LOOP_MS / max(1, len(payload["frames"]) or 1), + "scale": scale} @_pet_method("pet.cells", fail_open=_PET_OFF) def _(rid, params: dict) -> dict: - """Half-block cell frames for one pet state (TUI); each cell is ``[tr,tg,tb,ta, br,bg,bb,ba]``. - Params: ``state`` (idle/run/review/failed/wave/jump), ``cols``, ``graphics``.""" + """Half-block cell frames (``[tr,tg,tb,ta, br,bg,bb,ba]``) for one pet ``state``; ``cols``, ``graphics``.""" from agent.pet import constants, store from agent.pet.render import PetRenderer pet_cfg = _pet_display_cfg() - if not is_truthy_value(pet_cfg.get("enabled"), default=False): - return _ok(rid, {"enabled": False}) - pet = store.resolve_active_pet(str(pet_cfg.get("slug", "") or "")) + pet = None + if is_truthy_value(pet_cfg.get("enabled"), default=False): + pet = store.resolve_active_pet(str(pet_cfg.get("slug", "") or "")) if pet is None or not pet.exists: return _ok(rid, {"enabled": False}) state = str(params.get("state") or constants.PetState.IDLE.value) @@ -1389,32 +1180,26 @@ def _(rid, params: dict) -> dict: return _ok(rid, {**base, **kitty}) renderer = PetRenderer(str(pet.spritesheet), mode="unicode", scale=scale, unicode_cols=cols) count = renderer.frame_count(state) or 1 - frames = [ - [[[*top, *bottom] for (top, bottom) in row] for row in renderer.cells(state, i, cols=cols)] - for i in range(count)] - return _ok( - rid, - {**base, "cols": cols, "frameMs": constants.LOOP_MS / max(1, count), "frames": frames, "scale": scale}, - ) + frames = [[[[*top, *bottom] for (top, bottom) in row] for row in renderer.cells(state, i, cols=cols)] + for i in range(count)] + return _ok(rid, {**base, "cols": cols, "frameMs": constants.LOOP_MS / max(1, count), "frames": frames, + "scale": scale}) @_pet_method("pet.gallery", fail_open={"enabled": False, "active": "", "pets": []}) def _(rid, params: dict) -> dict: - """Petdex gallery merged with local install state; falls back to installed pets offline. - ``localOnly`` skips the remote manifest so the user's own pets render instantly.""" + """Petdex gallery + local install state (installed-only offline); ``localOnly`` skips the remote manifest.""" local_only = bool(params.get("localOnly")) from agent.pet import store pet_cfg = _pet_display_cfg() installed = {p.slug: p for p in store.installed_pets()} gallery: list[dict] = [] - seen: set[str] = set() try: from agent.pet.manifest import fetch_manifest, prefetch # Local-only still warms the manifest cache in the background. if local_only: prefetch() for entry in [] if local_only else fetch_manifest(): - seen.add(entry.slug) gallery.append({ "slug": entry.slug, "displayName": entry.display_name, "installed": entry.slug in installed, "spritesheetUrl": entry.spritesheet_url, @@ -1423,13 +1208,13 @@ def _(rid, params: dict) -> dict: "generated": entry.slug in installed and installed[entry.slug].generated}) except Exception as exc: # noqa: BLE001 - offline: fall back to installed logger.debug("pet.gallery manifest fetch failed: %s", exc) + seen = {item["slug"] for item in gallery} gallery.extend( {"slug": slug, "displayName": pet.display_name, "installed": True, "spritesheetUrl": "", "generated": pet.generated} for slug, pet in installed.items() if slug not in seen) - return _ok(rid, { - "enabled": is_truthy_value(pet_cfg.get("enabled"), default=False), - "active": str(pet_cfg.get("slug", "") or ""), "pets": gallery}) + return _ok(rid, {"enabled": is_truthy_value(pet_cfg.get("enabled"), default=False), + "active": str(pet_cfg.get("slug", "") or ""), "pets": gallery}) @_pet_method("pet.select", slug=True) @@ -1452,13 +1237,18 @@ def _(rid, params: dict, slug: str) -> dict: from agent.pet import store from hermes_cli.pets import _clear_active_if removed = store.remove_pet(slug) - try: - _clear_active_if(slug) - except Exception as exc: # noqa: BLE001 - removal already succeeded - logger.debug("pet.remove config update failed: %s", exc) + _pet_config_followup("pet.remove", _clear_active_if, slug) return _ok(rid, {"ok": removed, "slug": slug}) +def _pet_config_followup(what: str, fn, *args) -> None: + """Best-effort ``hermes_cli.pets`` active-slug update after a store op that already succeeded.""" + try: + fn(*args) + except Exception as exc: # noqa: BLE001 + logger.debug("%s config update failed: %s", what, exc) + + def _b64(data: bytes) -> str: import base64 return base64.standard_b64encode(data).decode("ascii") @@ -1475,29 +1265,22 @@ def _(rid, params: dict, slug: str) -> dict: @_pet_method("pet.rename", slug=True) def _(rid, params: dict, slug: str) -> dict: """Rename a pet's display name + realign its slug/dir; follows the active slug in config.""" - name = str(params.get("name") or "").strip() - if not name: + if not (name := _str_param(params, "name")): return _err(rid, 4004, "missing name") from agent.pet import store - new_slug = store.rename_pet(slug, name) - if not new_slug: + if not (new_slug := store.rename_pet(slug, name)): return _err(rid, 5031, "pet.rename failed") if new_slug != slug: - try: - from hermes_cli.pets import _rename_active_if - _rename_active_if(slug, new_slug) - except Exception as exc: # noqa: BLE001 - rename already succeeded - logger.debug("pet.rename config update failed: %s", exc) + from hermes_cli.pets import _rename_active_if + _pet_config_followup("pet.rename", _rename_active_if, slug, new_slug) return _ok(rid, {"ok": True, "slug": new_slug, "displayName": name}) -@_pet_method("pet.thumb", slug=True, fail_open=lambda params: {"ok": False, "slug": str(params.get("slug") or "").strip()}) +@_pet_method("pet.thumb", slug=True, fail_open=lambda params: {"ok": False, "slug": _str_param(params, "slug")}) def _(rid, params: dict, slug: str) -> dict: - """Idle-frame PNG data URI for the picker (desktop CSP / R2 hotlink rules break a CDN - ````). ``url`` serves not-yet-installed pets.""" + """Idle-frame PNG data URI for the picker (desktop CSP breaks CDN ````); ``url``: not-yet-installed.""" from agent.pet import store - data = store.thumbnail_png(slug, source_url=str(params.get("url") or "")) - if not data: + if not (data := store.thumbnail_png(slug, source_url=str(params.get("url") or ""))): return _ok(rid, {"ok": False, "slug": slug}) return _ok(rid, {"ok": True, "slug": slug, "dataUri": "data:image/png;base64," + _b64(data)}) @@ -1515,17 +1298,13 @@ def _(rid, params: dict) -> dict: """Persist ``display.pet.scale`` (clamped to engine bounds) from the desktop slider.""" from hermes_cli.pets import set_pet_scale scale, err = set_pet_scale(params.get("scale")) - if err: - return _err(rid, 4004, err) - return _ok(rid, {"ok": True, "scale": scale}) + return _err(rid, 4004, err) if err else _ok(rid, {"ok": True, "scale": scale}) @method("pet.cancel") def _(rid, params: dict) -> dict: - """Stop an in-flight ``pet.generate``/``pet.hatch`` by token. Idempotent; stays off the - worker pool so it lands while a generation occupies it.""" - token = str(params.get("token") or "").strip() - if token: + """Stop an in-flight generate/hatch by token (idempotent; off the pool so it lands mid-generation).""" + if token := _str_param(params, "token"): _pet_cancel_request(token) return _ok(rid, {"ok": True}) @@ -1534,33 +1313,37 @@ def _(rid, params: dict) -> dict: def _(rid, params: dict) -> dict: """Whether pet generation is possible: a reference-capable image backend is configured.""" from agent.pet.generate.imagegen import GenerationError, list_sprite_providers, resolve_provider + available, providers = True, [] try: resolve_provider(require_references=True) - available = True except GenerationError: available = False try: providers = list_sprite_providers() except Exception as exc: # noqa: BLE001 - picker is best-effort logger.debug("pet provider list failed: %s", exc) - providers = [] return _ok(rid, {"available": available, "providers": providers}) +def _pet_pick_provider(params: dict, *, require_references: bool): + """Picker-chosen ``params.provider`` resolved up front (a bad pick fails fast, not mid-fan-out).""" + from agent.pet.generate.imagegen import resolve_provider + name = _str_param(params, "provider") + return resolve_provider(require_references=require_references, prefer=name) if name else None + + @_pet_method("pet.generate", scoped=False) def _(rid, params: dict) -> dict: - """Candidate base looks for a new pet (draft step; worker pool). Params: ``prompt`` - (required unless ``referenceImage`` data URL), ``count`` (≤4), ``style``, ``provider``. - Returns ``{ok, token, drafts:[{index, dataUri}]}``; the token keys ``pet.hatch``.""" - prompt = str(params.get("prompt") or "").strip() - ref_raw = str(params.get("referenceImage") or "").strip() + """Candidate base looks for a new pet (draft step; worker pool): ``prompt`` (or a ``referenceImage`` + data URL), ``count`` (≤4), ``style``, ``provider`` → ``{ok, token, drafts:[{index, dataUri}]}``.""" + prompt = _str_param(params, "prompt") + ref_raw = _str_param(params, "referenceImage") if not prompt and not ref_raw: return _err(rid, 4004, "missing prompt") count = max(1, min(4, _int_param(params, "count", 4) or 4)) - style = str(params.get("style") or "auto").strip() or "auto" import shutil from agent.pet.generate import generate_base_drafts - from agent.pet.generate.imagegen import GenerationError, resolve_provider + from agent.pet.generate.imagegen import GenerationError root = _pet_gen_root() _pet_gen_sweep(root) # Token up front so each draft is staged + streamed the moment it lands. @@ -1574,15 +1357,10 @@ def _(rid, params: dict) -> dict: reference_images = _pet_reference_images_from_data_url(ref_raw, stage) except ValueError as exc: return _pet_gen_abort(rid, token, 4004, str(exc)) - # Resolve a picker-chosen provider up front so a bad pick fails fast, not mid-fan-out. - provider_name = str(params.get("provider") or "").strip() - sprite = None - if provider_name: - try: - sprite = resolve_provider(require_references=bool(reference_images), prefer=provider_name) - except GenerationError as exc: - return _pet_gen_abort(rid, token, 5031, str(exc)) - concept = prompt or "a pet based on the reference image" + try: + sprite = _pet_pick_provider(params, require_references=bool(reference_images)) + except GenerationError as exc: + return _pet_gen_abort(rid, token, 5031, str(exc)) out: list[dict] = [] # Token-only init event so a Stop fired before the first draft can target this run. _pet_emit("pet.generate.progress", {"token": token, "count": count}, "pet.generate init") @@ -1596,54 +1374,40 @@ def _(rid, params: dict) -> dict: logger.debug("pet.generate draft %d failed: %s", index, exc) return out.append({"index": index, "dataUri": data_uri}) - # Stream the draft so the grid fills live. - _pet_emit( - "pet.generate.progress", {"token": token, "index": index, "dataUri": data_uri, "count": count}, - "pet.generate progress") + _pet_emit("pet.generate.progress", {"token": token, "index": index, "dataUri": data_uri, "count": count}, + "pet.generate progress") try: - generate_base_drafts( - concept, n=count, style=style, reference_images=reference_images, provider=sprite, - on_draft=_on_draft, is_cancelled=lambda: _pet_is_cancelled(token)) + generate_base_drafts(prompt or "a pet based on the reference image", n=count, + style=_str_param(params, "style", "auto"), reference_images=reference_images, + provider=sprite, on_draft=_on_draft, is_cancelled=lambda: _pet_is_cancelled(token)) except GenerationError as exc: return _pet_gen_abort(rid, token, 5031, str(exc)) cancelled = _pet_is_cancelled(token) _pet_cancel_release(token) - if cancelled: - return _err(rid, 5031, "generation cancelled") - if not out: - return _err(rid, 5031, "generation produced no usable drafts") - out.sort(key=lambda d: d["index"]) - return _ok(rid, {"ok": True, "token": token, "drafts": out}) + if cancelled or not out: + return _err(rid, 5031, "generation cancelled" if cancelled else "generation produced no usable drafts") + return _ok(rid, {"ok": True, "token": token, "drafts": sorted(out, key=lambda d: d["index"])}) @_pet_method("pet.hatch", scoped=False) def _(rid, params: dict) -> dict: - """Turn a base draft into a full pet — installed but NOT active (``pet.select`` adopts, - ``pet.remove`` discards). Params: ``token`` + ``index``, ``name`` (required), ``description``, - ``prompt``, ``style``, ``cancelToken``. Returns ``{ok, slug, displayName, warnings, pet}``.""" - token = str(params.get("token") or "").strip() + """Turn a base draft (``token`` + ``index``) into a full pet — installed but NOT active (``pet.select`` + adopts, ``pet.remove`` discards) → ``{ok, slug, displayName, warnings, pet}``.""" + token, name = _str_param(params, "token"), _str_param(params, "name") + if not token or not name: + return _err(rid, 4004, "missing token" if not token else "missing name") # Own cancel key: pet.generate may still be releasing `token`. Falls back for old clients. - cancel_token = str(params.get("cancelToken") or "").strip() or token - name = str(params.get("name") or "").strip() - if not token: - return _err(rid, 4004, "missing token") - if not name: - return _err(rid, 4004, "missing name") - index = _int_param(params, "index", 0) + cancel_token = _str_param(params, "cancelToken") or token from agent.pet import store from agent.pet.generate import hatch_pet - from agent.pet.generate.imagegen import GenerationError, resolve_provider - base = _pet_gen_root() / token / f"draft-{index}.png" + from agent.pet.generate.imagegen import GenerationError + base = _pet_gen_root() / token / f"draft-{_int_param(params, 'index', 0)}.png" if not base.is_file(): return _err(rid, 4004, "draft expired — generate again") - # Picker override (rows always need reference grounding). - provider_name = str(params.get("provider") or "").strip() - sprite = None - if provider_name: - try: - sprite = resolve_provider(require_references=True, prefer=provider_name) - except GenerationError as exc: - return _err(rid, 5031, str(exc)) + try: + sprite = _pet_pick_provider(params, require_references=True) # rows always need reference grounding + except GenerationError as exc: + return _err(rid, 5031, str(exc)) _pet_cancel_arm(cancel_token) slug = store.unique_slug(name) @@ -1657,54 +1421,41 @@ def _(rid, params: dict) -> dict: try: result = hatch_pet( base_image=base, slug=slug, display_name=name, description=str(params.get("description") or ""), - concept=str(params.get("prompt") or name), - style=str(params.get("style") or "auto").strip() or "auto", provider=sprite, + concept=str(params.get("prompt") or name), style=_str_param(params, "style", "auto"), provider=sprite, on_progress=_on_progress, is_cancelled=lambda: _pet_is_cancelled(cancel_token)) except GenerationError as exc: return _err(rid, 5031, str(exc)) finally: _pet_cancel_release(cancel_token) pet = store.load_pet(result.slug) - return _ok(rid, { - "ok": True, "slug": result.slug, "displayName": result.display_name, - "warnings": result.validation.get("warnings", []), - "pet": _pet_sprite_payload(pet, scale=_pet_config_scale()) if pet else {}}) + return _ok(rid, {"ok": True, "slug": result.slug, "displayName": result.display_name, + "warnings": result.validation.get("warnings", []), + "pet": _pet_sprite_payload(pet, scale=_pet_config_scale()) if pet else {}}) # ── billing / subscription ─────────────────────────────────────────── # All fail-open: a logged-out / unreachable portal yields an ``ok`` envelope with a typed # ``error`` (not a JSON-RPC error) so the TUI maps it to copy. ``billing:manage`` routes # return error=insufficient_scope on 403, which drives the ``billing.step_up`` device flow. +def _billing_view(name: str, module: str, builder: str, serializer: str, fallback: dict) -> None: + """Read-only view RPC (no scope required): ``serializer(module.builder())``, ``fallback`` on any error. + The view module stays a lazy import (startup budget); the serializer is a server global.""" + @method(name) + def _(rid, params: dict) -> dict: + try: + from importlib import import_module + return _ok(rid, globals()[serializer](getattr(import_module(module), builder)())) + except Exception: + return _ok(rid, dict(fallback)) -@method("billing.state") -def _(rid, params: dict) -> dict: - """GET /api/billing/state → serialized BillingState. No scope required.""" - try: - from agent.billing_view import build_billing_state - return _ok(rid, _serialize_billing_state(build_billing_state())) - except Exception: - return _ok(rid, {"ok": True, "logged_in": False, "error": "could not load billing state"}) - - -@method("usage.bars") -def _(rid, params: dict) -> dict: - """Shared dollar usage model (two-bar view) for /usage + /subscription.""" - try: - from agent.billing_usage import build_usage_model - return _ok(rid, _serialize_usage_model(build_usage_model())) - except Exception: - return _ok(rid, {"ok": True, "available": False}) - - -@method("subscription.state") -def _(rid, params: dict) -> dict: - """GET /api/billing/subscription → serialized SubscriptionState (read-only).""" - try: - from agent.subscription_view import build_subscription_state - return _ok(rid, _serialize_subscription_state(build_subscription_state())) - except Exception: - return _ok(rid, {"ok": True, "logged_in": False, "error": "could not load subscription state"}) +_billing_view("billing.state", "agent.billing_view", "build_billing_state", "_serialize_billing_state", + {"ok": True, "logged_in": False, "error": "could not load billing state"}) +_billing_view("usage.bars", "agent.billing_usage", "build_usage_model", "_serialize_usage_model", # two-bar $ view + {"ok": True, "available": False}) +_billing_view("subscription.state", "agent.subscription_view", "build_subscription_state", + "_serialize_subscription_state", + {"ok": True, "logged_in": False, "error": "could not load subscription state"}) @method("subscription.preview") @@ -1712,107 +1463,71 @@ def _(rid, params: dict) -> dict: """POST /api/billing/subscription/preview → chargeless effect quote. billing:manage.""" from agent.subscription_view import subscription_change_preview_from_payload from hermes_cli.nous_billing import post_subscription_preview - tier_id = params.get("subscription_type_id") - if not tier_id: + if not (tier_id := params.get("subscription_type_id")): return _billing_invalid(rid, "subscription_type_id is required") return _billing_call(rid, lambda: _serialize_subscription_preview( - subscription_change_preview_from_payload(post_subscription_preview(subscription_type_id=tier_id)) - )) + subscription_change_preview_from_payload(post_subscription_preview(subscription_type_id=tier_id)))) -@method("subscription.change") -def _(rid, params: dict) -> dict: - """PUT /api/billing/subscription/pending-change: schedule a downgrade / same-price - change OR a period-end cancellation (chargeless). billing:manage.""" - from hermes_cli.nous_billing import put_subscription_pending_change - cancel = bool(params.get("cancel")) - tier_id = params.get("subscription_type_id") - if not cancel and not tier_id: - return _billing_invalid(rid, "subscription_type_id or cancel is required") - return _billing_call(rid, lambda: _billing_pending_change( - put_subscription_pending_change(subscription_type_id=tier_id, cancel=cancel))) +def _billing_route(name: str, call, *, invalid=None, message: str = "", error: str = "invalid_request", + idempotent: bool = False): + """Portal write route on ``hermes_cli.nous_billing`` (lazy; tests patch its functions): ``invalid(params)`` + → ``_billing_invalid(message, error)``; ``call(nb, params, key)`` performs the request. ``idempotent`` + mints ``idempotency_key`` if absent and echoes it (also on error) so the TUI retries the SAME operation.""" + @method(name) + def _(rid, params: dict) -> dict: + import hermes_cli.nous_billing as nb + if invalid is not None and invalid(params): + return _billing_invalid(rid, message, error=error) + key = extra = None + if idempotent: + from agent.billing_view import new_idempotency_key + key = params.get("idempotency_key") or new_idempotency_key() + extra = {"idempotency_key": key} + return _billing_call(rid, lambda: call(nb, params, key) | (extra or {}), extra=extra) -@method("subscription.resume") -def _(rid, params: dict) -> dict: - """DELETE /api/billing/subscription/pending-change: clear a scheduled downgrade / - cancellation. Re-enables recurring spend → billing:manage + kill-switch.""" - from hermes_cli.nous_billing import delete_subscription_pending_change - return _billing_call(rid, lambda: _billing_pending_change(delete_subscription_pending_change())) +# PUT pending-change: schedule a downgrade / same-price change OR a period-end cancellation. +_billing_route("subscription.change", lambda nb, p, _k: _billing_pending_change(nb.put_subscription_pending_change( + subscription_type_id=p.get("subscription_type_id"), cancel=bool(p.get("cancel")))), + invalid=lambda p: not p.get("cancel") and not p.get("subscription_type_id"), + message="subscription_type_id or cancel is required") +# DELETE pending-change: clear a scheduled downgrade / cancellation (re-enables recurring spend). +_billing_route("subscription.resume", + lambda nb, p, _k: _billing_pending_change(nb.delete_subscription_pending_change())) +# The money route (prorate + charge + flip plan). SCA / decline → status requires_action / payment_failed + +# recovery_url. +_billing_route("subscription.upgrade", lambda nb, p, key: _billing_pick( + nb.post_subscription_upgrade(subscription_type_id=p.get("subscription_type_id"), idempotency_key=key), + status="status", target_tier_name="targetTierName", recovery_url="recoveryUrl", reason="reason"), + invalid=lambda p: not p.get("subscription_type_id"), message="subscription_type_id is required", idempotent=True) +# POST /api/billing/charge → {ok, charge_id, idempotency_key}. +_billing_route("billing.charge", lambda nb, p, key: _billing_pick( + nb.post_charge(amount_usd=p.get("amount_usd"), idempotency_key=key), charge_id="chargeId"), + invalid=lambda p: p.get("amount_usd") is None, message="amount_usd is required", idempotent=True) +# GET /api/billing/charge/{id} — a single status read; the caller drives the poll cadence. +_billing_route("billing.charge_status", lambda nb, p, _k: _billing_pick( + nb.get_charge_status(p.get("charge_id")), status="status", amount_usd="amountUsd", settled_at="settledAt", + reason="reason"), invalid=lambda p: not p.get("charge_id"), message="charge_id is required", + error="invalid_charge_id") -@method("subscription.upgrade") -def _(rid, params: dict) -> dict: - """POST /api/billing/subscription/upgrade — the money route (prorate + charge + flip plan). - SCA / decline → status requires_action / payment_failed + recovery_url. Idempotency key - minted if absent and echoed (also on error) for retry of the SAME upgrade. billing:manage.""" - from agent.billing_view import new_idempotency_key - from hermes_cli.nous_billing import post_subscription_upgrade - tier_id = params.get("subscription_type_id") - if not tier_id: - return _billing_invalid(rid, "subscription_type_id is required") - key = params.get("idempotency_key") or new_idempotency_key() - - def call(): - result = post_subscription_upgrade(subscription_type_id=tier_id, idempotency_key=key) - return _billing_pick( - result, status="status", target_tier_name="targetTierName", recovery_url="recoveryUrl", - reason="reason", - ) | {"idempotency_key": key} - return _billing_call(rid, call, extra={"idempotency_key": key}) - - -@method("billing.charge") -def _(rid, params: dict) -> dict: - """POST /api/billing/charge → {ok, charge_id, idempotency_key}; key minted if absent - and echoed (also on error) so the TUI reuses it on retry of the SAME purchase.""" - from hermes_cli.nous_billing import post_charge - from agent.billing_view import new_idempotency_key - amount = params.get("amount_usd") - if amount is None: - return _billing_invalid(rid, "amount_usd is required") - key = params.get("idempotency_key") or new_idempotency_key() - return _billing_call( - rid, - lambda: _billing_pick(post_charge(amount_usd=amount, idempotency_key=key), charge_id="chargeId") - | {"idempotency_key": key}, - extra={"idempotency_key": key}) - - -@method("billing.charge_status") -def _(rid, params: dict) -> dict: - """GET /api/billing/charge/{id} — a single status read; the caller drives the poll cadence.""" - from hermes_cli.nous_billing import get_charge_status - charge_id = params.get("charge_id") - if not charge_id: - return _billing_invalid(rid, "charge_id is required", error="invalid_charge_id") - return _billing_call(rid, lambda: _billing_pick( - get_charge_status(charge_id), status="status", amount_usd="amountUsd", settled_at="settledAt", - reason="reason")) - - -@method("billing.auto_reload") -def _(rid, params: dict) -> dict: +def _auto_reload(nb, p: dict, _key) -> dict: """PATCH /api/billing/auto-top-up. params: {enabled, threshold, top_up_amount}.""" - from hermes_cli.nous_billing import patch_auto_top_up - enabled = bool(params.get("enabled")) - threshold = params.get("threshold") - top_up_amount = params.get("top_up_amount") - if threshold is None or top_up_amount is None: - return _billing_invalid(rid, "threshold and top_up_amount are required") + nb.patch_auto_top_up(enabled=bool(p.get("enabled")), threshold=p.get("threshold"), + top_up_amount=p.get("top_up_amount")) + return {"ok": True} - def call(): - patch_auto_top_up(enabled=enabled, threshold=threshold, top_up_amount=top_up_amount) - return {"ok": True} - return _billing_call(rid, call) + +_billing_route("billing.auto_reload", _auto_reload, message="threshold and top_up_amount are required", + invalid=lambda p: p.get("threshold") is None or p.get("top_up_amount") is None) @method("billing.step_up") def _(rid, params: dict) -> dict: - """billing:manage step-up device flow → {ok, granted} (false when the server downscopes). - Runs on the pool (_LONG_HANDLERS; blocks for minutes). URL/code reach the TUI via the - ``billing.step_up.verification`` event (stdout is the RPC pipe) and the browser opens - TUI-side, never via the gateway's headless webbrowser.open.""" + """billing:manage step-up device flow → {ok, granted} (false when the server downscopes). Pooled (blocks + for minutes); URL/code reach the TUI via ``billing.step_up.verification`` (stdout is the RPC pipe) and the + browser opens TUI-side, never via the gateway's headless webbrowser.open.""" sid = params.get("session_id") or "" def call(): @@ -1826,8 +1541,6 @@ def _(rid, params: dict) -> dict: # ── session status / history / undo / compress / save / close ──────── - - def _status_row(session: dict, params: dict, key: str) -> dict: """Stored row for ``key``: the live session's bound profile db first, else params.profile / launch.""" if not key: @@ -1852,18 +1565,15 @@ def _status_dt(value, fallback=None): return fallback or datetime.now() -@method("session.status") -@_with_session +@_session_method("session.status") def _(rid, params: dict, session: dict) -> dict: from hermes_constants import display_hermes_home key = session.get("session_key") or params.get("session_id") or "" agent = session.get("agent") meta = _status_row(session, params, key) created = _status_dt(meta.get("started_at")) - updated = next( - (_status_dt(meta[f], created) for f in ("updated_at", "last_updated_at", "last_activity_at") - if meta.get(f)), - created) + updated = next((_status_dt(meta[f], created) for f in ("updated_at", "last_updated_at", "last_activity_at") + if meta.get(f)), created) mirror = _metadata_mirror(session) provider = getattr(agent, "provider", None) or mirror.get("provider") or "unknown" model = getattr(agent, "model", None) or mirror.get("model") or "(unknown)" @@ -1879,23 +1589,21 @@ def _(rid, params: dict, session: dict) -> dict: return _ok(rid, {"output": "\n".join(lines)}) -@method("session.history") -@_with_session +@_session_method("session.history") def _(rid, params: dict, session: dict) -> dict: history = list(session.get("history", [])) if session.get("session_key"): with _session_db(session) as db: if db is not None: - # include_row_ids: the durable row id is how clients address a persisted - # turn (reactions, truncation targets); _history_to_messages forwards it. + # include_row_ids: the durable row id is how clients address a persisted turn (reactions, + # truncation targets); _history_to_messages forwards it. with contextlib.suppress(Exception): history = db.get_messages_as_conversation( session["session_key"], include_ancestors=True, include_row_ids=True) return _ok(rid, {"count": len(history), "messages": _history_to_messages(history)}) -@method("session.undo") -@_with_live_session +@_session_method("session.undo", live=True) def _(rid, params: dict, session: dict) -> dict: # Under a running turn the post-run write would clobber the undo — /interrupt first. busy = _err(rid, 4009, "session busy — /interrupt the current turn before /undo") @@ -1908,12 +1616,9 @@ def _(rid, params: dict, session: dict) -> dict: history = _history_without_ephemeral_scaffolding(session.get("history", [])) # Truncate from the last *real* user turn (not a timeline marker / compaction handoff). from agent.context_compressor import user_originated_turn_view - user_indices = [ - index for index, message in enumerate(history) if user_originated_turn_view(message) is not None - ] - if user_indices: + if user_turns := sum(1 for message in history if user_originated_turn_view(message) is not None): try: - removed = _rewind_active_session_history(session, len(user_indices) - 1)[2] + removed = _rewind_active_session_history(session, user_turns - 1)[2] except Exception as exc: return _err(rid, 5008, f"undo: {exc}") return _ok(rid, {"removed": removed}) @@ -1929,14 +1634,12 @@ def _compute_host_ack_error(rid, ack: dict, code: int, default: str): def _save_via_compute_host(rid, params: dict) -> dict: """``session.save`` for a turn-isolated session: the host owns the transcript file.""" try: - ack = _send_compute_host_control( - str(params.get("session_id") or ""), route_name="session.save", wait=True) + ack = _send_compute_host_control(str(params.get("session_id") or ""), route_name="session.save", wait=True) except Exception as exc: return _err(rid, 5011, f"compute-host session save failed: {exc}") if (resp := _compute_host_ack_error(rid, ack, 5011, "compute-host session save failed")) is not None: return resp - result = ack.get("result") - if not isinstance(result, dict): + if not isinstance(result := ack.get("result"), dict): return _err(rid, 5011, "compute-host session save returned an invalid response") return _ok(rid, result) @@ -1944,31 +1647,27 @@ def _save_via_compute_host(rid, params: dict) -> dict: def _compress_via_compute_host(rid, params: dict, session: dict) -> dict: """``session.compress`` for a turn-isolated session: forward ``/compress`` to the host.""" sid = str(params.get("session_id") or "") - focus_topic = str(params.get("focus_topic", "") or "").strip() - command = "/compress" + (f" {focus_topic}" if focus_topic else "") + focus_topic = _str_param(params, "focus_topic") def _on_late_ack(late: dict, _sid=sid) -> None: _adopt_late_compute_host_compress_ack(_sid, session, late, route_name="session.compress") try: ack = _send_compute_host_control( - sid, route_name="session.compress", command=command, wait=True, + sid, route_name="session.compress", command="/compress" + (f" {focus_topic}" if focus_topic else ""), # compression.context_total_ceiling_seconds: the host legitimately runs that long. - timeout=_compute_host_compress_wait_seconds(), on_late_ack=_on_late_ack) + wait=True, timeout=_compute_host_compress_wait_seconds(), on_late_ack=_on_late_ack) except queue.Empty: # Waiter gave up, host still compressing; the late-ack handler adopts the rotated session when it # lands. Not an error (a 5019 here reported timeouts that later succeeded). - return _ok(rid, { - "status": "pending", "turn_isolation": True, - "message": ( - "compression still running in the background; " - "the transcript will refresh when it finishes")}) + return _ok(rid, {"status": "pending", "turn_isolation": True, + "message": ("compression still running in the background; " + "the transcript will refresh when it finishes")}) except Exception as exc: return _err(rid, 5019, f"compute-host compress failed: {exc}") if (resp := _compute_host_ack_error(rid, ack, 4009, "compute-host compress failed")) is not None: return resp _apply_compute_host_metadata_mirror(session, ack) - host_result = ack.get("result") - if isinstance(host_result, dict): + if isinstance(host_result := ack.get("result"), dict): # Host-owned result verbatim (carries `status: aborted` / `summary.aborted`). return _ok(rid, {**host_result, "turn_isolation": True}) host_info = ack.get("session_info") if isinstance(ack.get("session_info"), dict) else {} @@ -1976,12 +1675,57 @@ def _compress_via_compute_host(rid, params: dict, session: dict) -> dict: "status": "compressed", "turn_isolation": True, # `messages` goes top-level for the transcript replacement; don't duplicate it in the ack. "host_ack": {key: value for key, value in ack.items() if key != "messages"}, "info": host_info, - "messages": ( - _history_to_messages(ack.get("messages")) if isinstance(ack.get("messages"), list) else [] - ), + "messages": _history_to_messages(ack.get("messages")) if isinstance(ack.get("messages"), list) else [], "usage": host_info.get("usage") if isinstance(host_info.get("usage"), dict) else {}}) +def _compress_live(rid, sid: str, session: dict, focus_topic: str) -> dict: + """In-process ``session.compress``: status pinned "compressing", then the before/after summary + messages.""" + 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 + with session["history_lock"]: + before_messages = list(session.get("history", [])) + history_version = int(session.get("history_version", 0)) + before_count = len(before_messages) + _agent = session["agent"] + _sys_prompt = getattr(_agent, "_cached_system_prompt", "") or "" + _tools = getattr(_agent, "tools", None) or None + + def _tokens(msgs) -> int: + # Re-reads prompt + tools each call: _compress_context may have rebuilt the system prompt. + sys_prompt = getattr(_agent, "_cached_system_prompt", "") or _sys_prompt + tools = getattr(_agent, "tools", None) or _tools + return estimate_request_tokens_rough(msgs, system_prompt=sys_prompt, tools=tools) if msgs else 0 + before_tokens = _tokens(before_messages) + if before_count >= 4: + focus_suffix = f', focus: "{focus_topic}"' if focus_topic else "" + _status_update(sid, "compressing", + f"⠋ compressing {before_count} messages (~{before_tokens:,} tok){focus_suffix}…") + try: + removed, usage = _compress_session_history( + session, focus_topic, approx_tokens=before_tokens, before_messages=before_messages, + history_version=history_version) + with session["history_lock"]: + messages = list(session.get("history", [])) + after_tokens = _tokens(messages) + agent = session["agent"] + _sync_session_key_after_compress(sid, session) + summary = summarize_manual_compression(before_messages, messages, before_tokens, after_tokens, + compression_state=getattr(agent, "context_compressor", None)) + info = _session_info(agent, session) + _emit("session.info", sid, info) + finalize_context_engine_compression_notification(agent, committed=True) + return _ok(rid, { + "status": "aborted" if summary["aborted"] else "compressed", "removed": removed, + "before_messages": before_count, "after_messages": len(messages), + "before_tokens": before_tokens, "after_tokens": after_tokens, "summary": summary, + "usage": usage, "info": info, "messages": _history_to_messages(messages)}) + finally: + # Always clear the pinned compressing status (success, no-op, or raise). + _status_update(sid, "ready") + + @method("session.compress") def _(rid, params: dict) -> dict: session, err = _sess_nowait(params, rid) @@ -1994,68 +1738,20 @@ def _(rid, params: dict) -> dict: return err if session.get("running"): return _err(rid, 4009, "session busy — /interrupt the current turn before /compress") - from agent.conversation_compression import finalize_context_engine_compression_notification sid = params.get("session_id", "") - focus_topic = str(params.get("focus_topic", "") or "").strip() try: - from agent.manual_compression_feedback import summarize_manual_compression - from agent.model_metadata import estimate_request_tokens_rough - with session["history_lock"]: - before_messages = list(session.get("history", [])) - history_version = int(session.get("history_version", 0)) - before_count = len(before_messages) - _agent = session["agent"] - _sys_prompt = getattr(_agent, "_cached_system_prompt", "") or "" - _tools = getattr(_agent, "tools", None) or None - - def _tokens(msgs, sys_prompt, tools) -> int: - return estimate_request_tokens_rough(msgs, system_prompt=sys_prompt, tools=tools) if msgs else 0 - before_tokens = _tokens(before_messages, _sys_prompt, _tools) - if before_count >= 4: - focus_suffix = f', focus: "{focus_topic}"' if focus_topic else "" - _status_update( - sid, "compressing", - f"⠋ compressing {before_count} messages (~{before_tokens:,} tok){focus_suffix}…") - try: - removed, usage = _compress_session_history( - session, focus_topic, approx_tokens=before_tokens, before_messages=before_messages, - history_version=history_version) - with session["history_lock"]: - messages = list(session.get("history", [])) - after_count = len(messages) - # Re-read prompt + tools: _compress_context may have rebuilt the system prompt. - after_tokens = _tokens( - messages, getattr(_agent, "_cached_system_prompt", "") or _sys_prompt, - getattr(_agent, "tools", None) or _tools) - agent = session["agent"] - _sync_session_key_after_compress(sid, session) - summary = summarize_manual_compression( - before_messages, messages, before_tokens, after_tokens, - compression_state=getattr(agent, "context_compressor", None)) - info = _session_info(agent, session) - _emit("session.info", sid, info) - finalize_context_engine_compression_notification(agent, committed=True) - return _ok(rid, { - "status": "aborted" if summary["aborted"] else "compressed", "removed": removed, - "before_messages": before_count, "after_messages": after_count, - "before_tokens": before_tokens, "after_tokens": after_tokens, "summary": summary, - "usage": usage, "info": info, - # Same projection as session.resume / session.history. - "messages": _history_to_messages(messages)}) - finally: - # Always clear the pinned compressing status (success, no-op, or raise). - _status_update(sid, "ready") + return _compress_live(rid, sid, session, _str_param(params, "focus_topic")) except CompressionLockHeld as e: _status_update(sid, "ready") from agent.manual_compression_feedback import describe_compression_lock_skip return _ok(rid, {"compressed": False, "lock_held": True, "message": describe_compression_lock_skip(e.holder)}) except Exception as e: + from agent.conversation_compression import finalize_context_engine_compression_notification finalize_context_engine_compression_notification(session["agent"], committed=False) return _err(rid, 5005, str(e)) -@method("session.save") -@_with_live_session +@_session_method("session.save", live=True) def _(rid, params: dict, session: dict) -> dict: if _session_uses_compute_host(session): return _save_via_compute_host(rid, params) @@ -2070,68 +1766,49 @@ def _(rid, params: dict, session: dict) -> dict: with session["history_lock"]: messages = list(session.get("history", [])) # Prefer the agent's session_start (classic CLI export); else the gateway created_at. - agent_start = getattr(agent, "session_start", None) - if isinstance(agent_start, datetime): - session_start = agent_start.isoformat() - else: + started = getattr(agent, "session_start", None) + if not isinstance(started, datetime): created_at = session.get("created_at") - session_start = datetime.fromtimestamp(created_at).isoformat() if isinstance(created_at, (int, float)) else "" + started = datetime.fromtimestamp(created_at) if isinstance(created_at, (int, float)) else None try: with open(path, "w", encoding="utf-8") as f: - json.dump({ - "model": getattr(agent, "model", ""), - "session_id": getattr(agent, "session_id", None) or session.get("session_key") or "", - "session_start": session_start, - "system_prompt": getattr(agent, "_cached_system_prompt", "") or "", - "messages": messages, - }, f, indent=2, ensure_ascii=False) - return _ok(rid, {"file": str(path)}) + json.dump({"model": getattr(agent, "model", ""), + "session_id": getattr(agent, "session_id", None) or session.get("session_key") or "", + "session_start": started.isoformat() if started else "", + "system_prompt": getattr(agent, "_cached_system_prompt", "") or "", + "messages": messages}, f, indent=2, ensure_ascii=False) except Exception as e: return _err(rid, 5011, str(e)) + return _ok(rid, {"file": str(path)}) @method("session.close") def _(rid, params: dict) -> dict: - sid = params.get("session_id", "") - # Lock only the ownership claim; finalization (plugin cleanup) must not block resumes. - with _session_resume_lock: - session = _pop_session_by_id(sid) - closed = _teardown_popped_session(session, end_reason="tui_close") - return _ok(rid, {"closed": closed}) + with _session_resume_lock: # lock only the ownership claim; finalization must not block resumes + session = _pop_session_by_id(params.get("session_id", "")) + return _ok(rid, {"closed": _teardown_popped_session(session, end_reason="tui_close")}) # ── session.branch ─────────────────────────────────────────────────── - - def _visible_branch_history(messages) -> list: - """user/assistant rows with visible text, as FULL row copies (reasoning + timeline-marker - tags must survive the branch).""" - return [ - dict(message) for message in messages or [] - if isinstance(message, dict) and message.get("role") in {"user", "assistant"} - and _coerce_message_text(message.get("content")).strip()] + """user/assistant rows with visible text, as FULL copies (reasoning + timeline-marker tags survive).""" + return [dict(message) for message in messages or [] + if isinstance(message, dict) and message.get("role") in {"user", "assistant"} + and _coerce_message_text(message.get("content")).strip()] def _build_branch_agent(session: dict, new_sid: str, new_key: str, history: list, source: str): - """Build + register the branched agent bound to the parent's profile (home + secret scope, - the profile's own state.db handle). The DEDICATED handle is ours until - ``_transfer_db_to_agent`` (unconditional drop, as session.resume); released here on failure.""" + """Build + register the branched agent in the parent's profile; the DEDICATED db handle is ours until + ``_transfer_db_to_agent`` (released here on failure).""" parent_home = session.get("profile_home") - branch_db = None - branch_owns_db = False + branch_db, branch_owns_db = _profile_session_db(parent_home) if parent_home else (None, False) try: - if parent_home: - from hermes_state import get_shared_session_db - branch_db = get_shared_session_db(Path(parent_home) / "state.db") - branch_owns_db = True with _profile_build_scope(parent_home): - agent = _make_agent_in_context( - new_sid, new_key, session_db=branch_db, platform_override=source, - context_cwd_is_launch_artifact=_context_cwd_is_launch_artifact(session)) - _init_session( - new_sid, new_key, agent, list(history), cols=session.get("cols", 80), - cwd=_session_cwd(session), session_db=branch_db, source=source, profile_home=parent_home, - explicit_cwd=bool(session.get("explicit_cwd"))) + agent = _make_agent_in_context(new_sid, new_key, session_db=branch_db, platform_override=source, + context_cwd_is_launch_artifact=_context_cwd_is_launch_artifact(session)) + _init_session(new_sid, new_key, agent, list(history), cols=session.get("cols", 80), + cwd=_session_cwd(session), session_db=branch_db, source=source, profile_home=parent_home, + explicit_cwd=bool(session.get("explicit_cwd"))) _transfer_db_to_agent(agent, branch_db) branch_owns_db = False if new_sid in _sessions: @@ -2139,93 +1816,79 @@ def _build_branch_agent(session: dict, new_sid: str, new_key: str, history: list return agent finally: if branch_owns_db and branch_db is not None: - with contextlib.suppress(Exception): - from hermes_state import release_or_close - release_or_close(branch_db) + _release_db(branch_db) _BRANCH_COPY_FIELDS = ( "reasoning", "reasoning_content", "reasoning_details", "codex_reasoning_items", "codex_message_items", - # Timeline markers ride as role=user; without the tag they become bare user turns after a restart, - # corrupting the truncate ordinal address space. + # Timeline markers ride as role=user; untagged they become bare user turns after a restart, corrupting + # the truncate ordinal address space. "display_kind", "display_metadata", # Branch copies are history, not new activity: keep the parent's timestamps. "timestamp") -@method("session.branch") -@_with_live_session +def _branch_source_history(db, session: dict, old_key: str) -> list: + """Rows a branch copies: the persisted DISPLAY projection reconciled with live memory (live history is + the MODEL projection — post-compaction summary + tail — the child would lose every archived turn).""" + with session["history_lock"]: + in_memory_history = [ + dict(msg) for msg in list(session.get("display_history_prefix") or []) + list(session.get("history", [])) + if isinstance(msg, dict)] + history = None + if callable(get_resume_conversations := getattr(db, "get_resume_conversations", None)): + try: + _, display_history = get_resume_conversations(old_key) + history = _visible_branch_history(_reconcile_display_with_live(display_history, in_memory_history)) + except Exception: + logger.debug("branch display projection read failed", exc_info=True) + return history or _visible_branch_history(in_memory_history) + + +@_session_method("session.branch", live=True) def _(rid, params: dict, session: dict) -> dict: # Write into the parent's profile-scoped state.db; the launch handle would orphan rows. with _session_db(session) as db: if db is None: return _db_unavailable_error(rid, code=5008) old_key = session["session_key"] - with session["history_lock"]: - in_memory_history = [ - dict(msg) - for msg in list(session.get("display_history_prefix") or []) + list(session.get("history", [])) - if isinstance(msg, dict)] - # Live history is the MODEL projection (post-compaction: summary + tail). Snapshot the persisted - # display projection or the child loses every archived turn. - history = None - get_resume_conversations = getattr(db, "get_resume_conversations", None) - if callable(get_resume_conversations): - try: - _, display_history = get_resume_conversations(old_key) - display_history = _reconcile_display_with_live(display_history, in_memory_history) - history = _visible_branch_history(display_history) - except Exception: - logger.debug("branch display projection read failed", exc_info=True) - history = history or _visible_branch_history(in_memory_history) + history = _branch_source_history(db, session, old_key) if not history: return _err(rid, 4008, "nothing to branch — send a message first") - count = params.get("count") - if isinstance(count, int) and count > 0: + if isinstance(count := params.get("count"), int) and count > 0: history = history[:count] - new_key = _new_session_key() - new_sid = uuid.uuid4().hex[:8] - source = _session_source(session) + new_key, new_sid, source = _new_session_key(), uuid.uuid4().hex[:8], _session_source(session) try: title = params.get("name", "") or _branch_title(db, old_key) - _create_branch_row( - db, new_key, old_key, source=source, cwd=_session_cwd(session), - profile_name=( - Path(session["profile_home"]).name if session.get("profile_home") else _current_profile_name() - )) - _copy_branch_transcript(db, new_key, title, history, _BRANCH_COPY_FIELDS) + home = session.get("profile_home") + _persist_branch(db, new_key, old_key, title, history, source=source, cwd=_session_cwd(session), + profile_name=Path(home).name if home else _current_profile_name(), + copy_fields=_BRANCH_COPY_FIELDS) except Exception as e: return _err(rid, 5008, f"branch failed: {e}") try: agent = _build_branch_agent(session, new_sid, new_key, history, source) except Exception as e: return _err(rid, 5000, f"agent init failed on branch: {e}") - return _ok(rid, { - "session_id": new_sid, "stored_session_id": new_key, "title": title, "parent": old_key, - "message_count": len(history), "messages": _history_to_messages(history), - "info": _session_info(agent, _sessions.get(new_sid))}) + return _ok(rid, {"session_id": new_sid, "stored_session_id": new_key, "title": title, "parent": old_key, + "message_count": len(history), "messages": _history_to_messages(history), + "info": _session_info(agent, _sessions.get(new_sid))}) # ── interrupt / steer / redirect ───────────────────────────────────── - - @method("session.interrupt") def _(rid, params: dict) -> dict: - # Keypress barge-in also silences streaming TTS (voice is process-global). - _tts_stream_stop() + _tts_stream_stop() # keypress barge-in also silences streaming TTS (voice is process-global) session, err = _sess_nowait(params, rid) if err: return err - expected_hosted_task_id = str(params.get("expected_hosted_task_id") or "").strip() - if expected_hosted_task_id: + if expected := _str_param(params, "expected_hosted_task_id"): with session["history_lock"]: - active_task = session.get("_hosted_room_task") - if not ( - session.get("running") and isinstance(active_task, dict) - and active_task.get("task_id") == expected_hosted_task_id): + task = session.get("_hosted_room_task") + if not (session.get("running") and isinstance(task, dict) and task.get("task_id") == expected): return _ok(rid, {"status": "not_interrupted", "interrupted": False}) + sid = str(params.get("session_id") or "") if _session_uses_compute_host(session): - sid = str(params.get("session_id") or "") try: _interrupt_session_turn(sid, session, request_id=f"interrupt-{rid}") except Exception as exc: @@ -2234,10 +1897,10 @@ def _(rid, params: dict) -> dict: session, err = _sess(params, rid) if err: return err - _interrupt_session_turn(str(params.get("session_id") or ""), session) - # Retire the crash-recovery marker NOW: until the run thread's finally, a backend exit looks like a - # crash and session.resume auto-continues the turn the user just stopped. The extra key covers - # compression rotating session_key mid-turn. + _interrupt_session_turn(sid, session) + # Retire the crash-recovery marker NOW: until the run thread's finally, a backend exit looks like a crash + # and session.resume auto-continues the turn the user just stopped (the extra key covers compression + # rotating session_key mid-turn). with session["history_lock"]: active_marker_key = str(session.pop("_active_turn_marker_key", "") or "") _retire_turn_marker(session, active_marker_key) @@ -2245,8 +1908,8 @@ def _(rid, params: dict) -> dict: def _apply_correction(rid, session: dict, verb: str, text: str, accepted_status: str) -> dict: - """Run ``agent.(text)``; on acceptance record it on the live turn (mid-turn resume rebuilds - the bubble) and purge queued self-copies so post-turn drain cannot re-fire the old prompt.""" + """``agent.(text)``; on acceptance record it on the live turn (mid-turn resume rebuilds the bubble) + and purge queued self-copies so post-turn drain cannot re-fire the old prompt.""" try: accepted = getattr(session["agent"], verb)(text) except Exception as exc: @@ -2259,53 +1922,45 @@ def _apply_correction(rid, session: dict, verb: str, text: str, accepted_status: return _ok(rid, {"status": accepted_status if accepted else "rejected", "text": text}) -@method("session.steer") -def _(rid, params: dict) -> dict: - """Inject text into the next tool result without interrupting (AIAgent.steer(): no new - user turn, no role alternation violation).""" - text = (params.get("text") or "").strip() - if not text: - return _err(rid, 4002, "text is required") - session, err = _sess_nowait(params, rid) - if err: - return err - if not hasattr(session.get("agent"), "steer"): - return _err(rid, 4010, "agent does not support steer") - return _apply_correction(rid, session, "steer", text, "queued") +def _correction_method(name: str, verb: str, accepted_status: str, supported, unsupported: str): + """steer/redirect RPC: ``params.text`` (4002, checked before the session) into a live session; + ``supported(agent)`` gates 4010.""" + @method(name) + def _(rid, params: dict) -> dict: + if not (text := (params.get("text") or "").strip()): + return _err(rid, 4002, "text is required") + session, err = _sess_nowait(params, rid) + if err: + return err + agent = session.get("agent") + # Redirect during the turn-build window (running=True, agent None): queue for the next turn instead of + # a misleading 4010 the client swallows into a lost follow-up. + if verb == "redirect" and agent is None and session.get("running"): + _enqueue_prompt(session, text, current_transport() or _stdio_transport) + session["last_active"] = time.time() + return _ok(rid, {"status": "queued", "text": text}) + if not supported(agent): + return _err(rid, 4010, unsupported) + return _apply_correction(rid, session, verb, text, accepted_status) -@method("session.redirect") -def _(rid, params: dict) -> dict: - """Redirect the active model turn while preserving valid work/context.""" - text = (params.get("text") or "").strip() - if not text: - return _err(rid, 4002, "text is required") - session, err = _sess_nowait(params, rid) - if err: - return err - agent = session.get("agent") - # Turn-build window (running=True, agent None): queue for the next turn instead of a misleading 4010 - # the client swallows into a lost follow-up. - if agent is None and session.get("running"): - _enqueue_prompt(session, text, current_transport() or _stdio_transport) - session["last_active"] = time.time() - return _ok(rid, {"status": "queued", "text": text}) - if getattr(agent, "_supports_active_turn_redirect", False) is not True or not hasattr(agent, "redirect"): - return _err(rid, 4010, "agent does not support active-turn redirect") - return _apply_correction(rid, session, "redirect", text, "redirected") +# Inject text into the next tool result without interrupting (AIAgent.steer(): no new user turn, no role +# alternation violation). +_correction_method("session.steer", "steer", "queued", lambda agent: hasattr(agent, "steer"), + "agent does not support steer") +# Redirect the active model turn while preserving valid work/context. +_correction_method("session.redirect", "redirect", "redirected", + lambda agent: getattr(agent, "_supports_active_turn_redirect", False) is True + and hasattr(agent, "redirect"), "agent does not support active-turn redirect") # ── delegation / spawn trees ───────────────────────────────────────── - - @method("delegation.status") def _(rid, params: dict) -> dict: - from tools.delegate_tool import ( - is_spawn_paused, list_active_subagents, _get_max_concurrent_children, _get_max_spawn_depth) - return _ok(rid, { - "active": list_active_subagents(), "paused": is_spawn_paused(), - "max_spawn_depth": _get_max_spawn_depth(), "max_concurrent_children": _get_max_concurrent_children(), - }) + from tools import delegate_tool as dt + return _ok(rid, {"active": dt.list_active_subagents(), "paused": dt.is_spawn_paused(), + "max_spawn_depth": dt._get_max_spawn_depth(), + "max_concurrent_children": dt._get_max_concurrent_children()}) @method("delegation.pause") @@ -2317,57 +1972,46 @@ def _(rid, params: dict) -> dict: @method("subagent.interrupt") def _(rid, params: dict) -> dict: from tools.delegate_tool import interrupt_subagent - subagent_id = str(params.get("subagent_id") or "").strip() - if not subagent_id: + if not (subagent_id := _str_param(params, "subagent_id")): return _err(rid, 4000, "subagent_id required") return _ok(rid, {"found": interrupt_subagent(subagent_id), "subagent_id": subagent_id}) @method("subagent.steer") def _(rid, params: dict) -> dict: - """Queue steering text into a live delegated child (AIAgent.steer(); the in-flight tool call - is never cut). "queued" is not "delivered": a child past its final tool batch surfaces - the race as ``missed_steer`` on the parent's completion entry.""" + """Queue steering text into a live delegated child (the in-flight tool call is never cut). "queued" + is not "delivered": a child past its final tool batch surfaces ``missed_steer`` on the parent entry.""" from tools.delegate_tool import steer_subagent - subagent_id = str(params.get("subagent_id") or "").strip() - if not subagent_id: + if not (subagent_id := _str_param(params, "subagent_id")): return _err(rid, 4000, "subagent_id required") - text = (params.get("text") or "").strip() - if not text: + if not (text := (params.get("text") or "").strip()): return _err(rid, 4002, "text is required") - _invoking_session, err = _sess_nowait(params, rid) - if err: + if (err := _sess_nowait(params, rid)[1]) is not None: return err - invoking_session_id = str(params.get("session_id") or "").strip() - invoking_transport, invoking_session = _current_session_steer_authority(invoking_session_id) - queued = invoking_transport is not None and invoking_session is not None and steer_subagent( - subagent_id, text, owner_session_id=invoking_session_id, owner_transport=invoking_transport, - owner_session_record=invoking_session) + owner_id = _str_param(params, "session_id") + transport, owner = _current_session_steer_authority(owner_id) + queued = transport is not None and owner is not None and steer_subagent( + subagent_id, text, owner_session_id=owner_id, owner_transport=transport, owner_session_record=owner) return _ok(rid, {"status": "queued" if queued else "rejected", "subagent_id": subagent_id, "text": text}) @method("spawn_tree.save") def _(rid, params: dict) -> dict: - session_id = str(params.get("session_id") or "").strip() + session_id = _str_param(params, "session_id") subagents = params.get("subagents") or [] if not isinstance(subagents, list) or not subagents: return _err(rid, 4000, "subagents list required") - started_at = params.get("started_at") - finished_at = params.get("finished_at") or time.time() - label = str(params.get("label") or "") - ts = datetime.utcfromtimestamp(float(finished_at)).strftime("%Y%m%dT%H%M%S") + started_at, label = params.get("started_at"), str(params.get("label") or "") + finished_at = float(params.get("finished_at") or time.time()) d = _spawn_tree_session_dir(session_id or "default") - path = d / f"{ts}.json" + path = d / f"{datetime.utcfromtimestamp(finished_at).strftime('%Y%m%dT%H%M%S')}.json" + meta = {"session_id": session_id, "started_at": float(started_at) if started_at else None, + "finished_at": finished_at, "label": label} try: - payload = { - "session_id": session_id, "started_at": float(started_at) if started_at else None, - "finished_at": float(finished_at), "label": label, "subagents": subagents} - path.write_text(json.dumps(payload, ensure_ascii=False), encoding="utf-8") + path.write_text(json.dumps({**meta, "subagents": subagents}, ensure_ascii=False), encoding="utf-8") except OSError as exc: return _err(rid, 5000, f"spawn_tree.save failed: {exc}") - _append_spawn_tree_index(d, { - "path": str(path), "session_id": session_id, "started_at": payload["started_at"], - "finished_at": payload["finished_at"], "label": label, "count": len(subagents)}) + _append_spawn_tree_index(d, {"path": str(path), **meta, "count": len(subagents)}) return _ok(rid, {"path": str(path), "session_id": session_id}) @@ -2377,51 +2021,41 @@ def _legacy_spawn_tree_entry(p, session_dir_name: str) -> dict | None: stat = p.stat() except OSError: return None - try: + raw = {} + with contextlib.suppress(Exception): raw = json.loads(p.read_text(encoding="utf-8")) - except Exception: - raw = {} subagents = raw.get("subagents") or [] - return { - "path": str(p), "session_id": raw.get("session_id") or session_dir_name, - "finished_at": raw.get("finished_at") or stat.st_mtime, "started_at": raw.get("started_at"), - "label": raw.get("label") or "", "count": len(subagents) if isinstance(subagents, list) else 0, - } + return {"path": str(p), "session_id": raw.get("session_id") or session_dir_name, + "finished_at": raw.get("finished_at") or stat.st_mtime, "started_at": raw.get("started_at"), + "label": raw.get("label") or "", "count": len(subagents) if isinstance(subagents, list) else 0} @method("spawn_tree.list") def _(rid, params: dict) -> dict: - session_id = str(params.get("session_id") or "").strip() - limit = int(params.get("limit") or 50) + session_id = _str_param(params, "session_id") if bool(params.get("cross_session")): roots = [p for p in _spawn_trees_root().iterdir() if p.is_dir()] else: roots = [_spawn_tree_session_dir(session_id or "default")] entries: list[dict] = [] for d in roots: - indexed = _read_spawn_tree_index(d) - if indexed: + if indexed := _read_spawn_tree_index(d): # Skip index entries whose snapshot file was manually deleted. entries.extend(e for e in indexed if (p := e.get("path")) and Path(p).exists()) - continue - # Legacy (pre-index) sessions: full scan, once per session until the next save. - for p in d.glob("*.json"): - if p.name != _SPAWN_TREE_INDEX and (entry := _legacy_spawn_tree_entry(p, d.name)) is not None: - entries.append(entry) + else: # Legacy (pre-index) sessions: full scan, once per session until the next save. + entries.extend( + entry for p in d.glob("*.json") + if p.name != _SPAWN_TREE_INDEX and (entry := _legacy_spawn_tree_entry(p, d.name)) is not None) entries.sort(key=lambda e: e.get("finished_at") or 0, reverse=True) - return _ok(rid, {"entries": entries[:limit]}) + return _ok(rid, {"entries": entries[:int(params.get("limit") or 50)]}) @method("spawn_tree.load") def _(rid, params: dict) -> dict: - raw_path = str(params.get("path") or "").strip() - if not raw_path: + if not (raw_path := _str_param(params, "path")): return _err(rid, 4000, "path required") - # Reject paths escaping the spawn-trees root. - root = _spawn_trees_root().resolve() try: - resolved = Path(raw_path).resolve() - resolved.relative_to(root) + (resolved := Path(raw_path).resolve()).relative_to(_spawn_trees_root().resolve()) except (ValueError, OSError) as exc: return _err(rid, 4030, f"path outside spawn-trees root: {exc}") try: @@ -2432,31 +2066,25 @@ def _(rid, params: dict) -> dict: # ── terminal / event replay ────────────────────────────────────────── - - -@method("terminal.resize") -@_with_session +@_session_method("terminal.resize") def _(rid, params: dict, session: dict) -> dict: - session["cols"] = int(params.get("cols", 80)) - return _ok(rid, {"cols": session["cols"]}) + session["cols"] = cols = int(params.get("cols", 80)) + return _ok(rid, {"cols": cols}) @method("session.events.since") def _(rid, params: dict) -> dict: - """Replay events newer than the client's last-seen seq (WS reconnect). Frames older than - the ring window report ``truncated`` so the client refetches instead of accepting a gap.""" + """Replay events after ``last_seen`` (WS reconnect); ``truncated`` past the ring window → client refetches.""" sid = str(params.get("session_id") or "") try: last_seen = int(params.get("last_seen", 0)) except (TypeError, ValueError): return _err(rid, -32602, "invalid params: last_seen must be an integer") - from tui_gateway import event_replay - frames = event_replay.events_since(sid, last_seen) - return _ok(rid, { - "events": frames, "latest_seq": event_replay.latest_seq(sid), - "truncated": event_replay.is_truncated(sid, last_seen), "count": len(frames), - # In-process seq: clients reset watermarks when this differs from gateway.ready's. - "epoch": event_replay.replay_epoch()}) + from tui_gateway import event_replay as er + frames = er.events_since(sid, last_seen) + # ``epoch``: in-process seq — clients reset watermarks when this differs from gateway.ready's. + return _ok(rid, {"events": frames, "latest_seq": er.latest_seq(sid), "truncated": er.is_truncated(sid, last_seen), + "count": len(frames), "epoch": er.replay_epoch()}) @method("session.events.stats") diff --git a/tui_gateway/methods_slash.py b/tui_gateway/methods_slash.py index 270da69700..4588aa9ca5 100644 --- a/tui_gateway/methods_slash.py +++ b/tui_gateway/methods_slash.py @@ -1,4 +1,4 @@ -"""slash.exec helpers: live-session command output and side-effect mirroring after a slash command ran in the worker. +"""slash.exec helpers: live-session command output + side-effect mirroring after a worker slash command. Bodies are rebound onto server.py's globals at install time (see method_ctx.bind_module), so they reference server.py globals bare. @@ -6,7 +6,6 @@ method_ctx.bind_module), so they reference server.py globals bare. from __future__ import annotations - import contextlib from .method_ctx import HandlerRegistry, bind_module @@ -16,9 +15,6 @@ _registry = HandlerRegistry() # ── Live-session slash output ──────────────────────────────────────── -_LIVE_SESSION_DIRECT_COMMANDS = frozenset( - {"clear", "compress", "effort", "history", "models", "prompt", "rename", "review", "status", "usage"} -) # Answered from the live session ONLY when the agent lives on a compute host. _ISOLATED_SESSION_READ_COMMANDS = frozenset({"context", "tools", "help"}) @@ -26,27 +22,21 @@ _NO_AGENT_USAGE = "(._.) No active agent -- send a message first." _NO_AGENT = "No active agent -- send a message first." -def _format_live_review_output(session: Optional[dict], arg: str) -> str: - """Dispatch /review against the live session's agent. - - The reviewer subagent runs on the async delegation rail; the TUI notification - poller drains its completion back into this chat. The dispatch stamps the - parent's durable session_id as the completion's session_key, which is what - ``_session_owns_notification_event`` matches against. - """ +def _format_live_review_output(sid: str, session: Optional[dict], arg: str) -> str: + """Dispatch /review against the live session's agent. The reviewer runs on the async + delegation rail; its completion is stamped with the parent's durable session_id, which + ``_session_owns_notification_event`` matches to drain it back into this chat.""" if session is None: return "Nothing to review yet — send a message first." if _session_uses_compute_host(session): return "/review runs on the local agent only for now — this session's agent lives on a remote compute host." - agent = session.get("agent") - if agent is None: + if (agent := session.get("agent")) is None: return "Nothing to review yet — send a message first." if session.get("running"): return "session busy — wait for the current turn to finish, then /review" with session.get("history_lock") or contextlib.nullcontext(): snapshot = list(session.get("history", [])) - if not snapshot: - snapshot = list(getattr(agent, "_session_messages", None) or []) + snapshot = snapshot or list(getattr(agent, "_session_messages", None) or []) try: from agent.review_engine import format_dispatch_note, start_review result = start_review(agent, snapshot, arg or "") @@ -57,7 +47,7 @@ def _format_live_review_output(session: Optional[dict], arg: str) -> str: return format_dispatch_note(result, arg or "") -def _format_live_usage_output(session: dict) -> str: +def _format_live_usage_output(sid: str, session: dict, arg: str) -> str: agent = session.get("agent") usage = _session_usage_snapshot(session) if agent is None and not usage: @@ -70,25 +60,18 @@ def _format_live_usage_output(session: dict) -> str: def n(key: str) -> str: return f"{int(usage.get(key) or 0):,}" - lines = [ - "Session Token Usage", - "────────────────────────────────────────", - f"Model: {usage.get('model') or _metadata_mirror(session).get('model') or getattr(agent, 'model', '') or '(unknown)'}", - f"Input tokens: {n('input')}", - f"Output tokens: {n('output')}"] + rows = [("Input tokens:", n("input")), ("Output tokens:", n("output"))] if int(usage.get("reasoning") or 0): - lines.append(f"Reasoning tokens: {n('reasoning')}") - lines += [ - f"Prompt tokens: {n('prompt')}", - f"Completion tokens: {n('completion')}", - f"Total tokens: {n('total')}", - f"API calls: {n('calls')}"] + rows.append(("Reasoning tokens:", n("reasoning"))) + rows += [("Prompt tokens:", n("prompt")), ("Completion tokens:", n("completion")), + ("Total tokens:", n("total")), ("API calls:", n("calls"))] if usage.get("context_max"): - lines.append( - f"Current context: {n('context_used')} / {n('context_max')} " - f"({int(usage.get('context_percent') or 0)}%)") - lines += [f"Messages: {message_count:,}", f"Compressions: {n('compressions')}"] - return "\n".join(lines) + pct = int(usage.get("context_percent") or 0) + rows.append(("Current context:", f"{n('context_used')} / {n('context_max')} ({pct}%)")) + rows += [("Messages:", f"{message_count:,}"), ("Compressions:", n("compressions"))] + model = usage.get("model") or _metadata_mirror(session).get("model") or getattr(agent, "model", "") or "(unknown)" + lines = ["Session Token Usage", "────────────────────────────────────────", f"Model: {model}"] + return "\n".join(lines + [f"{label:<30}{value}" for label, value in rows]) def _live_session_messages(session: dict) -> Optional[list]: @@ -97,15 +80,13 @@ def _live_session_messages(session: dict) -> Optional[list]: profile's state.db, and through the launch handle this read comes back empty.""" with _session_db(session) as db: if db is not None and session.get("session_key"): - try: + with contextlib.suppress(Exception): return db.get_messages_as_conversation( session["session_key"], include_ancestors=True, include_row_ids=True) - except Exception: - pass return None -def _format_live_history_output(session: dict) -> str: +def _format_live_history_output(sid: str, session: dict, arg: str) -> str: with session["history_lock"]: history = list(session.get("history", [])) db_history = _live_session_messages(session) @@ -115,30 +96,28 @@ def _format_live_history_output(session: dict) -> str: lines = ["Conversation History", "────────────────────────────────────────"] for idx, message in enumerate(messages, start=1): role = str(message.get("role") or "unknown") - label = "You" if role == "user" else "Hermes" if role == "assistant" else role.title() + label = {"user": "You", "assistant": "Hermes"}.get(role, role.title()) text = str(message.get("text") or message.get("context") or "").strip() - if len(text) > 400: - text = f"{text[:400]}..." + text = f"{text[:400]}..." if len(text) > 400 else text lines.append(f"[{label} #{idx}] {text or '(no text)'}") return "\n".join(lines) -def _format_live_prompt_output(session: dict) -> str: +def _format_live_prompt_output(sid: str, session: dict, arg: str) -> str: agent = session.get("agent") mirror = _metadata_mirror(session) if agent is None and "system_prompt" not in mirror: return _NO_AGENT prompt = ( - mirror.get("system_prompt") - or getattr(agent, "ephemeral_system_prompt", None) - or getattr(agent, "_cached_system_prompt", None) - or "") + mirror.get("system_prompt") or getattr(agent, "ephemeral_system_prompt", None) + or getattr(agent, "_cached_system_prompt", None) or "") if not prompt: return "Current system prompt is not built yet; send a message first." return f"Current system prompt:\n{prompt}" -def _format_live_context_output(session: dict) -> str: +def _format_live_context_output(sid: str, session: dict, arg: str) -> str: + from collections import Counter try: messages = _history_to_messages(_live_session_messages(session) or []) except Exception: @@ -149,23 +128,16 @@ def _format_live_context_output(session: dict) -> str: usage = _session_usage_snapshot(session) mirror = _metadata_mirror(session) lines = [f"Conversation: {len(messages)} messages" if messages else "Conversation is empty (no messages yet)."] - roles: dict[str, int] = {} - for msg in messages: - role = str(msg.get("role") or "unknown") - roles[role] = roles.get(role, 0) + 1 - lines.append( - f" user: {roles.get('user', 0)}, assistant: {roles.get('assistant', 0)}, " - f"tool: {roles.get('tool', 0)}, system: {roles.get('system', 0)}") - model = mirror.get("model") or usage.get("model") or "" - if model: + roles = Counter(str(msg.get("role") or "unknown") for msg in messages) + lines.append(" " + ", ".join(f"{r}: {roles.get(r, 0)}" for r in ("user", "assistant", "tool", "system"))) + if model := mirror.get("model") or usage.get("model") or "": lines.append(f"Model: {model}") lines.append(f"Provider: {mirror.get('provider') or 'auto'}") context_used = int(usage.get("context_used") or usage.get("total") or 0) context_max = int(usage.get("context_max") or 0) if context_used and context_max: lines.append( - f"Context usage: ~{context_used:,} / {context_max:,} tokens ({(context_used / context_max) * 100:.1f}%)" - ) + f"Context usage: ~{context_used:,} / {context_max:,} tokens ({(context_used / context_max) * 100:.1f}%)") elif context_used: lines.append(f"Context usage: ~{context_used:,} tokens") if usage.get("compressions"): @@ -173,7 +145,7 @@ def _format_live_context_output(session: dict) -> str: return "\n".join(lines) -def _format_live_tools_output(session: dict) -> str: +def _format_live_tools_output(sid: str, session: dict, arg: str) -> str: info = _session_info(session.get("agent"), session) groups = info.get("tools") if isinstance(info, dict) else {} if not isinstance(groups, dict) or not groups: @@ -184,7 +156,7 @@ def _format_live_tools_output(session: dict) -> str: return "Available tools ({}):\n{}".format(len(names), "\n".join(f" {name}" for name in names)) -def _format_live_help_output() -> str: +def _format_live_help_output(sid: str, session: dict, arg: str) -> str: try: from hermes_cli.commands import COMMANDS_BY_CATEGORY lines = ["Available commands:", ""] @@ -200,34 +172,35 @@ def _format_live_model_output(session: dict) -> str: agent = session.get("agent") model = getattr(agent, "model", "") if agent is not None else "" provider = getattr(agent, "provider", "") if agent is not None else "" - if model and provider: - return f"Current model: {model} ({provider})" - return f"Current model: {model}" if model else "Current model: (unknown)" + if not model: + return "Current model: (unknown)" + return f"Current model: {model}" + (f" ({provider})" if provider else "") -def _format_live_status_output(sid: str) -> str: +def _format_live_status_output(sid: str, session: dict, arg: str) -> str: response = _methods["session.status"]("status", {"session_id": sid}) if response.get("error"): return str(response["error"].get("message") or "status unavailable") return str(response.get("result", {}).get("output") or "") -# name → (reply when there is no session, formatter(sid, session, arg)). A None -# no-session reply means the formatter handles a missing session itself. +# name → (reply when there is no session, formatter(sid, session, arg) or a fixed reply). +# A None no-session reply means the formatter handles a missing session itself. _LIVE_SLASH_OUTPUT = { - "compress": ("no active session for /compress", lambda sid, s, a: _mirror_slash_side_effects(sid, s, f"/compress {a}".strip())), - "usage": (_NO_AGENT_USAGE, lambda sid, s, a: _format_live_usage_output(s)), - "review": (None, lambda sid, s, a: _format_live_review_output(s, a)), - "history": ("No conversation history yet.", lambda sid, s, a: _format_live_history_output(s)), - "prompt": (_NO_AGENT, lambda sid, s, a: _format_live_prompt_output(s)), - "status": (None, lambda sid, s, a: _format_live_status_output(sid)), - "context": ("Conversation is empty (no messages yet).", lambda sid, s, a: _format_live_context_output(s)), - "tools": ("No tools available.", lambda sid, s, a: _format_live_tools_output(s)), - "help": (None, lambda sid, s, a: _format_live_help_output()), - "clear": (None, lambda sid, s, a: "Screen clear is terminal-only; desktop/TUI chat left unchanged."), - "models": (None, lambda sid, s, a: "Use /model to view or switch the current model; desktop users can also open the model picker."), - "rename": (None, lambda sid, s, a: "Use /title to rename this session."), - "effort": (None, lambda sid, s, a: "Use /reasoning to change reasoning effort.")} + "compress": ("no active session for /compress", + lambda sid, session, arg: _mirror_slash_side_effects(sid, session, f"/compress {arg}".strip())), + "usage": (_NO_AGENT_USAGE, _format_live_usage_output), + "review": (None, _format_live_review_output), + "history": ("No conversation history yet.", _format_live_history_output), + "prompt": (_NO_AGENT, _format_live_prompt_output), + "status": (None, _format_live_status_output), + "context": ("Conversation is empty (no messages yet).", _format_live_context_output), + "tools": ("No tools available.", _format_live_tools_output), + "help": (None, _format_live_help_output), + "clear": (None, "Screen clear is terminal-only; desktop/TUI chat left unchanged."), + "models": (None, "Use /model to view or switch the current model; desktop users can also open the model picker."), + "rename": (None, "Use /title to rename this session."), + "effort": (None, "Use /reasoning to change reasoning effort.")} def _live_slash_command_output(sid: str, session: Optional[dict], name: str, arg: str) -> Optional[str]: @@ -236,10 +209,7 @@ def _live_slash_command_output(sid: str, session: Optional[dict], name: str, arg arg = arg or "" if name == "model" and not arg.strip(): return _format_live_model_output(session or {}) - if name in _ISOLATED_SESSION_READ_COMMANDS: - if not (session is not None and _session_uses_compute_host(session)): - return None - elif name not in _LIVE_SESSION_DIRECT_COMMANDS: + if name in _ISOLATED_SESSION_READ_COMMANDS and not (session is not None and _session_uses_compute_host(session)): return None entry = _LIVE_SLASH_OUTPUT.get(name) if entry is None: @@ -247,7 +217,7 @@ def _live_slash_command_output(sid: str, session: Optional[dict], name: str, arg no_session_reply, fmt = entry if session is None and no_session_reply is not None: return no_session_reply - return fmt(sid, session, arg) + return fmt(sid, session, arg) if callable(fmt) else fmt # ── Side-effect mirroring ──────────────────────────────────────────── @@ -259,17 +229,11 @@ _MUTATES_WHILE_RUNNING = frozenset({"model", "personality", "prompt", "compress" def _compress_live_with_feedback(sid: str, session: dict, agent, arg: str, *, snapshot_kwargs: bool) -> str: - """Compress the live session; return the user-facing feedback text. - - Shared by command.dispatch /compress and the slash mirror so every route shows - "compressed N → M messages / ~X → ~Y tokens". ``snapshot_kwargs`` forwards the - pre-read snapshot (approx_tokens/before_messages/history_version) to - ``_compress_session_history``; the slash mirror passes only the raw arg. The raw - arg goes through unparsed — the choke point parses ``here [N]`` / ``--keep N``. - CompressionLockHeld is a clean no-op (its skip note is returned; the choke point - already discarded the deferred context-engine notification); other errors propagate - to the caller, which finalizes that notification. - """ + """Compress the live session; return the user-facing feedback text (shared by command.dispatch + /compress and the slash mirror). ``snapshot_kwargs`` forwards the pre-read snapshot to + ``_compress_session_history``; the raw arg goes through unparsed (the choke point parses + ``here [N]`` / ``--keep N``). CompressionLockHeld is a clean no-op (skip note returned); + other errors propagate to the caller, which finalizes the context-engine notification.""" from agent.conversation_compression import finalize_context_engine_compression_notification from agent.manual_compression_feedback import describe_compression_lock_skip, summarize_manual_compression from agent.model_metadata import estimate_request_tokens_rough @@ -278,14 +242,14 @@ def _compress_live_with_feedback(sid: str, session: dict, agent, arg: str, *, sn history_version = int(session.get("history_version", 0)) sys_prompt = getattr(agent, "_cached_system_prompt", "") or "" tools = getattr(agent, "tools", None) or None - before_tokens = ( - estimate_request_tokens_rough(before_messages, system_prompt=sys_prompt, tools=tools) if before_messages else 0 - ) + + def estimate(messages, prompt, tool_defs) -> int: + return estimate_request_tokens_rough(messages, system_prompt=prompt, tools=tool_defs) if messages else 0 + before_tokens = estimate(before_messages, sys_prompt, tools) + snapshot = {"approx_tokens": before_tokens, "before_messages": before_messages, "history_version": history_version} try: if snapshot_kwargs: - _compress_session_history( - session, arg.strip() or None, approx_tokens=before_tokens, before_messages=before_messages, - history_version=history_version) + _compress_session_history(session, arg.strip() or None, **snapshot) else: _compress_session_history(session, arg) except CompressionLockHeld as e: @@ -293,11 +257,8 @@ def _compress_live_with_feedback(sid: str, session: dict, agent, arg: str, *, sn _sync_session_key_after_compress(sid, session) with session["history_lock"]: after_messages = list(session.get("history", [])) - after_tokens = ( - estimate_request_tokens_rough( - after_messages, system_prompt=getattr(agent, "_cached_system_prompt", "") or sys_prompt, - tools=getattr(agent, "tools", None) or tools) - if after_messages else 0) + after_tokens = estimate( + after_messages, getattr(agent, "_cached_system_prompt", "") or sys_prompt, getattr(agent, "tools", None) or tools) _emit("session.info", sid, _session_info(agent, session)) fb = summarize_manual_compression( before_messages, after_messages, before_tokens, after_tokens, @@ -306,95 +267,71 @@ def _compress_live_with_feedback(sid: str, session: dict, agent, arg: str, *, sn return "\n".join(filter(None, [fb["headline"], fb["token_line"], fb.get("note")])) -def _mirror_model(sid, session, agent, arg) -> str: - if arg and agent: - return _apply_model_switch(sid, session, arg).get("warning", "") - return "" - - -def _mirror_approvals(sid, session, agent, arg) -> str: - # The worker already persisted approvals.mode; the bare read-only form needs no repaint. - if arg: +def _mirror_approvals(sid, session, agent, arg) -> None: + if arg: # the worker already persisted approvals.mode; the bare read-only form needs no repaint broadcast_session_info() - return "" -def _mirror_personality(sid, session, agent, arg) -> str: +def _mirror_personality(sid, session, agent, arg) -> None: if arg and agent: pname, new_prompt = _validate_personality(arg, _load_cfg()) - # Persist through the single owner so this surface never drifts from the others. - from hermes_cli.personality import persist_personality + from hermes_cli.personality import persist_personality # single owner: no surface drift persist_personality(pname) _apply_personality_to_session(sid, session, new_prompt, pname) - return "" -def _mirror_prompt(sid, session, agent, arg) -> str: +def _mirror_prompt(sid, session, agent, arg) -> None: if agent: cfg = _load_cfg() agent.ephemeral_system_prompt = _prompt_text((cfg.get("agent") or {}).get("system_prompt", "")) or None agent._cached_system_prompt = None - return "" - - -def _mirror_compress(sid, session, agent, arg) -> str: - return _compress_live_with_feedback(sid, session, agent, arg, snapshot_kwargs=False) if agent else "" _FAST_TIERS = {"fast": "priority", "on": "priority", "normal": None, "off": None, "auto": "auto", "cold": "cold"} -def _mirror_fast(sid, session, agent, arg) -> str: +def _mirror_fast(sid, session, agent, arg) -> None: if agent: - mode = arg.lower() - if mode in _FAST_TIERS: - agent.service_tier = _FAST_TIERS[mode] + if arg.lower() in _FAST_TIERS: + agent.service_tier = _FAST_TIERS[arg.lower()] _emit("session.info", sid, _session_info(agent, session)) - return "" -def _mirror_reload_mcp(sid, session, agent, arg) -> str: +def _mirror_reload_mcp(sid, session, agent, arg) -> None: if agent and hasattr(agent, "reload_mcp_tools"): agent.reload_mcp_tools() - return "" -def _mirror_stop(sid, session, agent, arg) -> str: +def _mirror_stop(sid, session, agent, arg) -> None: from tools.process_registry import process_registry process_registry.kill_all() - return "" +# name → mirror(sid, session, agent, arg); a falsy return means "no warning". _SLASH_MIRRORS = { - "model": _mirror_model, - "approvals": _mirror_approvals, - "personality": _mirror_personality, - "prompt": _mirror_prompt, - "compress": _mirror_compress, + "model": lambda sid, session, agent, arg: ( + _apply_model_switch(sid, session, arg).get("warning", "") if arg and agent else ""), + "approvals": _mirror_approvals, "personality": _mirror_personality, "prompt": _mirror_prompt, + "compress": lambda sid, session, agent, arg: ( + _compress_live_with_feedback(sid, session, agent, arg, snapshot_kwargs=False) if agent else ""), "fast": _mirror_fast, - "reload-mcp": _mirror_reload_mcp, - "stop": _mirror_stop} + "reload-mcp": _mirror_reload_mcp, "stop": _mirror_stop} def _compute_host_slash(sid: str, session: dict, name: str, command: str) -> tuple[str, str]: - """Forward a mutating slash command to the session's compute host. - - Returns ``(status, text)``: ``pending`` (compress still running after the wait), - ``failed`` (transport error/timeout), ``rejected`` (host control.error), ``ok`` - (host output; metadata mirror already applied). Compress waits longer and installs - a late-ack adopter so a slow compression still lands in this session. - """ + """Forward a mutating slash command to the session's compute host → ``(status, text)``: + ``pending`` (compress still running after the wait), ``failed`` (transport error/timeout), + ``rejected`` (host control.error), ``ok`` (host output; metadata mirror applied). Compress + waits longer and installs a late-ack adopter so a slow compression still lands here.""" route_name = f"slash.{name}" is_compress = name == "compress" - _late_session = session def _on_late_ack(late: dict, _sid=sid) -> None: - _adopt_late_compute_host_compress_ack(_sid, _late_session, late, route_name=route_name) + _adopt_late_compute_host_compress_ack(_sid, session, late, route_name=route_name) try: ack = _send_compute_host_control( sid, route_name=route_name, command=command, wait=True, - **({"timeout": _compute_host_compress_wait_seconds(), "on_late_ack": _on_late_ack} if is_compress else {}), - ) + **({"timeout": _compute_host_compress_wait_seconds(), "on_late_ack": _on_late_ack} if is_compress else {})) except queue.Empty: if is_compress: return "pending", "compression still running in the background; the transcript will refresh when it finishes" @@ -413,19 +350,17 @@ def _mirror_slash_side_effects(sid: str, session: dict, command: str) -> str: if not parts: return "" name, arg, agent = parts[0], (parts[1].strip() if len(parts) > 1 else ""), session.get("agent") - if name == "compact": - # /compact aliases /compress everywhere; the compute-host control forwards the - # raw alias verbatim, so without this the child mirror silently no-ops. + if name == "compact": # /compact aliases /compress; the compute-host control forwards the raw alias name = "compress" - if _session_uses_compute_host(session) and name in _MUTATES_WHILE_RUNNING: - return _compute_host_slash(sid, session, name, command)[1] - if name in _MUTATES_WHILE_RUNNING and session.get("running"): - return f"session busy — /interrupt the current turn before running /{name}" - mirror = _SLASH_MIRRORS.get(name) - if mirror is None: + if name in _MUTATES_WHILE_RUNNING: + if _session_uses_compute_host(session): + return _compute_host_slash(sid, session, name, command)[1] + if session.get("running"): + return f"session busy — /interrupt the current turn before running /{name}" + if (mirror := _SLASH_MIRRORS.get(name)) is None: return "" try: - return mirror(sid, session, agent, arg) + return mirror(sid, session, agent, arg) or "" except Exception as e: if name == "compress" and agent: from agent.conversation_compression import finalize_context_engine_compression_notification diff --git a/tui_gateway/methods_tools.py b/tui_gateway/methods_tools.py index 5657537088..b31804a3fc 100644 --- a/tui_gateway/methods_tools.py +++ b/tui_gateway/methods_tools.py @@ -15,40 +15,41 @@ _profile_scoped = _registry.profile_scoped # ─── Shared helpers ────────────────────────────────────────────────────────── - - -def _profile_scoped_rpc(fail_code: int, *, required=(), catch_resolve: bool = True, prefix: str = "", scoped: bool = True): - """Wrap a handler body with the optional ``profile`` HERMES_HOME scope. - - Order: ``required`` params checked first (4063 `` required``), then the - profile resolved (4064 when its dir is missing), then the body; body exceptions - become ``fail_code`` (message prefixed with ``prefix``). ``catch_resolve`` also maps - resolve-time exceptions to ``fail_code`` (cron/skills/catalog); mcp.servers.* let - them propagate to dispatch(). The override is always reset afterwards. - ``scoped=False`` (see ``_guarded``) ignores ``profile`` entirely. - """ +def _profile_scoped_rpc( + fail_code: int, *, required=(), catch_resolve: bool = True, prefix: str = "", + scoped: bool = True, live_session: bool = False, +): + """Wrap a handler body with the optional ``profile`` HERMES_HOME scope. Order: ``required`` + params (4063 `` required``) → ``live_session`` resolution via ``_sess`` (waits for the + agent build; body gets ``session`` as 3rd arg) → profile (4064 when its dir is missing) → body; + body exceptions become ``fail_code`` (``prefix`` + message). ``catch_resolve`` also maps + resolve-time exceptions to ``fail_code``; mcp.servers.* let them propagate to dispatch(). + ``scoped=False`` ignores ``profile``. The override is always reset afterwards.""" def deco(body): def handler(rid, params: dict) -> dict: for key, present in required: if not present(params.get(key)): return _err(rid, 4063, f"{key} required") - profile = str(params.get("profile") or "").strip() if scoped else "" + args = (rid, params) + if live_session: + session, err = _sess(params, rid) + if err: + return err + args = (rid, params, session) token = None - if profile: + if profile := _str_arg(params, "profile") if scoped else "": try: - from hermes_cli.profiles import get_profile_dir - from hermes_constants import set_hermes_home_override - profile_dir = get_profile_dir(profile) + profile_dir = _tools_mod("hermes_cli.profiles").get_profile_dir(profile) if not profile_dir or not profile_dir.is_dir(): return _err(rid, 4064, f"profile '{profile}' not found") - token = set_hermes_home_override(str(profile_dir)) + token = _tools_mod("hermes_constants").set_hermes_home_override(str(profile_dir)) except Exception as e: if not catch_resolve: raise return _err(rid, fail_code, str(e)) try: - return body(rid, params) + return body(*args) except Exception as e: return _err(rid, fail_code, f"{prefix}{e}") finally: @@ -58,53 +59,48 @@ def _profile_scoped_rpc(fail_code: int, *, required=(), catch_resolve: bool = Tr return deco -def _guarded(fail_code: int, prefix: str = ""): - """Handler body exceptions → ``_err(rid, fail_code, prefix + str(e))``.""" - return _profile_scoped_rpc(fail_code, prefix=prefix, scoped=False) +def _guarded(fail_code: int, prefix: str = "", *, live_session: bool = False): + """Body exceptions → ``_err(rid, fail_code, prefix + str(e))``; no profile scope. ``live_session`` + resolves the session first and calls ``body(rid, params, session)``.""" + return _profile_scoped_rpc(fail_code, prefix=prefix, scoped=False, live_session=live_session) -def _live_session_guarded(fail_code: int): - """Resolve the session via ``_sess`` (waits for the agent build) and call - ``body(rid, params, session)``; body exceptions → ``fail_code``.""" - - def deco(body): - def handler(rid, params: dict) -> dict: - session, err = _sess(params, rid) - if err: - return err - try: - return body(rid, params, session) - except Exception as e: - return _err(rid, fail_code, str(e)) - handler.__doc__ = body.__doc__ - return handler - return deco +def _rpc(name: str, fail_code: int, prefix: str = "", *, live_session: bool = False): + """``@method(name)`` + ``_guarded``.""" + return lambda body: method(name)(_guarded(fail_code, prefix, live_session=live_session)(body)) -def _stripped(v) -> bool: - return bool(str(v or "").strip()) +def _scoped_rpc(name: str, fail_code: int = 5024, **kw): + """``@method(name)`` + ``_profile_scoped_rpc`` (optional ``profile`` HERMES_HOME scope).""" + return lambda body: method(name)(_profile_scoped_rpc(fail_code, **kw)(body)) -def _nonempty(v) -> bool: - return not (v is None or str(v) == "") +def _str_arg(params: dict, key: str) -> str: + return str(params.get(key) or "").strip() +def _tools_mod(module: str): + """Deferred module import for one-liner bodies (startup budget: never import at load).""" + import importlib + return importlib.import_module(module) + + +_stripped = lambda v: bool(str(v or "").strip()) # noqa: E731 — required-param predicates +_nonempty = lambda v: not (v is None or str(v) == "") # noqa: E731 _NAME = (("name", _stripped),) _NAME_SESSION = (("name", _stripped), ("session_id", _stripped)) -def _mcp_server_scoped(body): - """mcp.servers.* contract: ``name`` required, profile scope, body errors → 5024.""" - return _profile_scoped_rpc(5024, required=_NAME, catch_resolve=False)(body) +def _mcp_rpc(name: str, required=_NAME): + """mcp.servers.* contract: profile scope, ``required`` params (default ``name``), body errors → 5024, + profile-resolve errors propagate to dispatch().""" + return _scoped_rpc(f"mcp.servers.{name}", required=required, catch_resolve=False) def _mcp_named_server(rid, params): """(name, servers, None) for a configured server, else (name, servers, 4064 error).""" - from hermes_cli.mcp_config import _get_mcp_servers - name = str(params.get("name") or "").strip() - servers = _get_mcp_servers() - err = None if name in servers else _err(rid, 4064, f"server '{name}' not found") - return name, servers, err + name, servers = _str_arg(params, "name"), _tools_mod("hermes_cli.mcp_config")._get_mcp_servers() + return name, servers, None if name in servers else _err(rid, 4064, f"server '{name}' not found") def _busy_error(rid, session, cmd: str): @@ -113,130 +109,161 @@ def _busy_error(rid, session, cmd: str): return None -def _session_key_or_err(rid, session): - """(session_key, None) or (None, 4001 error) for the /goal and /loop managers.""" +def _session_key_or_err(rid, session, module: str, label: str): + """(session_key, module, None) for the /goal and /loop managers, else (None, None, error): + 4001 without a session/key, 5030 when ``module`` fails to import.""" if not session: - return None, _err(rid, 4001, "no active session") - sid_key = session.get("session_key") or "" - if not sid_key: - return None, _err(rid, 4001, "no session key") - return sid_key, None + return None, None, _err(rid, 4001, "no active session") + if not (sid_key := session.get("session_key") or ""): + return None, None, _err(rid, 4001, "no session key") + try: + return sid_key, _tools_mod(module), None + except Exception as exc: + return None, None, _err(rid, 5030, f"{label} unavailable: {exc}") def _user_turn_indices(session): """(history, indices of user-originated turns) minus ephemeral scaffolding. Call under history_lock.""" - from agent.context_compressor import user_originated_turn_view + is_user = _tools_mod("agent.context_compressor").user_originated_turn_view history = _history_without_ephemeral_scaffolding(session.get("history", [])) - return history, [i for i, m in enumerate(history) if user_originated_turn_view(m) is not None] + return history, [i for i, m in enumerate(history) if is_user(m) is not None] + + +def _rewind_prelude(rid, session, cmd: str, empty_msg: str): + """Under history_lock: re-check busy, then (history, user_indices, None) or (None, None, error).""" + if busy := _busy_error(rid, session, cmd): + return None, None, busy + history, user_indices = _user_turn_indices(session) + if not user_indices: + return None, None, _err(rid, 4018, empty_msg) + return history, user_indices, None + + +def _rewind_or_err(rid, session, keep: int, value_err: tuple, fail_prefix: str, **kw): + """``_rewind_active_session_history`` → (result, None); ValueError → ``value_err`` (code, prefix), + other exceptions → 5008 ``fail_prefix`` + message.""" + try: + return _rewind_active_session_history(session, keep, **kw), None + except ValueError as exc: + return None, _err(rid, value_err[0], f"{value_err[1]}{exc}") + except Exception as exc: + return None, _err(rid, 5008, f"{fail_prefix}{exc}") def _clip(text: str, n: int = 120) -> str: return text[:n] + ("…" if len(text) > n else "") +def _exec_out(rid, output: str) -> dict: + """command.dispatch display-only result.""" + return _ok(rid, {"type": "exec", "output": output}) + + def _capture_run_kwargs(timeout: int) -> dict: - """subprocess.run kwargs shared by cli.exec / shell.exec / quick commands: captured - text, UTF-8 + lossy decode (non-UTF-8 child output must not crash the gateway thread - on locale-mismatched Windows), no stdin, no console flash under the desktop parent.""" - from hermes_cli._subprocess_compat import windows_hide_flags + """Shared captured-text subprocess.run kwargs: UTF-8 + lossy decode (non-UTF-8 child output must + not crash the gateway thread on Windows), no stdin, no console flash under the desktop parent.""" return dict( - capture_output=True, - text=True, - encoding="utf-8", - errors="replace", - timeout=timeout, - stdin=subprocess.DEVNULL, - creationflags=windows_hide_flags()) + capture_output=True, text=True, encoding="utf-8", errors="replace", timeout=timeout, + stdin=subprocess.DEVNULL, creationflags=_tools_mod("hermes_cli._subprocess_compat").windows_hide_flags()) + + +def _captured_exec(rid, cmd, timeout: int, *, on_result, timeout_err: tuple, fail_code: int, **kw) -> dict: + """Run ``cmd`` captured (see ``_capture_run_kwargs``) and hand the CompletedProcess to + ``on_result``; TimeoutExpired → ``timeout_err`` (code, message), other errors → ``fail_code``.""" + try: + return on_result(subprocess.run(cmd, cwd=os.getcwd(), **kw, **_capture_run_kwargs(timeout))) + except subprocess.TimeoutExpired: + return _err(rid, *timeout_err) + except Exception as e: + return _err(rid, fail_code, str(e)) + + +def _joined_output(r) -> str: + """stdout + stderr of a CompletedProcess, non-empty parts only, newline-joined and stripped.""" + return "\n".join(p for p in (r.stdout or "", r.stderr or "") if p).strip() def _toolset_rows(params: dict, *, with_tools: bool) -> list[dict]: - from toolsets import get_all_toolsets, get_toolset_info + toolsets = _tools_mod("toolsets") session = _sessions.get(params.get("session_id", "")) - enabled = ( - set(getattr(session["agent"], "enabled_toolsets", []) or []) if session else set(_load_enabled_toolsets() or []) - ) + enabled = set((getattr(session["agent"], "enabled_toolsets", []) if session else _load_enabled_toolsets()) or []) items = [] - for name in sorted(get_all_toolsets().keys()): - info = get_toolset_info(name) - if not info: - continue - row = { - "name": name, - "description": info["description"], - "tool_count": info["tool_count"], - "enabled": name in enabled if enabled else True} - if with_tools: - row["tools"] = info["resolved_tools"] - items.append(row) + for name in sorted(toolsets.get_all_toolsets().keys()): + if info := toolsets.get_toolset_info(name): + row = { + "name": name, "description": info["description"], "tool_count": info["tool_count"], + "enabled": name in enabled if enabled else True} + if with_tools: + row["tools"] = info["resolved_tools"] + items.append(row) return items # ─── System / process ──────────────────────────────────────────────────────── - - @method("system.battery") def _(rid, params: dict) -> dict: """Host battery for the status bar. Always resolves; ``available: false`` = no battery or read failed.""" try: - from agent.battery import battery_category, read_battery - batt = read_battery() - return _ok( - rid, - { - "available": batt.available, - "percent": batt.percent, - "plugged": batt.plugged, - "category": battery_category(batt)}) + battery = _tools_mod("agent.battery") + batt = battery.read_battery() + return _ok(rid, { + "available": batt.available, "percent": batt.percent, "plugged": batt.plugged, + "category": battery.battery_category(batt)}) except Exception: return _ok(rid, {"available": False, "percent": None, "plugged": None, "category": "dim"}) -@method("process.stop") -@_guarded(5010) -def _(rid, params: dict) -> dict: - from tools.process_registry import process_registry - return _ok(rid, {"killed": process_registry.kill_all()}) +# One-expression handlers: name → (fail_code, payload builder(params)). +_SIMPLE_RPCS = { + # Session-scoped view of the background process registry (desktop status stack). + "process.stop": (5010, lambda params: {"killed": _tools_mod("tools.process_registry").process_registry.kill_all()}), + # Re-read ``~/.hermes/.env`` (CLI ``/reload`` parity); built agents keep their pool, ``/new`` resolves fresh. + "reload.env": (5015, lambda params: {"updated": int(_tools_mod("hermes_cli.config").reload_env())}), + "plugins.list": (5032, lambda params: {"plugins": [ + {"name": n, "version": getattr(i, "version", "?"), "enabled": getattr(i, "enabled", True)} + for n, i in _tools_mod("hermes_cli.plugins").get_plugin_manager()._plugins.items()]}), + "tools.list": (5031, lambda params: {"toolsets": _toolset_rows(params, with_tools=True)}), + "toolsets.list": (5032, lambda params: {"toolsets": _toolset_rows(params, with_tools=False)}), + "agents.list": (5033, lambda params: {"processes": [ + {"session_id": p["session_id"], "command": p["command"][:80], "status": p["status"], "uptime": p["uptime_seconds"]} + for p in _tools_mod("tools.process_registry").process_registry.list_sessions()]}), +} +for _name, (_code, _build) in _SIMPLE_RPCS.items(): + # Look the builder up at call time: bind_module rebinds the table's lambdas onto server globals. + _rpc(_name, _code)(lambda rid, params, _n=_name: _ok(rid, _SIMPLE_RPCS[_n][1](params))) +del _name, _code, _build +_rpc("process.list", 5010, live_session=True)( + lambda rid, params, session: _ok(rid, {"processes": _session_processes(session)})) -@method("process.list") -@_live_session_guarded(5010) -def _(rid, params: dict, session) -> dict: - """Session-scoped view of the background process registry (desktop status stack).""" - return _ok(rid, {"processes": _session_processes(session)}) - - -@method("process.kill") -@_live_session_guarded(5010) +@_rpc("process.kill", live_session=True, fail_code=5010) def _(rid, params: dict, session) -> dict: """Kill ONE background process, scoped to the caller's session (unlike process.stop's kill_all).""" proc_id = str(params.get("process_id") or "") if not proc_id: return _err(rid, 4012, "process_id required") - from tools.process_registry import process_registry - proc = process_registry.get(proc_id) + registry = _tools_mod("tools.process_registry").process_registry + proc = registry.get(proc_id) if proc is None or str(getattr(proc, "session_key", "") or "") != str(session.get("session_key") or ""): return _err(rid, 4044, f"no such process: {proc_id}") - return _ok(rid, process_registry.kill_process(proc_id)) + return _ok(rid, registry.kill_process(proc_id)) def _mcp_reload_confirm_required() -> bool: """``approvals.mcp_reload_confirm`` from disk config; True (safe) on any failure.""" try: - from hermes_cli.config import load_config - cfg = load_config() + cfg = _tools_mod("hermes_cli.config").load_config() approvals = cfg.get("approvals") if isinstance(cfg, dict) else None return bool(approvals.get("mcp_reload_confirm", True)) if isinstance(approvals, dict) else True except Exception: return True -@method("reload.mcp") -@_guarded(5015) +@_rpc("reload.mcp", 5015) def _(rid, params: dict) -> dict: session = _sessions.get(params.get("session_id", "")) - # /reload-mcp invalidates the prompt cache: without confirm=true, honour - # ``approvals.mcp_reload_confirm`` (default true) by returning confirm_required; - # Ink prints ``message`` and re-invokes with confirm=true (or flips the config). + # Prompt-cache invalidation gate: without confirm=true honour ``approvals.mcp_reload_confirm`` + # (Ink prints ``message`` and re-invokes with confirm=true, or flips the config). if not bool(params.get("confirm", False)) and _mcp_reload_confirm_required(): message = ( "⚠️ /reload-mcp invalidates the prompt cache (next message re-sends full input tokens). " @@ -250,41 +277,35 @@ def _(rid, params: dict) -> dict: except Exception as exc: return _err(rid, 5019, f"compute-host reload_mcp failed: {exc}") return _ok(rid, {"status": "reloaded", "turn_isolation": True, "host_ack": ack}) - from tools.mcp_tool import shutdown_mcp_servers, discover_mcp_tools, reprobe_tool_availability + mcp_tool = _tools_mod("tools.mcp_tool") + global _mcp_reload_gen, _mcp_reload_loaded_rev + # Revision the CALLER wants loaded; empty on legacy clients / manual /reload-mcp + # (generation-only coalescing). + req_rev = str(params.get("rev") or "") def _refresh_session_agent() -> None: - """Rebuild THIS session's cached tool snapshot from the live registry and push - session.info (the agent never re-reads the registry itself; mirrors - gateway/run.py::_execute_mcp_reload). Runs under _mcp_reload_lock so a - concurrent reload can't tear the registry down mid-refresh.""" + """Rebuild THIS session's cached tool snapshot + push session.info (the agent never + re-reads the registry). Runs under _mcp_reload_lock so a concurrent reload can't + tear the registry down mid-refresh.""" if not session: return agent = session["agent"] - try: - from tools.mcp_tool import refresh_agent_mcp_tools - - # enabled_override re-resolves toolsets so a server enabled in config this session is picked up. - refresh_agent_mcp_tools(agent, enabled_override=_load_enabled_toolsets(), quiet_mode=True) + try: # enabled_override re-resolves toolsets so a server enabled in config this session is picked up + mcp_tool.refresh_agent_mcp_tools(agent, enabled_override=_load_enabled_toolsets(), quiet_mode=True) except Exception as _exc: logger.warning("Failed to refresh cached agent tools after /reload-mcp: %s", _exc) _emit("session.info", params.get("session_id", ""), _session_info(agent, session)) - global _mcp_reload_gen, _mcp_reload_loaded_rev - - # Revision the CALLER wants loaded (the mcp_rev its poll observed); empty on - # legacy clients / manual /reload-mcp, which coalesce on generation alone. - req_rev = str(params.get("rev") or "") def _do_full_reload() -> None: - """shutdown+discover+refresh under the lock, then mark a completed generation. - The lock spans the refresh too, else a second reload could tear the registry - down mid-rebuild. Config can change WHILE discover connects: re-hash after - discovery and repeat until stable so the marked generation matches what loaded.""" + """shutdown+discover+refresh under the lock, then mark a completed generation. Config + can change WHILE discover connects: re-hash and repeat until stable so the marked + generation matches what loaded.""" global _mcp_reload_gen, _mcp_reload_loaded_rev loaded = _compute_mcp_rev() for _ in range(_MCP_RELOAD_MAX_PASSES): - shutdown_mcp_servers() - reprobe_tool_availability() - discover_mcp_tools() + mcp_tool.shutdown_mcp_servers() + mcp_tool.reprobe_tool_availability() + mcp_tool.discover_mcp_tools() after = _compute_mcp_rev() if after == loaded: break @@ -293,11 +314,9 @@ def _(rid, params: dict) -> dict: _mcp_reload_loaded_rev = loaded _mcp_reload_gen += 1 - # LEADER (won the non-blocking acquire) runs the full reload. FOLLOWER snapshots - # the generation, waits, then — still holding the lock — coalesces only if a - # reload COMPLETED meanwhile (generation advanced ⇒ leader didn't throw) AND it - # loaded the requested revision; otherwise it re-runs the full reload so a - # failed/stale leader never leaves a follower acking an unloaded revision. + # LEADER (won the non-blocking acquire) runs the full reload. FOLLOWER waits, then — still + # holding the lock — coalesces only if a reload COMPLETED meanwhile (generation advanced + # ⇒ leader didn't throw) AND it loaded the requested revision; otherwise it re-runs. if _mcp_reload_lock.acquire(blocking=False): try: _do_full_reload() @@ -306,125 +325,108 @@ def _(rid, params: dict) -> dict: return _finish_reload(rid, params, coalesced=False) gen_before = _mcp_reload_gen with _mcp_reload_lock: - leader_completed = _mcp_reload_gen > gen_before - rev_satisfied = not req_rev or req_rev == _mcp_reload_loaded_rev - if leader_completed and rev_satisfied: - _refresh_session_agent() - coalesced = True - else: - _do_full_reload() - coalesced = False + coalesced = _mcp_reload_gen > gen_before and (not req_rev or req_rev == _mcp_reload_loaded_rev) + _refresh_session_agent() if coalesced else _do_full_reload() return _finish_reload(rid, params, coalesced=coalesced) -@method("reload.env") -@_guarded(5015) -def _(rid, params: dict) -> dict: - """Re-read ``~/.hermes/.env`` (classic CLI ``/reload`` parity). Already-built agents - keep their credential pool / provider routing; ``/new`` resolves fresh.""" - from hermes_cli.config import reload_env - return _ok(rid, {"updated": int(reload_env())}) - - # ─── Command catalog / dispatch ────────────────────────────────────────────── +class _Catalog: + """Accumulator for commands.catalog: ``pairs`` (every [key, desc]), ``canon`` (lowercase + key/alias → canonical key), ``commands`` (key → desktop meta) and ordered categories.""" + + def __init__(self) -> None: + self.pairs: list[list[str]] = [] + self.canon: dict[str, str] = {} + self.commands: dict[str, dict[str, str | None]] = {} + self.cat_map: dict[str, list[list[str]]] = {} # insertion order = category order + + def add(self, key: str, desc: str, cat: str) -> None: + self.canon[key.lower()] = key + self.pairs.append([key, desc]) + self.cat_map.setdefault(cat, []).append([key, desc]) -@method("commands.catalog") -@_guarded(5020) -def _(rid, params: dict) -> dict: - """Registry-backed slash metadata for the TUI — categorized, no aliases.""" - from hermes_cli.commands import COMMAND_REGISTRY, SUBCOMMANDS, _build_description, command_desktop_meta - all_pairs: list[list[str]] = [] - canon: dict[str, str] = {} - commands: dict[str, dict[str, str | None]] = {} - cat_map: dict[str, list[list[str]]] = {} - cat_order: list[str] = [] - - def bucket(cat: str) -> list[list[str]]: - if cat not in cat_map: - cat_map[cat] = [] - cat_order.append(cat) - return cat_map[cat] - - def add(key: str, desc: str, rows: list[list[str]]) -> None: - canon[key.lower()] = key - all_pairs.append([key, desc]) - rows.append([key, desc]) - for cmd in COMMAND_REGISTRY: - meta = command_desktop_meta(cmd) - commands[f"/{cmd.name}"] = dict(meta) - for alias in cmd.aliases: - commands[f"/{alias}"] = dict(meta) +def _catalog_registry(cat: _Catalog) -> None: + commands = _tools_mod("hermes_cli.commands") + for cmd in commands.COMMAND_REGISTRY: + meta = commands.command_desktop_meta(cmd) + cat.commands.update({f"/{key}": dict(meta) for key in (cmd.name, *cmd.aliases)}) if cmd.name in _TUI_HIDDEN or cmd.gateway_only: continue - c = f"/{cmd.name}" - add(c, _build_description(cmd), bucket(cmd.category)) + cat.add(f"/{cmd.name}", commands._build_description(cmd), cmd.category) for a in cmd.aliases: - canon[f"/{a}".lower()] = c - for name, desc, cat in _TUI_EXTRA: + cat.canon[f"/{a}".lower()] = f"/{cmd.name}" + for name, desc, category in _TUI_EXTRA: # Registry command/alias wins over a colliding TUI extra (e.g. /compact, /sessions). - if name.lower() not in canon: - add(name, desc, bucket(cat)) + if name.lower() not in cat.canon: + cat.add(name, desc, category) + + +def _catalog_quick_commands(cat: _Catalog) -> None: + qcmds = _load_cfg().get("quick_commands", {}) or {} + if not (isinstance(qcmds, dict) and qcmds): + return + cat.cat_map.setdefault("User commands", []) # category exists even when every entry is malformed + for qname, qc in sorted(qcmds.items()): + if not isinstance(qc, dict): + continue + qtype = qc.get("type", "") + default_desc = {"exec": f"exec: {qc.get('command', '')}", "alias": f"alias → {qc.get('target', '')}"} + desc = str(qc.get("description") or default_desc.get(qtype, qtype or "quick command")) + cat.add(f"/{qname}", _clip(desc), "User commands") + + +def _catalog_plugin_commands(cat: _Catalog) -> None: + plugin_cmds = _tools_mod("hermes_cli.plugins").get_plugin_commands() or {} + if plugin_cmds: + cat.cat_map.setdefault("Plugin commands", []) + for pname, info in sorted(plugin_cmds.items()): + key = f"/{pname}" + if not isinstance(info, dict) or key.lower() in cat.canon: + continue + cat.add(key, _clip(str(info.get("description") or "Plugin command")), "Plugin commands") + mode = info.get("argument_mode") + if mode not in {"options", "text", "mixed"}: + mode = "text" if str(info.get("args_hint") or "").strip() else None + cat.commands[key] = {"argument_mode": mode, "desktop": None} + + +def _catalog_skills(cat: _Catalog, skills: dict[str, dict]) -> None: + """Append skill pairs and fill ``skills`` = ``{key: {usage, origin}}`` (every consumer ranks by them).""" + usage, origin_of = _skill_usage_lookup() + for k, info in sorted(_tools_mod("agent.skill_commands").scan_skill_commands().items()): + cat.pairs.append([k, _clip(str(info.get("description", "Skill")))]) + name = str(info.get("name") or k.lstrip("/")) + skills[k] = {"usage": usage(name), "origin": origin_of(name)} + + +@_rpc("commands.catalog", 5020) +def _(rid, params: dict) -> dict: + """Registry-backed slash metadata, categorized, no aliases. Discovery failures land in ``warning`` + (skills' message wins, then quick commands', then plugins').""" + cat = _Catalog() + _catalog_registry(cat) warning = "" try: - qcmds = _load_cfg().get("quick_commands", {}) or {} - if isinstance(qcmds, dict) and qcmds: - rows = bucket("User commands") - for qname, qc in sorted(qcmds.items()): - if not isinstance(qc, dict): - continue - qtype = qc.get("type", "") - default_desc = { - "exec": f"exec: {qc.get('command', '')}", - "alias": f"alias → {qc.get('target', '')}", - }.get(qtype, qtype or "quick command") - add(f"/{qname}", _clip(str(qc.get("description") or default_desc)), rows) + _catalog_quick_commands(cat) except Exception as e: warning = f"quick_commands discovery unavailable: {e}" try: - from hermes_cli.plugins import get_plugin_commands - plugin_cmds = get_plugin_commands() or {} - if plugin_cmds: - rows = bucket("Plugin commands") - for pname, info in sorted(plugin_cmds.items()): - if not isinstance(info, dict): - continue - key = f"/{pname}" - if key.lower() in canon: - continue - add(key, _clip(str(info.get("description") or "Plugin command")), rows) - hint = str(info.get("args_hint") or "").strip() - mode = info.get("argument_mode") - if mode not in {"options", "text", "mixed"}: - mode = "text" if hint else None - commands[key] = {"argument_mode": mode, "desktop": None} + _catalog_plugin_commands(cat) except Exception as e: - if not warning: - warning = f"plugin command discovery unavailable: {e}" - skill_count = 0 + warning = warning or f"plugin command discovery unavailable: {e}" skills: dict[str, dict] = {} try: - from agent.skill_commands import scan_skill_commands - - # Usage + origin ride along (not a second RPC): every catalog consumer also ranks it. - usage, origin_of = _skill_usage_lookup() - for k, info in sorted(scan_skill_commands().items()): - all_pairs.append([k, _clip(str(info.get("description", "Skill")))]) - name = str(info.get("name") or k.lstrip("/")) - skills[k] = {"usage": usage(name), "origin": origin_of(name)} - skill_count += 1 + _catalog_skills(cat, skills) except Exception as e: warning = f"skill discovery unavailable: {e}" - payload = { - "pairs": all_pairs, - "sub": {k: v[:] for k, v in SUBCOMMANDS.items()}, - "canon": canon, - "commands": commands, - "categories": [{"name": cat, "pairs": cat_map[cat]} for cat in cat_order], - "skills": skills, - "skill_count": skill_count, - "warning": warning} - return _ok(rid, payload) + return _ok(rid, { + "pairs": cat.pairs, "sub": {k: v[:] for k, v in _tools_mod("hermes_cli.commands").SUBCOMMANDS.items()}, + "canon": cat.canon, + "commands": cat.commands, + "categories": [{"name": c, "pairs": rows} for c, rows in cat.cat_map.items()], + "skills": skills, "skill_count": len(skills), "warning": warning}) @method("cli.exec") @@ -436,27 +438,19 @@ def _(rid, params: dict) -> dict: hint = _cli_exec_blocked(argv) if hint: return _ok(rid, {"blocked": True, "hint": hint, "code": -1, "output": ""}) - try: - r = subprocess.run( - [sys.executable, "-m", "hermes_cli.main", *argv], - cwd=os.getcwd(), - # Can drive the agent → needs provider credentials; tier-1 secrets still stripped. - env=hermes_subprocess_env(inherit_credentials=True), - **_capture_run_kwargs(min(int(params.get("timeout", 240)), 600))) - parts = [r.stdout or "", r.stderr or ""] - out = "\n".join(p for p in parts if p).strip() or "(no output)" - return _ok(rid, {"blocked": False, "code": r.returncode, "output": out[:48_000]}) - except subprocess.TimeoutExpired: - return _err(rid, 5016, "cli.exec: timeout") - except Exception as e: - return _err(rid, 5017, str(e)) + + # Can drive the agent → needs provider credentials; tier-1 secrets still stripped. + return _captured_exec( + rid, [sys.executable, "-m", "hermes_cli.main", *argv], min(int(params.get("timeout", 240)), 600), + on_result=lambda r: _ok(rid, { + "blocked": False, "code": r.returncode, "output": (_joined_output(r) or "(no output)")[:48_000]}), + timeout_err=(5016, "cli.exec: timeout"), fail_code=5017, + env=hermes_subprocess_env(inherit_credentials=True)) -@method("command.resolve") -@_guarded(5012) +@_rpc("command.resolve", 5012) def _(rid, params: dict) -> dict: - from hermes_cli.commands import resolve_command - r = resolve_command(params.get("name", "")) + r = _tools_mod("hermes_cli.commands").resolve_command(params.get("name", "")) if r: return _ok(rid, {"canonical": r.name, "description": r.description, "category": r.category}) return _err(rid, 4011, f"unknown command: {params.get('name')}") @@ -467,61 +461,50 @@ def _(rid, params: dict) -> dict: def _dispatch_quick(rid, params, session, name, arg): - qcmds = _load_cfg().get("quick_commands", {}) - if name not in qcmds: + qc = _load_cfg().get("quick_commands", {}).get(name) + if qc is None: return None - qc = qcmds[name] if qc.get("type") == "exec": # Sanitized env: the TUI server process holds every API key in os.environ. - from tools.environments.local import build_subprocess_env - sanitized_env = build_subprocess_env() - r = subprocess.run(qc.get("command", ""), shell=True, env=sanitized_env, **_capture_run_kwargs(30)) - output = ((r.stdout or "") + ("\n" if r.stdout and r.stderr else "") + (r.stderr or "")).strip()[:4000] - if output: - from agent.redact import redact_sensitive_text - output = redact_sensitive_text(output) + env = _tools_mod("tools.environments.local").build_subprocess_env() + r = subprocess.run(qc.get("command", ""), shell=True, env=env, **_capture_run_kwargs(30)) + output = _joined_output(r)[:4000] + output = _tools_mod("agent.redact").redact_sensitive_text(output) if output else output if r.returncode != 0: return _err(rid, 4018, output or f"quick command failed with exit code {r.returncode}") - return _ok(rid, {"type": "exec", "output": output}) - if qc.get("type") == "alias": - return _ok(rid, {"type": "alias", "target": qc.get("target", "")}) - return None + return _exec_out(rid, output) + return _ok(rid, {"type": "alias", "target": qc.get("target", "")}) if qc.get("type") == "alias" else None def _plugin_command_handler(name: str): try: - from hermes_cli.plugins import get_plugin_command_handler - return get_plugin_command_handler(name) + return _tools_mod("hermes_cli.plugins").get_plugin_command_handler(name) except Exception: return None def _run_plugin_command(handler, arg: str) -> str: - from hermes_cli.plugins import resolve_plugin_command_result - return str(resolve_plugin_command_result(handler(arg)) or "") + return str(_tools_mod("hermes_cli.plugins").resolve_plugin_command_result(handler(arg)) or "") def _is_profile_skill_command(session: dict, base: str) -> bool: - """True when ``/base`` is a skill command of the session's profile. HERMES_HOME is bound - to that profile so get_skill_commands() sees its skills.external_dirs: dispatch() runs on - the pool and nothing upstream binds the override. False on any failure.""" + """True when ``/base`` is a skill command of the session's profile (HERMES_HOME bound to it so + get_skill_commands() sees its skills.external_dirs; nothing upstream binds it). False on failure.""" try: - from agent.skill_commands import get_skill_commands - from hermes_constants import reset_hermes_home_override, set_hermes_home_override + hc = _tools_mod("hermes_constants") profile_home = session.get("profile_home") - token = set_hermes_home_override(profile_home) if profile_home else None + token = hc.set_hermes_home_override(profile_home) if profile_home else None try: - return f"/{base}" in get_skill_commands() + return f"/{base}" in _tools_mod("agent.skill_commands").get_skill_commands() finally: if token is not None: - reset_hermes_home_override(token) + hc.reset_hermes_home_override(token) except Exception: return False def _dispatch_plugin(rid, params, session, name, arg): - handler = _plugin_command_handler(name) - if handler: + if handler := _plugin_command_handler(name): with contextlib.suppress(Exception): return _ok(rid, {"type": "plugin", "output": _run_plugin_command(handler, arg)}) return None @@ -530,9 +513,9 @@ def _dispatch_plugin(rid, params, session, name, arg): def _bundle_key_for(name: str): """Skill-bundle key for ``name`` when it is NOT a registry command; None otherwise / on failure.""" try: - from agent.skill_bundles import resolve_bundle_command_key - from hermes_cli.commands import resolve_command - return resolve_bundle_command_key(name) if resolve_command(name) is None else None + if _tools_mod("hermes_cli.commands").resolve_command(name) is None: + return _tools_mod("agent.skill_bundles").resolve_bundle_command_key(name) + return None except Exception: return None @@ -541,39 +524,33 @@ def _dispatch_bundle(rid, params, session, name, arg): bundle_key = _bundle_key_for(name) if bundle_key is None: return None - from agent.skill_bundles import build_bundle_invocation_message, get_skill_bundles + bundles = _tools_mod("agent.skill_bundles") try: - bundle_result = build_bundle_invocation_message( - bundle_key, - arg, - task_id=session.get("session_key", "") if session else "", + bundle_result = bundles.build_bundle_invocation_message( + bundle_key, arg, task_id=session.get("session_key", "") if session else "", platform=_resolve_session_platform()) except Exception as exc: return _err(rid, 4018, f"bundle dispatch failed: {exc}") if not bundle_result: return _err(rid, 4018, f"failed to load bundle: {bundle_key}") msg, loaded_names, missing = bundle_result - bundle_name = get_skill_bundles().get(bundle_key, {}).get("name", bundle_key.lstrip("/")) + bundle_name = bundles.get_skill_bundles().get(bundle_key, {}).get("name", bundle_key.lstrip("/")) notice = f"⚡ Loading bundle: {bundle_name} ({len(loaded_names)} skills)" - if missing: - notice += f"\nSkipped missing skills: {', '.join(missing)}" + notice += f"\nSkipped missing skills: {', '.join(missing)}" if missing else "" # UIs render `display`, never `message`: the expanded body is model-facing scaffolding. return _ok(rid, {"type": "send", "message": msg, "notice": notice, "display": _skill_scaffold_projection(msg)}) def _dispatch_skill(rid, params, session, name, arg): - try: - from agent.skill_commands import scan_skill_commands, build_skill_invocation_message - cmds = scan_skill_commands() - key = f"/{name}" + with contextlib.suppress(Exception): + sc = _tools_mod("agent.skill_commands") + cmds, key = sc.scan_skill_commands(), f"/{name}" if key in cmds: - msg = build_skill_invocation_message(key, arg, task_id=session.get("session_key", "") if session else "") - if msg: - # UIs render `display`, never `message`. - display = _skill_scaffold_projection(msg) - return _ok(rid, {"type": "skill", "message": msg, "name": cmds[key].get("name", name), "display": display}) - except Exception: - pass + msg = sc.build_skill_invocation_message(key, arg, task_id=session.get("session_key", "") if session else "") + if msg: # UIs render `display`, never `message`. + return _ok(rid, { + "type": "skill", "message": msg, "name": cmds[key].get("name", name), + "display": _skill_scaffold_projection(msg)}) return None @@ -582,68 +559,51 @@ def _dispatch_skill(rid, params, session, name, arg): def _cmd_queue(rid, params, session, name, arg): - if not arg: - return _err(rid, 4004, "usage: /queue ") - return _ok(rid, {"type": "send", "message": arg}) + return _ok(rid, {"type": "send", "message": arg}) if arg else _err(rid, 4004, "usage: /queue ") -def _cmd_learn(rid, params, session, name, arg): - # Submitted as a normal turn; the live agent gathers sources and authors the skill via skill_manage. - from agent.learn_prompt import build_learn_prompt - return _ok(rid, {"type": "send", "message": build_learn_prompt(arg)}) +def _prompt_builtin(module: str, fn: str, kw: str = ""): + """/learn, /plan, /init: submit ``module.fn(arg)`` as a normal turn (the live agent does the work).""" + + def cmd(rid, params, session, name, arg): + build = getattr(_tools_mod(module), fn) + return _ok(rid, {"type": "send", "message": build(**{kw: arg}) if kw else build(arg)}) + return cmd -def _cmd_plan(rid, params, session, name, arg): - # Normal turn (as /learn); the agent saves the plan under .hermes/plans/ via write_file. - from agent.plan_prompt import build_plan_prompt - return _ok(rid, {"type": "send", "message": build_plan_prompt(arg)}) - - -def _cmd_init(rid, params, session, name, arg): - # Generate-or-update AGENTS.md as a normal turn (as /learn). - from hermes_cli.init_command import build_init_prompt_for_cwd - return _ok(rid, {"type": "send", "message": build_init_prompt_for_cwd(extra=arg)}) +_cmd_learn = _prompt_builtin("agent.learn_prompt", "build_learn_prompt") +_cmd_plan = _prompt_builtin("agent.plan_prompt", "build_plan_prompt") +_cmd_init = _prompt_builtin("hermes_cli.init_command", "build_init_prompt_for_cwd", kw="extra") def _cmd_moa(rid, params, session, name, arg): - # One prompt through the default MoA preset, then restore the prior model. Whole-session - # switching goes through the model picker (MoA presets = virtual "Mixture of Agents" provider). + # One prompt through the default MoA preset, then restore the prior model (whole-session + # switching goes through the model picker). try: - from hermes_cli.moa_config import moa_usage, normalize_moa_config + moa = _tools_mod("hermes_cli.moa_config") if not arg: - return _err(rid, 4004, moa_usage()) + return _err(rid, 4004, moa.moa_usage()) if not session: return _err(rid, 4001, "no active session") - sid = params.get("session_id", "") - preset = normalize_moa_config(_load_cfg().get("moa") or {})["default_preset"] + preset = moa.normalize_moa_config(_load_cfg().get("moa") or {})["default_preset"] # Record the live identity for post-turn restore, then swap the agent's client in # place: session["model_override"] alone never switches an already-built agent. agent = session.get("agent") session["moa_one_shot_restore"] = { - "override": session.get("model_override"), - "model": getattr(agent, "model", None) if agent else None, - "provider": getattr(agent, "provider", None) if agent else None} + "override": session.get("model_override"), "model": getattr(agent, "model", None), + "provider": getattr(agent, "provider", None)} if agent is not None: - try: + try: # persist_override=False: turn-scoped, never persist the MoA provider to config.yaml _apply_model_switch( - sid, - session, - f"{preset} --provider moa", - confirm_expensive_model=False, - pin_session_override=True, - persist_override=False, # turn-scoped: never persist the MoA provider to config.yaml - ) - except Exception as exc: + params.get("session_id", ""), session, f"{preset} --provider moa", + confirm_expensive_model=False, pin_session_override=True, persist_override=False) + except Exception: session.pop("moa_one_shot_restore", None) - return _err(rid, 5030, f"moa unavailable: {exc}") - else: - # Lazy/fresh session: the override is consumed by the first build. + raise + else: # lazy/fresh session: the override is consumed by the first build session["model_override"] = { - "provider": "moa", - "model": preset, - "base_url": "moa://local", - "api_key": "moa-virtual-provider", - "api_mode": "chat_completions"} + "provider": "moa", "model": preset, "base_url": "moa://local", + "api_key": "moa-virtual-provider", "api_mode": "chat_completions"} notice = f"MoA one-shot queued with preset {preset}; previous model will be restored after this turn." return _ok(rid, {"type": "send", "notice": notice, "message": arg}) except Exception as exc: @@ -652,23 +612,21 @@ def _cmd_moa(rid, params, session, name, arg): def _cmd_focus(rid, params, session, name, arg): # Display-only; routed through the config.set branch Ink uses so both surfaces share one state machine. - from hermes_cli.focus_view import format_focus_status, format_focus_toggle_message, resolve_focus_arg + fv = _tools_mod("hermes_cli.focus_view") display = _load_cfg().get("display") display = display if isinstance(display, dict) else {} - cur = bool(display.get("focus_view", False)) - action, target = resolve_focus_arg(arg, cur) + action, target = fv.resolve_focus_arg(arg, cur := bool(display.get("focus_view", False))) if action == "usage": return _err(rid, 4004, "usage: /focus [on|off|status]") if action == "status": saved = display.get("focus_saved_tool_progress") or _load_tool_progress_mode() - return _ok(rid, {"type": "exec", "output": format_focus_status(cur, saved)}) + return _exec_out(rid, fv.format_focus_status(cur, saved)) res = _methods["config.set"]( - rid, {"key": "focus", "value": "on" if target else "off", "session_id": params.get("session_id", "")} - ) + rid, {"key": "focus", "value": "on" if target else "off", "session_id": params.get("session_id", "")}) if "error" in res: return res - output = format_focus_toggle_message(bool(target), (res.get("result") or {}).get("tool_progress") or "all") - return _ok(rid, {"type": "exec", "output": output}) + tool_progress = (res.get("result") or {}).get("tool_progress") or "all" + return _exec_out(rid, fv.format_focus_toggle_message(bool(target), tool_progress)) def _cmd_retry(rid, params, session, name, arg): @@ -676,28 +634,25 @@ def _cmd_retry(rid, params, session, name, arg): return _err(rid, 4001, "no active session to retry") if busy := _busy_error(rid, session, "retry"): return busy - from agent.context_compressor import history_before_user_originated_turn, retryable_user_text + cc = _tools_mod("agent.context_compressor") with session["history_lock"]: if busy := _busy_error(rid, session, "retry"): return busy if session.get("attached_images"): return _err(rid, 4018, "retry cannot safely reconstruct or combine attached media") - history, user_indices = _user_turn_indices(session) - if not user_indices: - return _err(rid, 4018, "no previous user message to retry") - _prefix, live_view = history_before_user_originated_turn(history, user_indices[-1]) + history, user_indices, err = _rewind_prelude(rid, session, "retry", "no previous user message to retry") + if err: + return err + _prefix, live_view = cc.history_before_user_originated_turn(history, user_indices[-1]) try: - content = retryable_user_text(live_view.get("content")) + content = cc.retryable_user_text(live_view.get("content")) except ValueError as exc: return _err(rid, 4018, str(exc)) - try: - _active, durable_live_view, _rewound_count = _rewind_active_session_history( - session, len(user_indices) - 1, require_retryable=True) - except ValueError as exc: - return _err(rid, 4018, str(exc)) - except Exception as exc: - return _err(rid, 5008, f"retry: failed to persist history: {exc}") - content = retryable_user_text(durable_live_view.get("content")) + rewound, err = _rewind_or_err( + rid, session, len(user_indices) - 1, (4018, ""), "retry: failed to persist history: ", require_retryable=True) + if err: + return err + content = cc.retryable_user_text(rewound[1].get("content")) return _ok(rid, {"type": "send", "message": content}) @@ -706,52 +661,42 @@ def _cmd_steer(rid, params, session, name, arg): return _err(rid, 4004, "usage: /steer ") agent = session.get("agent") if session else None if agent and hasattr(agent, "steer"): - try: + with contextlib.suppress(Exception): if agent.steer(arg): shown = f"{arg[:80]}{'...' if len(arg) > 80 else ''}" - return _ok(rid, {"type": "exec", "output": f"⏩ Steer queued — arrives after the next tool call: {shown}"}) - except Exception: - pass - # No active run: treat as next-turn message. - return _ok(rid, {"type": "send", "message": arg}) + return _exec_out(rid, f"⏩ Steer queued — arrives after the next tool call: {shown}") + return _ok(rid, {"type": "send", "message": arg}) # no active run: next-turn message def _cmd_goal(rid, params, session, name, arg): - sid_key, err = _session_key_or_err(rid, session) + sid_key, goals, err = _session_key_or_err(rid, session, "hermes_cli.goals", "goals") if err: return err - try: - from hermes_cli.goals import GoalManager - except Exception as exc: - return _err(rid, 5030, f"goals unavailable: {exc}") try: max_turns = int((_load_cfg().get("goals") or {}).get("max_turns", 20) or 20) except Exception: max_turns = 20 - mgr = GoalManager(session_id=sid_key, default_max_turns=max_turns) + mgr = goals.GoalManager(session_id=sid_key, default_max_turns=max_turns) lower = arg.strip().lower() - if not arg.strip() or lower == "status": - return _ok(rid, {"type": "exec", "output": mgr.status_line()}) + if not lower or lower == "status": + return _exec_out(rid, mgr.status_line()) if lower == "pause": state = mgr.pause(reason="user-paused") - out = "No goal set." if state is None else f"⏸ Goal paused: {state.goal}" - return _ok(rid, {"type": "exec", "output": out}) + return _exec_out(rid, "No goal set." if state is None else f"⏸ Goal paused: {state.goal}") if lower == "resume": state = mgr.resume() if state is None: - return _ok(rid, {"type": "exec", "output": "No goal to resume."}) - # Resume must restart work: `exec` is display-only, so return a `send` with the - # continuation prompt; `display` keeps model-facing scaffolding out of the transcript. - prompt = mgr.next_continuation_prompt() - if not prompt: - return _ok(rid, {"type": "exec", "output": f"▶ Goal resumed: {state.goal}"}) + return _exec_out(rid, "No goal to resume.") + # Resume must restart work: `exec` is display-only, so return a `send`; `display` + # keeps model-facing scaffolding out of the transcript. + if not (prompt := mgr.next_continuation_prompt()): + return _exec_out(rid, f"▶ Goal resumed: {state.goal}") notice = f"▶ Goal resumed: {state.goal}\nContinuing now — taking the next step." return _ok(rid, {"type": "send", "notice": notice, "message": prompt, "display": "/goal resume"}) if lower in {"clear", "stop", "done"}: had = mgr.has_goal() mgr.clear() - return _ok(rid, {"type": "exec", "output": "✓ Goal cleared." if had else "No active goal."}) - + return _exec_out(rid, "✓ Goal cleared." if had else "No active goal.") # Remaining text = new goal. Client renders `notice`, submits `message`; the post-turn judge takes over. try: state = mgr.set(arg) @@ -765,85 +710,68 @@ def _cmd_goal(rid, params, session, name, arg): def _cmd_loop(rid, params, session, name, arg): - # Recurring in-session wakeups; the notification poller fires due ones while the session is idle. - sid_key, err = _session_key_or_err(rid, session) + sid_key, loops, err = _session_key_or_err(rid, session, "hermes_cli.loops", "loops") if err: return err - try: - from hermes_cli.loops import LoopManager, dispatch_loop_command - except Exception as exc: - return _err(rid, 5030, f"loops unavailable: {exc}") - result = dispatch_loop_command(LoopManager(session_id=sid_key), arg) + result = loops.dispatch_loop_command(loops.LoopManager(session_id=sid_key), arg) output = result.get("output") or "" if result.get("created"): with contextlib.suppress(Exception): - from hermes_cli.loops import goal_blocks_loop_tick - if goal_blocks_loop_tick(sid_key): - output += ( - "\nNote: an active /goal is driving this session — loop " - "wakeups defer until the goal finishes, pauses, or parks.") - return _ok(rid, {"type": "exec", "output": output}) + if loops.goal_blocks_loop_tick(sid_key): + output += ("\nNote: an active /goal is driving this session — loop " + "wakeups defer until the goal finishes, pauses, or parks.") + return _exec_out(rid, output) def _cmd_undo(rid, params, session, name, arg): - # /undo [N]: back up N user turns, soft-delete truncated rows on disk, prefill the composer. if not session: return _err(rid, 4001, "no active session to undo") if busy := _busy_error(rid, session, "undo"): return busy - session_key = session.get("session_key", "") - if not session_key: + if not (session_key := session.get("session_key", "")): return _err(rid, 4001, "no session key for undo") - n = 1 arg_str = (arg or "").strip() - if arg_str: - try: - n = int(arg_str.split()[0]) - except (ValueError, IndexError): - return _err(rid, 4004, f"undo: invalid count {arg_str!r} — use /undo or /undo N") - n = max(n, 1) - from agent.message_content import flatten_message_text + try: + n = max(int(arg_str.split()[0]), 1) if arg_str else 1 + except (ValueError, IndexError): + return _err(rid, 4004, f"undo: invalid count {arg_str!r} — use /undo or /undo N") with session["history_lock"]: - if busy := _busy_error(rid, session, "undo"): - return busy - _history, user_indices = _user_turn_indices(session) - if not user_indices: - return _err(rid, 4018, "no user messages to undo") + _history, user_indices, err = _rewind_prelude(rid, session, "undo", "no user messages to undo") + if err: + return err turns_undone = min(n, len(user_indices)) - try: - active, live_view, rewound_count = _rewind_active_session_history(session, len(user_indices) - turns_undone) - except ValueError as exc: - return _err(rid, 4004, f"undo: {exc}") - except Exception as exc: - return _err(rid, 5008, f"undo: {exc}") - target_text = flatten_message_text(live_view.get("content")) + rewound, err = _rewind_or_err(rid, session, len(user_indices) - turns_undone, (4004, "undo: "), "undo: ") + if err: + return err + active, live_view, rewound_count = rewound + target_text = _tools_mod("agent.message_content").flatten_message_text(live_view.get("content")) # Notify memory providers (same hook /branch fires) with rewound=True so cached per-turn state invalidates. agent = session.get("agent") if agent is not None: mm = getattr(agent, "_memory_manager", None) - if mm is not None: + for step in ( + lambda: mm is not None and mm.on_session_switch( + session_key, parent_session_id="", reset=False, rewound=True), + lambda: hasattr(agent, "_invalidate_system_prompt") and agent._invalidate_system_prompt(), + lambda: hasattr(agent, "_last_flushed_db_idx") and setattr(agent, "_last_flushed_db_idx", len(active)), + ): with contextlib.suppress(Exception): - mm.on_session_switch(session_key, parent_session_id="", reset=False, rewound=True) - if hasattr(agent, "_invalidate_system_prompt"): - with contextlib.suppress(Exception): - agent._invalidate_system_prompt() - if hasattr(agent, "_last_flushed_db_idx"): - with contextlib.suppress(Exception): - agent._last_flushed_db_idx = len(active) + step() turn_word = "turn" if turns_undone == 1 else "turns" notice = f"↶ Undid {turns_undone} {turn_word} ({rewound_count} message(s)). Edit and resubmit, or send a new message." return _ok(rid, {"type": "prefill", "message": target_text, "notice": notice}) +def _is_snapshot_restore(arg: str) -> bool: + return (arg.split(maxsplit=1)[0].lower() if arg else "") in {"restore", "rewind"} + + def _cmd_snapshot(rid, params, session, name, arg): - subcommand = arg.split(maxsplit=1)[0].lower() if arg else "" - if subcommand not in {"restore", "rewind"}: + if not _is_snapshot_restore(arg): return None - output = ( - "/snapshot restore is blocked in the TUI because it changes config/state on disk " - "while the live agent has cached settings. Run it in the classic CLI, then restart the TUI." - ) - return _ok(rid, {"type": "exec", "output": output}) + return _exec_out( + rid, "/snapshot restore is blocked in the TUI because it changes config/state on disk " + "while the live agent has cached settings. Run it in the classic CLI, then restart the TUI.") def _cmd_compress(rid, params, session, name, arg): @@ -851,23 +779,23 @@ def _cmd_compress(rid, params, session, name, arg): return _err(rid, 4001, "no active session to compress") if busy := _busy_error(rid, session, "compress"): return busy - from agent.conversation_compression import finalize_context_engine_compression_notification sid = params.get("session_id", "") if _session_uses_compute_host(session): status, text = _compute_host_slash(sid, session, "compress", f"/{name}" + (f" {arg}" if arg else "")) if status in {"failed", "rejected"}: return _err(rid, 5019 if status == "failed" else 4009, text) - payload = {"type": "exec", "status": "pending", "output": text} if status == "pending" else {"type": "exec", "output": text} - return _ok(rid, payload) + if status == "pending": + return _ok(rid, {"type": "exec", "status": "pending", "output": text}) + return _exec_out(rid, text) try: output = _compress_live_with_feedback(sid, session, session["agent"], arg, snapshot_kwargs=True) - return _ok(rid, {"type": "exec", "output": output}) + return _exec_out(rid, output) except Exception as exc: - finalize_context_engine_compression_notification(session["agent"], committed=False) + _tools_mod("agent.conversation_compression").finalize_context_engine_compression_notification( + session["agent"], committed=False) return _err(rid, 5009, f"compress failed: {exc}") -# name → built-in handler (values are rebound onto server globals by bind_module). _SLASH_BUILTINS = { "queue": _cmd_queue, "q": _cmd_queue, "learn": _cmd_learn, "plan": _cmd_plan, "init": _cmd_init, "moa": _cmd_moa, "focus": _cmd_focus, "retry": _cmd_retry, "steer": _cmd_steer, "goal": _cmd_goal, @@ -877,20 +805,15 @@ _SLASH_BUILTINS = { @method("command.dispatch") def _(rid, params: dict) -> dict: - name, arg = params.get("name", "").lstrip("/"), params.get("arg", "") - name = _resolve_name(name) + name, arg = _resolve_name(params.get("name", "").lstrip("/")), params.get("arg", "") session = _sessions.get(params.get("session_id", "")) # Stage order is load-bearing: quick > plugin > bundle > skill > built-in. - for stage in (_dispatch_quick, _dispatch_plugin, _dispatch_bundle, _dispatch_skill): + stages = (_dispatch_quick, _dispatch_plugin, _dispatch_bundle, _dispatch_skill, _SLASH_BUILTINS.get(name)) + for stage in filter(None, stages): res = stage(rid, params, session, name, arg) if res is not None: return res - builtin = _SLASH_BUILTINS.get(name) - if builtin is not None: - res = builtin(rid, params, session, name, arg) - if res is not None: - return res return _err(rid, 4018, f"not a quick/plugin/bundle/skill command: {name}") @@ -902,7 +825,6 @@ def _(rid, params: dict) -> dict: cmd = params.get("command", "").strip() if not cmd: return _err(rid, 4004, "empty command") - # Skill/bundle and _PENDING_INPUT_COMMANDS must NOT reach the slash worker. Plugin # commands also bypass it but return normal slash.exec output (TUI keeps the pager path). parts = cmd.lstrip("/").split(maxsplit=1) @@ -912,29 +834,24 @@ def _(rid, params: dict) -> dict: live_output = _live_slash_command_output(sid, session, base, arg) if live_output is not None: return _ok(rid, {"output": live_output or "(no output)"}) - if base in _PENDING_INPUT_COMMANDS: - # Route straight to command.dispatch: some clients fail the error-then-retry fallback ("empty command"). - return _methods["command.dispatch"](rid, {"name": base, "arg": arg, "session_id": sid}) - if base in _WORKER_BLOCKED_COMMANDS: - subcommand = arg.split(maxsplit=1)[0].lower() if arg else "" - if subcommand in {"restore", "rewind"}: - return _err(rid, 4018, "snapshot restore mutates live config/state; use command.dispatch for /snapshot restore") - bundle_key = _bundle_key_for(base) - if bundle_key is not None: - return _methods["command.dispatch"](rid, {"name": bundle_key.lstrip("/"), "arg": arg, "session_id": sid}) + if base in _WORKER_BLOCKED_COMMANDS and _is_snapshot_restore(arg): + return _err(rid, 4018, "snapshot restore mutates live config/state; use command.dispatch for /snapshot restore") + # Pending-input built-ins route straight to command.dispatch (some clients fail the + # error-then-retry fallback); bundles go the same way under their resolved key. + target = base if base in _PENDING_INPUT_COMMANDS else _bundle_key_for(base) + if target is not None: + return _methods["command.dispatch"](rid, {"name": target.lstrip("/"), "arg": arg, "session_id": sid}) if _is_profile_skill_command(session, base): return _err(rid, 4018, f"skill command: use command.dispatch for /{base}") - plugin_handler = _plugin_command_handler(base) if base else None - if plugin_handler: + if plugin_handler := _plugin_command_handler(base) if base else None: try: return _ok(rid, {"output": _run_plugin_command(plugin_handler, arg) or "(no output)"}) except Exception as e: return _ok(rid, {"output": f"Plugin command error: {e}"}) worker = session.get("slash_worker") if not worker: - # slash.exec runs on the RPC pool: two concurrent commands could both see - # slash_worker=None and each fork a full MCP-fleet worker (the _attach_worker - # loser leaks). Serialize first-use spawn per session. + # slash.exec runs on the RPC pool: two concurrent commands could both see slash_worker=None + # and each fork a full MCP-fleet worker (the loser leaks). Serialize first-use spawn. with _sessions_lock: spawn_lock = session.setdefault("_slash_spawn_lock", threading.Lock()) with spawn_lock: @@ -942,17 +859,14 @@ def _(rid, params: dict) -> dict: if not worker: try: worker = _SlashWorker( - session["session_key"], - getattr(session.get("agent"), "model", _resolve_model()), + session["session_key"], getattr(session.get("agent"), "model", _resolve_model()), profile_home=session.get("profile_home")) _attach_worker(sid, session, worker) except Exception as e: return _err(rid, 5030, f"slash worker start failed: {e}") try: - output = worker.run(cmd) - warning = _mirror_slash_side_effects(sid, session, cmd) - payload = {"output": output or "(no output)"} - if warning: + payload = {"output": worker.run(cmd) or "(no output)"} + if warning := _mirror_slash_side_effects(sid, session, cmd): payload["warning"] = warning return _ok(rid, payload) except Exception as e: @@ -963,40 +877,30 @@ def _(rid, params: dict) -> dict: # ─── Insights / rollback / browser / config ────────────────────────────────── - - -@method("insights.get") +@_rpc("insights.get", 5017) def _(rid, params: dict) -> dict: days = params.get("days", 30) - db = _get_db() - if db is None: + if (db := _get_db()) is None: return _db_unavailable_error(rid, code=5017) - try: - cutoff = time.time() - days * 86400 - rows = [s for s in db.list_sessions_rich(limit=500, compact_rows=True) if (s.get("started_at") or 0) >= cutoff] - return _ok(rid, {"days": days, "sessions": len(rows), "messages": sum(s.get("message_count", 0) for s in rows)}) - except Exception as e: - return _err(rid, 5017, str(e)) + cutoff = time.time() - days * 86400 + rows = [s for s in db.list_sessions_rich(limit=500, compact_rows=True) if (s.get("started_at") or 0) >= cutoff] + return _ok(rid, {"days": days, "sessions": len(rows), "messages": sum(s.get("message_count", 0) for s in rows)}) -@method("rollback.list") -@_live_session_guarded(5020) +@_rpc("rollback.list", live_session=True, fail_code=5020) def _(rid, params: dict, session) -> dict: def go(mgr, cwd): if not mgr.enabled: return _ok(rid, {"enabled": False, "checkpoints": []}) - rows = [ - {"hash": c.get("hash", ""), "timestamp": c.get("timestamp", ""), "message": c.get("message", "")} - for c in mgr.list_checkpoints(cwd)] + keys = ("hash", "timestamp", "message") + rows = [{k: c.get(k, "") for k in keys} for c in mgr.list_checkpoints(cwd)] return _ok(rid, {"enabled": True, "checkpoints": rows}) return _with_checkpoints(session, go) -@method("rollback.restore") -@_live_session_guarded(5021) +@_rpc("rollback.restore", live_session=True, fail_code=5021) def _(rid, params: dict, session) -> dict: - target = params.get("hash", "") - file_path = params.get("file_path", "") + target, file_path = params.get("hash", ""), params.get("file_path", "") if not target: return _err(rid, 4014, "hash required") # Full-history rollback mutates session history → rejected mid-turn (prompt.submit @@ -1005,15 +909,14 @@ def _(rid, params: dict, session) -> dict: return _err(rid, 4009, "session busy — /interrupt the current turn before full rollback.restore") def go(mgr, cwd): - resolved = _resolve_checkpoint_hash(mgr, cwd, target) - result = mgr.restore(cwd, resolved, file_path=file_path or None) + result = mgr.restore(cwd, _resolve_checkpoint_hash(mgr, cwd, target), file_path=file_path or None) if result.get("success") and not file_path: removed = 0 with session["history_lock"]: _history, user_indices = _user_turn_indices(session) if user_indices: try: - _active, _live_view, removed = _rewind_active_session_history(session, len(user_indices) - 1) + removed = _rewind_active_session_history(session, len(user_indices) - 1)[2] except Exception as exc: raise RuntimeError(f"checkpoint restored, but session history rewind failed: {exc}") from exc result["history_removed"] = removed @@ -1021,17 +924,14 @@ def _(rid, params: dict, session) -> dict: return _ok(rid, _with_checkpoints(session, go)) -@method("rollback.diff") -@_live_session_guarded(5022) +@_rpc("rollback.diff", live_session=True, fail_code=5022) def _(rid, params: dict, session) -> dict: - target = params.get("hash", "") - if not target: + if not (target := params.get("hash", "")): return _err(rid, 4014, "hash required") r = _with_checkpoints(session, lambda mgr, cwd: mgr.diff(cwd, _resolve_checkpoint_hash(mgr, cwd, target))) raw = r.get("diff", "")[:4000] payload = {"stat": r.get("stat", ""), "diff": raw} - rendered = render_diff(raw, session.get("cols", 80)) - if rendered: + if rendered := render_diff(raw, session.get("cols", 80)): payload["rendered"] = rendered return _ok(rid, payload) @@ -1049,73 +949,44 @@ def _(rid, params: dict) -> dict: return _err(rid, 4015, f"unknown action: {action}") -@method("plugins.list") -@_guarded(5032) -def _(rid, params: dict) -> dict: - from hermes_cli.plugins import get_plugin_manager - rows = [ - {"name": n, "version": getattr(i, "version", "?"), "enabled": getattr(i, "enabled", True)} - for n, i in get_plugin_manager()._plugins.items()] - return _ok(rid, {"plugins": rows}) - - -@method("config.show") -@_guarded(5030) +@_rpc("config.show", 5030) def _(rid, params: dict) -> dict: cfg = _load_cfg() - model = _resolve_model() - from agent.secret_scope import get_secret - api_key = get_secret("HERMES_API_KEY", "") or cfg.get("api_key", "") + api_key = _tools_mod("agent.secret_scope").get_secret("HERMES_API_KEY", "") or cfg.get("api_key", "") masked = f"****{api_key[-4:]}" if len(api_key) > 4 else "(not set)" base_url = os.environ.get("HERMES_BASE_URL", "") or cfg.get("base_url", "") - agent_rows = [ - ["Max Turns", str(_cfg_max_turns(cfg, 500))], - ["Toolsets", ", ".join(cfg.get("enabled_toolsets", [])) or "all"], - ["Verbose", str(cfg.get("verbose", False))]] sections = [ - {"title": "Model", "rows": [["Model", model], ["Base URL", base_url or "(default)"], ["API Key", masked]]}, - {"title": "Agent", "rows": agent_rows}, + {"title": "Model", "rows": [ + ["Model", _resolve_model()], ["Base URL", base_url or "(default)"], ["API Key", masked]]}, + {"title": "Agent", "rows": [ + ["Max Turns", str(_cfg_max_turns(cfg, 500))], + ["Toolsets", ", ".join(cfg.get("enabled_toolsets", [])) or "all"], + ["Verbose", str(cfg.get("verbose", False))]]}, {"title": "Environment", "rows": [["Working Dir", os.getcwd()], ["Config File", str(_hermes_home / "config.yaml")]]}, ] return _ok(rid, {"sections": sections}) # ─── Tools / toolsets / agents ─────────────────────────────────────────────── - - -@method("tools.list") -@_guarded(5031) +@_rpc("tools.show", 5034) def _(rid, params: dict) -> dict: - return _ok(rid, {"toolsets": _toolset_rows(params, with_tools=True)}) - - -@method("toolsets.list") -@_guarded(5032) -def _(rid, params: dict) -> dict: - return _ok(rid, {"toolsets": _toolset_rows(params, with_tools=False)}) - - -@method("tools.show") -@_guarded(5034) -def _(rid, params: dict) -> dict: - from model_tools import get_toolset_for_tool, get_tool_definitions + mt = _tools_mod("model_tools") session = _sessions.get(params.get("session_id", "")) enabled = getattr(session["agent"], "enabled_toolsets", None) if session else _load_enabled_toolsets() # Pre-assembly list: /tools must also show tools deferred behind the tool_search bridge (as the CLI). - tools = get_tool_definitions(enabled_toolsets=enabled, quiet_mode=True, skip_tool_search_assembly=True) + tools = mt.get_tool_definitions(enabled_toolsets=enabled, quiet_mode=True, skip_tool_search_assembly=True) sections = {} for tool in sorted(tools, key=lambda t: t["function"]["name"]): name = tool["function"]["name"] desc = str(tool["function"].get("description", "") or "").split("\n")[0] if ". " in desc: desc = desc[: desc.index(". ") + 1] - sections.setdefault(get_toolset_for_tool(name) or "unknown", []).append({"name": name, "description": desc}) - sections_out = [{"name": name, "tools": rows} for name, rows in sorted(sections.items())] + sections.setdefault(mt.get_toolset_for_tool(name) or "unknown", []).append({"name": name, "description": desc}) + sections_out = [{"name": n, "tools": rows} for n, rows in sorted(sections.items())] return _ok(rid, {"sections": sections_out, "total": len(tools)}) -@method("tools.configure") -@_guarded(5035) +@_rpc("tools.configure", 5035) def _(rid, params: dict) -> dict: action = str(params.get("action", "") or "").strip().lower() targets = [str(name).strip() for name in params.get("names", []) or [] if str(name).strip()] @@ -1123,178 +994,132 @@ def _(rid, params: dict) -> dict: return _err(rid, 4017, f"unknown tools action: {action}") if not targets: return _err(rid, 4018, "names required") - from hermes_cli.config import load_config, save_config - from hermes_cli.tools_config import ( - CONFIGURABLE_TOOLSETS, _apply_mcp_change, _apply_toolset_change, _get_platform_tools, - _get_plugin_toolset_keys) - cfg = load_config() - valid_toolsets = {ts_key for ts_key, _, _ in CONFIGURABLE_TOOLSETS} | _get_plugin_toolset_keys() - toolset_targets = [name for name in targets if ":" not in name] + hc, tc = _tools_mod("hermes_cli.config"), _tools_mod("hermes_cli.tools_config") + cfg = hc.load_config() + valid_toolsets = {ts_key for ts_key, _, _ in tc.CONFIGURABLE_TOOLSETS} | tc._get_plugin_toolset_keys() mcp_targets = [name for name in targets if ":" in name] - unknown = [name for name in toolset_targets if name not in valid_toolsets] - toolset_targets = [name for name in toolset_targets if name in valid_toolsets] + unknown = [name for name in targets if ":" not in name and name not in valid_toolsets] + toolset_targets = [name for name in targets if ":" not in name and name in valid_toolsets] if toolset_targets: - _apply_toolset_change(cfg, "cli", toolset_targets, action) - missing_servers = _apply_mcp_change(cfg, mcp_targets, action) if mcp_targets else set() - save_config(cfg) + tc._apply_toolset_change(cfg, "cli", toolset_targets, action) + missing_servers = tc._apply_mcp_change(cfg, mcp_targets, action) if mcp_targets else set() + hc.save_config(cfg) sid = params.get("session_id", "") session = _sessions.get(sid) info = _reset_session_agent(sid, session) if session else None - enabled = sorted(_get_platform_tools(load_config(), "cli", include_default_mcp_servers=False)) + enabled = sorted(tc._get_platform_tools(hc.load_config(), "cli", include_default_mcp_servers=False)) changed = [ - name - for name in targets - if name not in unknown and (":" not in name or name.split(":", 1)[0] not in missing_servers) - ] + name for name in targets + if name not in unknown and (":" not in name or name.split(":", 1)[0] not in missing_servers)] return _ok(rid, { - "changed": changed, - "enabled_toolsets": enabled, - "info": info, - "missing_servers": sorted(missing_servers), - "reset": bool(session), - "unknown": unknown}) - - -@method("agents.list") -@_guarded(5033) -def _(rid, params: dict) -> dict: - from tools.process_registry import process_registry - rows = [ - {"session_id": p["session_id"], "command": p["command"][:80], "status": p["status"], "uptime": p["uptime_seconds"]} - for p in process_registry.list_sessions()] - return _ok(rid, {"processes": rows}) + "changed": changed, "enabled_toolsets": enabled, "info": info, + "missing_servers": sorted(missing_servers), "reset": bool(session), "unknown": unknown}) # ─── Cron / learning / skills ──────────────────────────────────────────────── - - -@method("cron.manage") -@_profile_scoped_rpc(5023) +@_scoped_rpc("cron.manage", 5023) def _(rid, params: dict) -> dict: - """cronjob() keys off HERMES_HOME, so the optional ``profile`` scope reaches a - per-profile cron store even when that profile runs its own gateway.""" - from tools.cronjob_tools import cronjob + """cronjob() keys off HERMES_HOME, so ``profile`` reaches a per-profile cron store.""" + cronjob = _tools_mod("tools.cronjob_tools").cronjob action, jid = params.get("action", "list"), params.get("name", "") if action == "list": # Paused jobs are excluded by default (reads as deletion in a toggle UI) — forward the flag. - result = json.loads( - cronjob(action="list", include_disabled=is_truthy_value(params.get("include_disabled", False))) - ) - # ``scoped`` proves the profile scope was honored: new clients treat every job as that - # profile's; older gateways omit it and clients keep the safe [bot:] filter. - profile = str(params.get("profile") or "").strip() - if profile: + include_disabled = is_truthy_value(params.get("include_disabled", False)) + result = json.loads(cronjob(action="list", include_disabled=include_disabled)) + # ``scoped`` proves the profile scope was honored; older gateways omit it and clients + # keep the safe [bot:] filter. + if profile := _str_arg(params, "profile"): result["scoped"] = profile return _ok(rid, result) if action == "add": # Optional repeat / continuity / deliver ('bot-chat[:name]'): None keeps each cronjob() default. raw = cronjob( - action="create", - name=jid, - schedule=params.get("schedule", ""), - prompt=params.get("prompt", ""), + action="create", name=jid, schedule=params.get("schedule", ""), prompt=params.get("prompt", ""), repeat=int(params["repeat"]) if str(params.get("repeat", "")).strip().isdigit() else None, continuity=is_truthy_value(params.get("continuity")) if params.get("continuity") is not None else None, - deliver=str(params.get("deliver") or "").strip() or None) + deliver=_str_arg(params, "deliver") or None) return _ok(rid, json.loads(raw)) if action in {"remove", "pause", "resume"}: return _ok(rid, json.loads(cronjob(action=action, job_id=jid))) return _err(rid, 4016, f"unknown cron action: {action}") -@method("learning.frames") -@_guarded(5000, "learning.frames failed: ") +@_rpc("learning.frames", 5000, "learning.frames failed: ") def _(rid, params: dict) -> dict: - """Pre-render the ``/journey`` timeline: ``frames`` (reveal 0→1) plus legend/summary/ - bucket metadata so Ink walks the tree locally. Shares its renderer with ``hermes journey``.""" + """Pre-render the ``/journey`` timeline (frames + legend/summary metadata) so Ink walks it locally.""" try: - cols = int(params.get("cols", 80) or 80) - rows = int(params.get("rows", 24) or 24) - frames = int(params.get("frames", 48) or 48) + cols, rows, frames = ( + int(params.get(k, d) or d) for k, d in (("cols", 80), ("rows", 24), ("frames", 48))) except (TypeError, ValueError): cols, rows, frames = 80, 24, 48 - from agent.learning_graph import build_learning_graph - from agent.learning_graph_render import render_frames - return _ok(rid, render_frames(build_learning_graph(), cols=max(20, cols), rows=max(10, rows), frames=frames)) + graph = _tools_mod("agent.learning_graph").build_learning_graph() + render_frames = _tools_mod("agent.learning_graph_render").render_frames + return _ok(rid, render_frames(graph, cols=max(20, cols), rows=max(10, rows), frames=frames)) def _learning_mutation(fn_name: str, arg_keys: tuple): """learning.* body: ``agent.learning_mutations.(*str(params[k]) for k in arg_keys)``.""" def body(rid, params: dict) -> dict: - import agent.learning_mutations as mutations - return _ok(rid, getattr(mutations, fn_name)(*(str(params.get(k, "")) for k in arg_keys))) + fn = getattr(_tools_mod("agent.learning_mutations"), fn_name) + return _ok(rid, fn(*(str(params.get(k, "")) for k in arg_keys))) return body # detail → node content for an edit prefill; delete → skills archived (restorable), memories # removed; edit → rewrite a node's content (SKILL.md or memory chunk). -for _rpc, _fn, _keys in ( +for _name, _fn, _keys in ( ("detail", "node_detail", ("id",)), ("delete", "delete_node", ("id",)), ("edit", "edit_node", ("id", "content")), ): - method(f"learning.{_rpc}")(_guarded(5000, f"learning.{_rpc} failed: ")(_learning_mutation(_fn, _keys))) -del _rpc, _fn, _keys - - -def _skills_list(rid, params, query): - from hermes_cli.banner import get_available_skills - return _ok(rid, {"skills": get_available_skills()}) + _rpc(f"learning.{_name}", 5000, f"learning.{_name} failed: ")(_learning_mutation(_fn, _keys)) +del _name, _fn, _keys def _skills_search(rid, params, query): - from tools.skills_hub import GitHubAuth, create_source_router, unified_search - raw = unified_search(query, create_source_router(GitHubAuth()), source_filter="all", limit=20) or [] + hub = _tools_mod("tools.skills_hub") + raw = hub.unified_search(query, hub.create_source_router(hub.GitHubAuth()), source_filter="all", limit=20) or [] return _ok(rid, {"results": [{"name": r.name, "description": r.description} for r in raw]}) def _skills_install(rid, params, query): - from hermes_cli.skills_hub import do_install - - class _Q: - def print(self, *a, **k): - pass - do_install(query, skip_confirm=True, console=_Q()) + quiet = _tools_mod("types").SimpleNamespace(print=lambda *a, **k: None) + _tools_mod("hermes_cli.skills_hub").do_install(query, skip_confirm=True, console=quiet) return _ok(rid, {"installed": True, "name": query}) def _skills_browse(rid, params, query): - from hermes_cli.skills_hub import browse_skills pg = int(params.get("page", 0) or 0) or (int(query) if query.isdigit() else 1) - return _ok(rid, browse_skills(page=pg, page_size=int(params.get("page_size", 20)))) + browse = _tools_mod("hermes_cli.skills_hub").browse_skills + return _ok(rid, browse(page=pg, page_size=int(params.get("page_size", 20)))) -def _skills_inspect(rid, params, query): - from hermes_cli.skills_hub import inspect_skill - return _ok(rid, {"info": inspect_skill(query) or {}}) +_SKILLS_ACTIONS = { + "list": lambda rid, params, query: _ok(rid, {"skills": _tools_mod("hermes_cli.banner").get_available_skills()}), + "search": _skills_search, "install": _skills_install, "browse": _skills_browse, + "inspect": lambda rid, params, query: _ok( + rid, {"info": _tools_mod("hermes_cli.skills_hub").inspect_skill(query) or {}})} -@method("skills.manage") -@_profile_scoped_rpc(5024) +def _run_action(rid, params: dict, table: dict, label: str, *extra) -> dict: + """Dispatch ``params['action']`` (default ``list``) through ``table``; unknown → 4017.""" + action = params.get("action", "list") + handler = table.get(action) + if handler is None: + return _err(rid, 4017, f"unknown {label} action: {action}") + return handler(rid, params, *extra) + + +@_scoped_rpc("skills.manage") def _(rid, params: dict) -> dict: """list/install use the scoped profile's skills dir; search/browse/inspect hit the shared hub.""" - action, query = params.get("action", "list"), params.get("query", "") - handler = { - "list": _skills_list, - "search": _skills_search, - "install": _skills_install, - "browse": _skills_browse, - "inspect": _skills_inspect, - }.get(action) - if handler is None: - return _err(rid, 4017, f"unknown skills action: {action}") - return handler(rid, params, query) + return _run_action(rid, params, _SKILLS_ACTIONS, "skills", params.get("query", "")) -@method("skills.reload") -@_guarded(5025) +@_rpc("skills.reload", 5025) def _(rid, params: dict) -> dict: - from agent.skill_commands import reload_skills - result = reload_skills() - added = result.get("added") or [] - removed = result.get("removed") or [] - lines = ["Reloading skills..."] - if not added and not removed: - lines.append("No new skills detected.") + result = _tools_mod("agent.skill_commands").reload_skills() + added, removed = result.get("added") or [], result.get("removed") or [] + lines = ["Reloading skills..."] + ([] if added or removed else ["No new skills detected."]) for label, items in (("Added skills:", added), ("Removed skills:", removed)): if items: lines.append(label) @@ -1306,14 +1131,10 @@ def _(rid, params: dict) -> dict: # ─── MCP catalog + per-profile server lifecycle (mcp.servers.*) ───────────── # Gateway mirrors of the dashboard REST surface (hermes_cli/web_routers/mcp.py) so a # desktop plugin can manage MCP servers for ANY profile. Persistence: hermes_cli/mcp_config.py. - - -@method("mcp.catalog") -@_profile_scoped_rpc(5024) +@_scoped_rpc("mcp.catalog") def _(rid, params: dict) -> dict: - """``{servers: [{name, description, installed, enabled, requires: [env keys], transport}]}`` - — the `hermes mcp` menu with per-profile state, so UIs know which entries need setup.""" - from hermes_cli import mcp_catalog + """``{servers: [{name, description, installed, enabled, requires: [env keys], transport}]}`` per profile.""" + mcp_catalog = _tools_mod("hermes_cli.mcp_catalog") out = [] for entry in mcp_catalog.list_catalog(): try: @@ -1321,104 +1142,83 @@ def _(rid, params: dict) -> dict: except Exception: requires = [] transport = getattr(entry, "transport", None) # TransportSpec → its kind string - out.append( - { - "name": entry.name, - "description": getattr(entry, "description", "") or "", - "installed": bool(mcp_catalog.is_installed(entry.name)), - "enabled": bool(mcp_catalog.is_enabled(entry.name)), - "requires": requires, - "transport": str(getattr(transport, "kind", "") or transport or "stdio")}) + out.append({ + "name": entry.name, "description": getattr(entry, "description", "") or "", + "installed": bool(mcp_catalog.is_installed(entry.name)), + "enabled": bool(mcp_catalog.is_enabled(entry.name)), "requires": requires, + "transport": str(getattr(transport, "kind", "") or transport or "stdio")}) return _ok(rid, {"servers": out}) -@method("mcp.servers.list") -@_profile_scoped_rpc(5024, catch_resolve=False) +@_mcp_rpc("list", required=()) def _(rid, params: dict) -> dict: - """``{servers: [{name, transport, url, command, args, env (key names only), - auth, oauth_tokens_present, enabled, tools}]}`` for the scoped profile.""" - from hermes_cli.mcp_config import _get_mcp_servers - servers = _get_mcp_servers() + """``{servers: [{name, transport, url, command, args, env (key names), auth, oauth_tokens_present, + enabled, tools}]}``""" + servers = _tools_mod("hermes_cli.mcp_config")._get_mcp_servers() return _ok(rid, {"servers": [_mcp_summarize_server(name, cfg) for name, cfg in sorted(servers.items())]}) -@method("mcp.servers.add") -@_mcp_server_scoped +@_mcp_rpc("add") def _(rid, params: dict) -> dict: - """Add ``name`` with EITHER ``preset`` (catalog id) or ``config`` (url/command/args/env/ - headers/auth/tools). ``bearer_token`` goes to the profile's .env; only the - ``Authorization`` header template is persisted. Duplicate names → 4090.""" - from hermes_cli.mcp_config import _apply_mcp_preset, _get_mcp_servers, _save_bearer_auth_token, _save_mcp_server - name = str(params.get("name") or "").strip() - if name in _get_mcp_servers(): + """Add ``name`` from ``preset`` (catalog id) and/or ``config`` (url/command/args/env/headers/auth/ + tools); ``bearer_token`` goes to the profile's .env (only the header template persists). Dup → 4090.""" + mc = _tools_mod("hermes_cli.mcp_config") + name, preset = _str_arg(params, "name"), _str_arg(params, "preset") + if name in mc._get_mcp_servers(): return _err(rid, 4090, f"server '{name}' already exists") - preset = str(params.get("preset") or "").strip() raw_cfg = params.get("config") server_config: dict = dict(raw_cfg) if isinstance(raw_cfg, dict) else {} if preset: # fills url/command/args when omitted; mutates server_config in place - _apply_mcp_preset( - name, - preset_name=preset, - url=server_config.get("url"), - command=server_config.get("command"), - cmd_args=list(server_config.get("args") or []), - server_config=server_config) + mc._apply_mcp_preset( + name, preset_name=preset, url=server_config.get("url"), command=server_config.get("command"), + cmd_args=list(server_config.get("args") or []), server_config=server_config) if not server_config.get("url") and not server_config.get("command"): return _err(rid, 4063, "config must specify a 'url' (http) or 'command' (stdio), or a valid 'preset'") - bearer_token = params.get("bearer_token") - if bearer_token: - server_config["headers"] = _save_bearer_auth_token(name, str(bearer_token)) - if not _save_mcp_server(name, server_config): + if bearer_token := params.get("bearer_token"): + server_config["headers"] = mc._save_bearer_auth_token(name, str(bearer_token)) + if not mc._save_mcp_server(name, server_config): return _err(rid, 4001, f"server '{name}' rejected: suspicious command/args configuration") - saved = _get_mcp_servers().get(name, server_config) + saved = mc._get_mcp_servers().get(name, server_config) return _ok(rid, {"ok": True, "name": name, "server": _mcp_summarize_server(name, saved)}) -@method("mcp.servers.set_api_key") -@_profile_scoped_rpc(5024, required=(("name", _stripped), ("value", _nonempty)), catch_resolve=False) +@_mcp_rpc("set_api_key", (*_NAME, ("value", _nonempty))) def _(rid, params: dict) -> dict: - """Secret → profile .env under ``env_var`` (default ``MCP__API_KEY``); config.yaml - gets a reference: ``Authorization: Bearer ${ENV}`` header (http) or ``env: {VAR: "${ENV}"}`` - (stdio), matching ``cmd_mcp_configure`` / ``_save_bearer_auth_token``.""" - from hermes_cli.config import load_config, save_config, save_env_value - from hermes_cli.mcp_config import _bearer_auth_headers, _env_key_for_server, _strip_bearer_prefix + """Secret → profile .env under ``env_var`` (default ``MCP__API_KEY``); config.yaml gets only + a ``${ENV}`` reference (Bearer header for http, ``env`` entry for stdio).""" + hc, mc = _tools_mod("hermes_cli.config"), _tools_mod("hermes_cli.mcp_config") name, servers, err = _mcp_named_server(rid, params) if err: return err value = params.get("value") - env_var = str(params.get("env_var") or "").strip() or _env_key_for_server(name) + env_var = _str_arg(params, "env_var") or mc._env_key_for_server(name) entry = servers[name] if not isinstance(entry, dict): return _err(rid, 4001, "malformed server config") if entry.get("url"): - normalized = _strip_bearer_prefix(str(value)) + normalized = mc._strip_bearer_prefix(str(value)) if not normalized or normalized.lower() == "bearer": return _err(rid, 4063, "value is not a valid credential") - save_env_value(env_var, normalized) - if env_var == _env_key_for_server(name): - entry["headers"] = _bearer_auth_headers(name) - else: - entry["headers"] = {"Authorization": f"Bearer ${{{env_var}}}"} + hc.save_env_value(env_var, normalized) + is_default = env_var == mc._env_key_for_server(name) + entry["headers"] = ( + mc._bearer_auth_headers(name) if is_default else {"Authorization": f"Bearer ${{{env_var}}}"}) else: - save_env_value(env_var, str(value)) + hc.save_env_value(env_var, str(value)) env_block = entry.get("env") - if not isinstance(env_block, dict): - env_block = {} - env_block[env_var] = f"${{{env_var}}}" - entry["env"] = env_block - cfg = load_config() + entry["env"] = env_block if isinstance(env_block, dict) else {} + entry["env"][env_var] = f"${{{env_var}}}" + cfg = hc.load_config() cfg.setdefault("mcp_servers", {})[name] = entry - save_config(cfg) + hc.save_config(cfg) return _ok(rid, {"ok": True, "name": name, "env_var": env_var, "server": _mcp_summarize_server(name, entry)}) -@method("mcp.servers.test") -@_mcp_server_scoped +@_mcp_rpc("test") def _(rid, params: dict) -> dict: - """Connect, list tools, disconnect. Success: ``{ok, tools, prompts, resources, oauth_needed, - oauth_tokens_present}``; failure: ``{ok: false, error, tools: [], oauth_needed, ...}``. - Runs on the RPC pool (_LONG_HANDLERS): a cold stdio `npx` spawn can block for seconds.""" - from hermes_cli.mcp_config import _oauth_tokens_present, _probe_single_server + """Connect, list tools, disconnect → ``{ok, tools, prompts, resources, oauth_needed, + oauth_tokens_present}`` (``{ok: false, error, tools: []...}`` on failure). RPC pool: cold npx blocks.""" + mc = _tools_mod("hermes_cli.mcp_config") name, servers, err = _mcp_named_server(rid, params) if err: return err @@ -1429,50 +1229,38 @@ def _(rid, params: dict) -> dict: details: dict = {} def failure(error: str, oauth_needed: bool, tokens_present) -> dict: - payload = {"ok": False, "error": error, "tools": [], "oauth_needed": oauth_needed} - return _ok(rid, {**payload, "oauth_tokens_present": tokens_present}) + return _ok(rid, {"ok": False, "error": error, "tools": [], "oauth_needed": oauth_needed, + "oauth_tokens_present": tokens_present}) try: - tools = _probe_single_server(name, cfg, details=details) - token_present = _oauth_tokens_present(name) if needs_oauth_token else True + tools = mc._probe_single_server(name, cfg, details=details) + token_present = mc._oauth_tokens_present(name) if needs_oauth_token else True except Exception as exc: - return failure(str(exc), needs_oauth_token, _oauth_tokens_present(name) if needs_oauth_token else None) + return failure(str(exc), needs_oauth_token, mc._oauth_tokens_present(name) if needs_oauth_token else None) if not token_present: return failure("OAuth authentication required — no token found.", True, False) - payload = { - "ok": True, - "tools": [{"name": t, "description": d} for t, d in tools], - "prompts": details.get("prompts", 0), - "resources": details.get("resources", 0), - "oauth_needed": needs_oauth_token, - "oauth_tokens_present": True if needs_oauth_token else None} - return _ok(rid, payload) + return _ok(rid, { + "ok": True, "tools": [{"name": t, "description": d} for t, d in tools], + "prompts": details.get("prompts", 0), "resources": details.get("resources", 0), + "oauth_needed": needs_oauth_token, "oauth_tokens_present": True if needs_oauth_token else None}) -@method("mcp.servers.remove") -@_mcp_server_scoped +@_mcp_rpc("remove") def _(rid, params: dict) -> dict: """Remove a server from the profile's config.yaml → ``{ok: true, removed: true}``.""" - from hermes_cli.mcp_config import _remove_mcp_server - name = str(params.get("name") or "").strip() - if not _remove_mcp_server(name): + name = _str_arg(params, "name") + if not _tools_mod("hermes_cli.mcp_config")._remove_mcp_server(name): return _err(rid, 4064, f"server '{name}' not found") return _ok(rid, {"ok": True, "removed": True}) -@method("mcp.servers.oauth.start") -@_mcp_server_scoped +@_mcp_rpc("oauth.start") def _(rid, params: dict) -> dict: - """Begin a session-backed OAuth flow → ``{ok, session_id, auth_url, flow: "pkce"}``. - - The client opens ``auth_url`` and polls ``mcp.servers.oauth.poll`` until ``approved``. - A background worker drives the ``hermes mcp login`` machinery with a loopback - listener. With ``client_redirect_uri`` the CLIENT hosts the loopback and relays the - code via ``mcp.servers.oauth.callback`` — the only flow that works when desktop and - gateway are on different machines. Runs on the RPC pool (_LONG_HANDLERS).""" - client_redirect_uri = str(params.get("client_redirect_uri") or "").strip() or None + """Begin a session-backed OAuth flow → ``{ok, session_id, auth_url, flow: "pkce"}``; the client + opens ``auth_url`` and polls ``mcp.servers.oauth.poll``. With ``client_redirect_uri`` the CLIENT + hosts the loopback and relays the code via ``mcp.servers.oauth.callback`` (desktop and gateway + on different machines). Runs on the RPC pool (_LONG_HANDLERS).""" + client_redirect_uri = _str_arg(params, "client_redirect_uri") or None try: - from hermes_constants import get_hermes_home - from tui_gateway import mcp_oauth_sessions name, servers, err = _mcp_named_server(rid, params) if err: return err @@ -1482,72 +1270,45 @@ def _(rid, params: dict) -> dict: if cfg.get("headers") and cfg.get("auth") != "oauth": return _err(rid, 4001, "this server uses header/API-key auth, not OAuth") cfg["auth"] = "oauth" - hermes_home = str(get_hermes_home().expanduser().resolve(strict=False)) - result = mcp_oauth_sessions.start_flow(hermes_home, name, cfg, client_redirect_uri=client_redirect_uri) + hermes_home = str(_tools_mod("hermes_constants").get_hermes_home().expanduser().resolve(strict=False)) + result = _tools_mod("tui_gateway.mcp_oauth_sessions").start_flow( + hermes_home, name, cfg, client_redirect_uri=client_redirect_uri) except ValueError as e: return _err(rid, 4001, str(e)) - return _ok(rid, {"ok": True, "session_id": result["session_id"], "auth_url": result["auth_url"], "flow": result["flow"]}) + return _ok(rid, {"ok": True, **{k: result[k] for k in ("session_id", "auth_url", "flow")}}) -@method("mcp.servers.oauth.poll") -@_profile_scoped_rpc(5024, required=_NAME_SESSION, catch_resolve=False) +@_mcp_rpc("oauth.poll", _NAME_SESSION) def _(rid, params: dict) -> dict: - """Poll a flow → ``{ok, status: pending|approved|error, error_message?, auth_url?, tools?}``. - On ``approved`` tokens persist for that server/profile (profile scope applies here too).""" - from tui_gateway import mcp_oauth_sessions - name = str(params.get("name") or "").strip() - session_id = str(params.get("session_id") or "").strip() - result = mcp_oauth_sessions.poll_flow(session_id, name) - return _ok(rid, {"ok": True, **result}) + """Poll a flow → ``{ok, status: pending|approved|error, ...}``; ``approved`` persists tokens per profile.""" + poll = _tools_mod("tui_gateway.mcp_oauth_sessions").poll_flow + return _ok(rid, {"ok": True, **poll(_str_arg(params, "session_id"), _str_arg(params, "name"))}) -@method("mcp.servers.oauth.callback") -@_profile_scoped_rpc(5024, required=_NAME_SESSION, catch_resolve=False) +@_mcp_rpc("oauth.callback", _NAME_SESSION) def _(rid, params: dict) -> dict: - """Relay a client-captured redirect (``code``/``state``/``error``) into a flow started with - ``client_redirect_uri``. ``{ok: true}`` once accepted (state verified), else ``{ok: false, error_message}``.""" - from tui_gateway import mcp_oauth_sessions - name = str(params.get("name") or "").strip() - session_id = str(params.get("session_id") or "").strip() - result = mcp_oauth_sessions.deliver_callback_flow( - session_id, - name, - code=str(params.get("code") or "") or None, - state=str(params.get("state") or "") or None, - error=str(params.get("error") or "") or None) - return _ok(rid, result) + """Relay a client-captured redirect (``code``/``state``/``error``) into a ``client_redirect_uri`` flow.""" + code, state, error = (str(params.get(k) or "") or None for k in ("code", "state", "error")) + deliver = _tools_mod("tui_gateway.mcp_oauth_sessions").deliver_callback_flow + return _ok(rid, deliver( + _str_arg(params, "session_id"), _str_arg(params, "name"), code=code, state=state, error=error)) # ─── Plugins ───────────────────────────────────────────────────────────────── - - def _plugin_rows() -> list[dict]: - from hermes_cli.plugins_cmd import ( - _bundled_default_on, - _discover_all_plugins, - _get_disabled_set, - _get_enabled_set, - _is_portable_plugin_dir, - _plugin_status) - enabled = _get_enabled_set() - disabled = _get_disabled_set() + pc = _tools_mod("hermes_cli.plugins_cmd") + enabled, disabled = pc._get_enabled_set(), pc._get_disabled_set() out = [] - for name, version, desc, source, _dir, key in sorted(_discover_all_plugins()): - status = _plugin_status(name, enabled, disabled, key=key) + for name, version, desc, source, _dir, key in sorted(pc._discover_all_plugins()): + status = pc._plugin_status(name, enabled, disabled, key=key) # Bundled backends/platforms/providers run without an explicit enable: report the # truthful default instead of "not enabled" (reads as OFF). - if status == "not enabled" and source == "bundled" and _bundled_default_on(_dir): + if status == "not enabled" and source == "bundled" and pc._bundled_default_on(_dir): status = "enabled" - out.append( - { - "name": name, - "key": key, # canonical registry key (``image_gen/fal``): names collide across category dirs - "version": str(version or ""), - "description": desc or "", - "source": source, - "status": status, - "portable": _is_portable_plugin_dir(_dir), # Agent Plugins v1 package vs native Hermes plugin - }) + # key = canonical registry key (names collide across category dirs); portable = Agent Plugins v1. + out.append({ + "name": name, "key": key, "version": str(version or ""), "description": desc or "", + "source": source, "status": status, "portable": pc._is_portable_plugin_dir(_dir)}) return out @@ -1558,13 +1319,12 @@ def _plugins_list(rid, params): def _plugins_toggle(rid, params): - from hermes_cli.plugins_cmd import dashboard_set_agent_plugin_enabled - # Prefer the canonical key — bare names are ambiguous across categories. ident = (params.get("key") or params.get("name") or "").strip() if not ident: return _err(rid, 4019, "plugins.toggle requires a 'key' or 'name'") - result = dashboard_set_agent_plugin_enabled(ident, enabled=bool(params.get("enable"))) + toggle = _tools_mod("hermes_cli.plugins_cmd").dashboard_set_agent_plugin_enabled + result = toggle(ident, enabled=bool(params.get("enable"))) if not result.get("ok"): return _err(rid, 5026, result.get("error") or "toggle failed") row = next((r for r in _plugin_rows() if ident in (r["key"], r["name"])), None) @@ -1572,29 +1332,23 @@ def _plugins_toggle(rid, params): def _plugins_install(rid, params): - from hermes_cli.plugins_cmd import dashboard_install_plugin ident = (params.get("identifier") or params.get("repo") or "").strip() if not ident: return _err(rid, 4019, "plugins.install requires 'identifier' or 'repo'") - result = dashboard_install_plugin(ident, force=bool(params.get("force")), enable=params.get("enable", True)) - if not result.get("ok"): - return _err(rid, 5026, result.get("error") or "install failed") - return _ok(rid, result) + result = _tools_mod("hermes_cli.plugins_cmd").dashboard_install_plugin( + ident, force=bool(params.get("force")), enable=params.get("enable", True)) + return _ok(rid, result) if result.get("ok") else _err(rid, 5026, result.get("error") or "install failed") -@method("plugins.manage") -@_profile_scoped_rpc(5026, catch_resolve=False) +_PLUGINS_ACTIONS = {"list": _plugins_list, "toggle": _plugins_toggle, "install": _plugins_install} + + +@_scoped_rpc("plugins.manage", 5026, catch_resolve=False) def _(rid, params: dict) -> dict: - """TUI Plugins Hub backend (shares primitives with ``hermes plugins`` / the dashboard). - - ``list`` → {plugins: [{name, key, version, description, source, status, portable}], user_count, bundled_count} - - ``toggle`` → flip ``key`` (or ``name``) per ``enable``; returns the row + {ok, unchanged} - - ``install`` → git-clone ``identifier``/``repo`` into ~/.hermes/plugins/ (``force``, ``enable`` default True) - Optional ``profile`` scopes HERMES_HOME (mcp.servers.* contract).""" - action = params.get("action", "list") - handler = {"list": _plugins_list, "toggle": _plugins_toggle, "install": _plugins_install}.get(action) - if handler is None: - return _err(rid, 4017, f"unknown plugins action: {action}") - return handler(rid, params) + """TUI Plugins Hub backend (shares primitives with ``hermes plugins`` / the dashboard): + ``list`` → {plugins, user_count, bundled_count}; ``toggle`` flips ``key``/``name`` per ``enable``; + ``install`` git-clones ``identifier``/``repo`` (``force``, ``enable`` default True).""" + return _run_action(rid, params, _PLUGINS_ACTIONS, "plugins") @method("shell.exec") @@ -1603,22 +1357,18 @@ def _(rid, params: dict) -> dict: if not cmd: return _err(rid, 4004, "empty command") try: - from tools.approval import detect_dangerous_command, detect_hardline_command - is_hardline, hardline_desc = detect_hardline_command(cmd) + approval = _tools_mod("tools.approval") + is_hardline, hardline_desc = approval.detect_hardline_command(cmd) if is_hardline: return _err(rid, 4005, f"blocked (hardline): {hardline_desc}. Use the agent for dangerous commands.") - is_dangerous, _, desc = detect_dangerous_command(cmd) + is_dangerous, _, desc = approval.detect_dangerous_command(cmd) if is_dangerous: return _err(rid, 4005, f"blocked: {desc}. Use the agent for dangerous commands.") except ImportError: return _err(rid, 5001, "shell.exec unavailable: approval safety module not importable") - try: - r = subprocess.run(cmd, shell=True, cwd=os.getcwd(), **_capture_run_kwargs(30)) - return _ok(rid, {"stdout": r.stdout[-4000:], "stderr": r.stderr[-2000:], "code": r.returncode}) - except subprocess.TimeoutExpired: - return _err(rid, 5002, "command timed out (30s)") - except Exception as e: - return _err(rid, 5003, str(e)) + return _captured_exec( + rid, cmd, 30, shell=True, fail_code=5003, timeout_err=(5002, "command timed out (30s)"), + on_result=lambda r: _ok(rid, {"stdout": r.stdout[-4000:], "stderr": r.stderr[-2000:], "code": r.returncode})) def register(server) -> None: diff --git a/tui_gateway/methods_voice.py b/tui_gateway/methods_voice.py index 52cf5b87e2..00b2395917 100644 --- a/tui_gateway/methods_voice.py +++ b/tui_gateway/methods_voice.py @@ -1,6 +1,5 @@ -"""Voice / TTS / wake-word JSON-RPC handlers and their process-global state (one mic, one -speaker per process). Bodies are rebound onto server.py's globals (method_ctx.bind_module) -and reference them bare. +"""Voice / TTS / wake-word JSON-RPC handlers and their process-global state (one mic, one speaker +per process). Bodies are rebound onto server.py's globals (method_ctx.bind_module), used bare. """ from __future__ import annotations @@ -14,7 +13,8 @@ _registry = HandlerRegistry() method = _registry.method -# ── Voice state ────────────────────────────────────────────────────────── +# ── Voice state: HERMES_VOICE / HERMES_VOICE_TTS are runtime-only env flags (never config.yaml) +# so a prior session can't auto-start REC. _voice_sid_lock = threading.Lock() _voice_event_sid: str = "" @@ -41,19 +41,16 @@ def _resume_voice_wake() -> None: def _voice_mode_enabled() -> bool: - """Runtime-only flag (env, never config.yaml) so a prior session can't auto-start REC.""" return os.environ.get("HERMES_VOICE", "").strip() == "1" def _voice_tts_enabled() -> bool: - """Whether agent replies are spoken back via TTS (runtime only).""" return os.environ.get("HERMES_VOICE_TTS", "").strip() == "1" def _end_voice_chat(*, stop_loop: bool, stop_tts: bool) -> None: """Flip voice + TTS off; optionally halt the continuous loop / cut live TTS (best-effort).""" - os.environ["HERMES_VOICE"] = "0" - os.environ["HERMES_VOICE_TTS"] = "0" + os.environ["HERMES_VOICE"] = os.environ["HERMES_VOICE_TTS"] = "0" if stop_loop: with contextlib.suppress(Exception): from hermes_cli.voice import stop_continuous @@ -64,8 +61,8 @@ def _end_voice_chat(*, stop_loop: bool, stop_tts: bool) -> None: def _tts_lease_async(lease: str, active: bool) -> None: - """Acquire/release a TTS engine lease off the RPC thread: acquiring warms the provider - (local engines load a model) and must not block the toggle's reply. Best-effort.""" + """Acquire/release a TTS lease off the RPC thread (acquiring warms a local engine; must not + block the toggle's reply). Best-effort.""" def _run(): try: from tools.tts_tool import acquire_tts_lease, release_tts_lease @@ -75,17 +72,21 @@ def _tts_lease_async(lease: str, active: bool) -> None: threading.Thread(target=_run, name=f"tts-lease-{lease}", daemon=True).start() +def _running_sessions() -> list: + with _sessions_lock: + return [s for s in _sessions.values() if s.get("running")] + + def _any_session_running() -> bool: """Voice busy-probe: silent captures during a long turn don't count toward the no-speech limit.""" try: - with _sessions_lock: - return any(s.get("running") for s in _sessions.values()) + return bool(_running_sessions()) except Exception: return False -# ── Streaming TTS: one pipeline per process (one speaker); a new turn's pipeline barges in -# on the previous. Token deltas feed a sentence-buffering consumer (stream_tts_to_speaker). +# ── Streaming TTS: one pipeline per process (one speaker); a new turn's pipeline barges in on +# the previous. Token deltas feed a sentence-buffering consumer (stream_tts_to_speaker). _tts_stream_lock = threading.Lock() _tts_stream_state: Optional[dict] = None @@ -104,8 +105,7 @@ def _tts_stream_begin() -> Optional[queue.Queue]: _tts_stream_stop() text_queue: queue.Queue = queue.Queue() stop, done = threading.Event(), threading.Event() - threading.Thread(target=stream_tts_to_speaker, args=(text_queue, stop, done), - daemon=True).start() + threading.Thread(target=stream_tts_to_speaker, args=(text_queue, stop, done), daemon=True).start() global _tts_stream_state with _tts_stream_lock: _tts_stream_state = {"stop": stop, "done": done} @@ -115,17 +115,17 @@ def _tts_stream_begin() -> Optional[queue.Queue]: def _tts_stream_stop(user_barge: bool = True) -> None: """Cut in-flight streaming TTS. *user_barge* latches the interruption for the next turn's - model note — pass ``False`` for mode changes (/voice off).""" + model note; ``False`` for mode changes (/voice off).""" global _tts_stream_state with _tts_stream_lock: state, _tts_stream_state = _tts_stream_state, None if state is None: return if user_barge and not state["done"].is_set(): - import traceback as _tb - logger.debug("TTS CUT: _tts_stream_stop(user_barge=True) — new turn or " - "interrupt cutting in-flight TTS\n%s", "".join(_tb.format_stack())) + import traceback from tools.tts_streaming import mark_speech_interrupted + logger.debug("TTS CUT: _tts_stream_stop(user_barge=True) — new turn or " + "interrupt cutting in-flight TTS\n%s", "".join(traceback.format_stack())) mark_speech_interrupted() state["stop"].set() with contextlib.suppress(Exception): @@ -134,13 +134,13 @@ def _tts_stream_stop(user_barge: bool = True) -> None: # ── Full-duplex agent-turn listener: arms at utterance-submit, spans generation AND playback -# (per-playback monitors were deaf during generation and mis-calibrated against speaker -# bleed), disarms when no session runs, no TTS is pending, and no audio flows. +# (per-playback monitors were deaf during generation and mis-calibrated against speaker bleed), +# disarms when no session runs, no TTS is pending, and no audio flows. _fd_speak_pipelines holds +# (stop, done) pairs of fallback whole-reply speak paths: the listener cuts their private stop +# events too, and keeps listening while any is still speaking. _fd_listener_lock = threading.Lock() _fd_listener_active = False -# (stop, done) pairs of fallback whole-reply speak paths: the listener cuts their private stop -# events too, and keeps listening while any is still speaking. _fd_speak_pipelines: "set[tuple[threading.Event, threading.Event]]" = set() @@ -160,34 +160,25 @@ def _arm_barge_listener_if_enabled() -> None: _arm_full_duplex_listener() -def _fd_speak_pipelines_snapshot() -> list: - with _fd_listener_lock: - return list(_fd_speak_pipelines) - - def _fd_tts_pending() -> bool: """True while any TTS (streaming pipeline or fallback speak) is unfinished.""" with _tts_stream_lock: state = _tts_stream_state - if state is not None and not state["done"].is_set(): - return True - return any(not done.is_set() for _stop, done in _fd_speak_pipelines_snapshot()) + with _fd_listener_lock: + pending = ([state["done"]] if state is not None else []) + [done for _stop, done in _fd_speak_pipelines] + return any(not done.is_set() for done in pending) def _full_duplex_listener() -> None: - """Mic live from utterance-submit to turn-complete; a trip (see ``_fd_trip``) transcribes - the utterance and emits it as ``voice.transcript``.""" + """Mic live from utterance-submit to turn-complete; a trip transcribes -> ``voice.transcript``.""" global _fd_listener_active try: from tools.voice_mode import (full_duplex_listen, is_audio_output_active, transcribe_recording) def _should_stop() -> bool: - if not _voice_mode_enabled(): - return True - if _any_session_running() or _fd_tts_pending(): - return False - return not is_audio_output_active() + return not _voice_mode_enabled() or not ( + _any_session_running() or _fd_tts_pending() or is_audio_output_active()) tripped = threading.Event() def _on_trigger(phase: str) -> None: @@ -201,9 +192,8 @@ def _full_duplex_listener() -> None: return try: result = transcribe_recording(wav_path) - text = (result.get("transcript") or "").strip() if result.get("success") else "" - if text: - _deliver_fd_transcript(text) + if result.get("success") and (result.get("transcript") or "").strip(): + _deliver_fd_transcript(result["transcript"].strip()) finally: with contextlib.suppress(OSError): os.unlink(wav_path) @@ -216,43 +206,36 @@ def _full_duplex_listener() -> None: def _fd_barge_params(cfg: dict) -> tuple[float, int]: """``(threshold multiplier, grace ms)`` from the voice config; malformed -> defaults.""" - try: - mult = float(cfg.get("barge_in_threshold_multiplier", 0) or 0) - except (TypeError, ValueError): - mult = 0.0 - try: - grace_ms = int(float(cfg.get("barge_in_grace_seconds", 0.5)) * 1000) - except (TypeError, ValueError): - grace_ms = 500 - return mult, max(0, grace_ms) - - -def _cut_all_tts() -> None: - """Cut streaming TTS, every fallback speak pipeline, and the file player.""" - from tools.voice_mode import stop_playback - _tts_stream_stop(user_barge=True) - for _stop, _done in _fd_speak_pipelines_snapshot(): - _stop.set() - stop_playback() + def num(conv, key, default): + try: + return conv(cfg.get(key, default)) + except (TypeError, ValueError): + return conv(default) + mult = num(lambda v: float(v or 0), "barge_in_threshold_multiplier", 0) + return mult, max(0, num(lambda v: int(float(v) * 1000), "barge_in_grace_seconds", 0.5)) def _fd_trip(phase: str) -> None: - """Listener tripped: latch the interruption, cut TTS, and during generation also - interrupt every running turn (the ``agent.interrupt()`` seam ``session.interrupt`` uses).""" + """Listener tripped: latch the interruption, cut TTS FIRST (so a stale reply can never + speak), and during generation also interrupt every running turn (the ``agent.interrupt()`` + seam ``session.interrupt`` uses).""" from tools.tts_streaming import mark_speech_interrupted + from tools.voice_mode import stop_playback mark_speech_interrupted() if phase == "playback": logger.debug("TTS CUT: full-duplex listener tripped during playback") - _cut_all_tts() else: logger.debug("full-duplex listener tripped during generation — " "interrupting running turn(s)") - # Cut pending TTS FIRST so the stale reply can never speak. - _cut_all_tts() + # Cut streaming TTS, every fallback speak pipeline, and the file player. + _tts_stream_stop(user_barge=True) + with _fd_listener_lock: + for _stop, _done in _fd_speak_pipelines: + _stop.set() + stop_playback() + if phase != "playback": try: - with _sessions_lock: - running = [s for s in _sessions.values() if s.get("running")] - for s in running: + for s in _running_sessions(): agent = s.get("agent") if agent is not None and hasattr(agent, "interrupt"): with contextlib.suppress(Exception): @@ -263,25 +246,20 @@ def _fd_trip(phase: str) -> None: def _deliver_fd_transcript(text: str) -> None: - """Emit the captured interjection; a bare stop phrase also ends the voice chat.""" - # Stop-check must never break transcript delivery (stubbed voice_mode in tests, - # partial installs) — treat as not-a-stop. + """Emit the captured interjection; a bare stop phrase also ends the voice chat. The stop + check must never break delivery (stubbed voice_mode in tests, partial installs).""" try: from tools.voice_mode import is_voice_stop_phrase is_stop = is_voice_stop_phrase(text) except Exception: is_stop = False - if is_stop: - # Turn already interrupted / TTS cut at trip time; now end the chat. + if is_stop: # turn already interrupted / TTS cut at trip time; now end the chat _end_voice_chat(stop_loop=True, stop_tts=False) - _voice_emit("voice.transcript", {"stop_phrase": True, "text": text}) - else: - _voice_emit("voice.transcript", {"text": text}) + _voice_emit("voice.transcript", {"stop_phrase": True, "text": text} if is_stop else {"text": text}) def _speak_text_with_barge(text: str) -> None: - """Speak via hermes_cli.voice.speak_text, registered in ``_fd_speak_pipelines`` so the - full-duplex listener can cut it and keeps listening while it is pending.""" + """speak_text registered in ``_fd_speak_pipelines`` so the listener can cut it / waits for it.""" from hermes_cli.voice import speak_text stop, done = threading.Event(), threading.Event() with _fd_listener_lock: @@ -290,8 +268,7 @@ def _speak_text_with_barge(text: str) -> None: def _speak(): try: speak_text(text, stop) - except TypeError: - # Older wrapper without the stop_event parameter. + except TypeError: # older wrapper without the stop_event parameter speak_text(text) finally: done.set() @@ -313,10 +290,12 @@ def _voice_cfg_number(value, default): return value if isinstance(value, (int, float)) and not isinstance(value, bool) else default -def _voice_record_key() -> str: - """Current ``voice.record_key`` value, documented default on error.""" +def _voice_status_payload(**extra) -> dict: + """``{enabled, record_key, tts, **extra}``: record_key (default ``ctrl+b``) rides every voice.toggle + branch so a tts toggle never resets a custom binding.""" record_key = _voice_cfg_dict().get("record_key") - return str(record_key) if isinstance(record_key, str) and record_key else "ctrl+b" + record_key = record_key if isinstance(record_key, str) and record_key else "ctrl+b" + return {"enabled": _voice_mode_enabled(), "record_key": record_key, "tts": _voice_tts_enabled(), **extra} # ── Wake word ("Hey Hermes"): process-global detector (one mic). The first eligible transport @@ -333,12 +312,6 @@ def _wake_owner_snapshot(): return _wake_owner_transport, _wake_owner_surface -def _set_wake_owner(transport, surface: str) -> None: - global _wake_owner_transport, _wake_owner_surface - with _wake_lock: - _wake_owner_transport, _wake_owner_surface = transport, surface - - def _release_wake_for_transport(transport: "Transport") -> bool: """Release the wake lease iff ``transport`` is the current gateway owner.""" global _wake_owner_transport, _wake_owner_surface @@ -365,11 +338,10 @@ _wake_resume_retry_active = False def _wake_resume_if_owner(owner: "Transport", *, retry_seconds: float = 15.0, retry_interval: float = 1.0) -> bool: - """Resume the wake detector for ``owner``; self-heal a busy microphone. Reopening the mic - right after a voice turn can fail while the device is still being released (browser WebRTC - tracks release async): on an exception retry in a background thread until it sticks, the - lease changes hands, or ``retry_seconds`` elapses. ``False`` from ``resume_listening`` - (lease gone / other owner) is final — never retried, so this can't steal another's mic.""" + """Resume the wake detector for ``owner``, self-healing a busy microphone: reopening right after + a voice turn can fail while the device is still being released (browser WebRTC tracks release + async), so an exception retries in a background thread until it sticks, the lease changes hands, + or ``retry_seconds`` elapses. ``False`` (lease gone / other owner) is final — never retried.""" from tools.wake_word import resume_listening try: return resume_listening(owner=owner) @@ -387,14 +359,10 @@ def _wake_resume_if_owner(owner: "Transport", *, retry_seconds: float = 15.0, try: while time.monotonic() < deadline: time.sleep(retry_interval) - try: + with contextlib.suppress(Exception): if resume_listening(owner=owner): logger.info("wake: detector resumed after retry") - return - except Exception: - continue - # False — detector gone or lease moved: stop, don't fight it. - return + return # False — detector gone or lease moved: stop, don't fight it. logger.warning("wake: could not resume detector after voice turn " "(microphone still busy?) — toggle the wake word to re-arm") finally: @@ -414,18 +382,20 @@ def _persist_wake_enabled(enabled: bool) -> bool: return False -def _wake_prefers_client(params: dict, surface: str) -> bool: - """Desktop (gui) prefers client capture (Mac mic → wake.feed PCM); CLI/TUI stay local.""" - return surface in ("gui", "desktop") or bool(params.get("client_capture")) +def _owner_result(rid, field: str, ok, **extra) -> dict: + """``{field: ok, reason: None | "not_owner", **extra}`` for the owner-gated wake RPCs.""" + return _ok(rid, {field: ok, "reason": None if ok else "not_owner", **extra}) def _frame_fields(frame: dict) -> dict: return {"sample_rate": frame.get("sample_rate", 16000), "frame_length": frame.get("frame_length", 1280)} -def _wake_probe(cfg: dict, prefer_client: bool) -> tuple[str, dict]: - """``(capture_mode, requirements)``; capture stamped so the probe matches what would arm.""" +def _wake_probe(cfg: dict, params: dict, surface: str) -> tuple[str, dict]: + """``(capture_mode, requirements)``; capture stamped so the probe matches what would arm. + Desktop (gui) prefers client capture (Mac mic → wake.feed PCM); CLI/TUI stay local.""" from tools.wake_word import check_wake_word_requirements, resolve_capture_mode + prefer_client = surface in ("gui", "desktop") or bool(params.get("client_capture")) capture_mode = resolve_capture_mode(cfg, prefer_client=prefer_client) return capture_mode, check_wake_word_requirements({**cfg, "capture": capture_mode}) @@ -455,27 +425,30 @@ def _wake_detect_handler(transport, sid: str, phrase: str, new_session: bool): @method("gateway.capabilities") def _(rid, params: dict) -> dict: - """Advertise what THIS BUILD enforces (a client withholds unless the guarantee is advertised). - Sourced from the enforcing module, never config: a believed-but-absent capability is worse.""" + """What THIS BUILD enforces (a client withholds unless advertised), sourced from the enforcing + module, never config: a believed-but-absent capability is worse.""" from hermes_cli.active_sessions import PER_SESSION_EXCLUSIVE_SUBMIT return _ok(rid, {"per_session_exclusive_submit": bool(PER_SESSION_EXCLUSIVE_SUBMIT)}) @method("ping") def _(rid, params: dict) -> dict: - """Cheapest liveness probe, answered on the WS reader thread (works while every agent is - mid-turn) so the desktop can tell a half-open socket after sleep/wake and reconnect.""" + """Cheapest liveness probe, answered on the WS reader thread (works while every agent is mid-turn) + so the desktop can tell a half-open socket after sleep/wake.""" return _ok(rid, {"pong": True}) @method("wake.start") def _(rid, params: dict) -> dict: """Arm the wake-word listener for the calling surface ("tui" | "gui"); ``{started: False, - reason}`` when disabled, owned by another surface, or deps/mic aren't ready. ``persist: true`` - (explicit gesture) flips ``wake_word.enabled`` on before arming; auto-arm callers omit it.""" + reason}`` when disabled/owned/not ready. ``persist: true`` (explicit gesture) also flips + ``wake_word.enabled`` on; auto-arm callers omit it.""" + global _wake_owner_transport, _wake_owner_surface surface = str(params.get("surface") or "auto").strip().lower() - persist = bool(params.get("persist")) transport = _caller_transport() + + def refused(reason, **extra): + return _ok(rid, {"started": False, "reason": reason, **extra}) try: from tools.wake_word import ( WakeWordInUse, detector_frame_info, load_wake_word_config, owns_listener, @@ -483,16 +456,13 @@ def _(rid, params: dict) -> dict: except Exception as e: return _err(rid, 5026, f"wake module unavailable: {e}") cfg = load_wake_word_config() - capture_mode, reqs = _wake_probe(cfg, _wake_prefers_client(params, surface)) - external_audio = capture_mode == "client" + capture_mode, reqs = _wake_probe(cfg, params, surface) # Requirements first: a gesture on an un-armable setup must refuse WITHOUT flipping # wake_word.enabled — else config says on while nothing can arm. if not reqs["available"]: logger.warning("wake.start(%s): not available — %s", surface, reqs.get("hint")) - return _ok(rid, { - "started": False, "reason": "unavailable", "hint": reqs.get("hint") or "", - "capture": capture_mode}) - enabled_persisted = bool(persist and not cfg.get("enabled") and _persist_wake_enabled(True)) + return refused("unavailable", hint=reqs.get("hint") or "", capture=capture_mode) + enabled_persisted = bool(params.get("persist")) and not cfg.get("enabled") and _persist_wake_enabled(True) if enabled_persisted: cfg = {**cfg, "enabled": True} if not wake_surface_enabled(surface, cfg): @@ -501,31 +471,28 @@ def _(rid, params: dict) -> dict: reason = "disabled" if not cfg.get("enabled") else "disabled_for_surface" logger.info("wake.start(%s): %s (enabled=%s, surface=%s)", surface, reason, cfg.get("enabled"), cfg.get("surface")) - return _ok(rid, {"started": False, "reason": reason}) + return refused(reason) existing_owner, existing_surface = _wake_owner_snapshot() - if existing_owner is not None and ( - _transport_is_dead(existing_owner) or not owns_listener(existing_owner) - ): + if existing_owner is not None and (_transport_is_dead(existing_owner) or not owns_listener(existing_owner)): _release_wake_for_transport(existing_owner) existing_owner, existing_surface = None, "" if existing_owner is not None and existing_owner is not transport: - return _ok(rid, {"started": False, "reason": "owned", "owner_surface": existing_surface}) - sid = str(params.get("session_id") or "") + return refused("owned", owner_surface=existing_surface) try: - on_detect = _wake_detect_handler( - transport, sid, wake_phrase(cfg), bool(cfg.get("start_new_session", True))) - start_listening(on_detect, owner=transport, config=cfg, external_audio=external_audio) + on_detect = _wake_detect_handler(transport, str(params.get("session_id") or ""), + wake_phrase(cfg), bool(cfg.get("start_new_session", True))) + start_listening(on_detect, owner=transport, config=cfg, + external_audio=capture_mode == "client") except WakeWordInUse: - return _ok(rid, {"started": False, "reason": "owned", - "owner_surface": existing_surface or None}) + return refused("owned", owner_surface=existing_surface or None) except Exception as e: logger.warning("wake.start(%s): failed to start listener: %s", surface, e) return _err(rid, 5026, str(e)) - _set_wake_owner(transport, surface) + with _wake_lock: + _wake_owner_transport, _wake_owner_surface = transport, surface frame = detector_frame_info() - logger.info( - "wake.start(%s): listening for %r (%s) capture=%s frame=%s", - surface, reqs["phrase"], reqs["provider"], capture_mode, frame.get("frame_length")) + logger.info("wake.start(%s): listening for %r (%s) capture=%s frame=%s", + surface, reqs["phrase"], reqs["provider"], capture_mode, frame.get("frame_length")) return _ok(rid, { "started": True, "phrase": reqs["phrase"], "provider": reqs["provider"], "owner_surface": surface, "enabled_persisted": enabled_persisted, "capture": capture_mode, @@ -535,34 +502,29 @@ def _(rid, params: dict) -> dict: @method("wake.stop") def _(rid, params: dict) -> dict: """Stop this surface's listener; ``persist: true`` also writes ``wake_word.enabled: false``.""" - transport = _caller_transport() - stopped = _release_wake_for_transport(transport) + stopped = _release_wake_for_transport(_caller_transport()) disabled_persisted = False - if bool(params.get("persist")): + if params.get("persist"): try: from tools.wake_word import load_wake_word_config currently_enabled = bool(load_wake_word_config().get("enabled")) except Exception: currently_enabled = True - if currently_enabled: - disabled_persisted = _persist_wake_enabled(False) - return _ok(rid, { - "stopped": stopped, "reason": None if stopped else "not_owner", - "disabled_persisted": disabled_persisted}) + disabled_persisted = currently_enabled and _persist_wake_enabled(False) + return _owner_result(rid, "stopped", stopped, disabled_persisted=disabled_persisted) @method("wake.pause") def _(rid, params: dict) -> dict: """Release the mic (e.g. while the desktop's browser captures audio).""" - transport = _caller_transport() try: from tools.wake_word import pause_listening - paused = pause_listening(owner=transport) + paused = pause_listening(owner=_caller_transport()) logger.info("wake.pause: detector paused=%s", paused) except Exception as e: logger.debug("wake.pause failed: %s", e) paused = False - return _ok(rid, {"paused": paused, "reason": None if paused else "not_owner"}) + return _owner_result(rid, "paused", paused) @method("wake.resume") @@ -570,7 +532,7 @@ def _(rid, params: dict) -> dict: """Reclaim the mic after a pause; no-op if the listener isn't armed.""" resumed = _wake_resume_if_owner(_caller_transport()) logger.info("wake.resume: detector resumed=%s", resumed) - return _ok(rid, {"resumed": resumed, "reason": None if resumed else "not_owner"}) + return _owner_result(rid, "resumed", resumed) @method("wake.status") @@ -580,11 +542,9 @@ def _(rid, params: dict) -> dict: audio_is_silent, detector_frame_info, get_input_device_status, is_listening, load_wake_word_config, owns_listener, silent_audio_hint) cfg = load_wake_word_config() - surface = str(params.get("surface") or "").strip().lower() - probe_capture, reqs = _wake_probe(cfg, _wake_prefers_client(params, surface)) - transport = _caller_transport() + probe_capture, reqs = _wake_probe(cfg, params, str(params.get("surface") or "").strip().lower()) owner, owner_surface = _wake_owner_snapshot() - owned_by_caller = owns_listener(transport) + owned_by_caller = owns_listener(_caller_transport()) listening = owned_by_caller and is_listening() silent = listening and audio_is_silent() input_device = get_input_device_status(cfg) @@ -596,22 +556,19 @@ def _(rid, params: dict) -> dict: # Effective capture: prefer the *armed* detector over config/auto, else with capture:auto # a bare status probe reports "local" and the desktop never reattaches the PCM feeder. frame = detector_frame_info() - if owned_by_caller and frame.get("external_audio"): - capture = "client" - elif owned_by_caller and listening: - capture = "local" + if owned_by_caller and (frame.get("external_audio") or listening): + capture = "client" if frame.get("external_audio") else "local" else: capture = probe_capture or reqs.get("capture") or str(cfg.get("capture") or "auto") + # `enabled` is config truth (clients re-arm after a voice turn from it); `audio_silent` = + # armed but deaf despite an open stream (see the platform-specific hint). return _ok(rid, { "listening": listening, "owned_by_caller": owned_by_caller, "owner_surface": owner_surface if owner is not None else None, "phrase": reqs["phrase"], "provider": reqs["provider"], "configured_surface": str(cfg.get("surface") or "auto"), "input_device": input_device, "available": reqs["available"], "hint": hint, - # Config truth: clients re-arm after a voice turn ("permanent on") from this. - "enabled": bool(cfg.get("enabled")), - # Armed but deaf despite an open stream; see platform-specific hint. - "audio_silent": silent, "capture": capture, + "enabled": bool(cfg.get("enabled")), "audio_silent": silent, "capture": capture, "local_input_available": bool(reqs.get("local_input_available")), **_frame_fields(frame)}) except Exception as e: return _err(rid, 5026, str(e)) @@ -620,44 +577,38 @@ def _(rid, params: dict) -> dict: @method("wake.feed") def _(rid, params: dict) -> dict: """Push client-captured PCM (``pcm``/``pcm_b64``: base64 int16 mono LE, 16 kHz only) into the - armed detector (``capture: "client"``) so mic-less remote backends can run openWakeWord.""" - transport = _caller_transport() + armed detector (``capture: "client"``) — mic-less remote backends can run openWakeWord.""" raw_b64 = params.get("pcm") or params.get("pcm_b64") or "" if not isinstance(raw_b64, str) or not raw_b64.strip(): return _err(rid, 4001, "wake.feed requires base64 pcm") + import base64 try: - import base64 pcm = base64.b64decode(raw_b64, validate=False) except Exception as e: return _err(rid, 4001, f"invalid base64 pcm: {e}") if not pcm: return _ok(rid, {"fed": False, "reason": "empty"}) - # Soft size cap: 64000 bytes = 2s of 16 kHz int16 mono - if len(pcm) > 64000: + if len(pcm) > 64000: # soft cap: 2s of 16 kHz int16 mono return _err(rid, 4001, "pcm frame too large") - sr = params.get("sample_rate") - if sr is not None and int(sr) not in (0, 16000): + if params.get("sample_rate") is not None and int(params["sample_rate"]) not in (0, 16000): return _err(rid, 4001, "wake.feed only accepts 16 kHz PCM") try: from tools.wake_word import feed_audio - ok = feed_audio(owner=transport, pcm_int16=pcm) + ok = feed_audio(owner=_caller_transport(), pcm_int16=pcm) except Exception as e: logger.debug("wake.feed failed: %s", e) return _err(rid, 5026, str(e)) - return _ok(rid, {"fed": bool(ok), "reason": None if ok else "not_owner"}) + return _owner_result(rid, "fed", bool(ok)) def _voice_toggle_status(rid, params: dict) -> dict: # Mirrors CLI _show_voice_status: STT/TTS availability tells the user WHY voice isn't # working; record_key lets the TUI bind and display the shortcut. - payload: dict = {"enabled": _voice_mode_enabled(), "record_key": _voice_record_key(), - "tts": _voice_tts_enabled()} + payload = _voice_status_payload() try: from tools.voice_mode import check_voice_requirements reqs = check_voice_requirements() - payload.update(available=bool(reqs.get("available")), - audio_available=bool(reqs.get("audio_available")), - stt_available=bool(reqs.get("stt_available")), + payload.update({k: bool(reqs.get(k)) for k in ("available", "audio_available", "stt_available")}, details=reqs.get("details") or "") except Exception as e: # Optional transcription deps — /voice status must always answer. @@ -667,16 +618,13 @@ def _voice_toggle_status(rid, params: dict) -> dict: def _voice_toggle_mode(rid, params: dict) -> dict: enabled = params.get("action") == "on" - # Runtime-only flag — never persisted, so the next launch starts with voice OFF. os.environ["HERMES_VOICE"] = "1" if enabled else "0" stop_hint = "" if enabled: # Spoken-stop hint for the client; sourced from voice.stop_phrases, empty when disabled. - try: + with contextlib.suppress(Exception): from tools.voice_mode import voice_stop_hint stop_hint = voice_stop_hint() - except Exception: - stop_hint = "" # Speech output already on → warm the engine now, not on the first reply. if _voice_tts_enabled(): _tts_lease_async("tui:voice-tts", True) @@ -689,25 +637,23 @@ def _voice_toggle_mode(rid, params: dict) -> dict: pass except Exception as e: logger.warning("voice: stop_continuous failed during toggle off: %s", e) - # Clear TTS so it can be toggled independently later; silence live speech. - os.environ["HERMES_VOICE_TTS"] = "0" + _set_voice_tts(False) # TTS is toggled independently later + return _ok(rid, _voice_status_payload(stop_hint=stop_hint)) + + +def _set_voice_tts(on: bool) -> None: + """Flip TTS; off silences live speech. The lease pre-loads the engine (on) / releases it (off).""" + os.environ["HERMES_VOICE_TTS"] = "1" if on else "0" + if not on: _tts_stream_stop(user_barge=False) - _tts_lease_async("tui:voice-tts", False) - return _ok(rid, {"enabled": enabled, "record_key": _voice_record_key(), - "tts": _voice_tts_enabled(), "stop_hint": stop_hint}) + _tts_lease_async("tui:voice-tts", on) def _voice_toggle_tts(rid, params: dict) -> dict: if not _voice_mode_enabled(): return _err(rid, 4014, "enable voice mode first: /voice on") - new_value = not _voice_tts_enabled() - os.environ["HERMES_VOICE_TTS"] = "1" if new_value else "0" - if not new_value: - _tts_stream_stop(user_barge=False) - # on → pre-load the engine so the first reply starts hot; off → release the lease. - _tts_lease_async("tui:voice-tts", new_value) - # record_key on every branch so a tts toggle never resets a custom binding. - return _ok(rid, {"enabled": True, "record_key": _voice_record_key(), "tts": new_value}) + _set_voice_tts(not _voice_tts_enabled()) + return _ok(rid, _voice_status_payload()) _VOICE_TOGGLE_ACTIONS = { @@ -726,14 +672,10 @@ def _(rid, params: dict) -> dict: return handler(rid, params) -# voice.record callbacks (module-level: they touch only process-global state). -def _vr_on_transcript(t): - _voice_emit("voice.transcript", {"text": t}) - _resume_voice_wake() - - -def _vr_on_silent(): - _voice_emit("voice.transcript", {"no_speech_limit": True}) +# voice.record callbacks: each terminal capture event resumes the wake detector so wake-triggered +# and manual captures coexist. +def _vr_transcript(payload: dict) -> None: + _voice_emit("voice.transcript", payload) _resume_voice_wake() @@ -741,8 +683,7 @@ def _vr_on_stop_phrase(t): # A SPOKEN bare stop phrase: end the chat like /voice off and emit a distinct signal so # clients end the conversation instead of treating it as a no-speech timeout. _end_voice_chat(stop_loop=False, stop_tts=True) - _voice_emit("voice.transcript", {"stop_phrase": True, "text": t}) - _resume_voice_wake() + _vr_transcript({"stop_phrase": True, "text": t}) def _vr_on_status(state): @@ -775,36 +716,32 @@ def _(rid, params: dict) -> dict: _resume_voice_wake() return _ok(rid, {"status": "stopped"}) from hermes_cli.voice import start_continuous - # Busy probe holds the no-speech counter during long agent turns. - # Safe to re-register every start; older wrappers lack the setter. + # Busy probe holds the no-speech counter during long agent turns; safe to re-register every + # start (older wrappers lack the setter). with contextlib.suppress(Exception): from hermes_cli.voice import set_voice_busy_probe set_voice_busy_probe(_any_session_running) - # Shape-safe: malformed voice YAML falls back to documented defaults. - # max_recording_seconds: explicit numeric <= 0 disables the cap (0.0). + # Shape-safe: malformed voice YAML falls back to documented defaults; an explicit numeric + # max_recording_seconds <= 0 disables the cap (0.0). voice_cfg = _voice_cfg_dict() max_rec = _voice_cfg_number(voice_cfg.get("max_recording_seconds"), 120.0) - # Hand the mic to STT if the wake detector holds it; resume on a terminal - # capture event so wake-triggered and manual captures coexist. - try: + # Hand the mic to STT if the wake detector holds it; a terminal capture event resumes it. + with contextlib.suppress(Exception): from tools.wake_word import pause_listening wake_paused = pause_listening(owner=transport) - except Exception: - wake_paused = False if wake_paused: with _voice_sid_lock: _voice_wake_owner = transport started = start_continuous( - on_transcript=_vr_on_transcript, on_status=_vr_on_status, on_silent_limit=_vr_on_silent, + on_transcript=lambda t: _vr_transcript({"text": t}), on_status=_vr_on_status, + on_silent_limit=lambda: _vr_transcript({"no_speech_limit": True}), silence_threshold=_voice_cfg_number(voice_cfg.get("silence_threshold"), 200), silence_duration=_voice_cfg_number(voice_cfg.get("silence_duration"), 3.0), auto_restart=False, max_recording_seconds=max_rec if max_rec > 0 else 0.0, - on_stop_phrase=_vr_on_stop_phrase, - ) + on_stop_phrase=_vr_on_stop_phrase) if started is False: _resume_voice_wake() - return _ok(rid, {"status": "busy"}) - return _ok(rid, {"status": "recording"}) + return _ok(rid, {"status": "busy" if started is False else "recording"}) except Exception as e: if wake_paused or action == "stop": _resume_voice_wake() @@ -819,12 +756,9 @@ def _(rid, params: dict) -> dict: if not text: return _err(rid, 4020, "text required") try: - # Import check up front so a missing voice module returns 5026, not a silent thread death. - import hermes_cli.voice # noqa: F401 - except ImportError: - return _err(rid, 5026, "voice module not available") + import hermes_cli.voice # noqa: F401 (a missing module must answer 5026, not die in a thread) except Exception as e: - return _err(rid, 5026, str(e)) + return _err(rid, 5026, "voice module not available" if isinstance(e, ImportError) else str(e)) threading.Thread(target=_speak_text_with_barge, args=(text,), daemon=True).start() return _ok(rid, {"status": "speaking"}) diff --git a/tui_gateway/model_switch.py b/tui_gateway/model_switch.py index 86cd88aa6d..02c458105c 100644 --- a/tui_gateway/model_switch.py +++ b/tui_gateway/model_switch.py @@ -1,8 +1,6 @@ -"""Model switching for a live session: persist, snapshot/restore runtime, /model apply with guards, bot-capability + config sync. - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so they reference server.py globals bare. -""" +"""Model switching for a live session: persist, snapshot/restore runtime, /model apply with +guards, bot-capability + config sync. Bodies are rebound onto server.py's globals at install +time (method_ctx.bind_module), so they reference server.py globals bare.""" from __future__ import annotations @@ -17,7 +15,6 @@ def _persist_model_switch(result) -> None: # Targeted key writes: a full `model:` block rewrite via save_config() would destroy # sibling keys the user set there (`model_slots`, `model_fallback`, ...). from cli import save_config_value - save_config_value("model.default", result.new_model) save_config_value("model.provider", result.target_provider) # A provider without a base_url must clear the stale one (custom endpoint -> native) @@ -25,11 +22,13 @@ def _persist_model_switch(result) -> None: save_config_value("model.base_url", result.base_url or None) +_RUNTIME_KEYS = ("model", "provider", "api_key", "base_url", "api_mode") + + def _snapshot_agent_model_runtime(agent) -> dict: """Capture the current agent model runtime for a one-turn restore.""" - snap = {k: getattr(agent, k, "") for k in ("model", "provider", "api_key", "base_url", "api_mode")} - snap["primary_runtime"] = copy.deepcopy(getattr(agent, "_primary_runtime", None)) - return snap + return {**{k: getattr(agent, k, "") for k in _RUNTIME_KEYS}, + "primary_runtime": copy.deepcopy(getattr(agent, "_primary_runtime", None))} def _restore_agent_model_runtime(agent, snapshot: dict | None) -> None: @@ -47,10 +46,10 @@ def _restore_agent_model_runtime(agent, snapshot: dict | None) -> None: except Exception: logger.debug("TUI one-turn model restore via primary runtime failed", exc_info=True) if hasattr(agent, "switch_model"): + model, provider, api_key, base_url, api_mode = (snapshot.get(k, "") for k in _RUNTIME_KEYS) agent.switch_model( - new_model=snapshot.get("model", ""), new_provider=snapshot.get("provider", ""), - api_key=snapshot.get("api_key", ""), base_url=snapshot.get("base_url", ""), - api_mode=snapshot.get("api_mode", ""), capabilities=snapshot.get("capabilities")) + new_model=model, new_provider=provider, api_key=api_key, base_url=base_url, + api_mode=api_mode, capabilities=snapshot.get("capabilities")) @contextlib.contextmanager @@ -65,7 +64,6 @@ def _session_profile_runtime_scope(session: dict): # Same terminal policy the gateway binds per turn: a docker-configured profile # must never resolve the launch process's pinned env. Failure → refusal scope. from tools.terminal_scope import install_profile_terminal_scope, reset_terminal_scope - terminal_token = install_profile_terminal_scope(Path(profile_home)) try: yield @@ -79,17 +77,14 @@ def _restart_completed_failed_agent_build(sid: str, session: dict, failed_ready: """Replace one completed failed build generation and start its retry.""" if failed_ready is None: return False - build_lock = session.setdefault("agent_build_lock", threading.Lock()) - with build_lock: - if ( - session.get("agent") is not None or session.get("agent_error") is None - or session.get("agent_ready") is not failed_ready or not failed_ready.is_set()): + with session.setdefault("agent_build_lock", threading.Lock()): + if (session.get("agent") is not None or session.get("agent_error") is None + or session.get("agent_ready") is not failed_ready or not failed_ready.is_set()): return False model_override = session.get("model_override") resume_overrides = session.get("resume_runtime_overrides") if isinstance(model_override, dict) and isinstance(resume_overrides, dict): - resume_overrides = dict(resume_overrides) - resume_overrides["model_override"] = model_override + resume_overrides = {**resume_overrides, "model_override": model_override} if provider := model_override.get("provider"): resume_overrides["provider_override"] = provider else: @@ -107,18 +102,11 @@ def _switch_request(raw_input: str, parsed_flags, persist_override) -> tuple[str """Normalize /model flags → (model_input, explicit_provider, one_turn, persist_global).""" from hermes_cli.model_switch import ( MODEL_SWITCH_ERR_ONCE_WITH_GLOBAL, MODEL_SWITCH_ERROR_TEXT, parse_model_switch_args, - resolve_persist_behavior, - ) + resolve_persist_behavior) - if parsed_flags is None: - parsed_flags = parse_model_switch_args(raw_input) - if hasattr(parsed_flags, "model_input"): - model_input, explicit_provider = parsed_flags.model_input, parsed_flags.explicit_provider - is_global_flag, is_session = parsed_flags.is_global, parsed_flags.is_session - one_turn = parsed_flags.is_once - else: - model_input, explicit_provider, is_global_flag, _force_refresh, is_session = parsed_flags - one_turn = False + f = parse_model_switch_args(raw_input) if parsed_flags is None else parsed_flags + model_input, explicit_provider, is_global_flag, is_session, one_turn = ( + f.model_input, f.explicit_provider, f.is_global, f.is_session, f.is_once) # Conflict validation is the shared parser's; surface it with the canonical copy. if is_global_flag and one_turn: raise ValueError(MODEL_SWITCH_ERROR_TEXT[MODEL_SWITCH_ERR_ONCE_WITH_GLOBAL]) @@ -134,13 +122,11 @@ def _current_model_runtime(agent, explicit_provider: str) -> tuple: """(provider, model, base_url, api_key) to switch from: live agent, else configured runtime.""" if agent: return tuple( - getattr(agent, k, "") or "" for k in ("provider", "model", "base_url", "api_key") - ) + getattr(agent, k, "") or "" for k in ("provider", "model", "base_url", "api_key")) current_model = _resolve_model() if explicit_provider: return explicit_provider.strip(), current_model, "", "" from hermes_cli.runtime_provider import resolve_runtime_provider - runtime = resolve_runtime_provider(requested=None) # Keep a callable api_key (Azure Entra bearer) unchanged: ``str()`` would # yield "" and poison switch_model validation. @@ -151,30 +137,14 @@ def _current_model_runtime(agent, explicit_provider: str) -> tuple: return provider, current_model, str(runtime.get("base_url", "") or ""), key -def _provider_context() -> tuple: - """(user providers, compatible custom providers, cfg) from config; all None on load failure.""" - user_provs = custom_provs = cfg = None - try: - from hermes_cli.config import get_compatible_custom_providers, load_config - - cfg = load_config() - user_provs = cfg.get("providers") - custom_provs = get_compatible_custom_providers(cfg) - except Exception: - pass - return user_provs, custom_provs, cfg - - def _merge_preflight_warning(result, agent, session: dict, cfg, custom_provs) -> None: """Fold the context-compression preflight warning into ``result`` (best-effort).""" try: from hermes_cli.context_switch_guard import merge_preflight_compression_warning - cfg_ctx = None - if isinstance(cfg, dict): - mc = cfg.get("model", {}) - if isinstance(mc, dict) and mc.get("context_length") is not None: - cfg_ctx = int(mc["context_length"]) + mc = cfg.get("model", {}) if isinstance(cfg, dict) else None + if isinstance(mc, dict) and mc.get("context_length") is not None: + cfg_ctx = int(mc["context_length"]) merge_preflight_compression_warning( result, agent=agent, messages=list(session.get("history", [])), custom_providers=custom_provs, config_context_length=cfg_ctx) @@ -186,7 +156,6 @@ def _expensive_model_confirm(result, current_base_url: str, current_api_key) -> """Deferred-confirm response when the selection guards flag the target model, else None.""" try: from hermes_cli.model_selection_guards import combined_selection_warning - warning = combined_selection_warning( result.new_model, provider=result.target_provider, base_url=result.base_url or current_base_url, api_key=result.api_key or current_api_key, model_info=result.model_info) @@ -194,12 +163,9 @@ def _expensive_model_confirm(result, current_base_url: str, current_api_key) -> warning = None if warning is None: return None - confirm_msg = warning.message - if result.warning_message: - confirm_msg = f"{confirm_msg}\n\n{result.warning_message}" - # Same contract as _set_model's deferred branch: confirm_message is - # canonical, warning is the legacy alias — keep identical. - return {"value": result.new_model, "warning": confirm_msg, "confirm_required": True, "confirm_message": confirm_msg} + msg = f"{warning.message}\n\n{result.warning_message}" if result.warning_message else warning.message + # Same contract as _set_model's deferred branch: confirm_message is canonical, warning legacy. + return {"value": result.new_model, "warning": msg, "confirm_required": True, "confirm_message": msg} def _commit_agent_switch(sid: str, session: dict, agent, result, current_model: str, snapshot): @@ -213,10 +179,8 @@ def _commit_agent_switch(sid: str, session: dict, agent, result, current_model: # The in-place swap rolled the agent back and re-raised. Abort the whole commit (worker # restart, persist, marker, override, config write) or the session pins a broken model. logger.warning("In-place model switch failed for TUI agent: %s", exc) - raise ValueError( - f"Model switch to {result.new_model} failed ({exc}); " - f"staying on {getattr(agent, 'model', current_model)}." - ) from exc + raise ValueError(f"Model switch to {result.new_model} failed ({exc}); " + f"staying on {getattr(agent, 'model', current_model)}.") from exc _restart_slash_worker(sid, session) _persist_live_session_runtime(session) _persist_live_session_system_prompt(session) @@ -233,7 +197,6 @@ def _apply_model_switch( pin_session_override: bool = True, parsed_flags: Any | None = None, persist_override: bool | None = None) -> dict: from hermes_cli.model_switch import switch_model - model_input, explicit_provider, one_turn, persist_global = _switch_request( raw_input, parsed_flags, persist_override) agent = session.get("agent") @@ -243,12 +206,17 @@ def _apply_model_switch( agent, explicit_provider) # User-defined providers let switch_model resolve named custom endpoints # (e.g. "ollama-launch") and validate against saved model lists. - user_provs, custom_provs, cfg = _provider_context() + user_provs = custom_provs = cfg = None + with contextlib.suppress(Exception): + from hermes_cli.config import get_compatible_custom_providers, load_config + cfg = load_config() + user_provs = cfg.get("providers") + custom_provs = get_compatible_custom_providers(cfg) result = switch_model( raw_input=model_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, - explicit_provider=explicit_provider, user_providers=user_provs, custom_providers=custom_provs, - ) + explicit_provider=explicit_provider, user_providers=user_provs, + custom_providers=custom_provs) if not result.success: raise ValueError(result.error_message or "model switch failed") restore_snapshot = _snapshot_agent_model_runtime(agent) if (one_turn and agent) else None @@ -276,12 +244,10 @@ def _apply_model_switch( def _sync_bot_capabilities(sid: str, session: dict) -> None: - """Rebuild a Bot Chat session's agent when its capability surface changed. - - Bot Chats are eternal sessions with toolsets/MCP baked in at construction, so a capability - edit would otherwise wait for /new. At turn start, fingerprint the profile's capabilities - and on change swap in a fresh agent for the SAME session (history is DB-backed). - """ + """Rebuild a Bot Chat session's agent when its capability surface changed. Bot Chats are + eternal sessions with toolsets/MCP baked in at construction, so a capability edit would + otherwise wait for /new: fingerprint at turn start and on change swap in a fresh agent for + the SAME session (history is DB-backed).""" agent = session.get("agent") if agent is None: return @@ -293,7 +259,6 @@ def _sync_bot_capabilities(sid: str, session: dict) -> None: if title != "Bot Chat": return from tools.bot_mode_probe import capability_fingerprint - current = capability_fingerprint(session.get("profile_home") or None) if current == "unavailable": return @@ -303,29 +268,23 @@ def _sync_bot_capabilities(sid: str, session: dict) -> None: return except Exception: return - try: tokens = _set_session_context(sid, cwd=_session_cwd(session)) try: - new_agent = _make_agent( - sid, session["session_key"], session_id=session["session_key"], platform_override=_session_source(session) - ) + new_agent = _make_agent(sid, session["session_key"], session_id=session["session_key"], + platform_override=_session_source(session)) finally: _clear_session_context(tokens) new_agent._session_title_hint = "Bot Chat" - session["agent"] = new_agent - session["config_model_seen"] = _config_model_target() + session.update(agent=new_agent, config_model_seen=_config_model_target()) _emit("notice", sid, {"message": "Capabilities updated — this bot's tools and prompt were refreshed."}) except Exception as e: logger.warning("Bot capability sync failed for %s: %s", sid, e) def _sync_agent_model_with_config(sid: str, session: dict) -> None: - """Adopt a config.yaml model change at turn start (like gateways do per message). - - Sessions pinned with /model keep their choice; a failed switch keeps the current - model and never blocks the turn. - """ + """Adopt a config.yaml model change at turn start (like gateways do per message). Sessions + pinned with /model keep their choice; a failed switch keeps the current model.""" agent = session.get("agent") if agent is None or session.get("model_override"): return @@ -335,36 +294,31 @@ def _sync_agent_model_with_config(sid: str, session: dict) -> None: seen = session.get("config_model_seen") # Record first so a broken config gets one attempt per edit, not per turn. session["config_model_seen"] = target - if target == seen: - return model, provider = target # Already on the configured model (resumed before first sync, or a config revert after # a failed switch): adopt without switching. - if model == getattr(agent, "model", "") and (not provider or provider == getattr(agent, "provider", "")): + if target == seen or ( + model == getattr(agent, "model", "") and (not provider or provider == getattr(agent, "provider", ""))): return raw = f"{model} --provider {provider}" if provider else model try: # This sync ADOPTS a config.yaml change; it must never write config back (that is # how `hermes --tui -m` once leaked into config.yaml). _apply_model_switch( - sid, session, raw, confirm_expensive_model=True, pin_session_override=False, persist_override=False - ) + sid, session, raw, confirm_expensive_model=True, pin_session_override=False, + persist_override=False) except Exception as e: _emit("error", sid, {"message": f"Could not switch to configured model {model}: {e}"}) def _pending_switch_selection_warning(model: str, provider: str) -> str | None: - """Selection-guard message for a model queued mid-turn, or ``None``. - - Runs BEFORE the pick is stashed, while the client can still turn the response into a - confirm prompt. Only pre-resolution inputs exist, so this can only under-fire; - ``_apply_model_switch`` is the backstop. Exceptions mean "no warning". - """ + """Selection-guard message for a model queued mid-turn, or ``None``. Runs BEFORE the pick is + stashed (the client can still turn the response into a confirm prompt); only pre-resolution + inputs exist so it can only under-fire — ``_apply_model_switch`` is the backstop.""" if not model: return None try: from hermes_cli.model_selection_guards import combined_selection_warning - warning = combined_selection_warning(model, provider=provider or None) except Exception: return None diff --git a/tui_gateway/project_tree.py b/tui_gateway/project_tree.py index 31c6dcaffb..126b99f539 100644 --- a/tui_gateway/project_tree.py +++ b/tui_gateway/project_tree.py @@ -12,11 +12,10 @@ from __future__ import annotations import re from typing import Any, Callable, Optional -# cwd -> ``{"repo_root", "worktree_root"}`` (COMMON main root shared across worktrees / -# this cwd's own checkout root); ``None`` when not in git or unprobeable (remote backend). +# cwd -> ``{"repo_root", "worktree_root"}`` (COMMON main root / this cwd's checkout root); +# ``None`` when not in git or unprobeable (remote backend). Resolve = Callable[[str], Optional[dict]] -# "does this directory still exist?"; defaults to always-True so callers that can't -# stat (remote backends) don't hide a project living on the other host. +# "does this directory still exist?"; always-True default keeps remote-host projects visible. Exists = Callable[[str], bool] # Only KANBAN-TASK worktrees (`/.worktrees/t_`, the id kanban_db mints) @@ -25,23 +24,21 @@ _KANBAN_DIR_RE = re.compile(r"^(.*[/\\]\.worktrees)[/\\]t_[0-9a-f]+[/\\]?$") _TRUNK_BRANCHES = {"main", "master", "trunk", "develop"} DEFAULT_BRANCH_LABEL = "main" -# Synthetic bucket for every session no project claimed (no cwd, bare home, HERMES -# state, deleted workspace); the id/flag name what it MEANS since membership keys off them. +# Synthetic bucket for every session no project claimed (no cwd, bare home, deleted workspace). NO_PROJECT_ID = "__no_project__" NO_PROJECT_LABEL = "Home" -# Sibling probes when recovering a deleted worktree's parent repo; each miss is a git -# invocation and real suffixes are one or two segments. +# Sibling probes when recovering a deleted worktree's parent repo (each miss is a git call). _MAX_SIBLING_PROBES = 4 def stamp_profile(projects: list[dict], profile: str) -> None: - """Stamp every session row with the request-scope profile (authoritative even - for legacy rows whose ``profile_name`` is NULL) for cross-profile routing.""" + """Stamp every session row with the request-scope profile (authoritative even for legacy + rows whose ``profile_name`` is NULL) for cross-profile routing.""" for project in projects: lanes = [g for repo in project.get("repos") or [] for g in repo.get("groups") or []] - lane_rows = [s for g in lanes for s in g.get("sessions") or []] - for session in (project.get("previewSessions") or []) + lane_rows: + for session in (project.get("previewSessions") or []) + [ + s for g in lanes for s in g.get("sessions") or []]: session["profile"] = profile @@ -59,15 +56,14 @@ def _segments(path: str) -> list[str]: def _is_windows_path(path: str) -> bool: - # Drive-letter (`C:\…`), UNC (`\\srv`, `//srv`), or any backslash-rooted path - # (`\wsl.localhost\…`, `\Users\…`). A single leading `/` stays POSIX. + # Drive-letter, UNC (`\\srv`, `//srv`) or backslash-rooted; a single leading `/` stays POSIX. value = (path or "").strip() return bool(re.match(r"^[A-Za-z]:[/\\]", value)) or value.startswith(("\\", "//")) def _comparison_segments(path: str) -> list[str]: - """Segments for identity comparison: Windows paths casefold (even when running - on POSIX); display paths and emitted IDs keep their spelling.""" + """Segments for identity comparison: Windows paths casefold (even on POSIX); display + paths and emitted IDs keep their spelling.""" segs = _segments(path) return [s.casefold() for s in segs] if _is_windows_path(path) else segs @@ -78,8 +74,7 @@ def _path_key(path: str) -> str: def _lane_key(path_or_lane: str) -> str: - """Canonicalize only the path portion of a lane id; branch labels stay - byte-preserved so equivalent Windows spellings don't fork lanes.""" + """Canonicalize only the path portion of a lane id (branch labels stay byte-preserved).""" marker = next((m for m in ("::branch::", "::kanban") if m in path_or_lane), None) if marker is None: return _path_key(path_or_lane) @@ -98,25 +93,15 @@ def kanban_worktree_dir(path: str) -> Optional[str]: return m.group(1) if m else None -def _with_base_name(path: str, name: str) -> str: - return re.sub(r"[^/\\]+$", name, (path or "").rstrip("/\\")) - - -def _parent_dir(path: str) -> str: - """The containing directory of ``path`` (``""`` once the root is passed).""" - return _with_base_name(path, "").rstrip("/\\") +def _with_base_name(path: str, name: str = "") -> str: + """Swap the last segment for ``name`` (``""`` -> the parent dir, ``""`` past the root).""" + return re.sub(r"[^/\\]+$", name, (path or "").rstrip("/\\")).rstrip("" if name else "/\\") def _field(row: dict, key: str) -> str: return (row.get(key) or "").strip() -def _branch_label(branch: str) -> str: - # An unrecorded branch folds into the one trunk lane so a repo never shows two - # "main" lanes (recorded "main" + the empty-branch bucket). - return (branch or "").strip() or DEFAULT_BRANCH_LABEL - - def _session_time(session: dict) -> float: return float(session.get("last_active") or session.get("started_at") or 0) @@ -131,12 +116,12 @@ def _placement( return { "repo_key": repo_root, "repo_label": base_name(repo_root) or repo_root, "lane_key": lane_key, "lane_label": lane_label, "lane_path": lane_path, - "is_main": is_main, "is_kanban": is_kanban, - } + "is_main": is_main, "is_kanban": is_kanban} def _trunk_placement(repo_root: str, branch: str) -> dict: - b = _branch_label(branch) + # An unrecorded branch folds into the trunk lane so a repo never shows two "main" lanes. + b = (branch or "").strip() or DEFAULT_BRANCH_LABEL return _placement(repo_root, _branch_lane_id(repo_root, b), b, repo_root, True, False) @@ -145,13 +130,9 @@ def _kanban_placement(repo_root: str, kanban_dir: str) -> dict: def _probe_sibling_worktree(cwd: str, resolve: Resolve) -> str: - """The parent repo root of a deleted ``-`` worktree, else ``""``. - - A deleted dir can't be probed, so trim one ``-`` at a time off its name - and return the first sibling that resolves. The cwd is often a SUBDIR of the dead - worktree (``-/apps/desktop``), so the trim runs on each ANCESTOR, - deepest first. Probes are bounded in total (each is a git invocation). - """ + """The parent repo root of a deleted ``-`` worktree, else ``""``: trim one + ``-`` at a time off each ancestor's name (the cwd is often a SUBDIR of the dead + worktree), deepest first, returning the first sibling that resolves; probes are bounded.""" probes = 0 path = (cwd or "").rstrip("/\\") while path and probes < _MAX_SIBLING_PROBES: @@ -163,7 +144,7 @@ def _probe_sibling_worktree(cwd: str, resolve: Resolve) -> str: info = resolve(_with_base_name(path, "-".join(parts[:i]))) if info and info.get("repo_root"): return (info["repo_root"] or "").strip() - path = _parent_dir(path) + path = _with_base_name(path) return "" @@ -174,14 +155,15 @@ def _place_by_heuristic(path: str) -> Optional[dict]: return None kanban_dir = kanban_worktree_dir(path) if kanban_dir: - return _kanban_placement(_parent_dir(kanban_dir), kanban_dir) + return _kanban_placement(_with_base_name(kanban_dir), kanban_dir) m = re.match(r"^(.+)-wt-(.+)$", base) if m: return _placement(_with_base_name(path, m.group(1)), path, m.group(2), path, False, False) return _placement(path, _branch_lane_id(path, DEFAULT_BRANCH_LABEL), base, path, True, False) -def _place(cwd: str, branch: str, resolve: Optional[Resolve], persisted_root: str) -> Optional[dict]: +def _place( + cwd: str, branch: str, resolve: Optional[Resolve], persisted_root: str) -> Optional[dict]: info = resolve(cwd) if resolve else None if info and info.get("repo_root") and info.get("worktree_root"): repo_root, worktree_root = info["repo_root"], info["worktree_root"] @@ -193,16 +175,14 @@ def _place(cwd: str, branch: str, resolve: Optional[Resolve], persisted_root: st label = base_name(worktree_root) or worktree_root return _placement(repo_root, worktree_root, label, worktree_root, False, False) - # No live probe: trust the backend-persisted root (split main by the recorded - # branch). Kanban tasks still collapse by path shape. + # No live probe: trust the persisted root; kanban tasks still collapse by path shape. if persisted_root: kanban_dir = kanban_worktree_dir(cwd) if kanban_dir: return _kanban_placement(persisted_root, kanban_dir) return _trunk_placement(persisted_root, branch) - # Unresolvable cwd: a deleted ``-`` worktree still belongs to its - # parent; absorb it into the trunk lane rather than stranding a dead-path lane. + # Unresolvable cwd: a deleted ``-`` worktree still belongs to its parent. sibling_root = _probe_sibling_worktree(cwd, resolve) if resolve else "" if sibling_root: return _trunk_placement(sibling_root, branch) @@ -228,8 +208,7 @@ def _session_repo_root(session: dict, resolve: Optional[Resolve]) -> str: def _lane_sort_key(group: dict) -> tuple: - # Trunk pins to the top; the kanban aggregate sinks to the bottom; the rest - # (branches + linked worktrees) sort by most-recent activity, then label. + # Trunk pins to the top, the kanban aggregate to the bottom; the rest by recency, then label. is_trunk = bool(group.get("isMain")) and group["label"].lower() in _TRUNK_BRANCHES return (0 if is_trunk else 1, 1 if group.get("isKanban") else 0, -_last_active(group.get("sessions") or []), group["label"].lower()) @@ -240,7 +219,6 @@ def _disambiguate_labels(items: list[dict]) -> None: by_label: dict[str, list[dict]] = {} for item in items: by_label.setdefault(item["label"], []).append(item) - for bucket in by_label.values(): pathed = [g for g in bucket if g.get("path")] if len(pathed) < 2: @@ -277,7 +255,6 @@ def _build_repos(sessions: list[dict], resolve: Optional[Resolve], hydrate: bool group["sessions"] = [] lanes[lane_identity] = (group, placement) lanes[lane_identity][0]["sessions"].append(session) - repos: dict[str, dict] = {} for group, placement in lanes.values(): group["sessions"].sort(key=_session_time, reverse=True) @@ -285,13 +262,11 @@ def _build_repos(sessions: list[dict], resolve: Optional[Resolve], hydrate: bool repo = repos.setdefault(_path_key(repo_key), _repo_node(repo_key, placement["repo_label"])) repo["groups"].append(group) repo["sessionCount"] += len(group["sessions"]) - repo_list = list(repos.values()) for repo in repo_list: repo["groups"] = sorted(repo["groups"], key=_lane_sort_key) _disambiguate_labels(repo["groups"]) - # Drop per-lane rows only AFTER sorting: _lane_sort_key derives recency - # from them. Counts were captured above, so the overview payload stays slim. + # Drop per-lane rows only AFTER sorting (_lane_sort_key reads them); counts stay. if not hydrate: for group in repo["groups"]: group["sessions"] = [] @@ -299,11 +274,10 @@ def _build_repos(sessions: list[dict], resolve: Optional[Resolve], hydrate: bool return repo_list -def _seed_folder_repos(repos: list[dict], folders: list[dict], resolve: Optional[Resolve]) -> list[dict]: - """Ensure every declared project folder shows as a repo, even with 0 sessions: - otherwise the desktop's entered-project view renders blank and the optimistic - live-session overlay has no lane for a fresh session until a full refresh. - Folders already covered by a session-derived repo (same git root) are untouched.""" +def _seed_folder_repos( + repos: list[dict], folders: list[dict], resolve: Optional[Resolve]) -> list[dict]: + """Ensure every declared project folder shows as a repo, even with 0 sessions (else the + entered-project view renders blank); folders covered by a session-derived repo are untouched.""" seen = {_path_key(v) for repo in repos for v in (repo.get("id"), repo.get("path")) if v} seeded = list(repos) for folder in folders or []: @@ -323,8 +297,7 @@ def _seed_folder_repos(repos: list[dict], folders: list[dict], resolve: Optional class _FolderIndex: - """Normalized folder path -> (owning project, depth): a session is matched by - walking its cwd's ancestors instead of scanning every project x folder.""" + """Normalized folder path -> (owning project, depth); matched by walking cwd ancestors.""" def __init__(self, projects: list[dict]) -> None: self._by_path: dict[str, tuple[dict, int]] = {} @@ -345,7 +318,8 @@ class _FolderIndex: return None, -1 -def _project_for_session(session: dict, index: _FolderIndex, resolve: Optional[Resolve]) -> Optional[dict]: +def _project_for_session( + session: dict, index: _FolderIndex, resolve: Optional[Resolve]) -> Optional[dict]: cwd = _field(session, "cwd") if not cwd: return None @@ -355,73 +329,49 @@ def _project_for_session(session: dict, index: _FolderIndex, resolve: Optional[R return max((index.match(t) for t in candidates), key=lambda hit: hit[1])[0] -def _session_cost(session: dict) -> float: - """A session's spend, billed if the provider reported it, else estimated.""" - for key in ("actual_cost_usd", "estimated_cost_usd"): - if session.get(key): - return float(session[key]) - return 0.0 - - def _project_node( pid: str, label: str, path: Optional[str], repos: list[dict], session_count: int, last_active: float, preview_sessions: list[dict], sessions: Optional[list[dict]] = None, - **flags: Any, -) -> dict: - """``flags`` overrides ``color`` / ``icon`` / ``isAuto`` / ``isNoProject`` (key order is - fixed by the defaults below — the renderer's wire shape).""" + **flags: Any) -> dict: + """``flags`` overrides ``color``/``icon``/``isAuto``/``isNoProject``; key order = wire shape.""" + rows = sessions or [] node = { "id": pid, "label": label, "path": path, "color": None, "icon": None, "isAuto": False, "isNoProject": False, "sessionCount": session_count, "lastActive": last_active, - # Totals over the same sessions `sessionCount` counts, so a project header - # adds up to what its rows show. - "totalTokens": sum((s.get("input_tokens") or 0) + (s.get("output_tokens") or 0) for s in sessions or []), - "totalCostUsd": sum(_session_cost(s) for s in sessions or []), - "repos": repos, "previewSessions": preview_sessions, - } + # Totals over the same sessions `sessionCount` counts (billed cost, else estimated). + "totalTokens": sum( + (s.get("input_tokens") or 0) + (s.get("output_tokens") or 0) for s in rows), + "totalCostUsd": sum( + float(s.get("actual_cost_usd") or s.get("estimated_cost_usd") or 0) for s in rows), + "repos": repos, "previewSessions": preview_sessions} node.update(flags) return node def _auto_buckets( unowned: list[dict], resolve: Optional[Resolve], junk: Callable, junk_cwd: Callable, - exists: Callable, -) -> tuple[dict[str, dict], list[dict]]: - """Group leftover sessions by auto-project root; the rest go to the Home bucket. - Prefer the common git root, then the session cwd for non-git workspaces (the - pre-Projects desktop grouped every cwd; dropping that flattens them into Recents).""" + exists: Callable) -> tuple[dict[str, dict], list[dict]]: + """Group leftover sessions by auto-project root (common git root, else the session cwd + for non-git workspaces); the rest go to the Home bucket.""" by_auto_root: dict[str, dict] = {} homeless: list[dict] = [] - - def _add_auto(root: str, session: dict) -> None: - key = _path_key(root) - if not key: - homeless.append(session) - return - by_auto_root.setdefault(key, {"root": root, "sessions": []})["sessions"].append(session) - for session in unowned: root = _session_repo_root(session, resolve) if root: - # A real git root uses the stricter repo policy; never reinterpret a - # filtered internal repo as a cwd-only project. A root no longer on - # disk is a stale persisted value and must not resurrect as a project. - if not junk(root) and exists(root): - _add_auto(root, session) - else: - homeless.append(session) - continue - cwd = _field(session, "cwd") - if not cwd or junk_cwd(cwd): - homeless.append(session) - continue - placement = _place_session(session, resolve) - # A placement that only echoes back an unresolvable cwd is the path-only - # heuristic guessing. If that dir is also gone from disk, promoting it - # mints a phantom project that can only be dismissed by hand -> Home. - if placement and exists(placement["repo_key"]): - _add_auto(placement["repo_key"], session) + # Stricter repo policy for real git roots; a root gone from disk is stale and + # must not resurrect as a project (never reinterpret it as a cwd-only project). + if junk(root) or not exists(root): + root = "" + elif (cwd := _field(session, "cwd")) and not junk_cwd(cwd): + # A path-only heuristic placement whose dir is gone from disk would mint a phantom + # project that can only be dismissed by hand -> Home. + placement = _place_session(session, resolve) + if placement and exists(placement["repo_key"]): + root = placement["repo_key"] + key = _path_key(root) if root else "" + if key: + by_auto_root.setdefault(key, {"root": root, "sessions": []})["sessions"].append(session) else: homeless.append(session) return by_auto_root, homeless @@ -431,35 +381,26 @@ def _home_project(homeless: list[dict], hydrate: bool, previews: list[dict]) -> """The synthetic Home bucket: no folder => no repo/lane structure, one lane carries the rows.""" lane = { "id": NO_PROJECT_ID, "label": NO_PROJECT_LABEL, "path": None, "isMain": False, - "isKanban": False, "sessions": homeless if hydrate else [], - } + "isKanban": False, "sessions": homeless if hydrate else []} home_repo = { "id": NO_PROJECT_ID, "label": NO_PROJECT_LABEL, "path": None, "groups": [lane], - "sessionCount": len(homeless), - } + "sessionCount": len(homeless)} return _project_node( NO_PROJECT_ID, NO_PROJECT_LABEL, None, [home_repo], len(homeless), _last_active(homeless), previews, homeless, isNoProject=True) def build_tree( - projects: list[dict], - sessions: list[dict], - discovered_repos: list[dict], - resolve: Optional[Resolve] = None, - *, - preview_limit: int = 3, - hydrate: bool = False, + projects: list[dict], sessions: list[dict], discovered_repos: list[dict], + resolve: Optional[Resolve] = None, *, preview_limit: int = 3, hydrate: bool = False, is_junk_root: Optional[Callable[[str], bool]] = None, - is_junk_cwd: Optional[Callable[[str], bool]] = None, - exists: Optional[Exists] = None) -> dict: + is_junk_cwd: Optional[Callable[[str], bool]] = None, exists: Optional[Exists] = None) -> dict: """Build the authoritative project tree -> ``{"projects", "scoped_session_ids"}``. - ``is_junk_root`` flags git roots that must never become an AUTO project (bare home, - HERMES_HOME); ``is_junk_cwd`` is the narrower policy for non-git folders; explicit - projects are honored regardless. ``exists`` keeps a DELETED workspace from becoming - a phantom AUTO project (omit on remote backends). ``hydrate`` False (overview) empties - lane ``sessions`` but keeps counts + ``preview_limit`` ``previewSessions``. + ``is_junk_root`` flags git roots that must never become an AUTO project; ``is_junk_cwd`` + is the narrower non-git policy (explicit projects are honored regardless); ``exists`` + keeps a DELETED workspace from becoming a phantom AUTO project (omit on remote backends). + ``hydrate`` False empties lane ``sessions`` but keeps counts + ``previewSessions``. """ active_projects = [p for p in projects if not p.get("archived")] _junk = is_junk_root or (lambda _root: False) @@ -482,7 +423,6 @@ def build_tree( def _scope(project_sessions: list[dict]) -> None: scoped_ids.extend(s["id"] for s in project_sessions if s.get("id")) - # Tier 1: explicit, user-created projects (always shown, even with 0 sessions). for project in active_projects: psessions = by_project.get(project["id"], []) @@ -513,8 +453,7 @@ def build_tree( repo_node["sessionCount"], _last_active(auto_sessions), _previews(auto_sessions), auto_sessions, isAuto=True)) - # Tier 3: repos discovered from full history / disk scan with no loaded - # sessions, folded to their common root and not owned by an explicit project. + # Tier 3: discovered repos with no loaded sessions, folded to their common root. for repo in discovered_repos or []: raw_root = _field(repo, "root") if not raw_root: @@ -530,12 +469,10 @@ def build_tree( root, label, root, [_repo_node(root, label)], int(repo.get("sessions") or 0), float(repo.get("last_active") or 0), [], isAuto=True)) - # Auto projects are labelled by repo basename, which can collide; grow path - # prefixes so each is distinct. Explicit projects keep their user-chosen names. + # Auto-project basename labels can collide; explicit projects keep their user-chosen names. _disambiguate_labels([p for p in result if p.get("isAuto")]) - # Tier 0: everything above could not place, so the grouped view loses no - # session. Leads the list; omitted entirely when empty. + # Tier 0: whatever the tiers above could not place. Leads the list; omitted when empty. if homeless: homeless.sort(key=_session_time, reverse=True) _scope(homeless) diff --git a/tui_gateway/prompt_attachments.py b/tui_gateway/prompt_attachments.py index 60d1687ce3..bdb749e4ef 100644 --- a/tui_gateway/prompt_attachments.py +++ b/tui_gateway/prompt_attachments.py @@ -28,20 +28,17 @@ del _re # bodies are rebound onto server globals: import inside functions only def _b64_payload(raw: str, data_url_re: str, flags: int) -> bytes: - """Strip an optional ``data:...;base64,`` wrapper and all whitespace, then - strictly decode (raises ``binascii.Error``/``ValueError`` on bad base64).""" + """Strip an optional ``data:...;base64,`` wrapper and all whitespace, then strictly decode.""" import base64 as _base64 import re as _re cleaned = (raw or "").strip() - m = _re.match(data_url_re, cleaned, flags) - if m: + if m := _re.match(data_url_re, cleaned, flags): cleaned = m.group(1) return _base64.b64decode(_re.sub(r"\s+", "", cleaned), validate=True) def _decode_attach_base64(raw: str, *, mime_prefix: str) -> bytes | None: - """Decode a base64 payload, optionally ``data:...;base64,``-wrapped, - tolerating embedded whitespace. ``None`` when not valid base64.""" + """Decode a (``data:...;base64,``-wrapped) payload; None when invalid.""" import re as _re try: return _b64_payload( @@ -52,8 +49,7 @@ def _decode_attach_base64(raw: str, *, mime_prefix: str) -> bytes | None: def _decode_attach_payload( rid, raw_b64: str, *, mime_prefix: str, max_bytes: int, label: str, empty_msg: str): - """``(bytes, None)`` or ``(None, error)`` for an upload: 4017 on bad/empty - base64, 4018 over *max_bytes*.""" + """``(bytes, None)`` or ``(None, error)``: 4017 on bad/empty base64, 4018 over *max_bytes*.""" data = _decode_attach_base64(raw_b64, mime_prefix=mime_prefix) if data is None: return None, _err(rid, 4017, "data is not valid base64") @@ -66,8 +62,7 @@ def _decode_attach_payload( def _sniff_image_ext(img_bytes: bytes, filename: str = "") -> str: - """Extension from the filename hint, else magic bytes (WebP needs the RIFF/WEBP - container check), else ``.png``.""" + """Extension from the filename hint, else magic bytes (WebP: RIFF container), else ``.png``.""" if filename and (suffix := Path(filename).suffix.lower()): return suffix head = img_bytes[:16] @@ -85,26 +80,20 @@ def _allowed_image_extensions() -> frozenset[str]: def _session_home_dir(session: dict, name: str) -> Path: - """``/``, anchored on the session's stored ``profile_home``. - - Attach RPCs run BEFORE ``prompt.submit`` installs the profile HERMES_HOME - override, so ``get_hermes_home()`` would return the gateway's launch home — - while the sandbox mounts and the vision host-read allowlist resolve the - *session profile's* dirs at run time. Writing anywhere else means the agent - can never see the file. - """ + """``/``, anchored on the session's stored ``profile_home``: attach + RPCs run BEFORE ``prompt.submit`` installs the profile HERMES_HOME override, while + the sandbox mounts and the vision host-read allowlist resolve the *session profile's* + dirs at run time — writing anywhere else means the agent can never see the file.""" profile_home = session.get("profile_home") return (Path(profile_home) if profile_home else _hermes_home) / name def _session_images_dir(session: dict) -> Path: - """Uploads ``images/`` dir for the session (see ``_session_home_dir``).""" return _session_home_dir(session, "images") def _queue_attached_image(session: dict, img_bytes: bytes, ext: str, *, prefix: str) -> Path: - """Write image bytes into the session images dir and append to - ``session["attached_images"]`` so the next ``prompt.submit`` picks them up.""" + """Write image bytes into the session images dir and queue them for the next submit.""" session["image_counter"] = session.get("image_counter", 0) + 1 img_dir = _session_images_dir(session) img_dir.mkdir(parents=True, exist_ok=True) @@ -120,8 +109,7 @@ def _queue_attached_image(session: dict, img_bytes: bytes, ext: str, *, prefix: def _format_ref_value(value: str) -> str: - """Quote a context-ref value containing whitespace/brackets/quotes so the staged - ``@file:`` ref round-trips through ``agent.context_references``.""" + """Quote a value with whitespace/brackets/quotes so the ``@file:`` ref round-trips.""" if not value or not _ATTACHMENT_REF_NEEDS_QUOTING_RE.search(value): return value for q in ("`", '"', "'"): @@ -139,74 +127,33 @@ def _attachment_ref_path(session: dict, target: Path) -> str: return str(target.resolve()) -def _desktop_attachment_dir(session: dict) -> Path: - """File-attachment staging dir (``attachments/``, see ``_session_home_dir``); - registered in ``tools.credential_files._CACHE_DIRS`` and auto-mounted into - containers, so a staged file lands where the bind mount points.""" - root = _session_home_dir(session, "attachments") - root.mkdir(parents=True, exist_ok=True) - return root - - def _sanitize_attachment_name(name: str) -> str: import re as _re candidate = _re.sub(r"[\x00-\x1f]+", "_", Path(str(name or "").strip()).name) return candidate.strip().strip(".") or "attachment" -def _unique_attachment_path(root: Path, filename: str) -> Path: - candidate = root / filename - if not candidate.exists(): - return candidate - stem = Path(filename).stem or "attachment" - suffix = Path(filename).suffix - counter = 2 - while (next_candidate := root / f"{stem}-{counter}{suffix}").exists(): - counter += 1 - return next_candidate - - -def _resolve_gateway_attachment_path(raw: str) -> Path | None: - """Resolve a raw path token to a gateway-visible file, or None.""" - if not raw: - return None - try: - from cli import _detect_file_drop, _resolve_attachment_path, _split_path_input - except Exception: - return None - dropped = _detect_file_drop(raw) - if dropped: - return Path(dropped["path"]).resolve() - path_token, _remainder = _split_path_input(raw) - resolved = _resolve_attachment_path(path_token) - return Path(resolved).resolve() if resolved is not None else None - - -def _decode_attachment_data_url(data_url: str) -> bytes: - """Decode a ``data:;base64,`` payload (any media type, unlike the - image-specific ``_decode_attach_base64``); bare base64 also accepted.""" - import binascii as _binascii - import re as _re - try: - return _b64_payload( - data_url, r"^data:[^;,]*(?:;[^;,=]+=[^;,]+)*;base64,(.*)$", _re.DOTALL | _re.I) - except (ValueError, _binascii.Error) as exc: - raise ValueError("invalid data_url payload") from exc - - def _stage_session_file_attachment( session: dict, *, raw_path: str, data_url: str, name: str) -> tuple[Path, bool]: - """Make a desktop file attachment available to the gateway agent. - - 1. Path resolves INSIDE the session workspace -> use as-is (``uploaded=False``). - 2. Gateway-visible file OUTSIDE the workspace -> copy into ``attachments/`` - (bind-mounted into container backends) so ``@file:`` resolves in the sandbox. - 3. Not on the gateway (remote client disk) -> decode ``data_url`` bytes into - ``attachments/``. - Returns ``(stored_path, uploaded)``. - """ + """Make a desktop file attachment available to the gateway agent: ``(stored_path, uploaded)``. + Inside the workspace -> as-is; gateway-visible but outside -> copied into ``attachments/`` + (bind-mounted into container backends so ``@file:`` resolves in the sandbox); not on the + gateway -> ``data_url`` bytes decoded into ``attachments/``.""" workspace = Path(_session_cwd(session)).resolve() - resolved = _resolve_gateway_attachment_path(raw_path) + resolved = None + if raw_path: + try: + from cli import _detect_file_drop, _resolve_attachment_path, _split_path_input + except Exception: + _detect_file_drop = None + if _detect_file_drop is not None: + dropped = _detect_file_drop(raw_path) + if dropped: + resolved = Path(dropped["path"]).resolve() + else: + path_token, _remainder = _split_path_input(raw_path) + found = _resolve_attachment_path(path_token) + resolved = Path(found).resolve() if found is not None else None if resolved is not None: try: resolved.relative_to(workspace) @@ -217,10 +164,25 @@ def _stage_session_file_attachment( else: if not data_url: raise ValueError("file not found on gateway and no data_url provided") - payload = _decode_attachment_data_url(data_url) + # Any media type (unlike the image-specific decoder); bare base64 also accepted. + import binascii as _binascii + import re as _re + try: + payload = _b64_payload( + data_url, r"^data:[^;,]*(?:;[^;,=]+=[^;,]+)*;base64,(.*)$", _re.DOTALL | _re.I) + except (ValueError, _binascii.Error) as exc: + raise ValueError("invalid data_url payload") from exc filename = _sanitize_attachment_name(name or Path(str(raw_path or "")).name) - target = _unique_attachment_path( - _desktop_attachment_dir(session), _sanitize_attachment_name(filename)) + root = _session_home_dir(session, "attachments") + root.mkdir(parents=True, exist_ok=True) + filename = _sanitize_attachment_name(filename) + target = root / filename + if target.exists(): + stem = Path(filename).stem or "attachment" + suffix = Path(filename).suffix + counter = 2 + while (target := root / f"{stem}-{counter}{suffix}").exists(): + counter += 1 target.write_bytes(payload) return target.resolve(), True diff --git a/tui_gateway/prompt_turn.py b/tui_gateway/prompt_turn.py index 97fb49364f..3cc33ad224 100644 --- a/tui_gateway/prompt_turn.py +++ b/tui_gateway/prompt_turn.py @@ -21,41 +21,37 @@ def _hook_failure(what: str, exc: BaseException) -> None: def _is_successful_goal_turn(result: Any, status: str, raw: Any) -> bool: - """Return whether a turn produced a real response the goal judge can use.""" + """Whether a turn produced a real response the goal judge can use.""" return bool( status == "complete" and isinstance(raw, str) and raw.strip() and not (isinstance(result, dict) and result.get("failed")) and not (isinstance(result, dict) and result.get("completed") is False)) -def _goal_max_turns() -> int: +def _active_goal_manager(session: dict): + """The session's GoalManager when a goal is active, else None.""" + from hermes_cli.goals import GoalManager try: - goals_cfg = _load_cfg().get("goals") or {} - return int(goals_cfg.get("max_turns", 20) or 20) + max_turns = int((_load_cfg().get("goals") or {}).get("max_turns", 20) or 20) except Exception: - return 20 + max_turns = 20 + goal_mgr = GoalManager( + session_id=str(session.get("session_key") or ""), default_max_turns=max_turns) + return goal_mgr if goal_mgr.is_active() else None def _plan_goal_compression_recovery( session: dict, result: Any, *, status: str, raw: Any) -> tuple[str | None, str | None]: - """Plan a bounded active-goal retry after compression exhaustion. - - Exhaustion is a failed turn: never judge input, never a spent goal turn. One - fresh continuation is allowed; if that also exhausts, pause the goal instead of - spinning until a random user message wakes it. Returns - ``(continuation_prompt, status_notice)``; no active goal -> ``(None, None)``. - """ - compression_exhausted = bool(isinstance(result, dict) and result.get("compression_exhausted")) - if not compression_exhausted: + """Bounded active-goal retry after compression exhaustion: ``(continuation, notice)``. + Exhaustion is a failed turn (never judge input, never a spent goal turn); one fresh + continuation is allowed, a second exhaustion pauses the goal instead of spinning.""" + if not (isinstance(result, dict) and result.get("compression_exhausted")): if _is_successful_goal_turn(result, status, raw): session.pop(_GOAL_COMPRESSION_RECOVERY_ATTEMPTS, None) return None, None - from hermes_cli.goals import GoalManager - sid_key = str(session.get("session_key") or "") - if not sid_key: + if not str(session.get("session_key") or ""): return None, None - goal_mgr = GoalManager(session_id=sid_key, default_max_turns=_goal_max_turns()) - if not goal_mgr.is_active(): + if (goal_mgr := _active_goal_manager(session)) is None: session.pop(_GOAL_COMPRESSION_RECOVERY_ATTEMPTS, None) return None, None goal_created_at = float(getattr(goal_mgr.state, "created_at", 0.0) or 0.0) @@ -66,39 +62,29 @@ def _plan_goal_compression_recovery( isinstance(recovery_state, dict) and recovery_state.get("goal_created_at") == goal_created_at and recovery_state.get("goal") == goal_text): - try: + with contextlib.suppress(TypeError, ValueError): attempts = int(recovery_state.get("attempts", 0) or 0) - except (TypeError, ValueError): - attempts = 0 continuation_prompt = goal_mgr.next_continuation_prompt() if attempts < _GOAL_COMPRESSION_RECOVERY_LIMIT and continuation_prompt: session[_GOAL_COMPRESSION_RECOVERY_ATTEMPTS] = { "goal_created_at": goal_created_at, "goal": goal_text, "attempts": attempts + 1} return ( - continuation_prompt, "Context compression was exhausted. Retrying the active goal once." - ) + continuation_prompt, + "Context compression was exhausted. Retrying the active goal once.") goal_mgr.pause(reason="context compression exhausted twice consecutively") # A later explicit /goal resume gets a fresh bounded recovery cycle. session.pop(_GOAL_COMPRESSION_RECOVERY_ATTEMPTS, None) - return ( - None, + return None, ( "Goal paused after context compression was exhausted twice. " "Run /compress, then /goal resume to continue.") -# ── turn admission ─────────────────────────────────────────────────── - - def _admit_prompt_turn( sid: str, session: dict, text: Any, image_paths: list[str] | None, queued_prompt_generation: int | None) -> tuple[list[str], Any] | None: - """Ownership + liveness gate every fresh turn source must cross. - - prompt.submit claims the slot in its RPC handler, but auto-continue, wake-ups - and other synthesized turns call ``_run_prompt_submit`` directly — the bypass - that once let a second backend run a duplicate turn. Returns - ``(images, agent)`` or None when refused (``running`` already reset). - """ + """Ownership + liveness gate every turn source must cross; ``(images, agent)`` or None. + Synthesized turns (auto-continue, wake-ups) call ``_run_prompt_submit`` directly — the + bypass that once let a second backend run a duplicate turn.""" if (ownership_refusal := _ensure_active_session_slot(sid, session)) is not None: logger.info( "Refusing turn for session %s at _run_prompt_submit: %s", @@ -114,33 +100,25 @@ def _admit_prompt_turn( and int(session.get("_queued_prompt_generation", 0)) != queued_prompt_generation): session["running"] = False return None + images = list(session.get("attached_images", []) if image_paths is None else image_paths) if image_paths is None: - images = list(session.get("attached_images", [])) session["attached_images"] = [] - else: - images = list(image_paths) inflight = session.get("inflight_turn") # A retained failed turn (see _fail_inflight_turn) is a stale leftover # by the time a new turn starts — replace it, never append onto it. if not isinstance(inflight, dict) or inflight.get("status") == "error": _start_inflight_turn(session, text) agent = session["agent"] - if hasattr(agent, "clear_interrupt"): - with contextlib.suppress(Exception): - agent.clear_interrupt() + with contextlib.suppress(Exception): + agent.clear_interrupt() return images, agent def _record_turn_marker(session: dict, text: Any) -> str: - """Write the durable crash marker; returns the session key it was written under. - - Retired when the outcome reaches the client; a surviving marker means the - process died mid-turn and session.resume auto-continues from it. Compression - can rotate session_key mid-turn, so the caller keeps this key. The key is - published before the disk write so an interrupt racing startup can retire it; - the post-write cancel check closes the inverse race (Stop landed first, no - file to clear yet). - """ + """Write the durable crash marker; returns the session key it was written under (compression + can rotate session_key mid-turn). A surviving marker means the process died mid-turn. + The key is published before the disk write so an interrupt racing startup can retire + it; the post-write cancel check closes the inverse race (Stop landed first, no file).""" marker_home = _session_home(session) marker_key = str(session.get("session_key") or "") marker_attempt = int(session.pop("_auto_continue_attempt", 0) or 0) @@ -156,9 +134,6 @@ def _record_turn_marker(session: dict, text: Any) -> str: return marker_key -# ── per-turn scopes ────────────────────────────────────────────────── - - @dataclasses.dataclass(slots=True) class _TurnScopes: """Reset tokens for the thread/context scopes a turn binds (filled incrementally).""" @@ -170,67 +145,10 @@ class _TurnScopes: terminal: Any = None -def _bind_turn_scopes(sid: str, session: dict, scopes: _TurnScopes) -> None: - """Bind approval/session/profile/terminal scopes for this turn thread. - - Fills ``scopes`` field by field so a failure midway still leaves every bound - token for ``_release_turn_scopes``. The profile's COMPLETE terminal policy is - bound too: terminal_tool otherwise reads the launch process's pinned env, and - a failed install leaves a refusal scope so terminal tools fail closed. - """ - from tools.approval import set_current_session_key - scopes.approval = set_current_session_key(session["session_key"]) - scopes.session_tokens = _set_session_context(session["session_key"], ui_session_id=sid) - profile_home = session.get("profile_home") - if profile_home: - scopes.home = set_hermes_home_override(profile_home) - scopes.secret = set_secret_scope(build_profile_secret_scope(Path(profile_home))) - from tools.terminal_scope import install_profile_terminal_scope - scopes.terminal = install_profile_terminal_scope(Path(profile_home)) - # The sudo password callback is thread-local, so the build thread's wiring - # doesn't reach this turn thread — sudo prompts would fall through to - # /dev/tty and hang the headless gateway. Re-wire to the sudo.request - # overlay (secret capture is a module global; re-running is a no-op). - _wire_callbacks(sid) - - -def _release_turn_scopes(scopes: _TurnScopes) -> None: - with contextlib.suppress(Exception): - if scopes.approval is not None: - from tools.approval import reset_current_session_key - reset_current_session_key(scopes.approval) - if scopes.home is not None: - reset_hermes_home_override(scopes.home) - if scopes.secret is not None: - reset_secret_scope(scopes.secret) - if scopes.terminal is not None: - from tools.terminal_scope import reset_terminal_scope - reset_terminal_scope(scopes.terminal) - _clear_session_context(scopes.session_tokens) - - -# ── message resolution ─────────────────────────────────────────────── - - -def _expand_context_references(agent, prompt: str, cwd: str): - """Expand ``@file`` references; returns the preprocess result (``.blocked``/``.message``).""" - from agent.context_references import preprocess_context_references - from agent.model_metadata import get_model_context_length - ctx_len = get_model_context_length( - getattr(agent, "model", "") or _resolve_model(), - base_url=getattr(agent, "base_url", "") or "", api_key=getattr(agent, "api_key", "") or "", - provider=getattr(agent, "provider", "") or "", - config_context_length=getattr(agent, "_config_context_length", None)) - return preprocess_context_references(prompt, cwd=cwd, allowed_root=cwd, context_length=ctx_len) - - def _route_turn_images(agent, prompt: Any, images: list[str]) -> Any: - """Build the run message for a turn with attached images. - - "native" passes pixels as OpenAI-style content parts; "text" references the - paths so the agent analyzes them in-loop with vision_analyze, never blocking - the submit path on vision calls. Decision table: agent/image_routing.py. - """ + """Run message for a turn with attached images: "native" content parts, or "text" path + references the agent analyzes in-loop (never blocking submit on vision calls). + Decision table: agent/image_routing.py.""" try: from agent.image_routing import build_native_content_parts, decide_image_input_mode from hermes_cli.config import load_config as _tui_load_config @@ -241,9 +159,8 @@ def _route_turn_images(agent, prompt: Any, images: list[str]) -> Any: if getattr(agent, "api_mode", "") == "codex_app_server": mode = "text" except Exception as _img_exc: - print( - f"[tui_gateway] image_routing decision failed, defaulting to text: {_img_exc}", - file=sys.stderr) + print(f"[tui_gateway] image_routing decision failed, defaulting to text: {_img_exc}", + file=sys.stderr) mode = "text" if mode != "native": return _build_image_ref_message(prompt, images) @@ -256,22 +173,15 @@ def _route_turn_images(agent, prompt: Any, images: list[str]) -> Any: if any(p.get("type") == "image_url" for p in parts): return parts except Exception as _img_exc: - print( - f"[tui_gateway] native attach failed, falling back to text: {_img_exc}", - file=sys.stderr, - ) + print(f"[tui_gateway] native attach failed, falling back to text: {_img_exc}", + file=sys.stderr) return _build_image_ref_message(prompt, images) def _start_turn_voice() -> tuple[Any, bool]: - """Arm voice-mode turn audio; returns ``(tts_queue, thinking_started)``. - - ``_tts_stream_begin`` goes first: cutting a still-speaking previous turn IS - this turn's barge-in, so it must latch before the caller consumes the latch. - The full-duplex listener lets the user interject DURING generation. The - "thinking" sound keeps long silences from reading as a dead session; its - gate skips while TTS plays or the mic captures; stopped in the turn's finally. - """ + """Arm voice-mode turn audio; ``(tts_queue, thinking_started)``. ``_tts_stream_begin`` + goes first: cutting a still-speaking previous turn IS this turn's barge-in, so it must + latch before the caller consumes the latch.""" tts_queue = _tts_stream_begin() if not _voice_mode_enabled(): return tts_queue, False @@ -293,115 +203,12 @@ def _start_turn_voice() -> tuple[Any, bool]: return tts_queue, False -def _stop_thinking_sound() -> None: - with contextlib.suppress(Exception): - from tools.voice_mode import stop_thinking_sound - stop_thinking_sound() - - -def _apply_turn_notes(run_message: Any, session: dict) -> Any: - """Prepend the per-turn API-message notes (same enrichment channel as images): - barge mid-speech, reactions since the last turn, then which window the message - was typed into (HUD mode is per-turn state; it cannot live in the byte-stable - system prompt).""" - from tools.tts_streaming import SPEECH_INTERRUPTED_NOTE, take_speech_interrupted - if take_speech_interrupted(): - run_message = _prepend_note(run_message, SPEECH_INTERRUPTED_NOTE) - run_message = _prepend_note(run_message, _pending_reaction_notes(session)) - return _prepend_note(run_message, _hud_surface_note(session)) - - -def _build_run_kwargs( - agent, session: dict, history: list, prompt: Any, images: list[str], run_message: Any, - stream_cb, display_kind: str | None, display_metadata: dict | None) -> dict: - """Assemble ``run_conversation`` kwargs, feature-detecting optional parameters. - - A synthesized turn is typed at turn START so the crash persist writes its row - as a timeline event, not a raw user bubble (forever, if the turn never ends). - The post-turn stamp is the fallback for an older agent; re-stamping is a no-op. - """ - run_kwargs = { - "conversation_history": list(history), - "stream_callback": stream_cb, - "persist_user_message": ( - _build_persist_user_message(prompt, images, run_message) if images else prompt)} - try: - run_params = inspect.signature(agent.run_conversation).parameters - except (TypeError, ValueError): - run_params = {} - if "task_id" in run_params: - run_kwargs["task_id"] = session["session_key"] - if display_kind and "persist_user_display_kind" in run_params: - run_kwargs["persist_user_display_kind"] = display_kind - run_kwargs["persist_user_display_metadata"] = display_metadata - return run_kwargs - - -# ── post-run bookkeeping ───────────────────────────────────────────── - - -def _stamp_synthetic_display_kind( - agent, session: dict, result: Any, text: str, display_kind: str, display_metadata: dict | None -) -> None: - """Post-turn fallback stamp of a synthesized turn's display kind (DB row + result).""" - db = getattr(agent, "_session_db", None) - current_session_id = getattr(agent, "session_id", None) or session.get("session_key") - if db is not None: - try: - db.set_latest_matching_message_display_kind( - current_session_id, role="user", content=text, display_kind=display_kind, - display_metadata=display_metadata) - except Exception: - logger.debug("failed to stamp synthetic display kind", exc_info=True) - if isinstance(result, dict) and isinstance(result.get("messages"), list): - for message in reversed(result["messages"]): - if message.get("role") == "user" and message.get("content") == text: - message["display_kind"] = display_kind - if display_metadata: - message["display_metadata"] = display_metadata - break - - -def _restore_moa_one_shot(sid: str, session: dict) -> None: - """Undo a /moa one-shot after its turn — through the switch path, because the - one-shot did a real in-place ``agent.switch_model()``; resetting - ``model_override`` alone would leave the live client pinned to MoA.""" - _restore = session.pop("moa_one_shot_restore", None) - if isinstance(_restore, dict): - _prev_override = _restore.get("override") - _prev_model = _restore.get("model") - _prev_provider = _restore.get("provider") - if _prev_override is None: - session.pop("model_override", None) - else: - session["model_override"] = _prev_override - if _prev_model: - _raw = f"{_prev_model} --provider {_prev_provider}" if _prev_provider else _prev_model - try: - _apply_model_switch( - sid, session, _raw, confirm_expensive_model=False, - pin_session_override=bool(_prev_override), - persist_override=False, # session-internal restore, never config.yaml - ) - except Exception as _moa_restore_exc: - logger.warning("MoA one-shot model restore failed: %s", _moa_restore_exc) - elif _restore is None: - session.pop("model_override", None) - else: - session["model_override"] = _restore - - def _commit_turn_history( session: dict, result: dict, history: list, history_version: int) -> str | None: """Write the agent's messages back to session history; returns a client warning or None. - - Caller holds no lock. If history_version moved during the turn, the only - tolerated mutation is a pivot marker the gateway itself inserted mid-turn - (model switch, /personality); then the output is merged after the current - history. ``_append_model_switch_marker`` strips prior markers in place then - appends, so the delta is NOT a tail slice — compare content, not indices. - Any other desync (undo/compress/retry/rollback) is surfaced, never dropped. - """ + If history_version moved mid-turn, the only tolerated mutation is a gateway-inserted + pivot marker (compare content, not indices: ``_append_model_switch_marker`` strips prior + markers in place); any other desync is surfaced, never dropped.""" with session["history_lock"]: current_version = int(session.get("history_version", 0)) if current_version == history_version: @@ -411,15 +218,11 @@ def _commit_turn_history( current_history = list(session["history"]) history_no_markers = [e for e in history if not _is_pivot_marker(e)] current_no_markers = [e for e in current_history if not _is_pivot_marker(e)] - pivot_only = current_no_markers == history_no_markers and any( - _is_pivot_marker(e) for e in current_history) - if pivot_only: - # Auto-compression can make result["messages"] shorter than the - # turn-start history; then the full result is the base. - if len(result["messages"]) > len(history): - new_messages = result["messages"][len(history):] - else: - new_messages = list(result["messages"]) + if current_no_markers == history_no_markers and any( + _is_pivot_marker(e) for e in current_history): + # Auto-compression can leave the result shorter than the turn-start history. + msgs = result["messages"] + new_messages = msgs[len(history):] if len(msgs) > len(history) else list(msgs) session["history"] = current_history + new_messages session["history_version"] = current_version + 1 return None @@ -433,52 +236,37 @@ def _commit_turn_history( "but was not saved to session history.") +def _result_status(result: dict) -> str: + return ( + "interrupted" if result.get("interrupted") + else "error" if result.get("error") else "complete") + + def _turn_outcome(result: Any) -> tuple[Any, str, str | None]: """Reduce a run_conversation result to ``(raw_text, status, last_reasoning)``.""" if not isinstance(result, dict): return str(result), "complete", None raw = result.get("final_response", "") - status = ( - "interrupted" if result.get("interrupted") - else "error" if result.get("error") else "complete") - # No visible response AND a real error (e.g. invalid model slug -> provider - # 4xx): surface the error as the text (classic CLI parity) instead of an - # empty turn. An empty successful turn still renders as empty. + status = _result_status(result) + # No visible response AND a real error: surface the error as the text (classic CLI + # parity). An empty successful turn still renders as empty. if (not raw) and result.get("error") and (result.get("failed") or result.get("partial")): raw = f"Error: {result.get('error')}" # "Operation interrupted: waiting for model response (…)" is cancellation # metadata, not assistant prose (gateway/run.py and ACP suppress it too). if status == "interrupted" and isinstance(raw, str) and raw.strip().startswith( - INTERRUPT_WAITING_FOR_MODEL_PREFIX): + INTERRUPT_WAITING_FOR_MODEL_PREFIX): raw = "" lr = result.get("last_reasoning") last_reasoning = lr.strip() if isinstance(lr, str) and lr.strip() else None return raw, status, last_reasoning -def _turn_error_surface(agent, result: Any) -> Any: - """{layer, code, retryable} descriptor for an error result (advisory, never raises).""" - try: - from agent.error_surface import build_error_surface_from_result - return build_error_surface_from_result( - result, provider=str(getattr(agent, "provider", "") or ""), - model=str(getattr(agent, "model", "") or "")) - except Exception: - return None - - -# ── post-turn hooks ────────────────────────────────────────────────── - - def _goal_followup_after_turn( sid: str, session: dict, result: Any, status: str, raw: Any) -> str | None: - """/goal continuation (mirrors gateway/run._post_turn_goal_continuation). - - Asks the judge whether the goal is done and, if not and under budget, returns - the continuation prompt to chain once ``running`` is released. The verdict is - surfaced as a status line either way. Compression failures are never judge - input: the error text is not work toward the goal, and judging it spends a turn. - """ + """/goal continuation (mirrors gateway/run._post_turn_goal_continuation): the prompt to + chain once ``running`` is released, or None. Compression failures are never judge + input: the error text is not work toward the goal, and judging it spends a turn.""" goal_followup = None compression_exhausted = bool(isinstance(result, dict) and result.get("compression_exhausted")) try: @@ -486,18 +274,13 @@ def _goal_followup_after_turn( session, result, status=status, raw=raw) if recovery_notice: _emit("status.update", sid, {"kind": "goal", "text": recovery_notice}) - if recovery_prompt: - goal_followup = recovery_prompt + goal_followup = recovery_prompt or None except Exception as _goal_recovery_exc: _hook_failure("goal compression recovery", _goal_recovery_exc) if compression_exhausted or not _is_successful_goal_turn(result, status, raw): return goal_followup try: - from hermes_cli.goals import GoalManager - sid_key = session.get("session_key") or "" - if sid_key and ( - goal_mgr := GoalManager(session_id=sid_key, default_max_turns=_goal_max_turns()) - ).is_active(): + if session.get("session_key") and (goal_mgr := _active_goal_manager(session)) is not None: try: from hermes_cli.goals import gather_background_processes as _gather_bg _bg_procs = _gather_bg() @@ -515,8 +298,8 @@ def _goal_followup_after_turn( return goal_followup -def _complete_loop_tick(sid: str, session: dict, raw: Any) -> None: - """If this turn was a /loop wakeup, evaluate it (LOOP_COMPLETE, --until judge, caps, next).""" +def _after_complete_turn(sid: str, session: dict, st: _TurnRun, raw: Any) -> None: + """Hooks for a ``complete`` turn: /loop tick evaluation, pending title, voice fallback.""" try: from hermes_cli.loops import LoopManager loop_sid_key = session.get("session_key") or "" @@ -529,51 +312,34 @@ def _complete_loop_tick(sid: str, session: dict, raw: Any) -> None: _emit("status.update", sid, {"kind": "loop", "text": loop_msg}) except Exception as _loop_exc: _hook_failure("loop completion hook", _loop_exc) - - -def _apply_pending_title(sid: str, session: dict) -> None: - """Apply pending_title now that the DB row exists — in the session-owned profile store.""" - _pending = session.get("pending_title") - if not _pending: - return - _session_key = session.get("session_key") or sid - try: - with _session_db(session) as _pdb: - if _pdb and _pdb.set_session_title(_session_key, _pending): - session["pending_title"] = None - except ValueError as exc: - # Invalid/duplicate title — non-retryable, drop it; auto-title takes over. - session["pending_title"] = None - logger.info("Dropping pending title for session %s: %s", _session_key, exc) - except Exception: - pass # transient DB failure — keep pending_title for retry - - -def _speak_turn_fallback(raw: str) -> None: - """Voice TTS fallback when the streaming pipeline couldn't start: speak the final text whole.""" - try: - # Barge-aware: spoken interruptions must cut this playback too. - threading.Thread(target=_speak_text_with_barge, args=(raw,), daemon=True).start() - except ImportError: - logger.warning("voice TTS skipped: hermes_cli.voice unavailable") - except Exception as e: - logger.warning("voice TTS dispatch failed: %s", e) - - -def _append_turn_crash_log(sid: str, trace: str) -> None: - with contextlib.suppress(Exception): - os.makedirs(os.path.dirname(_CRASH_LOG), exist_ok=True) - with open(_CRASH_LOG, "a", encoding="utf-8") as f: - f.write( - f"\n=== turn-dispatcher exception · " - f"{time.strftime('%Y-%m-%d %H:%M:%S')} · sid={sid} ===\n") - f.write(trace) + # Apply pending_title now that the DB row exists — in the session-owned profile store. + if _pending := session.get("pending_title"): + _session_key = session.get("session_key") or sid + try: + with _session_db(session) as _pdb: + if _pdb and _pdb.set_session_title(_session_key, _pending): + session["pending_title"] = None + except ValueError as exc: + # Invalid/duplicate title — non-retryable, drop it; auto-title takes over. + session["pending_title"] = None + logger.info("Dropping pending title for session %s: %s", _session_key, exc) + except Exception: + pass # transient DB failure — keep pending_title for retry + # Voice fallback when the streaming pipeline couldn't start (tts_queue already spoke + # everything otherwise); barge-aware. + if st.tts_queue is None and isinstance(raw, str) and raw.strip() and _voice_tts_enabled(): + try: + threading.Thread(target=_speak_text_with_barge, args=(raw,), daemon=True).start() + except ImportError: + logger.warning("voice TTS skipped: hermes_cli.voice unavailable") + except Exception as e: + logger.warning("voice TTS dispatch failed: %s", e) def _dispatch_followup_turn(rid, sid: str, session: dict, prompt: Any, what: str, *, on_done=None, on_error=None) -> None: - """Chain one follow-up turn (caller already set ``running``); a dispatch failure - runs ``on_error``, logs, and releases ``running``.""" + """Chain one follow-up turn (caller set ``running``); on failure run ``on_error``, log, + release ``running``.""" try: _emit("message.start", sid) _run_prompt_submit(rid, sid, session, prompt) @@ -589,19 +355,14 @@ def _dispatch_followup_turn(rid, sid: str, session: dict, prompt: Any, what: str def _run_post_turn_followups( rid, sid: str, session: dict, result: Any, goal_followup: str | None) -> None: - """Chain whatever should run after ``running`` was released. - - Order: a user prompt that arrived mid-turn wins over every auto follow-up — - drain it and skip the rest this cycle. A leftover /steer the agent couldn't - inject is requeued first so it isn't dropped (a real queued prompt still wins: - ``_enqueue_prompt`` merges both). Then the goal continuation, then completion - notifications that arrived mid-turn. Each nested ``_run_prompt_submit`` checks - ``running`` under the lock first, so a racing user prompt wins. - """ - _leftover_steer = result.get("pending_steer") if isinstance(result, dict) else None - if isinstance(_leftover_steer, str) and _leftover_steer.strip(): + """Chain whatever should run after ``running`` was released. Order: a mid-turn user + prompt wins over every auto follow-up (drain it, skip the rest); a leftover /steer is + requeued first so it isn't dropped; then goal continuation, then completion + notifications. Each nested submit re-checks ``running`` under the lock.""" + steer = result.get("pending_steer") if isinstance(result, dict) else None + if isinstance(steer, str) and steer.strip(): with session["history_lock"]: - _enqueue_prompt(session, _leftover_steer, session.get("transport")) + _enqueue_prompt(session, steer, session.get("transport")) if _drain_queued_prompt(rid, sid, session): return if goal_followup: @@ -610,12 +371,9 @@ def _run_post_turn_followups( return # user already sent something — their turn wins session["running"] = True _dispatch_followup_turn(rid, sid, session, goal_followup, "goal continuation dispatch") - - # Safety net for completion events that arrived mid-turn (the poller handles - # between-turn delivery). Ownership is positive-proof and compression-chain - # aware (same fail-closed gate as the poller): session B must not consume - # session A's event; a post-compression session still claims its - # pre-compression dispatches. Unclaimable events are requeued for the poller. + # Safety net for completion events that arrived mid-turn. Ownership is positive-proof + # and compression-chain aware (same fail-closed gate as the poller): session B must + # not consume session A's event. Unclaimable events are requeued for the poller. try: from tools.process_registry import process_registry drained = process_registry.drain_notifications( @@ -642,28 +400,17 @@ def _run_post_turn_followups( _hook_failure("completion queue drain", _drain_exc) -# ── the turn ───────────────────────────────────────────────────────── - - @dataclasses.dataclass(slots=True) class _TurnRun: - """Mutable state the phase helpers of one turn thread share. - - ``agent`` is bound eagerly so except/finally always have one even if setup - throws; re-read after ``_sync_bot_capabilities`` (may swap in a rebuilt Bot - Chat agent). ``error_retained``: the finally skips the inflight clear (the - failed snapshot stays for resume replay). ``error_detail``: cause for the - "tui turn finished" bookend, stashed by both failure paths because the finally - sees neither ``result`` nor the exception reliably; ``prompt_text`` is what was - submitted (post @-expansion) so the cause can be checked for quoting it back. - """ + """Shared state of one turn thread. ``agent`` is bound eagerly so except/finally always + have one; ``error_retained`` makes the finally keep the failed inflight snapshot for + resume replay; ``error_detail`` is the "tui turn finished" failure cause.""" agent: Any one_turn_restore: Any terminal_callback: Any receipt_committed: bool scopes: _TurnScopes = dataclasses.field(default_factory=_TurnScopes) - goal_followup: Any = None result: Any = None # read after the finally for leftover /steer tts_queue: Any = None thinking_started: bool = False @@ -678,25 +425,33 @@ class _TurnRun: def _prepare_turn_input(sid: str, session: dict, st: _TurnRun, text: Any, images: list[str]): - """Bind scopes, sync the agent, snapshot history and build the run message. - - Returns ``(prompt, run_message, cols, streamer)``, or None when @-expansion - was refused (error already emitted). The config-model sync is skipped while - a /model --once override is active: the once-model is deliberately not pinned - as model_override, so the sync would clobber it (a config.yaml change is - adopted NEXT turn). A model picked mid-turn was queued, not applied — apply - it before the config sync so the explicit pick wins over a config change. - """ - _bind_turn_scopes(sid, session, st.scopes) + """Bind scopes, sync the agent, snapshot history, build the run message; returns + ``(prompt, run_message, cols, streamer)`` or None when @-expansion was refused. + Scopes fill field by field so a failure midway still leaves every bound token for the + finally; the profile's terminal policy is bound too (a failed install leaves a + fail-closed refusal scope). The config-model sync is skipped under a /model --once + override (not pinned as model_override, the sync would clobber it); a model picked + mid-turn is applied first so the explicit pick wins over a config change.""" + from tools.approval import set_current_session_key + scopes = st.scopes + scopes.approval = set_current_session_key(session["session_key"]) + scopes.session_tokens = _set_session_context(session["session_key"], ui_session_id=sid) + profile_home = session.get("profile_home") + if profile_home: + scopes.home = set_hermes_home_override(profile_home) + scopes.secret = set_secret_scope(build_profile_secret_scope(Path(profile_home))) + from tools.terminal_scope import install_profile_terminal_scope + scopes.terminal = install_profile_terminal_scope(Path(profile_home)) + # The sudo password callback is thread-local: without re-wiring here, sudo prompts + # fall through to /dev/tty and hang the headless gateway (re-run is a no-op). + _wire_callbacks(sid) if not st.one_turn_restore: _apply_pending_model_switch(sid, session) _sync_agent_model_with_config(sid, session) _sync_agent_compression_with_config(sid, session) - # Bot Chat: adopt Settings->Capabilities edits into the eternal bot session first. - _sync_bot_capabilities(sid, session) + _sync_bot_capabilities(sid, session) # Bot Chat: adopt Settings->Capabilities edits st.agent = agent = session["agent"] - # Snapshot after turn-start model sync: a deferred switch mutates history - # and its version, and that mutation belongs to this turn. + # Snapshot after the model sync: a deferred switch's history mutation belongs to this turn. with session["history_lock"]: st.history = list(session["history"]) st.history_version = int(session.get("history_version", 0)) @@ -706,7 +461,16 @@ def _prepare_turn_input(sid: str, session: dict, st: _TurnRun, text: Any, images streamer = make_stream_renderer(cols) prompt = text if isinstance(prompt, str) and "@" in prompt: - ctx = _expand_context_references(agent, prompt, cwd) + from agent.context_references import preprocess_context_references + from agent.model_metadata import get_model_context_length + ctx_len = get_model_context_length( + getattr(agent, "model", "") or _resolve_model(), + base_url=getattr(agent, "base_url", "") or "", + api_key=getattr(agent, "api_key", "") or "", + provider=getattr(agent, "provider", "") or "", + config_context_length=getattr(agent, "_config_context_length", None)) + ctx = preprocess_context_references( + prompt, cwd=cwd, allowed_root=cwd, context_length=ctx_len) if ctx.blocked: _emit( "error", sid, {"message": "\n".join(ctx.warnings) or "Context injection refused."}) @@ -715,7 +479,13 @@ def _prepare_turn_input(sid: str, session: dict, st: _TurnRun, text: Any, images st.prompt_text = prompt if isinstance(prompt, str) else "" run_message: Any = _route_turn_images(agent, prompt, images) if images else prompt st.tts_queue, st.thinking_started = _start_turn_voice() - return prompt, _apply_turn_notes(run_message, session), cols, streamer + # Per-turn API-message notes: barge mid-speech, reactions, HUD surface (per-turn state + # that must not touch the byte-stable system prompt). + from tools.tts_streaming import SPEECH_INTERRUPTED_NOTE, take_speech_interrupted + if take_speech_interrupted(): + run_message = _prepend_note(run_message, SPEECH_INTERRUPTED_NOTE) + run_message = _prepend_note(run_message, _pending_reaction_notes(session)) + return prompt, _prepend_note(run_message, _hud_surface_note(session)), cols, streamer def _invoke_agent( @@ -734,21 +504,29 @@ def _invoke_agent( st.tts_queue.put(delta) _emit("message.delta", sid, payload) - # Interim assistant text (commentary beside tool calls, or a pre-nudge final - # answer) is sealed by the desktop as its own segment instead of being lost - # when message.complete replaces the streaming buffer. Gated on - # display.interim_assistant_messages (default true). - if _load_interim_assistant_messages(): - def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: - _emit("message.interim", sid, {"text": text, "already_streamed": already_streamed}) - agent.interim_assistant_callback = _interim_assistant_cb - else: - agent.interim_assistant_callback = None - st.run_kwargs = _build_run_kwargs( - agent, session, st.history, prompt, images, run_message, _stream, display_kind, - display_metadata) - # Auto-titling fires inside the turn prologue; this live-rename hook - # repaints the sidebar the moment a title lands. + # Interim assistant text (commentary beside tool calls, pre-nudge final answer) is sealed + # by the desktop as its own segment instead of being lost to message.complete. + def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: + _emit("message.interim", sid, {"text": text, "already_streamed": already_streamed}) + agent.interim_assistant_callback = ( + _interim_assistant_cb if _load_interim_assistant_messages() else None) + # A synthesized turn is typed at turn START so a crash persist writes a timeline event, + # not a raw user bubble; the post-turn stamp is the fallback for an older agent. + st.run_kwargs = run_kwargs = { + "conversation_history": list(st.history), + "stream_callback": _stream, + "persist_user_message": ( + _build_persist_user_message(prompt, images, run_message) if images else prompt)} + try: + run_params = inspect.signature(agent.run_conversation).parameters + except (TypeError, ValueError): + run_params = {} + if "task_id" in run_params: + run_kwargs["task_id"] = session["session_key"] + if display_kind and "persist_user_display_kind" in run_params: + run_kwargs["persist_user_display_kind"] = display_kind + run_kwargs["persist_user_display_metadata"] = display_metadata + # Live-rename hook: auto-titling fires inside the turn prologue. _title_key = session.get("session_key") or sid agent._on_session_title = lambda t, _src, _k=_title_key: _emit( "session.title", sid, {"session_id": _k, "title": t}) @@ -756,11 +534,8 @@ def _invoke_agent( try: st.result = agent.run_conversation(run_message, **st.run_kwargs) finally: - # Stop AND join before anything below emits: a tick surviving past - # message.complete would roll the client's final usage back to a stale - # snapshot. The join is deliberately unbounded — once stop is set it only - # waits out one in-flight _get_usage/_emit, whose worst case (a stalled - # transport write) would stall the message.complete emit just the same. + # Stop AND join before anything emits: a tick surviving past message.complete would + # roll the client's usage back to a stale snapshot (unbounded join: same worst case). _usage_stop.set() _usage_thread.join() @@ -769,28 +544,66 @@ def _absorb_turn_result( sid: str, session: dict, st: _TurnRun, text: Any, display_kind: str | None, display_metadata ) -> str | None: """Stamp, restore /moa, commit history, re-sync the session key; returns the history warning.""" - result = st.result + result, agent = st.result, st.agent if display_kind and isinstance(text, str): - _stamp_synthetic_display_kind( - st.agent, session, result, text, display_kind, display_metadata) + # Post-turn fallback stamp of a synthesized turn's display kind (DB row + result). + db = getattr(agent, "_session_db", None) + current_session_id = getattr(agent, "session_id", None) or session.get("session_key") + if db is not None: + try: + db.set_latest_matching_message_display_kind( + current_session_id, role="user", content=text, display_kind=display_kind, + display_metadata=display_metadata) + except Exception: + logger.debug("failed to stamp synthetic display kind", exc_info=True) + if isinstance(result, dict) and isinstance(result.get("messages"), list): + for message in reversed(result["messages"]): + if message.get("role") == "user" and message.get("content") == text: + message["display_kind"] = display_kind + if display_metadata: + message["display_metadata"] = display_metadata + break if "moa_one_shot_restore" in session: - _restore_moa_one_shot(sid, session) + # Undo a /moa one-shot through the switch path: resetting model_override alone + # would leave the live client pinned to MoA after the in-place switch_model(). + _restore = session.pop("moa_one_shot_restore", None) + if isinstance(_restore, dict): + _prev_override = _restore.get("override") + _prev_model = _restore.get("model") + _prev_provider = _restore.get("provider") + if _prev_override is None: + session.pop("model_override", None) + else: + session["model_override"] = _prev_override + if _prev_model: + _raw = ( + f"{_prev_model} --provider {_prev_provider}" if _prev_provider else _prev_model) + try: + _apply_model_switch( + sid, session, _raw, confirm_expensive_model=False, + pin_session_override=bool(_prev_override), + persist_override=False) # session-internal restore, never config.yaml + except Exception as _moa_restore_exc: + logger.warning("MoA one-shot model restore failed: %s", _moa_restore_exc) + elif _restore is None: + session.pop("model_override", None) + else: + session["model_override"] = _restore status_note = None if isinstance(result, dict): if isinstance(result.get("messages"), list): status_note = _commit_turn_history(session, result, st.history, st.history_version) - # Auto-compression inside run_conversation() may have rotated - # agent.session_id: sync session_key before title/goal/finalize use it, - # keep pending_title (user intent), and restart the slash worker so - # worker-backed commands (/title etc.) target the live session. + # Auto-compression may have rotated agent.session_id: sync session_key before + # title/goal/finalize use it, keep pending_title (user intent), restart the slash + # worker so worker-backed commands target the live session. _sync_session_key_after_compress( sid, session, clear_pending_title=False, restart_slash_worker=True) return status_note def _complete_turn_payload(session: dict, st: _TurnRun, status_note: str | None, cols: int): - """Build the ``message.complete`` payload, retain/clear the inflight turn and - settle the hosted-room terminal receipt. Returns ``(payload, raw, status)``.""" + """``(payload, raw, status)`` for message.complete; retains/clears the inflight turn and + settles the hosted-room terminal receipt.""" result, agent = st.result, st.agent raw, status, last_reasoning = _turn_outcome(result) payload = {"text": raw, "usage": _get_usage(agent), "status": status} @@ -800,29 +613,32 @@ def _complete_turn_payload(session: dict, st: _TurnRun, status_note: str | None, payload["warning"] = status_note if result.get("response_previewed"): payload["response_previewed"] = True - # Structured billing-wall descriptor so the client renders a - # billing-specific recovery surface instead of re-parsing text. - _billing_block = result.get("billing_block") if isinstance(result, dict) else None - if _billing_block: + # Structured billing-wall descriptor: the client renders recovery without re-parsing text. + if _billing_block := result.get("billing_block"): payload["billing"] = _billing_block payload["failure_reason"] = result.get("failure_reason") if rendered := render_message(raw, cols): payload["rendered"] = rendered - # Layer descriptor computed before the retain below so resume replay - # carries the same one (advisory; older clients ignore it). - _error_surface = _turn_error_surface(agent, result) if status == "error" else None - _result_error = result.get("error") if isinstance(result, dict) else None - error_value = _result_error if isinstance(result, dict) else raw + # Advisory {layer, code, retryable} descriptor; computed before the retain so resume + # replay carries the same one. + _error_surface = None + if status == "error": + try: + from agent.error_surface import build_error_surface_from_result + _error_surface = build_error_surface_from_result( + result, provider=str(getattr(agent, "provider", "") or ""), + model=str(getattr(agent, "model", "") or "")) + except Exception: + _error_surface = None + error_value = result.get("error") with session["history_lock"]: if status == "error": - # Retain the failed turn for resume replay: if this terminal frame - # is lost to a disconnect, resume's inflight payload is the only - # carrier of the failure. + # Retain the failed turn: resume's inflight payload is the only carrier of the + # failure if this frame is lost to a disconnect. _fail_inflight_turn(session, error_value, error_surface=_error_surface) st.error_retained = True st.error_detail = _turn_failure_detail( - error_value, result.get("failure_reason") if isinstance(result, dict) else None, - st.prompt_text) + error_value, result.get("failure_reason"), st.prompt_text) else: _clear_inflight_turn(session) if status == "error": @@ -833,13 +649,9 @@ def _complete_turn_payload(session: dict, st: _TurnRun, status_note: str | None, if st.terminal_callback is not None: st.receipt_attempted = True st.terminal_callback({ - "status": ( - "cancelled" if status == "interrupted" - else "failed" if status == "error" else "settled"), + "status": {"interrupted": "cancelled", "error": "failed"}.get(status, "settled"), "text": raw if isinstance(raw, str) else str(raw), - **( - {"error": str(_result_error or raw)} - if status == "error" and isinstance(result, dict) else {})}) + **({"error": str(error_value or raw)} if status == "error" else {})}) st.receipt_committed = True if st.receipt_committed: _retire_turn_marker(session, st.marker_key) @@ -849,12 +661,15 @@ def _complete_turn_payload(session: dict, st: _TurnRun, status_note: str | None, def _recover_turn_exception(sid: str, session: dict, st: _TurnRun, e: BaseException) -> None: """Except-path of the turn: crash log, history restore, terminal error frame.""" import traceback - _append_turn_crash_log(sid, traceback.format_exc()) + with contextlib.suppress(Exception): + os.makedirs(os.path.dirname(_CRASH_LOG), exist_ok=True) + with open(_CRASH_LOG, "a", encoding="utf-8") as f: + f.write( + f"\n=== turn-dispatcher exception · " + f"{time.strftime('%Y-%m-%d %H:%M:%S')} · sid={sid} ===\n") + f.write(traceback.format_exc()) print(f"[gateway-turn] {type(e).__name__}: {e}", file=sys.stderr, flush=True) - # An exception in the agent's finalizer can leave the gateway's in-memory - # history at the turn-start snapshot; keep the partial turn available to - # the next prompt (the durable inflight record still carries the - # recoverable error state). + # A finalizer exception can leave in-memory history at the turn-start snapshot. _restore_agent_history_after_turn_error(session, st.agent) if st.terminal_callback is not None and not st.receipt_attempted: st.receipt_attempted = True @@ -864,8 +679,7 @@ def _recover_turn_exception(sid: str, session: dict, st: _TurnRun, e: BaseExcept except Exception: logger.exception("hosted room terminal receipt commit failed") try: - # Same terminal error frame shape as the returned-error path (uniform - # client handling), retaining the turn for replay. + # Same terminal error frame shape as the returned-error path. _emit_terminal_turn_error(sid, session, e, retire_marker=st.receipt_committed) st.error_retained = True st.error_detail = _turn_failure_detail(e, type(e).__name__, st.prompt_text) @@ -878,54 +692,44 @@ def _recover_turn_exception(sid: str, session: dict, st: _TurnRun, e: BaseExcept def _finish_turn(sid: str, session: dict, st: _TurnRun) -> None: """Finally-path of the turn: release everything, then the "tui turn finished" bookend.""" - agent, history, run_kwargs = st.agent, st.history, st.run_kwargs - # Drop both snapshots of the pre-turn history before asking glibc to return - # pages; session["history"] already points at the new/pruned result. + # Drop both pre-turn history snapshots before asking glibc to return pages (a test + # inspects these two locals by name). + history, run_kwargs = st.history, st.run_kwargs history.clear() if isinstance(run_kwargs, dict): run_kwargs.clear() - # While any profile-specific HERMES_HOME override is still active, so - # context.memory_trim resolves from the session's own config. - try: + try: # while the profile HERMES_HOME override is still active (session's own config) from hermes_cli.mem_trim import trim_memory trim_memory(reason="tui turn completion") except Exception: logger.debug("post-turn memory trim failed", exc_info=True) if st.thinking_started: - _stop_thinking_sound() + with contextlib.suppress(Exception): + from tools.voice_mode import stop_thinking_sound + stop_thinking_sound() if st.tts_queue is not None: st.tts_queue.put(None) # end-of-text sentinel — flush + finish speaking if st.one_turn_restore: try: - _restore_agent_model_runtime(agent, st.one_turn_restore) + _restore_agent_model_runtime(st.agent, st.one_turn_restore) _restart_slash_worker(sid, session) _persist_live_session_runtime(session) _persist_live_session_system_prompt(session) except Exception: logger.debug("TUI one-turn model restore failed", exc_info=True) - _release_turn_scopes(st.scopes) - - -def _log_turn_finished(sid: str, session: dict, st: _TurnRun, started_monotonic: float) -> None: - """Closing bookend of "tui prompt accepted" — fires on every path, so one - accepted prompt produces exactly one finished record. agent.session_id is - re-read because compression may have rotated it mid-turn (an accepted/finished - pair whose id changed IS a rotation trace). A missing finished record means - the thread died before the finally.""" - result = st.result - if isinstance(result, dict): - status = ( - result.get("interrupted") and "interrupted" - or result.get("error") and "error" or "complete" - ) - else: - status = "error" if st.error_retained else "complete" - logger.info( - "tui turn finished: ui_session=%s session_key=%s " - "agent_session_id=%s status=%s error_retained=%s duration=%.1fs" - "%s", - sid, session.get("session_key") or "", getattr(st.agent, "session_id", "") or "", status, - st.error_retained, time.monotonic() - started_monotonic, st.error_detail) + scopes = st.scopes + with contextlib.suppress(Exception): + if scopes.approval is not None: + from tools.approval import reset_current_session_key + reset_current_session_key(scopes.approval) + if scopes.home is not None: + reset_hermes_home_override(scopes.home) + if scopes.secret is not None: + reset_secret_scope(scopes.secret) + if scopes.terminal is not None: + from tools.terminal_scope import reset_terminal_scope + reset_terminal_scope(scopes.terminal) + _clear_session_context(scopes.session_tokens) def _run_prompt_submit( @@ -937,10 +741,8 @@ def _run_prompt_submit( if admitted is None: return False images, agent = admitted - # The ONE INFO record proving a Desktop/TUI prompt was accepted by THIS - # process; ties the UI session id, gateway session_key and the agent's live - # session_id (compression rotates the last independently) together for a - # rotation-mute trace. No prompt content is logged. + # The ONE INFO record proving a prompt was accepted by THIS process; ties ui sid, + # session_key and the agent's live session_id together. No prompt content is logged. _turn_started_monotonic = time.monotonic() logger.info( "tui prompt accepted: ui_session=%s session_key=%s agent_session_id=%s " @@ -950,15 +752,15 @@ def _run_prompt_submit( _emit("message.start", sid) def run(): - # ContextVars from the RPC dispatcher do not follow onto this thread: - # rebind the exact transport stored on this session generation before any - # tool can commission a child (delegate_task captures it as authority). + # RPC-dispatcher ContextVars do not follow onto this thread: rebind the transport + # before any tool can commission a child (delegate_task captures it as authority). transport_token = bind_transport(session.get("transport")) runtime_session_token = _current_runtime_session_record.set(session) st = _TurnRun( session["agent"], session.pop("one_turn_model_restore", None), terminal_callback, receipt_committed=terminal_callback is None) st.marker_key = _record_turn_marker(session, text) + goal_followup = None try: prepared = _prepare_turn_input(sid, session, st, text, images) if prepared is None: @@ -966,21 +768,14 @@ def _run_prompt_submit( prompt, run_message, cols, streamer = prepared _invoke_agent( sid, session, st, prompt, run_message, streamer, images, display_kind, - display_metadata, - ) + display_metadata) status_note = _absorb_turn_result( sid, session, st, text, display_kind, display_metadata) payload, raw, status = _complete_turn_payload(session, st, status_note, cols) _emit("message.complete", sid, payload) - st.goal_followup = _goal_followup_after_turn(sid, session, st.result, status, raw) + goal_followup = _goal_followup_after_turn(sid, session, st.result, status, raw) if status == "complete": - _complete_loop_tick(sid, session, raw) - _apply_pending_title(sid, session) - # The streaming path already spoke everything via tts_queue. - if ( - st.tts_queue is None and isinstance(raw, str) and raw.strip() - and _voice_tts_enabled()): - _speak_turn_fallback(raw) + _after_complete_turn(sid, session, st, raw) except Exception as e: _recover_turn_exception(sid, session, st, e) finally: @@ -994,7 +789,19 @@ def _run_prompt_submit( session["last_active"] = time.time() if not st.error_retained: _clear_inflight_turn(session) - _log_turn_finished(sid, session, st, _turn_started_monotonic) + # Closing bookend of "tui prompt accepted" — exactly one per accepted prompt. + # agent.session_id is re-read because compression may have rotated it (an + # accepted/finished pair whose id changed IS a rotation trace). + if isinstance(st.result, dict): + status = _result_status(st.result) + else: + status = "error" if st.error_retained else "complete" + logger.info( + "tui turn finished: ui_session=%s session_key=%s agent_session_id=%s status=%s " + "error_retained=%s duration=%.1fs%s", + sid, session.get("session_key") or "", getattr(st.agent, "session_id", "") or "", + status, st.error_retained, time.monotonic() - _turn_started_monotonic, + st.error_detail) # Backstop for turns that never reached a terminal frame. if st.receipt_committed: _retire_turn_marker(session, st.marker_key) @@ -1004,11 +811,11 @@ def _run_prompt_submit( session.pop("_hosted_room_task", None) session.pop("_auto_continue_scheduled", None) _emit_settled_session_info(sid, session, st.agent) - _run_post_turn_followups(rid, sid, session, st.result, st.goal_followup) + _run_post_turn_followups(rid, sid, session, st.result, goal_followup) run_thread = threading.Thread(target=run, daemon=True) with _sessions_lock: registered = _sessions.get(sid) - can_start = (not session.get("_closing") and (registered is None or registered is session)) + can_start = not session.get("_closing") and (registered is None or registered is session) if can_start: session["_run_thread"] = run_thread run_thread.start()