From 957147365a3ef3cf5c30705769f8747fff958955 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:47:13 -0700 Subject: [PATCH] refactor(tui_gateway): split _apply_model_switch into phase helpers, compact controller/browser/relay/watcher/billing helpers --- tui_gateway/_env.py | 22 +-- tui_gateway/_stdin_recovery.py | 46 ++--- tui_gateway/billing_view.py | 137 ++++++------- tui_gateway/change_watcher.py | 61 +++--- tui_gateway/event_publisher.py | 65 ++---- tui_gateway/event_replay.py | 69 +++---- tui_gateway/git_probe.py | 66 +++---- tui_gateway/loop_noise.py | 61 ++---- tui_gateway/mcp_rpc_helpers.py | 10 +- tui_gateway/method_ctx.py | 54 ++--- tui_gateway/methods_bot_relay.py | 131 ++++--------- tui_gateway/methods_browser.py | 129 ++++++------ tui_gateway/methods_browser_control.py | 262 +++++++------------------ tui_gateway/model_switch.py | 242 ++++++++++++----------- 14 files changed, 528 insertions(+), 827 deletions(-) diff --git a/tui_gateway/_env.py b/tui_gateway/_env.py index 481edfac63..167030e6e6 100644 --- a/tui_gateway/_env.py +++ b/tui_gateway/_env.py @@ -1,24 +1,22 @@ -"""Tolerant env-var knob parsing shared by the gateway entry points. - -A bare ``float(os.environ[...])`` would raise at import time on a typo -(``HERMES_SLASH_WATCHDOG_POLL_S=2s``) and kill a worker before it serves a -single command; these fall back to ``default`` on absent/empty/malformed values. -""" +"""Tolerant env-var knob parsing shared by the gateway entry points: a bare +``float(os.environ[...])`` would raise at import on a typo (``...POLL_S=2s``) and kill a +worker before it serves a command; these fall back to ``default`` on absent/empty/malformed.""" from __future__ import annotations import os -def env_float(name: str, default: float) -> float: +def _env_number(cast, name: str, default): try: - return float(os.environ.get(name, "") or default) + return cast(os.environ.get(name, "") or default) except (TypeError, ValueError): return default +def env_float(name: str, default: float) -> float: + return _env_number(float, name, default) + + def env_int(name: str, default: int) -> int: - try: - return int(os.environ.get(name, "") or default) - except (TypeError, ValueError): - return default + return _env_number(int, name, default) diff --git a/tui_gateway/_stdin_recovery.py b/tui_gateway/_stdin_recovery.py index 4b7e685f2c..e4149da831 100644 --- a/tui_gateway/_stdin_recovery.py +++ b/tui_gateway/_stdin_recovery.py @@ -1,12 +1,10 @@ """Shared spurious stdin-EOF recovery for the TUI gateway entry point and slash worker. -When a child inherits fd 0 and sets ``O_NONBLOCK``, the flag lands on the -**shared open file description**, not just the child's descriptor. The next -``read()`` returns ``EAGAIN``, which CPython's buffered ``TextIOWrapper`` -converts to ``b''`` (apparent EOF), killing the gateway. - -Recovery is **POSIX-only** (``fcntl``); on Windows the guard just reports a -genuine EOF and lets the caller exit. +When a child inherits fd 0 and sets ``O_NONBLOCK``, the flag lands on the SHARED open +file description, not just the child's descriptor. The next ``read()`` returns +``EAGAIN``, which CPython's buffered ``TextIOWrapper`` converts to ``b''`` (apparent +EOF), killing the gateway. Recovery is POSIX-only (``fcntl``); on Windows the guard +just reports a genuine EOF and lets the caller exit. """ from __future__ import annotations @@ -21,10 +19,8 @@ try: except ImportError: # Windows fcntl = None # type: ignore[assignment] - -# At most this many recoveries per 60s window. A child aggressively flipping -# the flag would otherwise create a tight busy-loop; exceeding the cap exits so -# the parent respawns us with fresh state. +# Recoveries per 60s window. A child aggressively flipping the flag would otherwise +# create a tight busy-loop; exceeding the cap exits so the parent respawns us fresh. MAX_RECOVERIES_PER_MINUTE = 10 @@ -55,9 +51,8 @@ def _stdin_sockopt(getter): def diagnose_stdin_state() -> str: """Diagnostic string (``O_NONBLOCK`` / ``SO_RCVTIMEO``) for crash-log forensics. - ``SO_RCVTIMEO`` is a socket option equally shared on the open file - description; a child's ``setsockopt`` launders into the same spurious-EOF - path with ``O_NONBLOCK`` clear, so it is reported alongside the flag. + ``SO_RCVTIMEO`` is equally shared on the open file description; a child's + ``setsockopt`` launders into the same spurious-EOF path with ``O_NONBLOCK`` clear. """ parts: list[str] = [] if fcntl is None: @@ -77,16 +72,15 @@ def diagnose_stdin_state() -> str: def handle_spurious_eof(recovery_times: list[float], log_fn: object) -> bool: """Check whether an empty ``readline()`` is spurious; recover if so. - Returns True if the caller should ``continue`` the read loop (recovered), - False if it should ``break`` (genuine peer-close or rate limit exceeded). - ``log_fn`` receives a diagnostic string. + Returns True if the caller should ``continue`` the read loop (recovered), False if it + should ``break`` (genuine peer-close or rate limit exceeded). ``log_fn`` receives a + diagnostic string. """ - # Without fcntl (Windows) we can't check the flag, and the issue is - # POSIX-specific anyway; a clear flag means a genuine peer-close. + # Without fcntl (Windows) we can't check the flag and the issue is POSIX-specific + # anyway; a clear flag means a genuine peer-close. if fcntl is None or not _stdin_nonblock(): log_fn("stdin EOF (peer closed)") # type: ignore[operator] return False - now = time.time() recovery_times.append(now) recovery_times[:] = [t for t in recovery_times if t > now - 60] @@ -96,16 +90,12 @@ def handle_spurious_eof(recovery_times: list[float], log_fn: object) -> bool: f"({len(recovery_times)}/min, cap {MAX_RECOVERIES_PER_MINUTE})" ) return False - log_fn(f"stdin spurious EOF (subprocess O_NONBLOCK flip), recovering: {diagnose_stdin_state()}") # type: ignore[operator] - - # Restore blocking mode on the shared description, and clear SO_RCVTIMEO too: - # a non-zero timeout would make the next readline() return '' again, looping - # until the rate limiter fires. + # Restore blocking mode on the shared description, and clear SO_RCVTIMEO too: a + # non-zero timeout would make the next readline() return '' again until the limiter fires. os.set_blocking(0, True) # "ll" = struct timeval {tv_sec, tv_usec}; zero timeval disables the timeout. _stdin_sockopt(lambda s: s.setsockopt(socket.SOL_SOCKET, socket.SO_RCVTIMEO, struct.pack("ll", 0, 0))) - - # TextIOWrapper.readline returns '' on EAGAIN but does NOT stick EOF; the - # next call blocks until data arrives or the peer truly closes. + # TextIOWrapper.readline returns '' on EAGAIN but does NOT stick EOF; the next call + # blocks until data arrives or the peer truly closes. return True diff --git a/tui_gateway/billing_view.py b/tui_gateway/billing_view.py index 8cfa92a026..3be9bcad12 100644 --- a/tui_gateway/billing_view.py +++ b/tui_gateway/billing_view.py @@ -3,10 +3,8 @@ STRUCTURED envelopes (result.ok / result.error) rather than JSON-RPC errors, so rpc() always resolves and the client branches on the typed billing code. Data-building lives in agent/billing_view.py + hermes_cli/nous_billing.py. - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so tests may still monkeypatch e.g. -``server._usage_payload``. +Bodies are rebound onto server.py's globals at install time (method_ctx.bind_module), +so tests may still monkeypatch e.g. ``server._usage_payload``. """ from __future__ import annotations @@ -16,6 +14,11 @@ from typing import Optional from .method_ctx import bind_module +def _wire_str(value): + """Decimal/number → wire string (None passes through).""" + return None if value is None else str(value) + + def _serialize_billing_error(exc) -> dict: """Map a BillingError into the result.error envelope the TUI branches on.""" from hermes_cli.nous_billing import ( @@ -47,12 +50,45 @@ def _serialize_billing_error(exc) -> dict: } +def _serialize_payment_method(pm) -> dict | None: + # Each kind sends only its own fields. Emitting every key with nulls would contradict + # the shared type — a client checking `'brand' in pm` would read every Link method as a card. + if pm is None: + return None + if pm.kind == "card": + return { + "kind": "card", "brand": pm.brand, "last4": pm.last4, "wallet": pm.wallet, + "resolved_via": pm.resolved_via, + } + if pm.kind == "link": + return {"kind": "link", "email": pm.email, "resolved_via": pm.resolved_via} + return {"kind": "unknown", "raw_kind": pm.raw_kind, "resolved_via": pm.resolved_via} + + +def _serialize_auto_reload(ar, format_money) -> dict | None: + if ar is None: + return None + card_out = None + if ar.card is not None: + if ar.card.kind == "distinct": + card_out = { + "kind": "distinct", "payment_method_id": ar.card.payment_method_id, + "brand": ar.card.brand, "last4": ar.card.last4, + } + else: + card_out = {"kind": ar.card.kind} + return { + "enabled": ar.enabled, "threshold_usd": _wire_str(ar.threshold_usd), + "threshold_display": format_money(ar.threshold_usd), + "reload_to_usd": _wire_str(ar.reload_to_usd), + "reload_to_display": format_money(ar.reload_to_usd), "card": card_out, + } + + def _serialize_billing_state(state) -> dict: """Serialize a BillingState for the wire (Decimals → strings, money-safe).""" from agent.billing_view import format_money - def _s(value): - return None if value is None else str(value) card = None if state.card is not None: card = { @@ -64,50 +100,15 @@ def _serialize_billing_state(state) -> dict: "display": state.card.display, "resolved_via": state.card.resolved_via, } - payment_method = None - if state.payment_method is not None: - pm = state.payment_method - # Each kind sends only its own fields. Emitting every key with nulls - # would contradict the shared type — a client checking `'brand' in pm` - # would read every Link method as a card. - if pm.kind == "card": - payment_method = { - "kind": "card", "brand": pm.brand, "last4": pm.last4, "wallet": pm.wallet, - "resolved_via": pm.resolved_via, - } - elif pm.kind == "link": - payment_method = {"kind": "link", "email": pm.email, "resolved_via": pm.resolved_via} - else: - payment_method = { - "kind": "unknown", "raw_kind": pm.raw_kind, "resolved_via": pm.resolved_via, - } monthly_cap = None if state.monthly_cap is not None: mc = state.monthly_cap monthly_cap = { - "limit_usd": _s(mc.limit_usd), "limit_display": format_money(mc.limit_usd), - "spent_this_month_usd": _s(mc.spent_this_month_usd), + "limit_usd": _wire_str(mc.limit_usd), "limit_display": format_money(mc.limit_usd), + "spent_this_month_usd": _wire_str(mc.spent_this_month_usd), "spent_display": format_money(mc.spent_this_month_usd), "is_default_ceiling": mc.is_default_ceiling, } - auto_reload = None - if state.auto_reload is not None: - ar = state.auto_reload - card_out = None - if ar.card is not None: - if ar.card.kind == "distinct": - card_out = { - "kind": "distinct", "payment_method_id": ar.card.payment_method_id, - "brand": ar.card.brand, "last4": ar.card.last4, - } - else: - card_out = {"kind": ar.card.kind} - auto_reload = { - "enabled": ar.enabled, "threshold_usd": _s(ar.threshold_usd), - "threshold_display": format_money(ar.threshold_usd), - "reload_to_usd": _s(ar.reload_to_usd), - "reload_to_display": format_money(ar.reload_to_usd), "card": card_out, - } return { "ok": True, "logged_in": state.logged_in, @@ -117,17 +118,17 @@ def _serialize_billing_state(state) -> dict: "is_admin": state.is_admin, "can_change_plan": state.can_change_plan, "can_charge": state.can_charge, - "balance_usd": _s(state.balance_usd), + "balance_usd": _wire_str(state.balance_usd), "balance_display": format_money(state.balance_usd), "cli_billing_enabled": state.cli_billing_enabled, - "charge_presets": [_s(p) for p in state.charge_presets], + "charge_presets": [_wire_str(p) for p in state.charge_presets], "charge_presets_display": [format_money(p) for p in state.charge_presets], - "min_usd": _s(state.min_usd), - "max_usd": _s(state.max_usd), + "min_usd": _wire_str(state.min_usd), + "max_usd": _wire_str(state.max_usd), "card": card, - "payment_method": payment_method, + "payment_method": _serialize_payment_method(state.payment_method), "monthly_cap": monthly_cap, - "auto_reload": auto_reload, + "auto_reload": _serialize_auto_reload(state.auto_reload, format_money), "portal_url": state.portal_url, "error": state.error, # Shared two-bar dollar usage model so /topup matches /usage and @@ -137,11 +138,7 @@ def _serialize_billing_state(state) -> dict: def _usage_payload(state) -> dict: - """Best-effort shared usage model for the /topup + /subscription overlay bars. - - Only fetched when logged in; fail-open to {available:false} so the overview - still renders if the account-info path is down. - """ + """Shared usage model for the /topup + /subscription bars: fetched only when logged in, fail-open.""" if not getattr(state, "logged_in", False): return {"available": False} try: @@ -164,14 +161,13 @@ def _serialize_usage_bar(bar) -> Optional[dict]: def _serialize_usage_model(model) -> dict: - """Serialize a UsageModel for the wire — the shared two-bar dollar view. - - Dollars-only (no 'credits'); fail-open shape mirrors the other billing RPCs - ({ok, available:false} when logged out / unreachable). - """ + """Serialize a UsageModel for the wire — the shared two-bar dollar view (fail-open {ok, available:false}).""" from agent.billing_usage import _fmt_usd, format_renews if model is None or not getattr(model, "available", False): return {"ok": True, "available": False} + + def _usd(value): + return None if value is None else _fmt_usd(value) return { "ok": True, "available": True, @@ -179,15 +175,9 @@ def _serialize_usage_model(model) -> dict: "plan_name": model.plan_name, "renews_at": model.renews_at, "renews_display": getattr(model, "renews_display", None) or format_renews(model.renews_at), - "subscription_remaining_display": ( - None if model.subscription_remaining_usd is None else _fmt_usd(model.subscription_remaining_usd) - ), - "topup_remaining_display": ( - None if model.topup_remaining_usd is None else _fmt_usd(model.topup_remaining_usd) - ), - "total_spendable_display": ( - None if model.total_spendable_usd is None else _fmt_usd(model.total_spendable_usd) - ), + "subscription_remaining_display": _usd(model.subscription_remaining_usd), + "topup_remaining_display": _usd(model.topup_remaining_usd), + "total_spendable_display": _usd(model.total_spendable_usd), "has_topup": model.has_topup, "plan_bar": _serialize_usage_bar(model.plan_bar), "topup_bar": _serialize_usage_bar(model.topup_bar), @@ -199,14 +189,13 @@ def _serialize_subscription_state(state) -> dict: from agent.billing_usage import format_renews from agent.billing_view import format_money - def _s(value): - return None if value is None else str(value) current = None if state.current is not None: c = state.current current = { "tier_id": c.tier_id, "tier_name": c.tier_name, - "monthly_credits": _s(c.monthly_credits), "credits_remaining": _s(c.credits_remaining), + "monthly_credits": _wire_str(c.monthly_credits), + "credits_remaining": _wire_str(c.credits_remaining), "cycle_ends_at": c.cycle_ends_at, "pending_downgrade_tier_name": c.pending_downgrade_tier_name, "pending_downgrade_at": c.pending_downgrade_at, @@ -221,7 +210,7 @@ def _serialize_subscription_state(state) -> dict: { "tier_id": t.tier_id, "name": t.name, "tier_order": t.tier_order, "dollars_per_month_display": format_money(t.dollars_per_month), - "monthly_credits": _s(t.monthly_credits), "is_current": t.is_current, + "monthly_credits": _wire_str(t.monthly_credits), "is_current": t.is_current, "is_enabled": t.is_enabled, } for t in state.tiers @@ -255,9 +244,7 @@ def _serialize_subscription_preview(p) -> dict: "current_tier_name": p.current_tier_name, "target_tier_id": p.target_tier_id, "target_tier_name": p.target_tier_name, - "monthly_credits_delta": ( - None if p.monthly_credits_delta is None else str(p.monthly_credits_delta) - ), + "monthly_credits_delta": _wire_str(p.monthly_credits_delta), "amount_due_now_cents": p.amount_due_now_cents, "effective_at": p.effective_at, } diff --git a/tui_gateway/change_watcher.py b/tui_gateway/change_watcher.py index 369eed7a64..40620a67bc 100644 --- a/tui_gateway/change_watcher.py +++ b/tui_gateway/change_watcher.py @@ -6,7 +6,6 @@ method_ctx.bind_module), so they reference server.py globals bare. from __future__ import annotations - from .method_ctx import HandlerRegistry, bind_module _registry = HandlerRegistry() @@ -53,13 +52,18 @@ def _watcher_mtime_ns(path: Path): return None +def _newest_mtime_ns(paths) -> int | None: + """Max ``st_mtime_ns`` across ``paths`` (unstat-able ones ignored); None when none could be 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) + + def _skin_sig() -> tuple[str, float | None]: """(active skin name, its user-file mtime). Built-ins have no file, so only 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") - home = _watcher_home() try: - mtime: float | None = (home / "skins" / f"{name}.yaml").stat().st_mtime + mtime: float | None = (_watcher_home() / "skins" / f"{name}.yaml").stat().st_mtime except OSError: mtime = None return name, mtime @@ -68,10 +72,8 @@ def _skin_sig() -> tuple[str, float | None]: def _note_skin_broadcast() -> None: """Sync the baseline after the /skin RPC emits so the watcher doesn't re-broadcast it.""" global _last_skin_sig - try: + with contextlib.suppress(Exception): _last_skin_sig = _skin_sig() - except Exception: - pass def _broadcast_skin_if_changed() -> None: @@ -85,10 +87,8 @@ def _broadcast_skin_if_changed() -> None: if sig == _last_skin_sig: return _last_skin_sig = sig - try: + with contextlib.suppress(Exception): _broadcast_global_event("skin.changed", resolve_skin()) - except Exception: - pass def _pet_sig() -> tuple: @@ -132,12 +132,11 @@ 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.""" - mtimes = [ - _watcher_mtime_ns(root / name) + return _newest_mtime_ns( + root / name for root in (_watcher_home(), *_served_profile_homes) for name in ("state.db", "state.db-wal") - ] - return max((m for m in mtimes if m is not None), default=None) + ) def _platforms_sig(): @@ -147,32 +146,22 @@ def _platforms_sig(): def _pairing_sig(): - """Newest mtime across every profile's pairing store (legacy ``pairing/`` and + """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.""" home = _watcher_home() roots = [home / "pairing", home / "platforms" / "pairing"] - try: + with contextlib.suppress(OSError): for profile_dir in (home / "profiles").iterdir(): - roots.append(profile_dir / "pairing") - roots.append(profile_dir / "platforms" / "pairing") - except OSError: - pass - - sig = None + roots += [profile_dir / "pairing", profile_dir / "platforms" / "pairing"] + entries = [] for root in roots: - try: - entries = list(root.iterdir()) - except OSError: - continue - for entry in entries: + with contextlib.suppress(OSError): # Only the ledgers: _rate_limits.json moves on every unauthorized DM. - if not entry.name.endswith(("-pending.json", "-approved.json")): - continue - mtime = _watcher_mtime_ns(entry) - if mtime is not None: - sig = mtime if sig is None else max(sig, mtime) - return sig + entries += [ + e for e in root.iterdir() if e.name.endswith(("-pending.json", "-approved.json")) + ] + return _newest_mtime_ns(entries) # Newest outbox-envelope mtime EVER seen (monotone): a drain empties the outbox, @@ -188,12 +177,10 @@ def _bot_relay_outbox_sig(): home = _watcher_home() root = home.parent.parent if home.parent.name == "profiles" else home newest = 0 - try: + 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) - except OSError: - pass if newest > _bot_relay_outbox_seen: _bot_relay_outbox_seen = newest return _bot_relay_outbox_seen or None @@ -243,10 +230,8 @@ def _broadcast_watched_changes(now: float | None = None) -> None: continue # floored: old signature stays so it re-fires when the window opens _change_sigs[event] = sig _change_broadcast_at[event] = now - try: + with contextlib.suppress(Exception): _broadcast_global_event(event, payload_fn()) - except Exception: # noqa: BLE001 - pass _skin_watcher_started = False diff --git a/tui_gateway/event_publisher.py b/tui_gateway/event_publisher.py index 8510b8eac9..eba2ff5deb 100644 --- a/tui_gateway/event_publisher.py +++ b/tui_gateway/event_publisher.py @@ -1,24 +1,18 @@ """Best-effort WebSocket publisher transport for the PTY-side gateway. -The dashboard's `/api/pty` spawns `hermes --tui` as a child process, which -spawns its own ``tui_gateway.entry``. Tool/reasoning/status events fire on -*that* gateway's transport — three processes removed from the dashboard -server itself. To surface them in the dashboard sidebar (`/api/events`), -the PTY-side gateway opens a back-WS to the dashboard at startup and -mirrors every emit through this transport. - -Wire protocol: newline-framed JSON dicts (the same shape the dispatcher -already passes to ``write``). No JSON-RPC envelope here — the dashboard's -``/api/pub`` endpoint just rebroadcasts the bytes verbatim to subscribers. - -Failure mode: silent. The agent loop must never block waiting for the -sidecar to drain. A dead WS short-circuits all subsequent writes. -Actual ``send`` calls run on a daemon thread so the TeeTransport's -``write`` returns after enqueueing (best-effort; drop when the queue is full). +The dashboard's `/api/pty` spawns `hermes --tui`, which spawns its own +``tui_gateway.entry`` — three processes removed from the dashboard server. To surface +tool/reasoning/status events in the sidebar (`/api/events`), that gateway opens a +back-WS to the dashboard at startup and mirrors every emit through this transport as +newline-framed JSON (no JSON-RPC envelope; ``/api/pub`` rebroadcasts bytes verbatim). +Failure mode: silent. The agent loop must never block on the sidecar — ``send`` runs on +a daemon thread, ``write`` returns after enqueueing (drop when full), a dead WS +short-circuits all subsequent writes. """ from __future__ import annotations +import contextlib import json import logging import queue @@ -33,7 +27,6 @@ except ImportError: # pragma: no cover - websockets is a required install path _log = logging.getLogger(__name__) _DRAIN_STOP = object() - _QUEUE_MAX = 256 @@ -47,26 +40,17 @@ class WsPublisherTransport: self._dead = False self._q: queue.Queue[object] = queue.Queue(maxsize=_QUEUE_MAX) self._worker: Optional[threading.Thread] = None - if ws_connect is None: self._dead = True - return - try: self._ws = ws_connect(url, open_timeout=connect_timeout, max_size=None) except Exception as exc: _log.debug("event publisher connect failed: %s", exc) self._dead = True self._ws = None - return - - self._worker = threading.Thread( - target=self._drain, - name="hermes-ws-pub", - daemon=True, - ) + self._worker = threading.Thread(target=self._drain, name="hermes-ws-pub", daemon=True) self._worker.start() def _drain(self) -> None: @@ -74,9 +58,7 @@ class WsPublisherTransport: item = self._q.get() if item is _DRAIN_STOP: return - if not isinstance(item, str): - continue - if self._ws is None: + if not isinstance(item, str) or self._ws is None: continue try: with self._lock: @@ -90,12 +72,8 @@ class WsPublisherTransport: def write(self, obj: dict) -> bool: if self._dead or self._ws is None or self._worker is None: return False - - line = json.dumps(obj, ensure_ascii=False) - try: - self._q.put_nowait(line) - + self._q.put_nowait(json.dumps(obj, ensure_ascii=False)) return True except queue.Full: return False @@ -104,23 +82,14 @@ class WsPublisherTransport: self._dead = True w = self._worker if w is not None and w.is_alive(): - try: + # Best-effort: if the queue is wedged, the daemon thread dies with the process. + with contextlib.suppress(queue.Full): self._q.put_nowait(_DRAIN_STOP) - except queue.Full: - # Best-effort: if the queue is wedged, the daemon thread - # will be torn down with the process. - pass w.join(timeout=3.0) self._worker = None - if self._ws is None: return - - try: - with self._lock: - if self._ws is not None: - self._ws.close() # type: ignore[union-attr] - except Exception: - pass - + with contextlib.suppress(Exception), self._lock: + if self._ws is not None: + self._ws.close() # type: ignore[union-attr] self._ws = None diff --git a/tui_gateway/event_replay.py b/tui_gateway/event_replay.py index 0ec4e2a86b..a86f1505b7 100644 --- a/tui_gateway/event_replay.py +++ b/tui_gateway/event_replay.py @@ -1,19 +1,13 @@ """Per-session event sequencing + bounded replay for WS reconnects. -Every gateway event frame that flows through :func:`server.write_json` (and -therefore ``_emit``) is stamped with a per-session monotonic ``seq`` and -appended to a small ring buffer keyed by session id. A reconnecting client -calls the ``session.events.since`` RPC with its last observed seq; the server -replays everything newer from the buffer, then live events resume seamlessly. - -Design constraints honored: -- stdio TUI path unaffected: frames gain a ``seq`` field only on event frames; - Ink ignores unknown params keys. -- Thread safety: a single module lock guards counters + buffers; write_json - already serializes per-transport writes, so stamping under the lock cannot - reorder frames relative to each other. -- Memory bound: _REPLAY_BUFFER_MAX events / _REPLAY_SESSIONS_MAX sessions, - oldest session evicted FIFO. +Every event frame through :func:`server.write_json` (hence ``_emit``) is stamped with +a per-session monotonic ``seq`` and appended to a small ring per session id; a +reconnecting client calls ``session.events.since`` with its last seen seq and gets +everything newer, then live events resume. Invariants: stdio TUI unaffected (``seq`` +only on event frames; Ink ignores unknown keys); one module lock guards counters + +buffers, and write_json already serializes per-transport writes so stamping cannot +reorder frames; memory bound = _REPLAY_BUFFER_MAX events x _REPLAY_SESSIONS_MAX +sessions, oldest session evicted FIFO. """ from __future__ import annotations @@ -22,24 +16,19 @@ import threading import uuid from collections import OrderedDict, deque -# Process identity for the replay contract. Seq counters live in-process, so -# a gateway restart silently resets them to 1 while clients still hold high -# watermarks — events_since(sid, 97) then returns [] with truncated=False and -# the client believes it missed nothing (and its stale watermark makes every -# future replay empty too). The epoch lets clients detect the restart and -# reset their watermarks. +# Seq counters live in-process, so a restart resets them to 1 while clients hold high +# watermarks — events_since(sid, 97) would return [] with truncated=False forever. The +# epoch lets clients detect the restart and reset their watermarks. _REPLAY_EPOCH = uuid.uuid4().hex -# Replay ring per session. A long turn emits ~hundreds of token events; this -# covers several minutes of streaming plus all control events. +# A long turn emits ~hundreds of token events; 512 covers minutes of streaming plus +# all control events. Desktop users rarely exceed a dozen live chats. _REPLAY_BUFFER_MAX = 512 -# Distinct sessions remembered. Desktop users rarely exceed a dozen live chats. _REPLAY_SESSIONS_MAX = 64 _replay_lock = threading.Lock() -# sid -> deque of (seq, event_object) where event_object is the frame's -# ``params`` dict (bare event: type/session_id/seq/payload) — the exact shape -# the client's dispatch path consumes. +# sid -> deque of (seq, params dict) — the bare event (type/session_id/seq/payload), +# the exact shape the client's dispatch path consumes. _replay_buffers: "OrderedDict[str, deque]" = OrderedDict() _replay_next_seq: dict[str, int] = {} @@ -58,8 +47,7 @@ def _stamp_event(obj: dict) -> None: return sid = params.get("session_id") or "" if not sid: - # Session-less global events (skin.changed etc.) are re-fetchable via - # their own RPCs; no replay contract for them. + # Session-less global events (skin.changed etc.) are re-fetchable via their own RPCs. return with _replay_lock: seq = _replay_next_seq.get(sid, 0) + 1 @@ -67,8 +55,7 @@ def _stamp_event(obj: dict) -> None: params["seq"] = seq buf = _replay_buffers.get(sid) if buf is None: - buf = deque(maxlen=_REPLAY_BUFFER_MAX) - _replay_buffers[sid] = buf + buf = _replay_buffers[sid] = deque(maxlen=_REPLAY_BUFFER_MAX) while len(_replay_buffers) > _REPLAY_SESSIONS_MAX: _oldest_sid, _oldest_buf = _replay_buffers.popitem(last=False) _replay_next_seq.pop(_oldest_sid, None) @@ -76,30 +63,22 @@ def _stamp_event(obj: dict) -> None: def events_since(sid: str, last_seen: int) -> list[dict]: - """Return recorded EVENT OBJECTS with seq > last_seen for *sid*, in order. + """Recorded EVENT OBJECTS (each frame's ``params`` dict) with seq > last_seen for *sid*, in order. - Shape contract: each element is the frame's ``params`` dict — a bare event - object with top-level ``type`` / ``session_id`` / ``seq`` — because that is - exactly what the client's dispatch path consumes. Returning the full - JSON-RPC envelope here would make every replayed event fail the client's - ``event.type`` gate and be silently dropped. + Returning the full JSON-RPC envelope would make every replayed event fail the + client's ``event.type`` gate and be silently dropped. """ with _replay_lock: buf = _replay_buffers.get(sid or "") - if not buf: - return [] - return [event for seq, event in buf if seq > last_seen] + return [event for seq, event in buf if seq > last_seen] if buf else [] def is_truncated(sid: str, last_seen: int) -> bool: - """True when events between *last_seen* and the ring's oldest retained - seq were evicted — the client must refetch history instead of trusting - the replay to be gap-free.""" + """True when events between *last_seen* and the ring's oldest retained seq were + evicted — the client must refetch history instead of trusting the replay.""" with _replay_lock: buf = _replay_buffers.get(sid or "") - if not buf: - return False - return last_seen + 1 < buf[0][0] + return bool(buf) and last_seen + 1 < buf[0][0] def latest_seq(sid: str) -> int: diff --git a/tui_gateway/git_probe.py b/tui_gateway/git_probe.py index 754a3bbd5d..2cc6d5cce9 100644 --- a/tui_gateway/git_probe.py +++ b/tui_gateway/git_probe.py @@ -1,19 +1,12 @@ -"""Git working-tree probing for the gateway: run git, resolve repo roots, fold -linked worktrees under their common root. +"""Git working-tree probing for the gateway: run git, resolve repo roots, fold linked +worktrees under their common root. -Probing runs where the gateway runs, so it covers local and remote backends -(the desktop's electron probe only sees the local fs). Roots go through a -thread-safe single-flight cache: gateway handlers run on worker threads, so -concurrent identical probes share one ``git`` spawn instead of racing a dict. - -Positive results are cached for the process lifetime; negatives (not a repo, or -a deleted dir) only for ``_NEG_TTL``. Caching negatives matters: ``build_tree`` -resolves a cwd once *per session*, so hundreds of sessions in non-git/deleted -dirs would otherwise re-spawn ``git`` on every sidebar open (multi-second -"Projects" load). The TTL keeps a not-yet-repo cwd re-probable — we ``git init`` -a new project's folder on its first worktree, and a frozen "" would mislabel its -main lane by dir basename. ``invalidate()`` drops everything after a known -mutation. +Probing runs where the gateway runs, so it covers local and remote backends. Roots go +through a thread-safe single-flight cache: concurrent identical probes from worker +threads 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 not-yet-``git init``-ed folder re-probable. """ from __future__ import annotations @@ -28,23 +21,21 @@ from hermes_cli._subprocess_compat import bounded_git_probe _GIT_TIMEOUT = 1.5 _WARM_WORKERS = 8 - -# "Not a git repo" cache TTL: short enough that a freshly `git init`-ed folder -# shows correctly within seconds, long enough to collapse a tree build's -# hundreds of redundant probes. +# "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 def run_git(cwd: str, *args: str) -> str: """``git -C `` → stripped stdout, or ``""`` on any failure. - Uses :func:`bounded_git_probe` so post-kill cleanup is bounded on Windows — - a plain ``subprocess.run(timeout=...)`` deadlocked Desktop session readiness - when a killed git left a suspended descendant holding the pipe handles. + ``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 pipe handles. """ + # `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. if not cwd or not os.path.isdir(cwd): - # `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. return "" return bounded_git_probe(["git", "-C", cwd, *args], timeout=_GIT_TIMEOUT) @@ -84,12 +75,10 @@ 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. gate.wait(timeout=_GIT_TIMEOUT + 0.5) continue - value = "" try: value = probe() @@ -122,19 +111,14 @@ def repo_root(cwd: str) -> str: 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, splitting every - worktree into its own "repo"; the parent of the shared ``--git-common-dir`` - is the one true root (fallback: the toplevel root). - - The result is normalized to git's forward-slash spelling so it compares - equal to :func:`repo_root` (raw ``--show-toplevel``). ``os.path.realpath`` - uses native ``\\`` on Windows, so without this the main checkout compared - unequal to its own common root, was misread as a linked worktree, and the + ``--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 desktop 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 that the - # parallel warm can't absorb (`resolve()` only reaches here for repos). + # 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. if not cwd or not repo_root(cwd): return "" @@ -150,12 +134,8 @@ def common_repo_root(cwd: str) -> str: def resolve(cwd: str) -> dict | None: - """Inject-able resolver for ``project_tree.build_tree``. - - Returns ``{"repo_root": , "worktree_root": }`` - or ``None`` when ``cwd`` is not in a git repo. ``build_tree`` treats - ``worktree_root == repo_root`` as the main checkout. - """ + """Inject-able resolver for ``project_tree.build_tree``: ``{repo_root: , + worktree_root: }`` or None outside a repo (equal roots = main checkout).""" worktree_root = repo_root(cwd) if not worktree_root: return None diff --git a/tui_gateway/loop_noise.py b/tui_gateway/loop_noise.py index 321509747e..bf6944ba36 100644 --- a/tui_gateway/loop_noise.py +++ b/tui_gateway/loop_noise.py @@ -1,73 +1,52 @@ """Suppress benign event-loop teardown noise on the gateway serving loop. -When the Desktop client forcibly closes its WebSocket while the gateway still -has pending socket operations, asyncio's transport teardown logs a full -traceback for every pending ``_call_connection_lost`` callback. On Windows this -surfaces as ``ConnectionResetError: [WinError 10054]`` (and the rarer -``ConnectionAbortedError: [WinError 10053]``); on POSIX it is the equivalent -``ConnectionResetError``/``BrokenPipeError``. A single client disconnect can -emit 50+ identical tracebacks into ``errors.log`` (#50005). - -These are not actionable — they are the expected side effect of the peer -hanging up before our writes drained. We install a loop exception handler that -collapses exactly this class of teardown error to one debug line and forwards -everything else to asyncio's default handler unchanged, so genuine loop bugs -still surface. +When the Desktop client forcibly closes its WebSocket while the gateway still has +pending socket operations, asyncio logs a full traceback for every pending +``_call_connection_lost`` callback — ``ConnectionResetError`` (WinError 10054), +``ConnectionAbortedError`` (10053), or ``BrokenPipeError`` on POSIX; one disconnect can +emit 50+ identical tracebacks. They are the expected side effect of the peer hanging up +before our writes drained, so the loop exception handler installed here collapses exactly +that class to one debug line and forwards everything else to the previous handler. """ from __future__ import annotations import asyncio +import contextlib import logging from typing import Any _log = logging.getLogger(__name__) -# Connection-teardown errors that mean "the peer hung up mid-write". WinError -# 10054 (connection reset) and 10053 (connection aborted) raise as these. -_BENIGN_TEARDOWN_ERRORS = ( - ConnectionResetError, - ConnectionAbortedError, - BrokenPipeError, -) +# Connection-teardown errors that mean "the peer hung up mid-write". +_BENIGN_TEARDOWN_ERRORS = (ConnectionResetError, ConnectionAbortedError, BrokenPipeError) def _is_benign_teardown(context: dict[str, Any]) -> bool: """True when the loop error is a peer-hangup during transport teardown. - Gated on BOTH the exception type AND the ``_call_connection_lost`` - callback so we only swallow the disconnect flood — any other place these - errors surface (a real handler, a custom callback) still goes to the - default handler. + Gated on BOTH the exception type AND the ``_call_connection_lost`` callback (matched + on repr) so the same error type raised elsewhere still reaches the default handler. """ - exc = context.get("exception") - if not isinstance(exc, _BENIGN_TEARDOWN_ERRORS): + if not isinstance(context.get("exception"), _BENIGN_TEARDOWN_ERRORS): return False - # The flood originates from the transport's connection-lost callback. Match - # on its repr so we don't suppress the same error type raised elsewhere. - callback = context.get("callback") - handle = context.get("handle") marker = "_call_connection_lost" - return marker in repr(callback) or marker in repr(handle) + return marker in repr(context.get("callback")) or marker in repr(context.get("handle")) def install_loop_noise_filter(loop: asyncio.AbstractEventLoop) -> None: """Chain a teardown-noise filter ahead of the loop's existing handler. - Idempotent: re-installing on a loop that already has the filter is a no-op, - so it's safe to call on every reconnect/serve entry. + Idempotent: a loop already carrying the filter is left alone, so it's safe to call + on every reconnect/serve entry without stacking handlers. """ if getattr(loop, "_hermes_noise_filter_installed", False): return - previous = loop.get_exception_handler() def _handler(loop: asyncio.AbstractEventLoop, context: dict[str, Any]) -> None: if _is_benign_teardown(context): - _log.debug( - "ws peer hangup during teardown (suppressed): %s", - context.get("exception"), - ) + _log.debug("ws peer hangup during teardown (suppressed): %s", context.get("exception")) return if previous is not None: previous(loop, context) @@ -75,9 +54,5 @@ def install_loop_noise_filter(loop: asyncio.AbstractEventLoop) -> None: loop.default_exception_handler(context) loop.set_exception_handler(_handler) - # Mark on the loop instance so a second install (reconnect, re-serve) is a - # no-op rather than stacking handlers. - try: + with contextlib.suppress(AttributeError, TypeError): # pragma: no cover - exotic loop impls loop._hermes_noise_filter_installed = True # type: ignore[attr-defined] - except (AttributeError, TypeError): # pragma: no cover - exotic loop impls - pass diff --git a/tui_gateway/mcp_rpc_helpers.py b/tui_gateway/mcp_rpc_helpers.py index bd3f11161e..41bfebe9b3 100644 --- a/tui_gateway/mcp_rpc_helpers.py +++ b/tui_gateway/mcp_rpc_helpers.py @@ -6,26 +6,24 @@ Published onto ``tui_gateway.server`` as ``_mcp_reset_profile`` / from __future__ import annotations +import contextlib from typing import Any, Dict def reset_profile(token) -> None: if token is None: return - try: + with contextlib.suppress(Exception): from hermes_constants import reset_hermes_home_override reset_hermes_home_override(token) - except Exception: - pass def summarize_server(name: str, cfg: dict) -> Dict[str, Any]: """Serialize one server's config for a UI (no secret values). - Mirrors web_server._mcp_server_summary plus ``oauth_tokens_present`` so a UI - can tell an OAuth server that still needs authentication from one already - authenticated. + Mirrors web_server._mcp_server_summary plus ``oauth_tokens_present`` so a UI can + tell an OAuth server that still needs authentication from one already authenticated. """ from hermes_cli.mcp_config import _oauth_tokens_present diff --git a/tui_gateway/method_ctx.py b/tui_gateway/method_ctx.py index 84ba57bd86..3c6c20e589 100644 --- a/tui_gateway/method_ctx.py +++ b/tui_gateway/method_ctx.py @@ -1,15 +1,11 @@ """Seam for the server.py handler/helper split. -server.py's JSON-RPC handlers and many helpers close over its module globals -(``_sessions``, ``_ok``, ``_err``, config helpers, ...). To move them out -without rewriting a single body, each split module defines its code normally -and server.py calls :func:`bind_module` at the end of its own import, once -every global the code closes over exists. Function 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 -imports them and passes itself in. +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. """ import contextlib @@ -21,11 +17,8 @@ _CM_HELPER_CODE = contextlib.contextmanager(lambda: (yield)).__code__ def rebind(fn, g: dict, _seen=None): - """Copy ``fn`` with globals ``g``. - - Closure cells holding functions from the same module are rebound too, so - handlers produced by import-time decorator factories keep working. - """ + """Copy ``fn`` with globals ``g``; closure cells holding same-module functions are rebound too + (so handlers produced by import-time decorator factories keep working).""" _seen = {} if _seen is None else _seen if id(fn) in _seen: return _seen[id(fn)] @@ -90,14 +83,13 @@ _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; - other values (constants, ``global``-mutated state seeds) are copied as-is. - Imported modules/functions, dunders and registry plumbing are skipped, so a - split module needs no hand-maintained export list. Dispatch tables (dicts - whose values are this module's functions) get their values rebound too. + ``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. """ g = vars(server) @@ -108,7 +100,6 @@ def bind_module(module_globals: dict, server, *, skip=()) -> None: return isinstance(v, types.FunctionType) and v.__module__ == mod_name def _rebind_in(v): - """Rebind own functions nested in dict/tuple/list constants (dispatch tables).""" if _own_fn(v): return rebind(v, g, seen) if isinstance(v, dict): @@ -118,13 +109,11 @@ def bind_module(module_globals: dict, server, *, skip=()) -> None: return v def _has_own_fn(v): - if _own_fn(v): - return True if isinstance(v, dict): return any(_has_own_fn(x) for x in v.values()) if isinstance(v, (tuple, list)): return any(_has_own_fn(x) for x in v) - return False + return _own_fn(v) for name, obj in list(module_globals.items()): if name.startswith("__") or name in _PLUMBING or name in skip: @@ -132,15 +121,12 @@ def bind_module(module_globals: dict, server, *, skip=()) -> None: if isinstance(obj, (types.ModuleType, HandlerRegistry)): continue if isinstance(obj, types.FunctionType): - if obj.__module__ != mod_name: - if name == obj.__name__: - continue # plain import; server already has its own - # ``_alias = other_module.fn`` — publish as-is, no rebind - else: + if obj.__module__ == mod_name: obj = rebind(obj, g, seen) + elif name == obj.__name__: + continue # plain import; server has its own (an ``_alias = other.fn`` publishes as-is) elif isinstance(obj, (dict, tuple, list)) and _has_own_fn(obj): - obj = _rebind_in(obj) - module_globals[name] = obj # keep the split module's own view consistent + obj = module_globals[name] = _rebind_in(obj) # keep the split module's own view in sync elif isinstance(obj, type): if obj.__module__ != mod_name: continue diff --git a/tui_gateway/methods_bot_relay.py b/tui_gateway/methods_bot_relay.py index a633c7eea7..f4d25bfecd 100644 --- a/tui_gateway/methods_bot_relay.py +++ b/tui_gateway/methods_bot_relay.py @@ -1,28 +1,13 @@ """Bot-relay JSON-RPC handlers — the gateway side of cross-connection A2A. -Connections ARE the peer set: every gateway the Desktop holds a socket to -(local, remote URL, SSH, Hermes Cloud, docker) must be able to find every -other connection's agents and message them. The Desktop is the relay — it -owns every socket — and these four methods are the door it uses on EACH -connected gateway: - -- ``bot_relay.roster.sync`` — Desktop pushes the union roster of agents on - the OTHER connections into this gateway's ``bot_relay/roster.json``, so - ``message_agent`` can resolve cross-connection targets and Bot Chat - prompts list them (capability-epoch refresh picks up changes). -- ``bot_relay.outbox.drain`` — Desktop collects envelopes queued here by - ``message_agent`` for targets on other connections. -- ``bot_relay.deliver`` — Desktop hands an envelope to the TARGET - gateway; this method runs the same one-turn Bot Chat delivery local DMs - use and returns the reply text. -- ``bot_relay.reply`` — Desktop writes the reply (or a delivery - error) back on the SENDER gateway; the waiter spawned at send time picks - it up and wakes the sending agent via the standard completion path. - -Storage/validation plumbing lives in ``tools/bot_relay.py``. Handlers are -rebound onto server.py's globals at install time (see method_ctx.py) and may -reference server module globals (``_ok``, ``_err``) not imported here; this -module's own helpers reach them via keyword defaults. +Connections ARE the peer set: the Desktop owns every gateway socket (local, remote, +SSH, Cloud, docker) and relays between them through these four doors on EACH gateway: +``roster.sync`` (push the union roster of OTHER connections' agents so ``message_agent`` +resolves them), ``outbox.drain`` (collect envelopes queued here for other connections), +``deliver`` (run a one-turn Bot Chat delivery on the TARGET gateway, return the reply), +``reply`` (write the reply/error back on the SENDER gateway for the waiter to pick up). +Storage/validation plumbing lives in ``tools/bot_relay.py``. Handlers are rebound onto +server.py's globals (method_ctx.py) and reference ``_ok``/``_err`` etc. bare. """ import os @@ -45,12 +30,8 @@ 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, + local_delivery_command(profile, tmp), capture_output=True, text=True, encoding="utf-8", + errors="replace", timeout=600, ) @@ -58,9 +39,8 @@ def _run_delivery(profile: str, tmp: str) -> subprocess.CompletedProcess: def _(rid, params: dict, _root=_relay_root) -> dict: """Replace this gateway's view of agents on OTHER connections. - Params: ``agents`` — list of rows ``{profile, handle, connection_id, - connection_label?, title?, description?}``. Rows failing validation are - dropped, not fatal. Result: ``{count}`` (accepted rows). + Params: ``agents`` — rows ``{profile, handle, connection_id, connection_label?, title?, + description?}``; rows failing validation are dropped, not fatal. Result: ``{count}``. """ try: from tools.bot_relay import write_remote_roster @@ -72,10 +52,9 @@ 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. + """Claim every pending cross-connection envelope queued on this gateway. Result: ``{envelopes}``. - Claimed envelopes move to ``claimed/`` atomically, so concurrent drains - (two Desktop windows) can't double-deliver. Result: ``{envelopes}``. + Claimed envelopes move to ``claimed/`` atomically, so concurrent drains can't double-deliver. """ try: from tools.bot_relay import claim_pending_envelopes @@ -89,12 +68,10 @@ def _(rid, params: dict, _root=_relay_root) -> dict: def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: """Deliver a relayed DM into a profile's Bot Chat ON THIS GATEWAY. - Params: ``profile`` (target on this install), ``message`` (already - attribution-prefixed by the sender gateway). Runs the same one-turn - ``hermes -p chat -c "Bot Chat"`` transport local DMs use and - returns ``{reply}`` — the target agent's response text. Blocking by - design (the Desktop calls it from its relay worker, off any UI path; - the RPC pool keeps it off the WS reader thread). + Params: ``profile`` (target on this install), ``message`` (already attribution-prefixed). + Runs the same one-turn ``hermes -p chat -c "Bot Chat"`` transport local DMs + use and returns ``{reply}``. Blocking by design (the Desktop calls it from its relay + worker; the RPC pool keeps it off the WS reader thread). """ import os import subprocess @@ -110,7 +87,6 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: if len(message) > MESSAGE_MAX_CHARS + 200: # + attribution headroom return _err(rid, 4091, "message too long") - root = _root() known = {"default"} profiles_dir = root / "profiles" @@ -120,13 +96,11 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: if resolved not in known: return _err(rid, 4092, f"no profile '{profile}' on this gateway") - # When THIS gateway already hosts the target's Bot Chat live (the - # Desktop has it open), the subprocess transport is fenced out by the - # single-owner lease and the payload dropped. Land the DM in the live - # session as a normal user turn via prompt.submit instead — the - # composer's choke point, so role alternation, persistence and - # streaming behave as a typed message would. (Nested: needs server - # globals via method_ctx rebinding.) + # 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 method_ctx rebinding.) def _live_bot_chat_sid(profile_name: str) -> str: from tools.bot_mode_probe import BOT_CHAT_TITLE @@ -144,58 +118,46 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: live_sid = _live_bot_chat_sid(resolved) if live_sid: - # queued=True: a teammate's DM runs as the NEXT turn. It must never - # interrupt or steer a turn already in flight (the default busy - # mode does); hundreds of arrivals simply queue in arrival order. + # 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. submitted = _methods["prompt.submit"](rid, {"session_id": live_sid, "text": message, "queued": True}) if "error" in submitted: return submitted return _ok( - rid, - {"reply": f"Delivered into @{resolved}'s open Bot Chat; the reply will appear there."}, + rid, {"reply": f"Delivered into @{resolved}'s open Bot Chat; the reply will appear there."} ) fd, tmp = tempfile.mkstemp(prefix="hermes-relay-dm-", suffix=".txt", text=True) try: with os.fdopen(fd, "w", encoding="utf-8") as f: f.write(message) - # Per-profile turn lock serializes with any other delivery turn into - # this profile (relay or local message_agent) and covers only the - # turn execution window. Worst-case handler hold is lock wait - # (bot_mode.turn_wait_seconds, default 120s) + the 600s turn timeout, - # doubled when the retry policy grants one bounded re-run — callers - # must tolerate ~1320s before assuming failure. + # Per-profile turn lock serializes with any other delivery turn into this + # profile and covers only the turn window. Worst-case hold is lock wait + # (bot_mode.turn_wait_seconds, default 120s) + the 600s turn timeout, doubled + # when the retry policy grants one re-run — callers must tolerate ~1320s. with acquire_turn_lock(root, resolved): proc = _run(resolved, tmp) if proc.returncode != 0: - # Retry policy: transient classes re-run the SAME session - # once; context_overflow also re-runs the same session — the - # retried turn's pre-API compaction pass compacts the - # over-threshold Bot Chat transcript first (the sanctioned - # compression lever; no fresh session is ever minted). - # Auth/quota/config classes never retry. + # Retry policy: transient classes re-run the SAME session once; + # context_overflow too — the retried turn's pre-API compaction pass + # compacts the over-threshold transcript first (no fresh session is + # ever minted). Auth/quota/config classes never retry. from tools.bot_failure_reasons import ( - RETRY_NONE, - classify_agent_error, - retry_action, + 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: proc = _run(resolved, tmp) finally: - try: + with contextlib.suppress(OSError): os.unlink(tmp) - except OSError: - pass 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}", + rid, 5092, f"delivery turn failed: {detail or proc.returncode}", data={"reason": classify_agent_error(detail)}, ) return _ok(rid, {"reply": (proc.stdout or "").strip()}) @@ -212,8 +174,8 @@ def _(rid, params: dict, _root=_relay_root, _run=_run_delivery) -> dict: def _(rid, params: dict, _root=_relay_root) -> dict: """Write a relayed reply (or delivery error) for a sender-side waiter. - Params: ``id`` (envelope id), ``reply`` and/or ``error``, optional - ``reason`` (typed failure code, see ``tools.bot_failure_reasons``). + Params: ``id`` (envelope id), ``reply`` and/or ``error``, optional ``reason`` + (typed failure code, see ``tools.bot_failure_reasons``). """ envelope_id = str(params.get("id") or "").strip() if not envelope_id: @@ -222,11 +184,8 @@ def _(rid, params: dict, _root=_relay_root) -> dict: 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 ""), + _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: @@ -241,12 +200,8 @@ def register(server) -> None: server._LONG_HANDLERS = server._LONG_HANDLERS | methods_groups.LONG_HANDLERS for name in ( - "get_hosted_room_service", - "_WORKER_UNAVAILABLE", - "_profile_name", - "_requested_profile", - "_api_server_key", - "_room_link_run_storage_durable", + "get_hosted_room_service", "_WORKER_UNAVAILABLE", "_profile_name", "_requested_profile", + "_api_server_key", "_room_link_run_storage_durable", ): setattr(server, name, getattr(methods_groups, name)) methods_groups.bind_server(server) diff --git a/tui_gateway/methods_browser.py b/tui_gateway/methods_browser.py index c7d64a49c2..c06641f9b7 100644 --- a/tui_gateway/methods_browser.py +++ b/tui_gateway/methods_browser.py @@ -6,11 +6,12 @@ method_ctx.bind_module), so they reference server.py globals bare. from __future__ import annotations - from .method_ctx import HandlerRegistry, bind_module _registry = HandlerRegistry() +_CDP_SCHEMES = {"http", "https", "ws", "wss"} + def _resolve_browser_cdp_url() -> str: """Configured browser CDP override without network I/O. @@ -52,20 +53,20 @@ def _is_default_local_cdp(parsed) -> bool: ) -def _http_ok(url: str, timeout: float) -> bool: +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 - try: - with urllib.request.urlopen(url, timeout=timeout) as resp: - return 200 <= getattr(resp, "status", 200) < 300 - except Exception: - return False - - -def _probe_urls(parsed) -> list[str]: scheme = {"ws": "http", "wss": "https"}.get(parsed.scheme, parsed.scheme) root = f"{scheme}://{parsed.netloc}".rstrip("/") - return [f"{root}/json/version", f"{root}/json"] + 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 + return False def _normalize_cdp_url(parsed) -> str: @@ -76,7 +77,7 @@ def _normalize_cdp_url(parsed) -> str: return parsed._replace(path="", params="", query="", fragment="").geturl() -def _failure_messages(url: str, port: int, system: str) -> list[str]: +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) @@ -89,12 +90,53 @@ def _failure_messages(url: str, port: int, system: str) -> list[str]: ] ) return [ - f"Browser CDP is not reachable at {url}.", *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 → its CDP URL, or None after announcing failure.""" + from hermes_cli.browser_connect import ( + discover_local_cdp_url, find_free_debug_port, launch_chrome_debug, local_port_in_use, + ) + + # 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 against squatters that accept TCP but never answer HTTP. + discovered = discover_local_cdp_url(port, timeout=2.0) + if discovered is not None: + announce(f"Chromium-family browser is already listening at {discovered}") + return discovered + 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..." + ) + else: + announce("Chromium-family browser isn't running with remote debugging — attempting to launch...") + launch = launch_chrome_debug(launch_port, system) + if launch.launched: + # 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: + break + time.sleep(0.5) + if discovered: + announce(f"Chromium-family browser launched and listening on port {launch_port}") + return discovered + if launch.hint: + announce(launch.hint, level="error") + for line in _launch_failure_hints(launch_port, system): + announce(line, level="error") + return None + + def _browser_connect(rid, params: dict) -> dict: import platform @@ -106,7 +148,6 @@ def _browser_connect(rid, params: dict) -> dict: 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__}") url = (raw_url or "").strip() or DEFAULT_BROWSER_CDP_URL - sid = params.get("session_id") or "" system = platform.system() messages: list[str] = [] @@ -118,7 +159,7 @@ def _browser_connect(rid, params: dict) -> dict: _emit("browser.progress", sid, {"message": message, "level": level}) parsed = urlparse(url if "://" in url else f"http://{url}") - if parsed.scheme not in {"http", "https", "ws", "wss"}: + if parsed.scheme not in _CDP_SCHEMES: return _err(rid, 4015, f"unsupported browser url: {url}") if not parsed.hostname: return _err(rid, 4015, f"missing host in browser url: {url}") @@ -126,13 +167,11 @@ def _browser_connect(rid, params: dict) -> dict: port = parsed.port or (443 if parsed.scheme in {"https", "wss"} else 80) except ValueError: return _err(rid, 4015, f"invalid port in browser url: {url}") - # Normalize default-local to 127.0.0.1:9222 so comparisons + messaging match what we persist. if _is_default_local_cdp(parsed): url = DEFAULT_BROWSER_CDP_URL parsed = urlparse(url) port = parsed.port or 9222 - try: # Hosted ws[s]://.../devtools/browser/ endpoints don't serve the HTTP discovery # path: check TCP reachability only and let browser_navigate handshake. @@ -145,60 +184,15 @@ def _browser_connect(rid, params: dict) -> dict: except OSError as e: return _err(rid, 5031, f"could not reach browser CDP at {url}: {e}") elif _is_default_local_cdp(parsed): - from hermes_cli.browser_connect import ( - discover_local_cdp_url, - find_free_debug_port, - launch_chrome_debug, - local_port_in_use, - ) - - # 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 against squatters that accept TCP but never answer HTTP. - discovered = discover_local_cdp_url(port, timeout=2.0) - launch_port = port - + discovered = _connect_local_default(port, system, announce) if discovered is None: - 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..." - ) - else: - announce("Chromium-family browser isn't running with remote debugging — attempting to launch...") - - launch = launch_chrome_debug(launch_port, system) - if launch.launched: - # 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: - break - time.sleep(0.5) - - if discovered: - announce(f"Chromium-family browser launched and listening on port {launch_port}") - else: - hint = launch.hint - if hint: - announce(hint, level="error") - for line in _failure_messages(url, launch_port, system)[1:]: - announce(line, level="error") - return _ok(rid, {"connected": False, "url": url, "messages": messages}) - else: - announce(f"Chromium-family browser is already listening at {discovered}") - + return _ok(rid, {"connected": False, "url": url, "messages": messages}) # Adopt whatever loopback/port answered ([::1] and/or an alternate port when 9222 was squatted). url = discovered parsed = urlparse(url) - elif not any(_http_ok(p, timeout=2.0) for p in _probe_urls(parsed)): + elif not _cdp_http_reachable(parsed): return _err(rid, 5031, f"could not reach browser CDP at {url}") - normalized = _normalize_cdp_url(parsed) - # 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() @@ -206,7 +200,6 @@ 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 @@ -216,12 +209,10 @@ def _browser_connect(rid, params: dict) -> dict: def _browser_disconnect(rid) -> dict: # Reap, drop the override, reap again — same swap window as ``_browser_connect``. def reap() -> None: - try: + with contextlib.suppress(Exception): from tools.browser_tool import cleanup_all_browsers cleanup_all_browsers() - except Exception: - pass reap() os.environ.pop("BROWSER_CDP_URL", None) diff --git a/tui_gateway/methods_browser_control.py b/tui_gateway/methods_browser_control.py index f1a3547ab2..f42ad94ade 100644 --- a/tui_gateway/methods_browser_control.py +++ b/tui_gateway/methods_browser_control.py @@ -1,34 +1,15 @@ """Browser controller registration and result routing for the dashboard. -The dashboard's browser controller (the extension that physically drives a -browser) registers itself over the authenticated ``/api/ws`` JSON-RPC -gateway. Everything here is bound to the **server-minted identity** that the -dashboard auth layer stamped onto the WS connection: ``hermes_cli.web_server`` -consumes the single-use ticket and records ``ws._hermes_auth_identity``, the -WS transport carries it as ``WSTransport.auth_identity``, and the client can -never name its own principal (a spoofed ``principal_id`` param is ignored and -replaced by a server-derived digest of the authenticated identity). - -Registration attaches the shared transport-neutral broker -(:mod:`gateway.browser_control_broker`) with the calling transport as owner; -broker command/cancel frames are wrapped as standard Gateway ``event`` frames -(``type`` = broker method name, ``payload`` = broker params, plus the owning -``session_id``) so the dashboard consumes the same envelope as every other -gateway event. ``browser.controller.result`` resolves a pending command only -when the request arrives on the same transport that owns the session, and -only for the exact attached scope — the broker's exact-scope ``complete`` is -the last line of defense against cross-tenant completion. - -Both dashboard and local API transports use the broker's shared, explicit -capability allowlist. Raw CDP, script evaluation, console access, uploads, and -other privileged surfaces are not controller capabilities. - -Note on handler globals: ``HandlerRegistry.install`` (method_ctx.py) rebinds -each handler's ``__globals__`` onto server.py's namespace, so handler bodies -may only reference names server.py defines/imports (``_ok``, ``_err``, -``_sessions``, ``_sessions_lock``, ``current_transport``, ``logger``, ...). -This module's own helpers and constants reach the handlers through closure -cells of :func:`_controller_method` (rebind preserves and re-targets them). +The dashboard's browser controller (the extension driving a browser) registers over +the authenticated ``/api/ws`` gateway. Everything binds to the SERVER-MINTED identity +(``WSTransport.auth_identity``, stamped by ``hermes_cli.web_server`` from the single-use +ticket); a client-supplied ``principal_id`` is ignored and replaced by a digest of it. +Broker command/cancel frames are re-enveloped as standard Gateway ``event`` frames; +``browser.controller.result`` resolves a command only on the owning transport and only +for the exact attached scope (the broker's exact-scope ``complete`` is the backstop). +Both transports share the broker's explicit capability allowlist (no raw CDP/eval/uploads). +Handler bodies are rebound onto server.py's globals (method_ctx.bind_module publishes +this module's helpers/constants there too), so they reference both bare. """ from __future__ import annotations @@ -36,14 +17,8 @@ from __future__ import annotations import hashlib import logging -from gateway.browser_control_broker import ( - BROWSER_CONTROL_PROTOCOL_VERSION, - browser_control_protocol_supported, - filter_browser_control_capabilities, -) from hermes_cli.dashboard_auth.ws_tickets import ( - INTERNAL_PROVIDER as _INTERNAL_PROVIDER, - INTERNAL_USER_ID as _INTERNAL_USER_ID, + INTERNAL_PROVIDER as _INTERNAL_PROVIDER, INTERNAL_USER_ID as _INTERNAL_USER_ID, ) from .method_ctx import HandlerRegistry, bind_module @@ -53,14 +28,10 @@ logger = logging.getLogger(__name__) _registry = HandlerRegistry() method = _registry.method -#: Transport family stamped into every scope attached from this gateway. The -#: broker's exact-match contract treats it as an identity field, so an API -#: transport can never address a dashboard controller (and vice versa). +# Transport family stamped into every scope attached here; the broker treats it as an +# identity field, so an API transport can never address a dashboard controller. _CLOUD_TRANSPORT_FAMILY = "cloud-ticket-ws" - -#: JSON-RPC error code for identity / session / flag denials (forbidden). -_ERR_FORBIDDEN = 4403 - +_ERR_FORBIDDEN = 4403 # identity / session / flag denials _IDENTITY_REQUIRED = "authenticated controller identity required" _NOT_OWNED = "controller is not owned by this transport" _NO_CONTROLLER = "no controller registered for this session" @@ -70,8 +41,7 @@ def _is_authenticated_identity(identity: object) -> bool: """True for a server-minted, non-internal ``{user_id, provider}`` identity.""" if not isinstance(identity, dict): return False - user_id = identity.get("user_id") - provider = identity.get("provider") + user_id, provider = identity.get("user_id"), identity.get("provider") if not isinstance(user_id, str) or not user_id.strip(): return False if not isinstance(provider, str) or not provider.strip(): @@ -80,45 +50,26 @@ def _is_authenticated_identity(identity: object) -> bool: def _principal_digest(identity: dict) -> str: - """Server-derived principal id: a digest of the server-minted identity. - - The client-supplied ``principal_id`` RPC param is never trusted; the - digest is deterministic (stable across reconnects for the same user) but - unspoofable by a peer that does not hold the authenticated identity. - """ + """Server-derived principal id: stable per user, unspoofable without the authenticated identity.""" raw = f"{identity.get('provider')}\x00{identity.get('user_id')}" - digest = hashlib.sha256(raw.encode("utf-8")).hexdigest() - return f"principal:dashboard:{digest[:32]}" + return f"principal:dashboard:{hashlib.sha256(raw.encode('utf-8')).hexdigest()[:32]}" def _broker_event_writer(transport: object, session_id: str): - """Wrap broker command/cancel frames as standard Gateway event frames. - - The broker's send callback receives transport-neutral envelopes like - ``{"method": "browser.controller.command", "params": {...}}``; the - dashboard speaks Gateway events, so we re-envelope them: - ``{"jsonrpc": "2.0", "method": "event", "params": {"type": , - "session_id": , "payload": }}``. - """ + """Broker send callback: re-envelope ``{method, params}`` as a Gateway ``event`` frame + (``type`` = method, ``payload`` = params, plus the owning ``session_id``).""" def send(frame: dict) -> None: try: - accepted = transport.write( - { - "jsonrpc": "2.0", - "method": "event", - "params": { - "type": frame.get("method"), - "session_id": session_id, - "payload": frame.get("params"), - }, - } - ) + accepted = transport.write({ + "jsonrpc": "2.0", "method": "event", + "params": { + "type": frame.get("method"), "session_id": session_id, "payload": frame.get("params"), + }, + }) except Exception: logger.exception( - "browser controller event write failed session=%s frame=%s", - session_id, - frame.get("method"), + "browser controller event write failed session=%s frame=%s", session_id, frame.get("method") ) raise if accepted is False: @@ -128,28 +79,17 @@ def _broker_event_writer(transport: object, session_id: str): def _controller_method( - name: str, - *, - identity_message: str = _IDENTITY_REQUIRED, - lookup_scope: bool = True, - missing_scope_message: str = _NO_CONTROLLER, - precheck=None, + name: str, *, identity_message: str = _IDENTITY_REQUIRED, lookup_scope: bool = True, + missing_scope_message: str = _NO_CONTROLLER, precheck=None, ): """Register a handler behind the shared fail-closed (4403) controller gates. - Order: ``precheck(rid, params)`` (may return an error envelope) → the - calling transport holds a server-authenticated, non-internal identity - (``WSTransport.auth_identity``, never the RPC params) → the named session - exists and its ``transport`` is exactly the caller → when ``lookup_scope``, - a controller scope is attached for this session/principal/family and the - caller owns it. ``fn(rid, params, transport, identity, session_id, broker, - scope)`` then runs (``scope`` is ``None`` when ``lookup_scope`` is off). + Order: ``precheck(rid, params)`` (may return an error envelope) → caller holds a + server-authenticated, non-internal identity → the named session exists and its + ``transport`` is exactly the caller → when ``lookup_scope``, a scope is attached for + this session/principal/family and the caller owns it. Then + ``fn(rid, params, transport, identity, session_id, broker, scope, session)`` runs. """ - forbidden = _ERR_FORBIDDEN - family = _CLOUD_TRANSPORT_FAMILY - identity_ok = _is_authenticated_identity - digest = _principal_digest - not_owned = _NOT_OWNED def dec(fn): def handler(rid, params: dict) -> dict: @@ -161,28 +101,26 @@ def _controller_method( return denied transport = current_transport() identity = getattr(transport, "auth_identity", None) - if not identity_ok(identity): - return _err(rid, forbidden, identity_message) + if not _is_authenticated_identity(identity): + return _err(rid, _ERR_FORBIDDEN, identity_message) session_id = str(params.get("session_id") or "") with _sessions_lock: session = _sessions.get(session_id) if session is None or session.get("transport") is not transport: - return _err(rid, forbidden, "session is not owned by this transport") + return _err(rid, _ERR_FORBIDDEN, "session is not owned by this transport") broker = browser_control_broker.get_browser_control_broker() scope = None if lookup_scope: scope = broker.scope_for_session( - session_id=session_id, - principal_id=digest(identity), - transport_family=family, + session_id=session_id, principal_id=_principal_digest(identity), + transport_family=_CLOUD_TRANSPORT_FAMILY, ) if scope is None: - return _err(rid, forbidden, missing_scope_message) - # Defense in depth: the broker's exact-scope operations already - # reject foreign scopes; the owner check makes the "same - # transport" rule explicit at this layer too. + return _err(rid, _ERR_FORBIDDEN, missing_scope_message) + # Defense in depth: the broker's exact-scope ops already reject foreign + # scopes; the owner check makes the same-transport rule explicit here too. if not broker.is_owner(scope, transport): - return _err(rid, forbidden, not_owned) + return _err(rid, _ERR_FORBIDDEN, _NOT_OWNED) return fn(rid, params, transport, identity, session_id, broker, scope, session) handler.__doc__ = fn.__doc__ @@ -191,54 +129,28 @@ def _controller_method( return dec -def _register_precheck( - rid, - params: dict, - _forbidden=_ERR_FORBIDDEN, - _protocol_version=BROWSER_CONTROL_PROTOCOL_VERSION, - _protocol_supported=browser_control_protocol_supported, -): +def _register_precheck(rid, params: dict): from gateway import browser_control_broker if not browser_control_broker.browser_control_enabled(): - return _err(rid, _forbidden, "browser.extension_control.enabled is not set") - if not _protocol_supported(params.get("protocol_version")): - return _err( - rid, - _forbidden, - f"unsupported browser-control protocol version; expected {_protocol_version}", - ) + return _err(rid, _ERR_FORBIDDEN, "browser.extension_control.enabled is not set") + if not browser_control_broker.browser_control_protocol_supported(params.get("protocol_version")): + expected = browser_control_broker.BROWSER_CONTROL_PROTOCOL_VERSION + return _err(rid, _ERR_FORBIDDEN, f"unsupported browser-control protocol version; expected {expected}") return None @_controller_method( "browser.controller.register", identity_message="browser.controller.register requires an authenticated non-internal identity", - lookup_scope=False, - precheck=_register_precheck, + lookup_scope=False, precheck=_register_precheck, ) -def _( - rid, - params: dict, - transport, - identity, - session_id, - broker, - _scope, - session, - _family=_CLOUD_TRANSPORT_FAMILY, - _forbidden=_ERR_FORBIDDEN, - _filter_capabilities=filter_browser_control_capabilities, - _digest=_principal_digest, - _event_writer=_broker_event_writer, -) -> dict: +def _(rid, params: dict, transport, identity, session_id, broker, _scope, session) -> dict: """Attach this connection as the browser controller for one session. - Fails closed (4403) unless the ``browser.extension_control.enabled`` flag - is on, the protocol version is supported, the shared identity/session - gates pass, and at least one requested capability survives the allowlist. - The returned ``scope`` names a server-derived ``principal_id``, the - ``cloud-ticket-ws`` transport family, and the filtered capability set. + Fails closed (4403) unless the ``browser.extension_control.enabled`` flag is on, the + protocol version is supported, the identity/session gates pass, and at least one + requested capability survives the allowlist. """ from gateway import browser_control_broker @@ -247,60 +159,43 @@ def _( profile_id = str(session.get("profile") or "").strip() if not controller_id or not browser_profile_id or not profile_id: return _err( - rid, - _forbidden, - "controller_id, browser_profile_id, and server session profile are required", + rid, _ERR_FORBIDDEN, "controller_id, browser_profile_id, and server session profile are required" ) - - capabilities = _filter_capabilities(params.get("capabilities")) + capabilities = browser_control_broker.filter_browser_control_capabilities(params.get("capabilities")) if not capabilities: - return _err(rid, _forbidden, "no permitted controller capabilities requested") - + return _err(rid, _ERR_FORBIDDEN, "no permitted controller capabilities requested") scope = browser_control_broker.ControllerScope( - principal_id=_digest(identity), - profile_id=profile_id, - session_id=session_id, - controller_id=controller_id, - browser_profile_id=browser_profile_id, - transport_family=_family, - capabilities=capabilities, - ) - broker.attach(scope, _event_writer(transport, session_id), owner=transport) - return _ok( - rid, - { - "scope": { - "principal_id": scope.principal_id, - "profile_id": scope.profile_id, - "session_id": scope.session_id, - "controller_id": scope.controller_id, - "browser_profile_id": scope.browser_profile_id, - "transport_family": scope.transport_family, - "capabilities": sorted(scope.capabilities), - } - }, + principal_id=_principal_digest(identity), profile_id=profile_id, session_id=session_id, + controller_id=controller_id, browser_profile_id=browser_profile_id, + transport_family=_CLOUD_TRANSPORT_FAMILY, capabilities=capabilities, ) + broker.attach(scope, _broker_event_writer(transport, session_id), owner=transport) + return _ok(rid, { + "scope": { + "principal_id": scope.principal_id, + "profile_id": scope.profile_id, + "session_id": scope.session_id, + "controller_id": scope.controller_id, + "browser_profile_id": scope.browser_profile_id, + "transport_family": scope.transport_family, + "capabilities": sorted(scope.capabilities), + } + }) @_controller_method("browser.controller.result") -def _(rid, params: dict, _transport, _identity, _session_id, broker, scope, _session, _forbidden=_ERR_FORBIDDEN) -> dict: +def _(rid, params: dict, _transport, _identity, _session_id, broker, scope, _session) -> dict: """Deliver one controller command result back to the broker. - Only the transport that owns the session may resolve its commands, and - only against the exact scope attached for that session (the broker's - exact-scope ``complete`` rejects any other scope). ``accepted`` is - ``False`` for unknown / already-resolved / cancelled command ids — the - broker's idempotent answer, surfaced verbatim. + ``accepted`` is ``False`` for unknown / already-resolved / cancelled command ids — + the broker's idempotent answer, surfaced verbatim. """ command_id = str(params.get("command_id") or "") if not command_id: - return _err(rid, _forbidden, "command_id required") + return _err(rid, _ERR_FORBIDDEN, "command_id required") ok = params.get("ok") is True accepted = broker.complete( - command_id, - scope=scope, - ok=ok, - result=params.get("result") if ok else params.get("error"), + command_id, scope=scope, ok=ok, result=params.get("result") if ok else params.get("error") ) return _ok(rid, {"accepted": accepted}) @@ -319,10 +214,5 @@ def _(rid, params: dict, transport, _identity, _session_id, broker, scope, _sess def register(server) -> None: - """Publish this module's helpers/constants onto ``server`` and install handlers. - - ``rebind`` re-targets closure cells that hold this module's functions, so - the helpers (and the constants they read) must exist in server.py's - namespace too — ``bind_module`` publishes them. - """ + """Publish helpers/constants onto ``server`` and install handlers (rebound to its globals).""" bind_module(globals(), server, skip=("_",)) diff --git a/tui_gateway/model_switch.py b/tui_gateway/model_switch.py index e9ad77844e..8df0d1ad55 100644 --- a/tui_gateway/model_switch.py +++ b/tui_gateway/model_switch.py @@ -83,10 +83,8 @@ def _restart_completed_failed_agent_build(sid: str, session: dict, failed_ready: 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() + 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") @@ -107,29 +105,18 @@ def _restart_completed_failed_agent_build(sid: str, session: dict, failed_ready: return True -def _apply_model_switch( - sid: str, - session: dict, - raw_input: str, - *, - confirm_expensive_model: bool = False, - pin_session_override: bool = True, - parsed_flags: Any | None = None, - persist_override: bool | None = None, -) -> dict: +def _switch_request(raw_input: str, parsed_flags, persist_override) -> tuple[str, str, bool, bool]: + """Normalize /model flags → (model_input, explicit_provider, one_turn, persist_global); raises on conflict.""" from hermes_cli.model_switch import ( - parse_model_switch_args, resolve_persist_behavior, switch_model, - MODEL_SWITCH_ERR_ONCE_WITH_GLOBAL, MODEL_SWITCH_ERROR_TEXT, + MODEL_SWITCH_ERR_ONCE_WITH_GLOBAL, MODEL_SWITCH_ERROR_TEXT, parse_model_switch_args, + resolve_persist_behavior, ) - from hermes_cli.runtime_provider import resolve_runtime_provider if parsed_flags is None: parsed_flags = parse_model_switch_args(raw_input) if hasattr(parsed_flags, "model_input"): - model_input = parsed_flags.model_input - explicit_provider = parsed_flags.explicit_provider - is_global_flag = parsed_flags.is_global - is_session = parsed_flags.is_session + 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 @@ -137,39 +124,36 @@ def _apply_model_switch( # 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]) - persist_global = ( - persist_override - if persist_override is not None - else resolve_persist_behavior(is_global_flag, is_session, is_once=one_turn, explicit_provider=explicit_provider) - ) + if persist_override is None: + persist_override = resolve_persist_behavior( + is_global_flag, is_session, is_once=one_turn, explicit_provider=explicit_provider + ) if not model_input: raise ValueError("model value required") + return model_input, explicit_provider, one_turn, persist_override - agent = session.get("agent") - if one_turn and not agent: - raise ValueError("/model --once requires a live session") + +def _current_model_runtime(agent, explicit_provider: str) -> tuple: + """(provider, model, base_url, api_key) the switch starts from: the live agent's, else the configured runtime.""" if agent: - current_provider = getattr(agent, "provider", "") or "" - current_model = getattr(agent, "model", "") or "" - current_base_url = getattr(agent, "base_url", "") or "" - current_api_key = getattr(agent, "api_key", "") or "" - else: - current_model = _resolve_model() - current_provider = explicit_provider.strip() - current_base_url = "" - current_api_key = "" - if not explicit_provider: - runtime = resolve_runtime_provider(requested=None) - current_provider = str(runtime.get("provider", "") or "") - current_base_url = str(runtime.get("base_url", "") or "") - # Keep a callable api_key (Azure Entra bearer) unchanged: ``str()`` would - # yield "" and poison switch_model validation. - _runtime_key = runtime.get("api_key", "") - is_bearer = callable(_runtime_key) and not isinstance(_runtime_key, str) - current_api_key = _runtime_key if is_bearer else str(_runtime_key or "") + return tuple(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 - # User-defined providers let switch_model resolve named custom endpoints - # (e.g. "ollama-launch") and validate against saved model lists. + runtime = resolve_runtime_provider(requested=None) + # Keep a callable api_key (Azure Entra bearer) unchanged: ``str()`` would + # yield "" and poison switch_model validation. + key = runtime.get("api_key", "") + if not (callable(key) and not isinstance(key, str)): + key = str(key or "") + provider = str(runtime.get("provider", "") or "") + 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 when config fails to load.""" user_provs = custom_provs = cfg = None try: from hermes_cli.config import get_compatible_custom_providers, load_config @@ -179,7 +163,95 @@ def _apply_model_switch( 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"]) + merge_preflight_compression_warning( + result, agent=agent, messages=list(session.get("history", [])), + custom_providers=custom_provs, config_context_length=cfg_ctx, + ) + except Exception as exc: + logger.debug("preflight-compression switch warning failed: %s", exc) + + +def _expensive_model_confirm(result, current_base_url: str, current_api_key) -> dict | None: + """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, + ) + except Exception: + 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} + + +def _commit_agent_switch(sid: str, session: dict, agent, result, current_model: str, restore_snapshot): + """Swap the live agent in place, then restart/persist/mark/announce; a failed swap aborts it all.""" + try: + agent.switch_model( + new_model=result.new_model, new_provider=result.target_provider, api_key=result.api_key, + base_url=result.base_url, api_mode=result.api_mode, + capabilities=getattr(result, "runtime_capabilities", None), + ) + except Exception as exc: + # 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 stays pinned to a broken model. A failed switch is a no-op. + 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 + _restart_slash_worker(sid, session) + _persist_live_session_runtime(session) + _persist_live_session_system_prompt(session) + _append_model_switch_marker(session, model=result.new_model, provider=result.target_provider) + _emit_session_info(sid, session) + if restore_snapshot is not None: + session["one_turn_model_restore"] = restore_snapshot + else: + session.pop("one_turn_model_restore", None) + + +def _apply_model_switch( + sid: str, session: dict, raw_input: str, *, confirm_expensive_model: bool = False, + 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") + if one_turn and not agent: + raise ValueError("/model --once requires a live session") + current_provider, current_model, current_base_url, current_api_key = _current_model_runtime( + 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() 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, @@ -187,69 +259,15 @@ def _apply_model_switch( ) if not result.success: raise ValueError(result.error_message or "model switch failed") - restore_snapshot = _snapshot_agent_model_runtime(agent) if (one_turn and agent) else None - if agent: - try: - from hermes_cli.context_switch_guard import merge_preflight_compression_warning - - _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"]) - merge_preflight_compression_warning( - result, agent=agent, messages=list(session.get("history", [])), - custom_providers=custom_provs, config_context_length=_cfg_ctx, - ) - except Exception as exc: - logger.debug("preflight-compression switch warning failed: %s", exc) - + _merge_preflight_warning(result, agent, session, cfg, custom_provs) if not confirm_expensive_model: - 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, - ) - except Exception: - warning = None - if warning is not 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} - + confirm = _expensive_model_confirm(result, current_base_url, current_api_key) + if confirm is not None: + return confirm if agent: - try: - agent.switch_model( - new_model=result.new_model, new_provider=result.target_provider, api_key=result.api_key, - base_url=result.base_url, api_mode=result.api_mode, - capabilities=getattr(result, "runtime_capabilities", None), - ) - except Exception as exc: - # 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 stays pinned to a broken model. A failed switch is a no-op. - 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 - _restart_slash_worker(sid, session) - _persist_live_session_runtime(session) - _persist_live_session_system_prompt(session) - _append_model_switch_marker(session, model=result.new_model, provider=result.target_provider) - _emit_session_info(sid, session) - if one_turn: - session["one_turn_model_restore"] = restore_snapshot - else: - session.pop("one_turn_model_restore", None) - + _commit_agent_switch(sid, session, agent, result, current_model, restore_snapshot) # PER-SESSION override so a rebuild of THIS session (/new, resume) re-derives # the chosen model. Deliberately NOT written to process-global env vars # (HERMES_MODEL & co.): the desktop hosts every same-profile session in one @@ -317,9 +335,10 @@ def _sync_bot_capabilities(sid: str, session: dict) -> None: 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 and never blocks the turn. """ agent = session.get("agent") if agent is None or session.get("model_override"): @@ -363,7 +382,6 @@ def _pending_switch_selection_warning(model: str, provider: str) -> str | None: warning = combined_selection_warning(model, provider=provider or None) except Exception: return None - return warning.message if warning is not None else None