From b5e1d307d69e02a86ae53f42a2a19e93468ed360 Mon Sep 17 00:00:00 2001 From: lsaether <25539605+lsaether@users.noreply.github.com> Date: Sun, 31 May 2026 12:28:43 -0500 Subject: [PATCH 1/2] feat(tui): add opt-in remote bridge listener Add a disabled-by-default remote TUI bridge that can mirror and control live TUI sessions from a separate client without stealing the local TUI transport. The bridge uses explicit config/env enablement, loopback-first defaults, token validation for non-loopback binds, host/origin checks, and a bridge-scoped RPC allowlist. Also add reconnect hydration from the display journal, active-turn peer prompt handling, and regression coverage across backend gateway and TUI event handling paths. --- cli-config.yaml.example | 29 + hermes_cli/config.py | 13 + tests/test_tui_gateway_server.py | 3 + tests/tui_gateway/test_remote_bridge.py | 871 ++++++++++++++++++ tools/approval.py | 6 + tui_gateway/entry.py | 12 +- tui_gateway/remote_bridge.py | 411 +++++++++ tui_gateway/server.py | 519 +++++++++-- tui_gateway/ws.py | 80 +- .../createGatewayEventHandler.test.ts | 59 ++ ui-tui/src/app/createGatewayEventHandler.ts | 48 + ui-tui/src/app/turnController.ts | 13 +- ui-tui/src/app/useInputHandlers.ts | 6 +- ui-tui/src/app/useMainApp.ts | 8 +- ui-tui/src/gatewayTypes.ts | 20 + 15 files changed, 2025 insertions(+), 73 deletions(-) create mode 100644 tests/tui_gateway/test_remote_bridge.py create mode 100644 tui_gateway/remote_bridge.py diff --git a/cli-config.yaml.example b/cli-config.yaml.example index fb6912642ae9..1e95e0dff5f3 100644 --- a/cli-config.yaml.example +++ b/cli-config.yaml.example @@ -899,6 +899,35 @@ delegation: # Hermes-specific overrides (optional — most config comes from ~/.honcho/config.json): # honcho: {} +# ============================================================================= +# TUI Remote Bridge (backend-only live WebSocket attach) +# ============================================================================= +# Starts an opt-in WebSocket listener inside a running `hermes --tui` gateway. +# Remote clients speak the same JSON-RPC protocol as ui-tui's stdio transport: +# connect, wait for gateway.ready, call session.active_list, then +# session.activate to mirror/control an in-memory live TUI session. +# +# Defaults are loopback-only and disabled. To reach it from a phone over +# Tailscale/LAN, bind to 0.0.0.0 (or a specific host) AND set a token; Hermes +# refuses non-loopback binds without one. Native clients may pass the token as +# `?token=...`, `Authorization: Bearer YOUR_TOKEN`, or `X-Hermes-TUI-Remote-Token`. +# Browser clients should also set trusted_origins for their app origin. +# +# Equivalent environment overrides: +# HERMES_TUI_REMOTE_BRIDGE=1 +# HERMES_TUI_REMOTE_BRIDGE_HOST=0.0.0.0 +# HERMES_TUI_REMOTE_BRIDGE_PORT=8769 +# HERMES_TUI_REMOTE_BRIDGE_TOKEN=... +# HERMES_TUI_REMOTE_BRIDGE_ORIGINS=http://localhost:5174,https://app.example +# +# tui_remote_bridge: +# enabled: false +# host: "127.0.0.1" +# port: 8769 +# path: "/api/tui/ws" +# token: "" +# trusted_origins: [] + # ============================================================================= # Display # ============================================================================= diff --git a/hermes_cli/config.py b/hermes_cli/config.py index bb004d9445ad..aac1029dee59 100644 --- a/hermes_cli/config.py +++ b/hermes_cli/config.py @@ -1198,6 +1198,19 @@ def _ensure_hermes_home_managed(home: Path): "extra_body": {}, }, }, + + # Backend-only WebSocket listener that lets a mobile/web/native client attach + # to the live in-memory sessions owned by a running `hermes --tui` process. + # Off by default. Non-loopback binds (0.0.0.0, LAN/Tailscale hostnames, etc.) + # require a bearer token and still enforce Host/Origin guardrails. + "tui_remote_bridge": { + "enabled": False, + "host": "127.0.0.1", + "port": 8769, + "path": "/api/tui/ws", + "token": "", + "trusted_origins": [], + }, "display": { "compact": False, diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index 4524fb88cb68..3cf483510f6d 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -3980,7 +3980,10 @@ def run_conversation(self, prompt, conversation_history=None, stream_callback=No monkeypatch.setattr(server, "_get_db", lambda: None) monkeypatch.setattr(server, "_session_info", lambda agent: {"model": agent.model}) + original_emit = server._emit + def _emit(event, sid, payload=None): + original_emit(event, sid, payload) if event == "message.complete": done.set() diff --git a/tests/tui_gateway/test_remote_bridge.py b/tests/tui_gateway/test_remote_bridge.py new file mode 100644 index 000000000000..b0be4e020469 --- /dev/null +++ b/tests/tui_gateway/test_remote_bridge.py @@ -0,0 +1,871 @@ +from __future__ import annotations + +import asyncio +import json +import threading +import types + +import pytest + +from tui_gateway import server +from tui_gateway import remote_bridge +from tui_gateway import ws as ws_module + + +class _RecordingTransport: + def __init__(self, *, ok: bool = True) -> None: + self.frames: list[dict] = [] + self.ok = ok + self.closed = False + + def write(self, obj: dict) -> bool: + self.frames.append(obj) + return self.ok + + def close(self) -> None: + self.closed = True + + +class _FakeWS: + def __init__(self, *, headers=None, query_params=None) -> None: + self.headers = headers or {} + self.query_params = query_params or {} + self.close_codes: list[int] = [] + + async def close(self, code: int = 1000) -> None: + self.close_codes.append(code) + + +def _minimal_live_session(transport: _RecordingTransport | None = None, *, session_key: str = "session-key") -> dict: + return { + "agent": None, + "created_at": 123.0, + "history": [], + "history_lock": threading.Lock(), + "last_active": 123.0, + "running": True, + "session_key": session_key, + "transport": transport or _RecordingTransport(), + } + + +def test_session_event_mirrors_to_remote_bridge_without_replacing_primary(): + previous = dict(server._sessions) + primary = _RecordingTransport() + remote = _RecordingTransport() + try: + server._sessions.clear() + server._sessions["sid"] = {"transport": primary} + + assert server.attach_bridge_transport("sid", remote) is True + assert server.write_json( + { + "jsonrpc": "2.0", + "method": "event", + "params": {"type": "message.delta", "session_id": "sid"}, + } + ) is True + + assert len(primary.frames) == 1 + assert len(remote.frames) == 1 + assert server._sessions["sid"]["transport"] is primary + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_attach_bridge_transport_is_copy_on_write(): + """Attaching must rebind a fresh list, never mutate the one in place. + + write_json / _session_event_transports iterate bridge_transports from the + streaming thread; an in-place append would risk 'list changed size during + iteration'. A reader holding an earlier snapshot must see it unchanged. + """ + previous = dict(server._sessions) + primary = _RecordingTransport() + first = _RecordingTransport() + second = _RecordingTransport() + try: + server._sessions.clear() + server._sessions["sid"] = {"transport": primary} + + server.attach_bridge_transport("sid", first) + snapshot = server._sessions["sid"]["bridge_transports"] # a reader's view + assert snapshot == [first] + + server.attach_bridge_transport("sid", second) + assert snapshot == [first] # the snapshot the reader holds is untouched + assert server._sessions["sid"]["bridge_transports"] is not snapshot + assert server._sessions["sid"]["bridge_transports"] == [first, second] + + server.attach_bridge_transport("sid", second) # re-attach is a no-op + assert server._sessions["sid"]["bridge_transports"] == [first, second] + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_failed_remote_bridge_transport_is_pruned_but_primary_still_wins(): + previous = dict(server._sessions) + primary = _RecordingTransport(ok=True) + remote = _RecordingTransport(ok=False) + try: + server._sessions.clear() + server._sessions["sid"] = {"transport": primary} + server.attach_bridge_transport("sid", remote) + + assert server.write_json( + { + "jsonrpc": "2.0", + "method": "event", + "params": {"type": "message.delta", "session_id": "sid"}, + } + ) is True + + assert primary.frames + assert remote.frames + assert server._sessions["sid"].get("bridge_transports") == [] + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_remote_prompt_submitted_event_mirrors_user_turn_to_peers_except_sender(): + previous = dict(server._sessions) + primary = _RecordingTransport() + remote_sender = _RecordingTransport() + remote_peer = _RecordingTransport() + try: + server._sessions.clear() + session = {"transport": primary, "bridge_transports": [remote_sender, remote_peer]} + server._sessions["sid"] = session + token = server.bind_transport(remote_sender) + try: + server._emit_prompt_submitted_to_peers( + "sid", + session, + "hello from phone", + {"client_id": "mobile-1", "source": "mobile"}, + ) + finally: + server.reset_transport(token) + + assert remote_sender.frames == [] + assert primary.frames == remote_peer.frames + frame = primary.frames[0] + assert frame["method"] == "event" + assert frame["params"] == { + "type": "prompt.submitted", + "session_id": "sid", + "payload": {"text": "hello from phone", "client_id": "mobile-1", "source": "mobile"}, + } + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_primary_prompt_submitted_event_mirrors_user_turn_to_remote_clients_only(): + previous = dict(server._sessions) + primary = _RecordingTransport() + remote = _RecordingTransport() + try: + server._sessions.clear() + session = {"transport": primary, "bridge_transports": [remote]} + server._sessions["sid"] = session + token = server.bind_transport(primary) + try: + server._emit_prompt_submitted_to_peers("sid", session, "local tui prompt", {}) + finally: + server.reset_transport(token) + + assert primary.frames == [] + frame = remote.frames[0] + assert frame["method"] == "event" + assert frame["params"] == { + "type": "prompt.submitted", + "session_id": "sid", + "payload": {"text": "local tui prompt"}, + } + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_session_activate_rehydrates_live_display_journal_as_if_mobile_had_been_connected(): + previous = dict(server._sessions) + primary = _RecordingTransport() + attached_mobile = _RecordingTransport() + late_mobile = _RecordingTransport() + try: + server._sessions.clear() + session = _minimal_live_session(primary) + session["bridge_transports"] = [attached_mobile] + session["history"] = [ + {"role": "user", "content": "canonical prompt"}, + {"role": "assistant", "content": "canonical answer"}, + ] + server._sessions["sid"] = session + + server._emit_prompt_submitted_to_peers("sid", session, "local tui prompt", {}) + server._emit( + "tool.start", + "sid", + {"name": "web_search", "context": "web_search(query=remote control)"}, + ) + server._emit( + "tool.complete", + "sid", + {"name": "web_search", "summary": "Did 1 search"}, + ) + server._emit("status.update", "sid", {"kind": "goal", "text": "✓ Goal achieved"}) + server._emit("message.complete", "sid", {"status": "complete", "text": "final answer"}) + + resp = server.dispatch( + {"id": "activate", "method": "session.activate", "params": {"session_id": "sid"}}, + late_mobile, + ) + + assert resp is not None + assert resp["result"]["messages"] == [ + {"role": "user", "text": "local tui prompt"}, + {"role": "tool", "name": "web_search", "text": "web_search(query=remote control)"}, + {"role": "tool", "name": "web_search", "text": "Did 1 search"}, + {"role": "event", "name": "goal", "text": "✓ Goal achieved"}, + {"role": "assistant", "text": "final answer"}, + ] + assert resp["result"]["message_count"] == 5 + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_mobile_prompt_submit_interrupts_running_turn_and_runs_after_current_turn(): + previous = dict(server._sessions) + primary = _RecordingTransport() + mobile = _RecordingTransport() + + class _FakeAgent: + model = "fake/model" + provider = "fake" + base_url = "" + session_id = "session-key" + + def __init__(self) -> None: + self.calls: list[str] = [] + self.interrupts: list[str | None] = [] + self.first_started = threading.Event() + self.second_started = threading.Event() + self.release_first = threading.Event() + self.release_second = threading.Event() + + def interrupt(self, message: str | None = None) -> None: + self.interrupts.append(message) + self.release_first.set() + + def run_conversation(self, user_message, conversation_history=None, stream_callback=None): + self.calls.append(user_message) + if len(self.calls) == 1: + self.first_started.set() + assert self.release_first.wait(1) + return { + "final_response": "interrupted first turn", + "interrupted": True, + "messages": [{"role": "user", "content": "first prompt"}], + } + self.second_started.set() + assert self.release_second.wait(1) + return { + "completed": True, + "final_response": "second turn complete", + "messages": [{"role": "user", "content": user_message}], + } + + agent = _FakeAgent() + session = { + "agent": agent, + "attached_images": [], + "bridge_transports": [mobile], + "history": [], + "history_lock": threading.Lock(), + "history_version": 0, + "inflight_turn": None, + "last_active": 123.0, + "running": False, + "session_key": "session-key", + "transport": primary, + } + + try: + server._sessions.clear() + server._sessions["sid"] = session + + first = server.dispatch( + {"id": "first", "method": "prompt.submit", "params": {"session_id": "sid", "text": "first prompt"}}, + primary, + ) + assert first == {"jsonrpc": "2.0", "id": "first", "result": {"status": "streaming"}} + assert agent.first_started.wait(1) + + second = server.dispatch( + { + "id": "second", + "method": "prompt.submit", + "params": { + "client_id": "mobile-1", + "on_busy": "interrupt", + "session_id": "sid", + "source": "mobile", + "text": "mobile followup", + }, + }, + mobile, + ) + + assert second == {"jsonrpc": "2.0", "id": "second", "result": {"status": "interrupting"}} + assert agent.interrupts == ["mobile followup"] + prompt_frames = [f for f in primary.frames if f.get("params", {}).get("type") == "prompt.submitted"] + assert prompt_frames == [ + { + "jsonrpc": "2.0", + "method": "event", + "params": { + "type": "prompt.submitted", + "session_id": "sid", + "payload": {"text": "mobile followup", "client_id": "mobile-1", "source": "mobile"}, + }, + } + ] + mobile_echoes = [ + f + for f in mobile.frames + if f.get("params", {}).get("type") == "prompt.submitted" + and f.get("params", {}).get("payload", {}).get("text") == "mobile followup" + ] + assert mobile_echoes == [] + assert agent.second_started.wait(1) + assert agent.calls[:2] == ["first prompt", "mobile followup"] + finally: + agent.release_first.set() + agent.release_second.set() + server._sessions.clear() + server._sessions.update(previous) + + +def test_session_activate_attaches_current_transport_to_remote_created_session(): + previous = dict(server._sessions) + mobile_primary = _RecordingTransport() + local_tui = _RecordingTransport() + try: + server._sessions.clear() + session = _minimal_live_session(mobile_primary, session_key="persisted-session") + server._sessions["mobile-live"] = session + + resp = server.dispatch( + {"id": "activate", "method": "session.activate", "params": {"session_id": "mobile-live"}}, + local_tui, + ) + + assert resp is not None + assert resp["result"]["session_id"] == "mobile-live" + assert local_tui in session.get("bridge_transports", []) + + token = server.bind_transport(local_tui) + try: + server._emit_prompt_submitted_to_peers("mobile-live", session, "hello from desktop", {}) + finally: + server.reset_transport(token) + + assert local_tui.frames == [] + frame = mobile_primary.frames[0] + assert frame["params"] == { + "type": "prompt.submitted", + "session_id": "mobile-live", + "payload": {"text": "hello from desktop"}, + } + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_session_resume_title_reuses_existing_live_session_and_attaches_current_transport(monkeypatch): + previous = dict(server._sessions) + mobile_primary = _RecordingTransport() + local_tui = _RecordingTransport() + + class _FakeDB: + def get_session(self, target): + return None + + def get_session_by_title(self, target): + return {"id": "persisted-session"} if target == "mobile-title" else None + + try: + server._sessions.clear() + session = _minimal_live_session(mobile_primary, session_key="persisted-session") + server._sessions["mobile-live"] = session + monkeypatch.setattr(server, "_get_db", lambda: _FakeDB()) + + token = server.bind_transport(local_tui) + try: + resp = server.handle_request( + {"id": "resume", "method": "session.resume", "params": {"session_id": "mobile-title"}} + ) + finally: + server.reset_transport(token) + + assert resp is not None + assert resp["result"]["session_id"] == "mobile-live" + assert resp["result"]["resumed"] == "persisted-session" + assert resp["result"]["live"] is True + assert list(server._sessions) == ["mobile-live"] + assert local_tui in session.get("bridge_transports", []) + + token = server.bind_transport(local_tui) + try: + server._emit_prompt_submitted_to_peers("mobile-live", session, "hello after resume", {}) + finally: + server.reset_transport(token) + + assert local_tui.frames == [] + frame = mobile_primary.frames[0] + assert frame["params"] == { + "type": "prompt.submitted", + "session_id": "mobile-live", + "payload": {"text": "hello after resume"}, + } + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_prompt_resolved_event_mirrors_prompt_answer_to_peers_except_responder(): + previous = dict(server._sessions) + primary = _RecordingTransport() + remote_sender = _RecordingTransport() + remote_peer = _RecordingTransport() + try: + server._sessions.clear() + session = {"transport": primary, "bridge_transports": [remote_sender, remote_peer]} + server._sessions["sid"] = session + token = server.bind_transport(remote_sender) + try: + server._emit_prompt_resolved_to_peers( + "sid", + session, + "clarify", + {"answer": "do not mirror this", "client_id": "mobile-1", "source": "mobile"}, + request_id="rid-1", + ) + finally: + server.reset_transport(token) + + assert remote_sender.frames == [] + assert primary.frames == remote_peer.frames + frame = primary.frames[0] + assert frame["method"] == "event" + assert frame["params"] == { + "type": "prompt.resolved", + "session_id": "sid", + "payload": { + "kind": "clarify", + "resolved": 1, + "request_id": "rid-1", + "client_id": "mobile-1", + "source": "mobile", + }, + } + assert "answer" not in frame["params"]["payload"] + finally: + server._sessions.clear() + server._sessions.update(previous) + + +def test_clarify_respond_mirrors_resolution_to_peer_clients(): + previous_sessions = dict(server._sessions) + previous_pending = dict(server._pending) + previous_answers = dict(server._answers) + primary = _RecordingTransport() + remote_sender = _RecordingTransport() + remote_peer = _RecordingTransport() + ev = threading.Event() + try: + server._sessions.clear() + server._pending.clear() + server._answers.clear() + server._sessions["sid"] = {"transport": primary, "bridge_transports": [remote_sender, remote_peer]} + server._pending["rid-1"] = ("sid", ev) + + resp = server.dispatch( + { + "id": "respond", + "method": "clarify.respond", + "params": { + "answer": "do not mirror this", + "client_id": "mobile-1", + "request_id": "rid-1", + "source": "mobile", + }, + }, + remote_sender, + ) + + assert resp == {"jsonrpc": "2.0", "id": "respond", "result": {"status": "ok"}} + assert ev.is_set() + assert server._answers["rid-1"] == "do not mirror this" + assert remote_sender.frames == [] + assert primary.frames == remote_peer.frames + frame = primary.frames[0] + assert frame["params"]["type"] == "prompt.resolved" + assert frame["params"]["payload"] == { + "kind": "clarify", + "resolved": 1, + "request_id": "rid-1", + "client_id": "mobile-1", + "source": "mobile", + } + assert "answer" not in frame["params"]["payload"] + finally: + server._sessions.clear() + server._sessions.update(previous_sessions) + server._pending.clear() + server._pending.update(previous_pending) + server._answers.clear() + server._answers.update(previous_answers) + + +def test_approval_respond_mirrors_resolution_to_remote_clients(monkeypatch): + from tools import approval + + previous = dict(server._sessions) + primary = _RecordingTransport() + remote = _RecordingTransport() + try: + server._sessions.clear() + server._sessions["sid"] = { + "session_key": "sess-key", + "transport": primary, + "bridge_transports": [remote], + } + monkeypatch.setattr(approval, "resolve_gateway_approval", lambda key, choice, resolve_all=False: 1) + token = server.bind_transport(primary) + try: + resp = server.handle_request( + { + "id": "approve", + "method": "approval.respond", + "params": {"choice": "once", "session_id": "sid", "source": "tui"}, + } + ) + finally: + server.reset_transport(token) + + assert resp == {"jsonrpc": "2.0", "id": "approve", "result": {"resolved": 1}} + assert primary.frames == [] + frame = remote.frames[0] + assert frame["params"] == { + "type": "prompt.resolved", + "session_id": "sid", + "payload": {"kind": "approval", "resolved": 1, "choice": "once", "source": "tui"}, + } + finally: + server._sessions.clear() + server._sessions.update(previous) + + +@pytest.mark.parametrize( + ("event", "payload"), + [ + ("clarify.request", {"choices": ["yes", "no"], "question": "Proceed?", "request_id": "rid-clarify"}), + ("sudo.request", {"request_id": "rid-sudo"}), + ("secret.request", {"env_var": "API_KEY", "prompt": "Enter API key", "request_id": "rid-secret"}), + ], +) +def test_session_activate_includes_pending_blocking_prompt_for_late_attach(event, payload): + previous_sessions = dict(server._sessions) + previous_pending = dict(server._pending) + previous_prompt_payloads = dict(server._pending_prompt_payloads) + try: + server._sessions.clear() + server._pending.clear() + server._pending_prompt_payloads.clear() + server._sessions["sid"] = _minimal_live_session() + request_id = str(payload["request_id"]) + server._pending[request_id] = ("sid", threading.Event()) + server._pending_prompt_payloads[request_id] = (event, dict(payload)) + + resp = server.dispatch( + {"id": "activate", "method": "session.activate", "params": {"session_id": "sid"}}, + _RecordingTransport(), + ) + + assert resp is not None + assert resp["result"]["status"] == "waiting" + assert resp["result"]["pending_prompt"] == {"type": event, "payload": payload} + finally: + server._sessions.clear() + server._sessions.update(previous_sessions) + server._pending.clear() + server._pending.update(previous_pending) + server._pending_prompt_payloads.clear() + server._pending_prompt_payloads.update(previous_prompt_payloads) + + +def test_session_activate_includes_pending_approval_for_late_attach(): + from tools import approval + + previous_sessions = dict(server._sessions) + previous_queues = {key: list(value) for key, value in approval._gateway_queues.items()} + approval_payload = { + "command": "rm -rf /tmp/nope", + "description": "dangerous command", + "pattern_key": "rm-rf", + "pattern_keys": ["rm-rf"], + } + try: + server._sessions.clear() + approval._gateway_queues.clear() + server._sessions["sid"] = _minimal_live_session(session_key="session-key") + approval._gateway_queues["session-key"] = [approval._ApprovalEntry(dict(approval_payload))] + + resp = server.dispatch( + {"id": "activate", "method": "session.activate", "params": {"session_id": "sid"}}, + _RecordingTransport(), + ) + + assert resp is not None + assert resp["result"]["status"] == "waiting" + assert resp["result"]["pending_prompt"] == { + "type": "approval.request", + "payload": approval_payload, + } + finally: + server._sessions.clear() + server._sessions.update(previous_sessions) + approval._gateway_queues.clear() + approval._gateway_queues.update(previous_queues) + + +def test_remote_bridge_config_env_enables_and_overrides(): + cfg = remote_bridge.resolve_remote_bridge_config( + cfg={"tui_remote_bridge": {"enabled": False, "port": 1111}}, + environ={ + "HERMES_TUI_REMOTE_BRIDGE": "1", + "HERMES_TUI_REMOTE_BRIDGE_HOST": "0.0.0.0", + "HERMES_TUI_REMOTE_BRIDGE_PORT": "9999", + "HERMES_TUI_REMOTE_BRIDGE_TOKEN": "secret", + "HERMES_TUI_REMOTE_BRIDGE_ORIGINS": "https://mobile.example,http://localhost:5174/", + }, + ) + + assert cfg.enabled is True + assert cfg.host == "0.0.0.0" + assert cfg.port == 9999 + assert cfg.token == "secret" + assert cfg.trusted_origins == ("https://mobile.example", "http://localhost:5174") + cfg.validate() + + +def test_remote_bridge_refuses_non_loopback_without_token(): + cfg = remote_bridge.RemoteBridgeConfig(enabled=True, host="0.0.0.0", token="") + + with pytest.raises(remote_bridge.RemoteBridgeConfigError, match="token is required"): + cfg.validate() + + +def test_remote_bridge_host_guard_blocks_rebinding_on_loopback(): + assert remote_bridge.is_accepted_host("localhost:8769", "127.0.0.1") + assert remote_bridge.is_accepted_host("[::1]:8769", "::1") + assert not remote_bridge.is_accepted_host("evil.example", "127.0.0.1") + assert not remote_bridge.is_accepted_host("127.0.0.1.evil.example", "127.0.0.1") + + +def test_remote_bridge_origin_guard_accepts_native_or_same_origin_only(): + assert remote_bridge.is_accepted_origin( + "", + bound_host="127.0.0.1", + host_header="localhost:8769", + ) + assert remote_bridge.is_accepted_origin( + "http://localhost:5174", + bound_host="127.0.0.1", + host_header="localhost:8769", + ) + assert remote_bridge.is_accepted_origin( + "https://mobile.example", + bound_host="0.0.0.0", + host_header="tailnet-host:8769", + trusted_origins=("https://mobile.example",), + ) + assert not remote_bridge.is_accepted_origin( + "http://evil.example", + bound_host="127.0.0.1", + host_header="localhost:8769", + ) + + +def test_remote_bridge_authorize_ws_requires_configured_token(): + cfg = remote_bridge.RemoteBridgeConfig( + enabled=True, + host="127.0.0.1", + token="secret", + ) + + ok = _FakeWS( + headers={"host": "localhost:8769", "authorization": "Bearer secret"}, + ) + bad = _FakeWS( + headers={"host": "localhost:8769", "authorization": "Bearer wrong"}, + ) + + assert asyncio.run(remote_bridge.authorize_ws(ok, cfg)) is True + assert ok.close_codes == [] + assert asyncio.run(remote_bridge.authorize_ws(bad, cfg)) is False + assert bad.close_codes == [4401] + + +def test_start_remote_bridge_uses_daemon_thread_without_real_uvicorn(monkeypatch): + class _Config: + def __init__(self, app, host, port, log_level, lifespan): + self.app = app + self.host = host + self.port = port + self.log_level = log_level + self.lifespan = lifespan + + class _Server: + def __init__(self, config): + self.config = config + self.ran = threading.Event() + self.should_exit = False + + def run(self): + self.ran.set() + + fake_uvicorn = types.SimpleNamespace(Config=_Config, Server=_Server) + monkeypatch.setattr(remote_bridge, "_ensure_server_deps", lambda: fake_uvicorn) + monkeypatch.setattr(remote_bridge, "build_app", lambda config: {"path": config.path}) + + handle = remote_bridge.start_remote_bridge( + remote_bridge.RemoteBridgeConfig(enabled=True, host="127.0.0.1", port=9876) + ) + + assert handle is not None + assert handle.thread.daemon is True + assert handle.thread.name == "hermes-tui-remote-bridge" + assert handle.server.ran.wait(1) + assert handle.public_info()["url"] == "ws://127.0.0.1:9876/api/tui/ws" + handle.stop() + assert handle.server.should_exit is True + + +class _ScriptedWS: + """Minimal async WebSocket double: yields scripted requests, then disconnects.""" + + def __init__(self, requests: list[dict]) -> None: + self._queue = [json.dumps(r) for r in requests] + self.sent: list[dict] = [] + + async def accept(self) -> None: + pass + + async def receive_text(self) -> str: + if self._queue: + return self._queue.pop(0) + raise ws_module._WebSocketDisconnect() + + async def send_text(self, line: str) -> None: + self.sent.append(json.loads(line)) + + async def close(self, code: int = 1000) -> None: + pass + + +def test_dashboard_ws_default_does_not_apply_remote_bridge_allowlist(monkeypatch): + """Shared dashboard /api/ws keeps the full authenticated gateway surface.""" + dispatched: list[str] = [] + + def _fake_dispatch(req, transport=None): + dispatched.append(req.get("method")) + return {"jsonrpc": "2.0", "id": req.get("id"), "result": {"ok": True}} + + monkeypatch.setattr(server, "dispatch", _fake_dispatch) + + dashboard = _ScriptedWS( + [{"jsonrpc": "2.0", "id": 1, "method": "slash.exec", "params": {}}] + ) + asyncio.run(ws_module.handle_ws(dashboard)) + + assert dispatched == ["slash.exec"] + assert not [ + m + for m in dashboard.sent + if isinstance(m.get("error"), dict) and m["error"].get("code") == 4403 + ] + + +def test_remote_bridge_ws_passes_bridge_allowlist(monkeypatch): + seen: dict[str, object] = {} + + async def _allow(_ws, _config): + return True + + async def _handle_ws(_ws, *, allowed_methods=None): + seen["allowed_methods"] = allowed_methods + + monkeypatch.setattr(remote_bridge, "authorize_ws", _allow) + monkeypatch.setattr(ws_module, "handle_ws", _handle_ws) + + cfg = remote_bridge.RemoteBridgeConfig( + enabled=True, + host="127.0.0.1", + port=8769, + ) + asyncio.run(remote_bridge.handle_remote_ws(object(), cfg)) + + assert seen["allowed_methods"] is ws_module.BRIDGE_ALLOWED_METHODS + + +def test_remote_bridge_allowlist_blocks_dangerous_methods_but_forwards_allowed(monkeypatch): + """Only allowlisted methods reach dispatch over the remote bridge.""" + assert "prompt.submit" in ws_module.BRIDGE_ALLOWED_METHODS + assert "session.close" in ws_module.BRIDGE_ALLOWED_METHODS + assert "shell.exec" not in ws_module.BRIDGE_ALLOWED_METHODS + assert "config.set" not in ws_module.BRIDGE_ALLOWED_METHODS + assert "slash.exec" not in ws_module.BRIDGE_ALLOWED_METHODS + + dispatched: list[str] = [] + + def _fake_dispatch(req, transport=None): + dispatched.append(req.get("method")) + return {"jsonrpc": "2.0", "id": req.get("id"), "result": {"ok": True}} + + monkeypatch.setattr(server, "dispatch", _fake_dispatch) + + blocked = _ScriptedWS( + [{"jsonrpc": "2.0", "id": 1, "method": "shell.exec", "params": {}}] + ) + asyncio.run( + ws_module.handle_ws(blocked, allowed_methods=ws_module.BRIDGE_ALLOWED_METHODS) + ) + assert "shell.exec" not in dispatched + rejections = [ + m + for m in blocked.sent + if isinstance(m.get("error"), dict) and m["error"].get("code") == 4403 + ] + assert rejections and rejections[0]["id"] == 1 + + allowed = _ScriptedWS( + [ + { + "jsonrpc": "2.0", + "id": 2, + "method": "prompt.submit", + "params": {"session_id": "x", "text": "hi"}, + } + ] + ) + asyncio.run( + ws_module.handle_ws(allowed, allowed_methods=ws_module.BRIDGE_ALLOWED_METHODS) + ) + assert dispatched == ["prompt.submit"] diff --git a/tools/approval.py b/tools/approval.py index 1dbb6eb6e4f2..aeeaf094141b 100644 --- a/tools/approval.py +++ b/tools/approval.py @@ -580,6 +580,12 @@ def resolve_gateway_approval(session_key: str, choice: str, return len(targets) +def pending_gateway_approvals(session_key: str) -> list[dict]: + """Return a snapshot of pending blocking approval payloads for a session.""" + with _lock: + return [dict(entry.data) for entry in _gateway_queues.get(session_key, [])] + + def has_blocking_approval(session_key: str) -> bool: """Check if a session has one or more blocking gateway approvals waiting.""" with _lock: diff --git a/tui_gateway/entry.py b/tui_gateway/entry.py index 7069ec97605f..4cb3ec346d6f 100644 --- a/tui_gateway/entry.py +++ b/tui_gateway/entry.py @@ -19,6 +19,7 @@ from tui_gateway import server from tui_gateway.server import _CRASH_LOG, dispatch, resolve_skin, write_json +from tui_gateway.remote_bridge import start_remote_bridge_if_enabled from tui_gateway.transport import TeeTransport logger = logging.getLogger(__name__) @@ -157,11 +158,11 @@ def _hard_exit() -> None: # ``hermes --tui``) imports cleanly there. SIGBREAK (Windows' Ctrl+Break) # is installed when available as a weaker equivalent of SIGHUP. if hasattr(signal, "SIGPIPE"): - signal.signal(signal.SIGPIPE, signal.SIG_IGN) + signal.signal(signal.SIGPIPE, signal.SIG_IGN) # windows-footgun: ok - guarded by hasattr if hasattr(signal, "SIGTERM"): signal.signal(signal.SIGTERM, _log_signal) if hasattr(signal, "SIGHUP"): - signal.signal(signal.SIGHUP, _log_signal) + signal.signal(signal.SIGHUP, _log_signal) # windows-footgun: ok - guarded by hasattr elif hasattr(signal, "SIGBREAK"): # Windows-only: Ctrl+Break in a console window delivers SIGBREAK. # Route it through the same handler so kills are diagnosable. @@ -212,6 +213,7 @@ def wait_for_mcp_discovery(timeout: float = 0.75) -> None: def main(): _install_sidecar_publisher() + remote_bridge = start_remote_bridge_if_enabled() # MCP tool discovery — runs in a background daemon thread so a slow or # unreachable MCP server can't freeze TUI startup. Previously this ran @@ -263,10 +265,14 @@ def _discover_mcp_background() -> None: global _mcp_discovery_thread _mcp_discovery_thread = _mcp_thread + ready_payload = {"skin": resolve_skin()} + if remote_bridge is not None: + ready_payload["remote_bridge"] = remote_bridge.public_info() + if not write_json({ "jsonrpc": "2.0", "method": "event", - "params": {"type": "gateway.ready", "payload": {"skin": resolve_skin()}}, + "params": {"type": "gateway.ready", "payload": ready_payload}, }): _log_exit("startup write failed (broken stdout pipe before first event)") sys.exit(0) diff --git a/tui_gateway/remote_bridge.py b/tui_gateway/remote_bridge.py new file mode 100644 index 000000000000..df0764d89d50 --- /dev/null +++ b/tui_gateway/remote_bridge.py @@ -0,0 +1,411 @@ +"""Opt-in WebSocket listener for remote control of a live TUI gateway. + +The normal TUI process model is Node/Ink talking to ``tui_gateway.entry`` over +stdio. This module adds a *second* backend-only listener inside that same +Python process so another client (mobile/web/native) can attach to the live +in-memory sessions over the same JSON-RPC protocol used by stdio. + +The listener is deliberately off by default. Binding anywhere other than +loopback requires an explicit bearer token, and WebSocket upgrades enforce Host +and Origin checks to keep browser DNS-rebinding attacks out of localhost-bound +bridges. +""" + +from __future__ import annotations + +import hmac +import importlib +import os +import sys +import threading +from dataclasses import dataclass, field +from typing import Any, Mapping, Sequence +from urllib.parse import urlsplit + +_DEFAULT_HOST = "127.0.0.1" +_DEFAULT_PORT = 8769 +_DEFAULT_PATH = "/api/tui/ws" +_TOKEN_HEADER = "x-hermes-tui-remote-token" + +_LOOPBACK_HOST_VALUES: frozenset[str] = frozenset({"localhost", "127.0.0.1", "::1"}) +_ALL_INTERFACE_HOST_VALUES: frozenset[str] = frozenset({"0.0.0.0", "::"}) + + +class RemoteBridgeConfigError(RuntimeError): + """Configuration refused because it would expose the TUI unsafely.""" + + +@dataclass(frozen=True) +class RemoteBridgeConfig: + enabled: bool = False + host: str = _DEFAULT_HOST + port: int = _DEFAULT_PORT + path: str = _DEFAULT_PATH + token: str = "" + trusted_origins: tuple[str, ...] = field(default_factory=tuple) + + def validate(self) -> None: + if not self.enabled: + return + if not self.host.strip(): + raise RemoteBridgeConfigError("tui_remote_bridge.host must not be empty") + if not (1 <= int(self.port) <= 65535): + raise RemoteBridgeConfigError("tui_remote_bridge.port must be between 1 and 65535") + if not self.path.startswith("/") or "?" in self.path or "#" in self.path: + raise RemoteBridgeConfigError( + "tui_remote_bridge.path must be an absolute path without query or fragment" + ) + if not _is_loopback_bind(self.host) and not self.token: + raise RemoteBridgeConfigError( + "tui_remote_bridge.token is required when host is not loopback" + ) + + def public_info(self) -> dict[str, Any]: + return { + "enabled": self.enabled, + "host": self.host, + "port": self.port, + "path": self.path, + "requires_token": bool(self.token), + "trusted_origins": list(self.trusted_origins), + "url": f"ws://{_url_host(self.host)}:{self.port}{self.path}", + } + + +@dataclass +class RemoteBridgeHandle: + config: RemoteBridgeConfig + server: Any + thread: threading.Thread + + def public_info(self) -> dict[str, Any]: + return self.config.public_info() + + def stop(self) -> None: + try: + self.server.should_exit = True + except Exception: + pass + + +def _url_host(host: str) -> str: + return f"[{host}]" if ":" in host and not host.startswith("[") else host + + +def _truthy(value: Any) -> bool: + return str(value or "").strip().lower() in {"1", "true", "yes", "on"} + + +def _as_int(value: Any, default: int) -> int: + try: + return int(str(value).strip()) + except Exception: + return default + + +def _as_origins(value: Any) -> tuple[str, ...]: + if value is None: + return () + if isinstance(value, str): + parts = value.split(",") + elif isinstance(value, Sequence): + parts = [str(v) for v in value] + else: + parts = [str(value)] + normalized = [] + for part in parts: + origin = _normalize_origin(part) + if origin: + normalized.append(origin) + return tuple(dict.fromkeys(normalized)) + + +def _cfg_bool(node: Mapping[str, Any], key: str, default: bool) -> bool: + if key not in node: + return default + return _truthy(node.get(key)) + + +def resolve_remote_bridge_config( + cfg: Mapping[str, Any] | None = None, + environ: Mapping[str, str] | None = None, +) -> RemoteBridgeConfig: + """Resolve bridge settings from config + environment overrides. + + Config surface:: + + tui_remote_bridge: + enabled: false + host: 127.0.0.1 + port: 8769 + path: /api/tui/ws + token: "" + trusted_origins: [] + + Environment overrides use ``HERMES_TUI_REMOTE_BRIDGE_*``. The legacy-ish + short toggle ``HERMES_TUI_REMOTE=1`` is also accepted for quick manual + testing, but the full name is preferred for scripts. + """ + if cfg is None: + try: + from hermes_cli.config import read_raw_config + + cfg = read_raw_config() or {} + except Exception: + cfg = {} + environ = environ or os.environ + + node = cfg.get("tui_remote_bridge", {}) if isinstance(cfg, Mapping) else {} + if not isinstance(node, Mapping): + node = {} + + enabled = _cfg_bool(node, "enabled", False) + env_enabled = ( + environ.get("HERMES_TUI_REMOTE_BRIDGE") + or environ.get("HERMES_TUI_REMOTE_CONTROL") + or environ.get("HERMES_TUI_REMOTE") + ) + if env_enabled is not None and str(env_enabled).strip(): + enabled = _truthy(env_enabled) + + host = str(environ.get("HERMES_TUI_REMOTE_BRIDGE_HOST") or node.get("host") or _DEFAULT_HOST).strip() + port = _as_int(environ.get("HERMES_TUI_REMOTE_BRIDGE_PORT") or node.get("port"), _DEFAULT_PORT) + path = str(environ.get("HERMES_TUI_REMOTE_BRIDGE_PATH") or node.get("path") or _DEFAULT_PATH).strip() + token = str(environ.get("HERMES_TUI_REMOTE_BRIDGE_TOKEN") or node.get("token") or "").strip() + origins = _as_origins( + environ.get("HERMES_TUI_REMOTE_BRIDGE_ORIGINS") + if environ.get("HERMES_TUI_REMOTE_BRIDGE_ORIGINS") is not None + else node.get("trusted_origins") + ) + + return RemoteBridgeConfig( + enabled=enabled, + host=host, + port=port, + path=path or _DEFAULT_PATH, + token=token, + trusted_origins=origins, + ) + + +def _host_only(host_header: str) -> str: + """Strip optional port/brackets from a Host header-like value.""" + h = (host_header or "").strip() + if not h: + return "" + if h.startswith("["): + close = h.find("]") + return h[1:close].lower() if close != -1 else h.strip("[]").lower() + return (h.rsplit(":", 1)[0] if ":" in h else h).lower() + + +def _is_loopback_bind(host: str) -> bool: + return host.strip().lower() in _LOOPBACK_HOST_VALUES + + +def is_accepted_host(host_header: str, bound_host: str) -> bool: + """Return True when the Host header targets the bound interface. + + This mirrors the dashboard's Host guard: loopback binds only accept loopback + hostnames, explicit non-loopback binds require an exact host match, and + 0.0.0.0/:: all-interface binds accept any Host because the operator has + explicitly opted into network exposure (token still required by validate()). + """ + host = _host_only(host_header) + if not host: + return False + + bound = (bound_host or "").strip().lower() + if bound in _ALL_INTERFACE_HOST_VALUES: + return True + if bound in _LOOPBACK_HOST_VALUES: + return host in _LOOPBACK_HOST_VALUES + return host == bound + + +def _normalize_origin(raw: str) -> str: + raw = (raw or "").strip().rstrip("/") + if not raw: + return "" + try: + parsed = urlsplit(raw) + except Exception: + return "" + if not parsed.scheme or not parsed.netloc: + return "" + return f"{parsed.scheme.lower()}://{parsed.netloc.lower()}" + + +def _origin_host(raw: str) -> str: + try: + return (urlsplit(raw).hostname or "").lower() + except Exception: + return "" + + +def is_accepted_origin( + origin: str, + *, + bound_host: str, + host_header: str, + trusted_origins: Sequence[str] = (), +) -> bool: + """Validate browser Origin for WebSocket upgrades. + + Native/mobile WebSocket stacks commonly omit ``Origin``; those are accepted + after Host/token checks. Browser upgrades with Origin are accepted only + when the origin is explicitly trusted or matches the endpoint host boundary. + """ + if not origin: + return True + + normalized = _normalize_origin(origin) + if not normalized: + return False + trusted = {_normalize_origin(o) for o in trusted_origins if _normalize_origin(o)} + if normalized in trusted: + return True + + origin_host = _origin_host(normalized) + if not origin_host: + return False + + bound = (bound_host or "").strip().lower() + if bound in _LOOPBACK_HOST_VALUES: + return origin_host in _LOOPBACK_HOST_VALUES + if bound in _ALL_INTERFACE_HOST_VALUES: + # For 0.0.0.0/::, the actual request Host is the only useful browser + # boundary we can validate at this layer. + return origin_host == _host_only(host_header) + return origin_host == bound + + +def _mapping_get(mapping: Any, name: str, default: str = "") -> str: + try: + return str(mapping.get(name, default) or "") + except Exception: + return default + + +def _extract_token(ws: Any) -> str: + query = getattr(ws, "query_params", {}) or {} + headers = getattr(ws, "headers", {}) or {} + + query_token = _mapping_get(query, "token") + if query_token: + return query_token + + header_token = _mapping_get(headers, _TOKEN_HEADER) + if header_token: + return header_token + + auth = _mapping_get(headers, "authorization") + if auth.lower().startswith("bearer "): + return auth[7:].strip() + return "" + + +async def authorize_ws(ws: Any, config: RemoteBridgeConfig) -> bool: + """Validate Host/Origin/token before accepting the WebSocket.""" + headers = getattr(ws, "headers", {}) or {} + host_header = _mapping_get(headers, "host") + origin = _mapping_get(headers, "origin") + + close_code = 4403 + ok = is_accepted_host(host_header, config.host) and is_accepted_origin( + origin, + bound_host=config.host, + host_header=host_header, + trusted_origins=config.trusted_origins, + ) + + if ok and config.token: + close_code = 4401 + ok = hmac.compare_digest(_extract_token(ws).encode(), config.token.encode()) + + if not ok: + try: + await ws.close(code=close_code) + except Exception: + pass + return False + return True + + +async def handle_remote_ws(ws: Any, config: RemoteBridgeConfig) -> None: + if not await authorize_ws(ws, config): + return + + from tui_gateway.ws import BRIDGE_ALLOWED_METHODS, handle_ws + + await handle_ws(ws, allowed_methods=BRIDGE_ALLOWED_METHODS) + + +def build_app(config: RemoteBridgeConfig) -> Any: + """Build a tiny FastAPI app exposing only health + the bridge WS.""" + fastapi = importlib.import_module("fastapi") + FastAPI = getattr(fastapi, "FastAPI") + # Resolve the class eagerly so FastAPI can inspect the endpoint signature, + # but keep the import dynamic so base installs without dashboard deps do not + # trigger static import diagnostics. + _websocket_type = getattr(fastapi, "WebSocket") + + app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None) + + @app.get("/healthz") + async def healthz() -> dict[str, Any]: + return {"ok": True, "remote_bridge": config.public_info()} + + async def remote_ws(ws: Any) -> None: + await handle_remote_ws(ws, config) + + remote_ws.__annotations__["ws"] = _websocket_type + app.websocket(config.path)(remote_ws) + + return app + + +def _ensure_server_deps() -> Any: + try: + importlib.import_module("fastapi") + return importlib.import_module("uvicorn") + except ImportError: + from tools.lazy_deps import ensure + + ensure("tool.dashboard") + importlib.import_module("fastapi") + return importlib.import_module("uvicorn") + + +def start_remote_bridge(config: RemoteBridgeConfig | None = None) -> RemoteBridgeHandle | None: + """Start the bridge listener in a daemon thread when enabled.""" + config = config or resolve_remote_bridge_config() + if not config.enabled: + return None + config.validate() + + uvicorn = _ensure_server_deps() + app = build_app(config) + server_config = uvicorn.Config( + app, + host=config.host, + port=int(config.port), + log_level="warning", + lifespan="off", + ) + uvicorn_server = uvicorn.Server(server_config) + thread = threading.Thread( + target=uvicorn_server.run, + name="hermes-tui-remote-bridge", + daemon=True, + ) + thread.start() + return RemoteBridgeHandle(config=config, server=uvicorn_server, thread=thread) + + +def start_remote_bridge_if_enabled() -> RemoteBridgeHandle | None: + try: + return start_remote_bridge() + except Exception as exc: + print(f"[tui-remote-bridge] not started: {exc}", file=sys.stderr, flush=True) + return None diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 4af8e2887e4a..34b3c63dca43 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -362,14 +362,69 @@ def _db_unavailable_error(rid, *, code: int): return _err(rid, code, f"state.db unavailable: {detail}") +def _session_event_transports(session: dict) -> list[Transport]: + """Return the primary transport plus any remote bridge mirrors.""" + transports: list[Transport] = [] + seen: set[int] = set() + + def add(t: Transport | None) -> None: + if t is None: + return + ident = id(t) + if ident in seen: + return + seen.add(ident) + transports.append(t) + + add(session.get("transport")) + for t in session.get("bridge_transports") or (): + add(t) + return transports + + +def _attach_transport_to_session(session: dict, transport: Transport | None = None) -> None: + """Mirror future events for a live session to *transport* without stealing primary ownership.""" + transport = transport or current_transport() or _stdio_transport + if session.get("transport") is transport: + return + bridges = session.get("bridge_transports") or [] + if any(t is transport for t in bridges): + return + # Copy-on-write: never mutate the list in place. write_json / + # _session_event_transports iterate it from the streaming thread, so an + # in-place append would risk "list changed size during iteration". + # detach_bridge_transport rebinds the same way, so readers always iterate a + # stable snapshot. + session["bridge_transports"] = [*bridges, transport] + + +def attach_bridge_transport(session_id: str, transport: Transport) -> bool: + """Mirror future events for *session_id* to a remote WS transport.""" + session = _sessions.get(session_id or "") + if session is None: + return False + _attach_transport_to_session(session, transport) + return True + + +def detach_bridge_transport(transport: Transport) -> None: + """Remove a remote mirror transport from every live session.""" + for session in list(_sessions.values()): + bridges = session.get("bridge_transports") + if not bridges: + continue + session["bridge_transports"] = [t for t in bridges if t is not transport] + + def write_json(obj: dict) -> bool: """Emit one JSON frame. Routes via the most-specific transport available. Precedence: - 1. Event frames with a session id → the transport stored on that session, - so async events land with the client that owns the session even if - the emitting thread has no contextvar binding. + 1. Event frames with a session id → the transport stored on that session + plus any remote bridge mirrors. The primary transport owns the return + value so stdio peer loss still behaves as before; bridge mirrors are + best-effort and never kill the TUI. 2. Otherwise the transport bound on the current context (set by :func:`dispatch` for the lifetime of a request). 3. Otherwise the module-level stdio transport, matching the historical @@ -377,19 +432,195 @@ def write_json(obj: dict) -> bool: """ if obj.get("method") == "event": sid = ((obj.get("params") or {}).get("session_id")) or "" - if sid and (t := (_sessions.get(sid) or {}).get("transport")) is not None: - return t.write(obj) + if sid and (session := _sessions.get(sid)) is not None: + transports = _session_event_transports(session) + if transports: + primary_ok = transports[0].write(obj) + stale: list[Transport] = [] + for t in transports[1:]: + try: + if not t.write(obj): + stale.append(t) + except Exception: + stale.append(t) + for t in stale: + detach_bridge_transport(t) + return primary_ok return (current_transport() or _stdio_transport).write(obj) +_DISPLAY_JOURNAL_LIMIT = 500 + + +def _compact_display_text(value: Any, limit: int = 220) -> str: + text = _content_display_text(value).strip() + if len(text) <= limit: + return text + return f"{text[: max(0, limit - 1)].rstrip()}…" + + +def _display_message_for_event(event: str, payload: dict | None) -> dict | None: + payload = payload or {} + if event == "prompt.submitted": + text = _compact_display_text(payload.get("text"), limit=16_000) + return {"role": "user", "text": text} if text else None + if event == "message.complete": + text = _compact_display_text(payload.get("text"), limit=16_000) + return {"role": "assistant", "text": text} if text else None + if event == "tool.start": + name = str(payload.get("name") or "tool") + text = _compact_display_text( + payload.get("context") or payload.get("args_text") or name or "tool started" + ) + return {"role": "tool", "name": name, "text": text} if text else None + if event == "tool.complete": + name = str(payload.get("name") or "tool") + text = _compact_display_text( + payload.get("summary") + or payload.get("result_text") + or ("error" if payload.get("error") else "complete") + ) + return {"role": "tool", "name": name, "text": text} if text else None + if event == "status.update": + text = _compact_display_text(payload.get("text")) + kind = str(payload.get("kind") or "").strip() + if text and kind and kind != "status": + return {"role": "event", "name": kind, "text": text} + return None + if event == "error": + text = _compact_display_text(payload.get("message") or "unknown error") + return {"role": "system", "name": "error", "text": text} if text else None + return None + + +def _append_display_message(session: dict, message: dict | None) -> None: + if not message: + return + + def _append() -> None: + journal = session.setdefault("display_messages", []) + if not isinstance(journal, list): + journal = [] + session["display_messages"] = journal + journal.append(dict(message)) + if len(journal) > _DISPLAY_JOURNAL_LIMIT: + del journal[: len(journal) - _DISPLAY_JOURNAL_LIMIT] + + lock = session.get("history_lock") + if lock is not None: + with lock: + _append() + else: + _append() + + +def _append_display_event(sid: str, event: str, payload: dict | None = None) -> None: + session = _sessions.get(sid) + if not isinstance(session, dict): + return + _append_display_message(session, _display_message_for_event(event, payload)) + + def _emit(event: str, sid: str, payload: dict | None = None): + _append_display_event(sid, event, payload) params = {"type": event, "session_id": sid} if payload is not None: params["payload"] = payload write_json({"jsonrpc": "2.0", "method": "event", "params": params}) +def _prompt_submitted_payload(text: Any, params: dict) -> dict: + payload = {"text": _content_display_text(text).strip()} + client_id = str(params.get("client_id") or "").strip() + if client_id: + payload["client_id"] = client_id + source = str(params.get("source") or "").strip() + if source: + payload["source"] = source + return payload + + +def _emit_to_peer_transports(sid: str, session: dict, event_type: str, payload: dict) -> None: + """Emit a session event to every attached client except the current caller.""" + origin = current_transport() or _stdio_transport + frame = { + "jsonrpc": "2.0", + "method": "event", + "params": {"type": event_type, "session_id": sid, "payload": payload}, + } + stale: list[Transport] = [] + for transport in _session_event_transports(session): + if transport is origin: + continue + try: + if not transport.write(frame): + stale.append(transport) + except Exception: + stale.append(transport) + for transport in stale: + detach_bridge_transport(transport) + + +def _emit_prompt_submitted_to_peers(sid: str, session: dict, text: Any, params: dict) -> None: + """Mirror a user prompt to every attached client except the submitter. + + Each submitting client optimistically renders its own user prompt before + calling ``prompt.submit``. Peer clients do not see that local echo, so the + gateway emits a small user-turn event to the other transports only: + + - desktop TUI submit → mobile mirrors see the prompt, desktop does not + duplicate it; + - mobile submit → desktop and any other mirrors see the prompt, the sending + phone does not duplicate it. + """ + payload = _prompt_submitted_payload(text, params) + if not payload.get("text"): + return + _append_display_message(session, _display_message_for_event("prompt.submitted", payload)) + _emit_to_peer_transports(sid, session, "prompt.submitted", payload) + + +def _prompt_resolved_payload(kind: str, params: dict, *, request_id: str = "", resolved: int = 1) -> dict: + payload = {"kind": kind, "resolved": int(resolved)} + if request_id: + payload["request_id"] = request_id + choice = str(params.get("choice") or "").strip() + if kind == "approval" and choice: + payload["choice"] = choice + if params.get("all"): + payload["all"] = True + client_id = str(params.get("client_id") or "").strip() + if client_id: + payload["client_id"] = client_id + source = str(params.get("source") or "").strip() + if source: + payload["source"] = source + return payload + + +def _emit_prompt_resolved_to_peers( + sid: str, + session: dict | None, + kind: str, + params: dict, + *, + request_id: str = "", + resolved: int = 1, +) -> None: + """Tell peer clients a blocking prompt was answered elsewhere. + + The submitting client clears its own prompt after the RPC response. Peers + need a small event to clear stale approval/clarify/sudo/secret panels. Do + not include free-form answers/passwords/secrets in this payload; only the + non-sensitive resolution metadata crosses clients. + """ + if not sid or session is None or resolved <= 0: + return + payload = _prompt_resolved_payload(kind, params, request_id=request_id, resolved=resolved) + _emit_to_peer_transports(sid, session, "prompt.resolved", payload) + + def _status_update(sid: str, kind: str, text: str | None = None): body = (text if text is not None else kind).strip() if not body: @@ -755,6 +986,33 @@ def _clear_pending(sid: str | None = None) -> None: ev.set() +def _request_session_interrupt(sid: str, session: dict, message: Any = None) -> None: + """Best-effort interrupt for a single live session. + + ``message`` is passed through when available so agent implementations that + record the follow-up prompt (classic CLI parity) can attach it to the + interrupt result. Prompt/approval waiters are also released with the same + session scoping as the explicit ``session.interrupt`` RPC. + """ + agent = session.get("agent") + interrupt = getattr(agent, "interrupt", None) + if callable(interrupt): + try: + if message is None: + interrupt() + else: + interrupt(message) + except TypeError: + interrupt() + _clear_pending(sid) + try: + from tools.approval import resolve_gateway_approval + + resolve_gateway_approval(session.get("session_key", ""), "deny", resolve_all=True) + except Exception: + pass + + # ── Agent factory ──────────────────────────────────────────────────── @@ -2077,6 +2335,7 @@ def _init_session(sid: str, key: str, agent, history: list, cols: int = 80): "agent": agent, "session_key": key, "history": history, + "display_messages": _history_to_messages(list(history or [])), "history_lock": threading.Lock(), "history_version": 0, "inflight_turn": None, @@ -2285,6 +2544,37 @@ def _clear_inflight_turn(session: dict) -> None: session["inflight_turn"] = None +def _prompt_submit_interrupts_on_busy(params: dict) -> bool: + if params.get("interrupt") is True: + return True + raw = ( + params.get("on_busy") + or params.get("busy_mode") + or params.get("busy_input_mode") + or "" + ) + return str(raw).strip().lower() == "interrupt" + + +def _queue_interrupt_prompt(session: dict, text: Any) -> None: + queued = session.setdefault("interrupt_queue", []) + if not isinstance(queued, list): + queued = [] + session["interrupt_queue"] = queued + queued.append(text) + + +def _take_queued_interrupt_prompt(session: dict) -> Any | None: + queued = session.pop("interrupt_queue", []) + if not isinstance(queued, list) or not queued: + return None + if len(queued) == 1: + return queued[0] + parts = [_content_display_text(item).strip() for item in queued] + combined = "\n".join(part for part in parts if part) + return combined or None + + def _inflight_snapshot(session: dict) -> dict | None: turn = session.get("inflight_turn") if not isinstance(turn, dict): @@ -2321,6 +2611,7 @@ def _(rid, params: dict) -> dict: "attached_images": [], "cols": cols, "created_at": now, + "display_messages": [], "edit_snapshots": {}, "history": [], "history_lock": threading.Lock(), @@ -2462,6 +2753,16 @@ def _(rid, params: dict) -> dict: target = params.get("session_id", "") if not target: return _err(rid, 4006, "session_id required") + + live = _live_session_for_target(str(target)) + if live is not None: + sid, session = live + _attach_transport_to_session(session) + payload = _session_activate_payload(sid, session) + payload["resumed"] = payload.get("session_key") or str(target) + payload["live"] = True + return _ok(rid, payload) + db = _get_db() if db is None: return _db_unavailable_error(rid, code=5000) @@ -2472,6 +2773,18 @@ def _(rid, params: dict) -> dict: target = found["id"] else: return _err(rid, 4007, "session not found") + else: + target = found.get("id") or target + + live = _live_session_for_target(str(target)) + if live is not None: + sid, session = live + _attach_transport_to_session(session) + payload = _session_activate_payload(sid, session) + payload["resumed"] = str(target) + payload["live"] = True + return _ok(rid, payload) + sid = uuid.uuid4().hex[:8] _enable_gateway_prompts() try: @@ -2487,6 +2800,9 @@ def _(rid, params: dict) -> dict: finally: _clear_session_context(tokens) _init_session(sid, target, agent, history, cols=int(params.get("cols", 80))) + live_session = _sessions.get(sid) + if isinstance(live_session, dict): + live_session["display_messages"] = list(messages) except Exception as e: return _err(rid, 5000, f"resume failed: {e}") return _ok( @@ -2501,17 +2817,39 @@ def _(rid, params: dict) -> dict: ) -def _session_pending_kind(sid: str) -> str: +def _pending_prompt_for_session(sid: str, session: dict | None = None) -> dict | None: + """Return the oldest pending prompt payload a late-attaching client can render.""" for rid, (owner_sid, _ev) in list(_pending.items()): if owner_sid != sid: continue - event, _payload = _pending_prompt_payloads.get(rid, ("input.request", {})) - return str(event).removesuffix(".request") - return "" + event, payload = _pending_prompt_payloads.get(rid, ("input.request", {})) + return {"type": str(event), "payload": dict(payload)} + + if session is None: + session = _sessions.get(sid) + session_key = str((session or {}).get("session_key") or sid) + if not session_key: + return None + try: + from tools.approval import pending_gateway_approvals + + approvals = pending_gateway_approvals(session_key) + except Exception: + approvals = [] + if approvals: + return {"type": "approval.request", "payload": dict(approvals[0])} + return None + + +def _session_pending_kind(sid: str, session: dict | None = None) -> str: + pending = _pending_prompt_for_session(sid, session) + if not pending: + return "" + return str(pending.get("type") or "input.request").removesuffix(".request") def _session_live_status(sid: str, session: dict) -> str: - if _session_pending_kind(sid): + if _session_pending_kind(sid, session): return "waiting" ready = session.get("agent_ready") if ready is not None and not ready.is_set(): @@ -2565,6 +2903,23 @@ def _session_live_item(sid: str, session: dict, current_sid: str = "") -> dict: } +def _live_session_for_target(target: str) -> tuple[str, dict] | None: + """Return an in-process live session addressed by live id or persistent key.""" + needle = str(target or "").strip() + if not needle: + return None + try: + items = list(_sessions.items()) + except Exception: + return None + for sid, session in items: + key = str(session.get("session_key") or "").strip() + agent_session_id = str(getattr(session.get("agent"), "session_id", "") or "").strip() + if needle == sid or needle == key or (agent_session_id and needle == agent_session_id): + return sid, session + return None + + def _fallback_session_info(session: dict) -> dict: agent = session.get("agent") if agent is not None: @@ -2578,6 +2933,45 @@ def _fallback_session_info(session: dict) -> dict: } +def _session_activate_payload(sid: str, session: dict) -> dict: + with session["history_lock"]: + session["last_active"] = time.time() + display_messages = list(session.get("display_messages") or []) + history = list(session.get("display_history") or session.get("history") or []) + messages = display_messages if display_messages else _history_to_messages(history) + inflight = _inflight_snapshot(session) + if inflight: + inflight_user = str(inflight.get("user") or "").strip() + if inflight_user: + for index in range(len(messages) - 1, -1, -1): + message = messages[index] + if ( + isinstance(message, dict) + and message.get("role") == "user" + and str(message.get("text") or "").strip() == inflight_user + ): + messages = messages[:index] + messages[index + 1 :] + break + running = bool(session.get("running")) + pending_prompt = _pending_prompt_for_session(sid, session) + status = "waiting" if pending_prompt else _session_live_status(sid, session) + payload = { + "info": _fallback_session_info(session), + "message_count": len(messages), + "messages": messages, + "running": running, + "session_id": sid, + "session_key": session.get("session_key") or sid, + "started_at": float(session.get("created_at") or time.time()), + "status": status, + } + if inflight: + payload["inflight"] = inflight + if pending_prompt: + payload["pending_prompt"] = pending_prompt + return payload + + @method("session.active_list") def _(rid, params: dict) -> dict: """Return live TUI sessions in this gateway process. @@ -2610,25 +3004,10 @@ def _(rid, params: dict) -> dict: session, err = _sess_nowait({"session_id": sid}, rid) if err: return err + assert session is not None - with session["history_lock"]: - session["last_active"] = time.time() - history = list(session.get("display_history") or session.get("history") or []) - inflight = _inflight_snapshot(session) - running = bool(session.get("running")) - status = _session_live_status(sid, session) - payload = { - "info": _fallback_session_info(session), - "message_count": len(history), - "messages": _history_to_messages(history), - "running": running, - "session_id": sid, - "session_key": session.get("session_key") or sid, - "started_at": float(session.get("created_at") or time.time()), - "status": status, - } - if inflight: - payload["inflight"] = inflight + _attach_transport_to_session(session) + payload = _session_activate_payload(sid, session) return _ok( rid, payload, @@ -3070,19 +3449,8 @@ def _(rid, params: dict) -> dict: session, err = _sess(params, rid) if err: return err - if hasattr(session["agent"], "interrupt"): - session["agent"].interrupt() - # Scope the pending-prompt release to THIS session. A global - # _clear_pending() would collaterally cancel clarify/sudo/secret - # prompts on unrelated sessions sharing the same tui_gateway - # process, silently resolving them to empty strings. - _clear_pending(params.get("session_id", "")) - try: - from tools.approval import resolve_gateway_approval - - resolve_gateway_approval(session["session_key"], "deny", resolve_all=True) - except Exception: - pass + assert session is not None + _request_session_interrupt(params.get("session_id", ""), session) return _ok(rid, {"status": "interrupted"}) @@ -3354,13 +3722,26 @@ def _(rid, params: dict) -> dict: session, err = _sess_nowait(params, rid) if err: return err + assert session is not None + interrupting = False with session["history_lock"]: if session.get("running"): - return _err(rid, 4009, "session busy") - session["running"] = True - session["last_active"] = time.time() - _start_inflight_turn(session, text) + if not _prompt_submit_interrupts_on_busy(params): + return _err(rid, 4009, "session busy") + _queue_interrupt_prompt(session, text) + session["last_active"] = time.time() + interrupting = True + else: + session["running"] = True + session["last_active"] = time.time() + _start_inflight_turn(session, text) + if interrupting: + _emit_prompt_submitted_to_peers(sid, session, text, params) + _request_session_interrupt(sid, session, text) + return _ok(rid, {"status": "interrupting"}) + + _emit_prompt_submitted_to_peers(sid, session, text, params) _start_agent_build(sid, session) def run_after_agent_ready() -> None: @@ -3499,6 +3880,7 @@ def run(): approval_token = None session_tokens = [] goal_followup = None # set by the post-turn goal hook below + queued_interrupt_prompt = None try: from tools.approval import ( reset_current_session_key, @@ -3826,6 +4208,20 @@ def _stream(delta): session["running"] = False session["last_active"] = time.time() _clear_inflight_turn(session) + queued_interrupt_prompt = _take_queued_interrupt_prompt(session) + if queued_interrupt_prompt is not None: + session["running"] = True + session["last_active"] = time.time() + _start_inflight_turn(session, queued_interrupt_prompt) + + if queued_interrupt_prompt is not None: + _run_prompt_submit( + f"__interrupt__{int(time.time() * 1000)}", + sid, + session, + queued_interrupt_prompt, + ) + return # Chain a goal-continuation turn if the judge said so. We do # this AFTER the finally releases session["running"], so the @@ -4064,30 +4460,31 @@ def run(): # ── Methods: respond ───────────────────────────────────────────────── -def _respond(rid, params, key): +def _respond(rid, params, key, kind): r = params.get("request_id", "") entry = _pending.get(r) if not entry: return _err(rid, 4009, f"no pending {key} request") - _, ev = entry + sid, ev = entry _answers[r] = params.get(key, "") + _emit_prompt_resolved_to_peers(sid, _sessions.get(sid), kind, params, request_id=str(r)) ev.set() return _ok(rid, {"status": "ok"}) @method("clarify.respond") def _(rid, params: dict) -> dict: - return _respond(rid, params, "answer") + return _respond(rid, params, "answer", "clarify") @method("sudo.respond") def _(rid, params: dict) -> dict: - return _respond(rid, params, "password") + return _respond(rid, params, "password", "sudo") @method("secret.respond") def _(rid, params: dict) -> dict: - return _respond(rid, params, "value") + return _respond(rid, params, "value", "secret") @method("approval.respond") @@ -4095,19 +4492,23 @@ def _(rid, params: dict) -> dict: session, err = _sess(params, rid) if err: return err + assert session is not None try: from tools.approval import resolve_gateway_approval - return _ok( - rid, - { - "resolved": resolve_gateway_approval( - session["session_key"], - params.get("choice", "deny"), - resolve_all=params.get("all", False), - ) - }, + resolved = resolve_gateway_approval( + session["session_key"], + params.get("choice", "deny"), + resolve_all=params.get("all", False), + ) + _emit_prompt_resolved_to_peers( + params.get("session_id", ""), + session, + "approval", + params, + resolved=resolved, ) + return _ok(rid, {"resolved": resolved}) except Exception as e: return _err(rid, 5004, str(e)) diff --git a/tui_gateway/ws.py b/tui_gateway/ws.py index a5879ef3a1c7..8ed0df32c503 100644 --- a/tui_gateway/ws.py +++ b/tui_gateway/ws.py @@ -113,7 +113,51 @@ def close(self) -> None: self._closed = True -async def handle_ws(ws: Any) -> None: +# Methods a remote bridge client may invoke. Dashboard/internal WebSocket +# callers use ``handle_ws(..., allowed_methods=None)`` to preserve the full +# gateway surface; only the opt-in remote bridge passes this allowlist. Anything +# not listed is rejected before dispatch, so a bridge token never reaches +# shell.exec / config mutation / key management / scheduling / slash-command +# exec. Grow this set deliberately as the mobile client gains features. +BRIDGE_ALLOWED_METHODS: frozenset[str] = frozenset( + { + # live session control + "session.active_list", + "session.activate", + "session.create", + "session.interrupt", + "session.close", + "prompt.submit", + "approval.respond", + "clarify.respond", + "sudo.respond", + "secret.respond", + # read-only / informational + "session.status", + "session.usage", + "session.history", + "session.list", + "session.most_recent", + "commands.catalog", + "model.options", + "tools.list", + "tools.show", + "toolsets.list", + "agents.list", + "plugins.list", + "insights.get", + "setup.status", + "rollback.list", + "rollback.diff", + } +) + + +async def handle_ws( + ws: Any, + *, + allowed_methods: frozenset[str] | None = None, +) -> None: """Run one WebSocket session. Wire-compatible with ``tui_gateway.entry``.""" await ws.accept() @@ -155,6 +199,39 @@ async def handle_ws(ws: Any) -> None: break continue + method = req.get("method") if isinstance(req, dict) else None + params = req.get("params") if isinstance(req, dict) else None + + # Remote bridge callers may pass an allowlist. Dashboard/internal + # callers intentionally leave it as None to preserve the full + # authenticated gateway surface. + if ( + allowed_methods is not None + and method is not None + and method not in allowed_methods + ): + ok = await transport.write_async( + { + "jsonrpc": "2.0", + "id": req.get("id") if isinstance(req, dict) else None, + "error": { + "code": 4403, + "message": f"method '{method}' is not permitted over the remote bridge", + }, + } + ) + if not ok: + break + continue + + if method == "session.activate" and isinstance(params, dict): + # A remote client that activates an existing stdio-owned TUI + # session needs to receive future live events without stealing + # them from Ink. Register this WS as a best-effort mirror; + # session-owned writes still go to the original primary + # transport first. + server.attach_bridge_transport(str(params.get("session_id") or ""), transport) + # dispatch() may schedule long handlers on the pool; it returns # None in that case and the worker writes the response itself via # the transport we pass in (a separate thread, so transport.write @@ -168,6 +245,7 @@ async def handle_ws(ws: Any) -> None: # Detach the transport from any sessions it owned so later emits # fall back to stdio instead of crashing into a closed socket. + server.detach_bridge_transport(transport) for _, sess in list(server._sessions.items()): if sess.get("transport") is transport: sess["transport"] = server._stdio_transport diff --git a/ui-tui/src/__tests__/createGatewayEventHandler.test.ts b/ui-tui/src/__tests__/createGatewayEventHandler.test.ts index afebc4d10aca..d2ebc11cad18 100644 --- a/ui-tui/src/__tests__/createGatewayEventHandler.test.ts +++ b/ui-tui/src/__tests__/createGatewayEventHandler.test.ts @@ -132,6 +132,65 @@ describe('createGatewayEventHandler', () => { expect(ctx.system.sys).toHaveBeenCalledWith('compressing 968 messages (~123,400 tok)…') }) + it('renders remote prompt.submitted events as user messages', () => { + const appended: Msg[] = [] + const onEvent = createGatewayEventHandler(buildCtx(appended)) + + onEvent({ payload: { client_id: 'mobile-1', source: 'mobile', text: 'hello from phone' }, type: 'prompt.submitted' } as any) + + expect(appended).toEqual([{ role: 'user', text: 'hello from phone' }]) + }) + + it('archives the current live turn before rendering a remote interrupt prompt', () => { + const appended: Msg[] = [] + const ctx = buildCtx(appended) + const onEvent = createGatewayEventHandler(ctx) + + onEvent({ payload: {}, type: 'message.start' } as any) + onEvent({ payload: { text: 'old thinking' }, type: 'reasoning.delta' } as any) + expect(getTurnState().streamSegments).toEqual([ + expect.objectContaining({ kind: 'trail', role: 'system', thinking: 'old thinking' }) + ]) + + onEvent({ payload: { client_id: 'mobile-1', source: 'mobile', text: 'wait' }, type: 'prompt.submitted' } as any) + + expect(ctx.gateway.gw.request).not.toHaveBeenCalledWith('session.interrupt', expect.anything()) + expect(appended).toEqual([ + expect.objectContaining({ kind: 'trail', role: 'system', thinking: 'old thinking' }), + { role: 'user', text: 'wait' } + ]) + expect(getUiState().busy).toBe(false) + expect(turnController.reasoningText).toBe('') + expect(getTurnState().streamSegments).toEqual([]) + + onEvent({ payload: { text: 'late old chunk' }, type: 'reasoning.delta' } as any) + expect(turnController.reasoningText).toBe('') + + onEvent({ payload: {}, type: 'message.start' } as any) + onEvent({ payload: { text: 'fresh thinking' }, type: 'reasoning.delta' } as any) + expect(turnController.reasoningText).toBe('fresh thinking') + }) + + it('clears pending prompt overlays when a peer resolves them', () => { + const onEvent = createGatewayEventHandler(buildCtx([])) + + onEvent({ payload: { choices: ['yes'], question: 'Proceed?', request_id: 'rid-1' }, type: 'clarify.request' } as any) + expect(getOverlayState().clarify).toMatchObject({ requestId: 'rid-1' }) + + onEvent({ payload: { kind: 'clarify', request_id: 'other', source: 'mobile' }, type: 'prompt.resolved' } as any) + expect(getOverlayState().clarify).toMatchObject({ requestId: 'rid-1' }) + + onEvent({ payload: { kind: 'clarify', request_id: 'rid-1', source: 'mobile' }, type: 'prompt.resolved' } as any) + expect(getOverlayState().clarify).toBeNull() + expect(getUiState().status).toBe('clarify answered on mobile') + + onEvent({ payload: { command: 'rm -rf /tmp/nope', description: 'dangerous command' }, type: 'approval.request' } as any) + expect(getOverlayState().approval).toMatchObject({ description: 'dangerous command' }) + + onEvent({ payload: { choice: 'deny', kind: 'approval', source: 'mobile' }, type: 'prompt.resolved' } as any) + expect(getOverlayState().approval).toBeNull() + }) + it('keeps goal verdict text in transcript but shows a brief idle status (#goal statusbar)', () => { const appended: Msg[] = [] const ctx = buildCtx(appended) diff --git a/ui-tui/src/app/createGatewayEventHandler.ts b/ui-tui/src/app/createGatewayEventHandler.ts index 987518a4460c..4ab01ff24777 100644 --- a/ui-tui/src/app/createGatewayEventHandler.ts +++ b/ui-tui/src/app/createGatewayEventHandler.ts @@ -401,6 +401,20 @@ export function createGatewayEventHandler(ctx: GatewayEventHandlerContext): (ev: return } + case 'prompt.submitted': { + const text = String(ev.payload?.text ?? '').trim() + + if (text) { + if (getUiState().busy) { + turnController.archiveInterruptedTurn({ appendMessage, sys }) + } + + appendMessage({ role: 'user', text }) + } + + return + } + case 'thinking.delta': { if (!getUiState().busy) { return @@ -683,6 +697,40 @@ export function createGatewayEventHandler(ctx: GatewayEventHandlerContext): (ev: return + case 'prompt.resolved': { + const kind = String(ev.payload.kind ?? '') + const requestId = String(ev.payload.request_id ?? '') + let cleared = false + + patchOverlayState(state => { + if (kind === 'approval' && state.approval) { + cleared = true + return { ...state, approval: null } + } + if (kind === 'clarify' && state.clarify && (!requestId || state.clarify.requestId === requestId)) { + cleared = true + return { ...state, clarify: null } + } + if (kind === 'sudo' && state.sudo && (!requestId || state.sudo.requestId === requestId)) { + cleared = true + return { ...state, sudo: null } + } + if (kind === 'secret' && state.secret && (!requestId || state.secret.requestId === requestId)) { + cleared = true + return { ...state, secret: null } + } + return state + }) + + if (cleared) { + const source = String(ev.payload.source ?? '').trim() + const where = source ? ` on ${source}` : ' elsewhere' + setStatus(`${kind || 'prompt'} answered${where}`) + } + + return + } + case 'background.complete': dropBgTask(ev.payload.task_id) sys(`[bg ${ev.payload.task_id}] ${ev.payload.text}`) diff --git a/ui-tui/src/app/turnController.ts b/ui-tui/src/app/turnController.ts index 5f11145b0102..3d6f96abfc07 100644 --- a/ui-tui/src/app/turnController.ts +++ b/ui-tui/src/app/turnController.ts @@ -182,10 +182,8 @@ class TurnController { resetFlowOverlays() } - interruptTurn({ appendMessage, gw, sid, sys }: InterruptDeps) { + private archiveInterruptedState({ appendMessage, sys }: Pick) { this.interrupted = true - gw.request('session.interrupt', { session_id: sid }).catch(() => {}) - this.closeReasoningSegment() const segments = this.segmentMessages @@ -217,6 +215,15 @@ class TurnController { } else { sys('interrupted') } + } + + archiveInterruptedTurn(deps: Pick) { + this.archiveInterruptedState(deps) + } + + interruptTurn({ appendMessage, gw, sid, sys }: InterruptDeps) { + gw.request('session.interrupt', { session_id: sid }).catch(() => {}) + this.archiveInterruptedState({ appendMessage, sys }) patchUiState({ status: 'interrupted' }) this.clearStatusTimer() diff --git a/ui-tui/src/app/useInputHandlers.ts b/ui-tui/src/app/useInputHandlers.ts index 2cbb745b8fe2..4782430d200c 100644 --- a/ui-tui/src/app/useInputHandlers.ts +++ b/ui-tui/src/app/useInputHandlers.ts @@ -127,19 +127,19 @@ export function useInputHandlers(ctx: InputHandlerContext): InputHandlerResult { if (overlay.approval) { return gateway - .rpc('approval.respond', { choice: 'deny', session_id: getUiState().sid }) + .rpc('approval.respond', { choice: 'deny', session_id: getUiState().sid, source: 'tui' }) .then(r => r && (patchOverlayState({ approval: null }), patchTurnState({ outcome: 'denied' }))) } if (overlay.sudo) { return gateway - .rpc('sudo.respond', { password: '', request_id: overlay.sudo.requestId }) + .rpc('sudo.respond', { password: '', request_id: overlay.sudo.requestId, source: 'tui' }) .then(r => r && (patchOverlayState({ sudo: null }), actions.sys('sudo cancelled'))) } if (overlay.secret) { return gateway - .rpc('secret.respond', { request_id: overlay.secret.requestId, value: '' }) + .rpc('secret.respond', { request_id: overlay.secret.requestId, source: 'tui', value: '' }) .then(r => r && (patchOverlayState({ secret: null }), actions.sys('secret entry cancelled'))) } diff --git a/ui-tui/src/app/useMainApp.ts b/ui-tui/src/app/useMainApp.ts index 43e8a2ed628c..a5fc9a5ae109 100644 --- a/ui-tui/src/app/useMainApp.ts +++ b/ui-tui/src/app/useMainApp.ts @@ -580,7 +580,7 @@ export function useMainApp(gw: GatewayClient) { turnController.turnTools = turnController.turnTools.filter(line => !sameToolTrailGroup(label, line)) patchTurnState({ turnTrail: turnController.turnTools }) - rpc('clarify.respond', { answer, request_id: clarify.requestId }).then(r => { + rpc('clarify.respond', { answer, request_id: clarify.requestId, source: 'tui' }).then(r => { if (!r) { return } @@ -844,7 +844,7 @@ export function useMainApp(gw: GatewayClient) { const answerApproval = useCallback( (choice: string) => - respondWith('approval.respond', { choice, session_id: ui.sid }, () => { + respondWith('approval.respond', { choice, session_id: ui.sid, source: 'tui' }, () => { patchOverlayState({ approval: null }) patchTurnState({ outcome: choice === 'deny' ? 'denied' : `approved (${choice})` }) patchUiState({ status: 'running…' }) @@ -858,7 +858,7 @@ export function useMainApp(gw: GatewayClient) { return } - return respondWith('sudo.respond', { password: pw, request_id: overlay.sudo.requestId }, () => { + return respondWith('sudo.respond', { password: pw, request_id: overlay.sudo.requestId, source: 'tui' }, () => { patchOverlayState({ sudo: null }) patchUiState({ status: 'running…' }) }) @@ -872,7 +872,7 @@ export function useMainApp(gw: GatewayClient) { return } - return respondWith('secret.respond', { request_id: overlay.secret.requestId, value }, () => { + return respondWith('secret.respond', { request_id: overlay.secret.requestId, source: 'tui', value }, () => { patchOverlayState({ secret: null }) patchUiState({ status: 'running…' }) }) diff --git a/ui-tui/src/gatewayTypes.ts b/ui-tui/src/gatewayTypes.ts index 447dec3ea492..cbfdbbc48d8a 100644 --- a/ui-tui/src/gatewayTypes.ts +++ b/ui-tui/src/gatewayTypes.ts @@ -153,11 +153,17 @@ export interface SessionInflightTurn { user?: string } +export interface PendingPromptSnapshot { + payload?: Record + type: 'approval.request' | 'clarify.request' | 'sudo.request' | 'secret.request' | string +} + export interface SessionActivateResponse { inflight?: null | SessionInflightTurn info?: SessionInfo message_count?: number messages: GatewayTranscriptMessage[] + pending_prompt?: null | PendingPromptSnapshot running?: boolean session_id: string session_key?: string @@ -505,6 +511,7 @@ export type GatewayEvent = | { payload?: GatewaySkin; session_id?: string; type: 'skin.changed' } | { payload: SessionInfo; session_id?: string; type: 'session.info' } | { payload?: { text?: string }; session_id?: string; type: 'thinking.delta' } + | { payload?: { client_id?: string; source?: string; text?: string }; session_id?: string; type: 'prompt.submitted' } | { payload?: undefined; session_id?: string; type: 'message.start' } | { payload?: { kind?: string; text?: string }; session_id?: string; type: 'status.update' } | { payload?: { state?: 'idle' | 'listening' | 'transcribing' }; session_id?: string; type: 'voice.status' } @@ -551,6 +558,19 @@ export type GatewayEvent = | { payload: { command: string; description: string }; session_id?: string; type: 'approval.request' } | { payload: { request_id: string }; session_id?: string; type: 'sudo.request' } | { payload: { env_var: string; prompt: string; request_id: string }; session_id?: string; type: 'secret.request' } + | { + payload: { + all?: boolean + choice?: string + client_id?: string + kind: 'approval' | 'clarify' | 'sudo' | 'secret' | string + request_id?: string + resolved?: number + source?: string + } + session_id?: string + type: 'prompt.resolved' + } | { payload: { task_id: string; text: string }; session_id?: string; type: 'background.complete' } | { payload?: { text?: string }; session_id?: string; type: 'review.summary' } | { payload: SubagentEventPayload; session_id?: string; type: 'subagent.spawn_requested' } From a61e549644adf5f37a7081e937872320eae0354c Mon Sep 17 00:00:00 2001 From: lsaether <25539605+lsaether@users.noreply.github.com> Date: Sun, 31 May 2026 13:14:43 -0500 Subject: [PATCH 2/2] feat(tui): add remote control launch flag --- cli-config.yaml.example | 3 + hermes_cli/_parser.py | 22 +++++++ hermes_cli/main.py | 8 +++ .../test_argparse_flag_propagation.py | 29 +++++++++ tests/hermes_cli/test_tui_resume_flow.py | 65 +++++++++++++++++++ 5 files changed, 127 insertions(+) diff --git a/cli-config.yaml.example b/cli-config.yaml.example index 1e95e0dff5f3..da42ad41df6b 100644 --- a/cli-config.yaml.example +++ b/cli-config.yaml.example @@ -920,6 +920,9 @@ delegation: # HERMES_TUI_REMOTE_BRIDGE_TOKEN=... # HERMES_TUI_REMOTE_BRIDGE_ORIGINS=http://localhost:5174,https://app.example # +# One-shot loopback-only convenience: +# hermes --tui --remote-control # alias: --rc +# # tui_remote_bridge: # enabled: false # host: "127.0.0.1" diff --git a/hermes_cli/_parser.py b/hermes_cli/_parser.py index cf4ffc34e5c1..866eaebbbc02 100644 --- a/hermes_cli/_parser.py +++ b/hermes_cli/_parser.py @@ -218,6 +218,17 @@ def build_top_level_parser(): default=False, help="Launch the modern TUI instead of the classic REPL", ) + _inherited_flag( + parser, + "--remote-control", + "--rc", + action="store_true", + default=False, + help=( + "With --tui: start the opt-in loopback Remote Control bridge " + "for live clients" + ), + ) _inherited_flag( parser, "--dev", @@ -369,6 +380,17 @@ def build_top_level_parser(): default=False, help="Launch the modern TUI instead of the classic REPL", ) + _inherited_flag( + chat_parser, + "--remote-control", + "--rc", + action="store_true", + default=argparse.SUPPRESS, + help=( + "With --tui: start the opt-in loopback Remote Control bridge " + "for live clients" + ), + ) _inherited_flag( chat_parser, "--dev", diff --git a/hermes_cli/main.py b/hermes_cli/main.py index 1941dc2af313..ca23fad650e1 100644 --- a/hermes_cli/main.py +++ b/hermes_cli/main.py @@ -1543,6 +1543,7 @@ def _launch_tui( pass_session_id: bool = False, max_turns: Optional[int] = None, accept_hooks: bool = False, + remote_control: bool = False, ): """Replace current process with the TUI.""" tui_dir = PROJECT_ROOT / "ui-tui" @@ -1622,6 +1623,12 @@ def _launch_tui( env["HERMES_TUI_TOOL_PROGRESS"] = "off" if accept_hooks: env["HERMES_ACCEPT_HOOKS"] = "1" + if remote_control: + # Friendly CLI switch for the backend TUI Remote Control bridge. The + # bridge itself keeps the secure defaults: loopback bind, default port, + # and a token requirement for any non-loopback host configured by env or + # config. + env["HERMES_TUI_REMOTE_BRIDGE"] = "1" # Guarantee an 8GB V8 heap for the TUI. Default node cap is ~1.5–4GB # depending on version and can fatal-OOM on long sessions with large # transcripts / reasoning blobs. Token-level merge: respect any @@ -1852,6 +1859,7 @@ def cmd_chat(args): pass_session_id=getattr(args, "pass_session_id", False), max_turns=getattr(args, "max_turns", None), accept_hooks=getattr(args, "accept_hooks", False), + remote_control=getattr(args, "remote_control", False), ) # Import and run the CLI diff --git a/tests/hermes_cli/test_argparse_flag_propagation.py b/tests/hermes_cli/test_argparse_flag_propagation.py index 87db493850c0..1e545dbf0f66 100644 --- a/tests/hermes_cli/test_argparse_flag_propagation.py +++ b/tests/hermes_cli/test_argparse_flag_propagation.py @@ -109,6 +109,35 @@ def fake_main(**kwargs): assert "verbose" not in captured +class TestRemoteControlArg: + """Verify --remote-control/--rc parse on both top-level and chat forms.""" + + @pytest.mark.parametrize( + "argv", + [ + ["--tui", "--remote-control"], + ["--tui", "--rc"], + ["chat", "--tui", "--remote-control"], + ["chat", "--tui", "--rc"], + ], + ) + def test_remote_control_flag_sets_attribute(self, argv): + from hermes_cli._parser import build_top_level_parser + + parser, _subparsers, _chat_parser = build_top_level_parser() + args = parser.parse_args(argv) + + assert args.remote_control is True + + def test_chat_without_remote_control_preserves_parent_default(self): + from hermes_cli._parser import build_top_level_parser + + parser, _subparsers, _chat_parser = build_top_level_parser() + args = parser.parse_args(["--tui", "chat"]) + + assert args.remote_control is False + + class TestYoloEnvVar: """Verify --yolo sets HERMES_YOLO_MODE regardless of flag position. diff --git a/tests/hermes_cli/test_tui_resume_flow.py b/tests/hermes_cli/test_tui_resume_flow.py index d15d67c00718..1bbba0eadfce 100644 --- a/tests/hermes_cli/test_tui_resume_flow.py +++ b/tests/hermes_cli/test_tui_resume_flow.py @@ -16,6 +16,7 @@ def _args(**overrides): "toolsets": None, "tui": True, "tui_dev": False, + "remote_control": False, } base.update(overrides) return Namespace(**base) @@ -201,6 +202,7 @@ def fake_launch(resume_session_id=None, **kwargs): pass_session_id=True, max_turns=7, accept_hooks=True, + remote_control=True, ) ) @@ -214,6 +216,7 @@ def fake_launch(resume_session_id=None, **kwargs): assert captured["pass_session_id"] is True assert captured["max_turns"] == 7 assert captured["accept_hooks"] is True + assert captured["remote_control"] is True def test_main_top_level_tui_accepts_toolsets(monkeypatch, main_mod): @@ -896,6 +899,68 @@ def fake_call(argv, cwd=None, env=None): assert env["NODE_ENV"] == "production" +def test_launch_tui_remote_control_sets_bridge_env(monkeypatch, main_mod): + captured = {} + + monkeypatch.setattr( + main_mod, + "_make_tui_argv", + lambda tui_dir, tui_dev: (["node", "dist/entry.js"], Path(".")), + ) + monkeypatch.setattr( + main_mod.subprocess, + "call", + lambda argv, cwd=None, env=None: captured.update({"env": env}) or 1, + ) + + with pytest.raises(SystemExit): + main_mod._launch_tui(remote_control=True) + + assert captured["env"]["HERMES_TUI_REMOTE_BRIDGE"] == "1" + + +def test_launch_tui_without_remote_control_leaves_bridge_env_unset(monkeypatch, main_mod): + captured = {} + + monkeypatch.delenv("HERMES_TUI_REMOTE_BRIDGE", raising=False) + monkeypatch.setattr( + main_mod, + "_make_tui_argv", + lambda tui_dir, tui_dev: (["node", "dist/entry.js"], Path(".")), + ) + monkeypatch.setattr( + main_mod.subprocess, + "call", + lambda argv, cwd=None, env=None: captured.update({"env": env}) or 1, + ) + + with pytest.raises(SystemExit): + main_mod._launch_tui() + + assert "HERMES_TUI_REMOTE_BRIDGE" not in captured["env"] + + +def test_launch_tui_remote_control_overrides_disabled_bridge_env(monkeypatch, main_mod): + captured = {} + + monkeypatch.setenv("HERMES_TUI_REMOTE_BRIDGE", "0") + monkeypatch.setattr( + main_mod, + "_make_tui_argv", + lambda tui_dir, tui_dev: (["node", "dist/entry.js"], Path(".")), + ) + monkeypatch.setattr( + main_mod.subprocess, + "call", + lambda argv, cwd=None, env=None: captured.update({"env": env}) or 1, + ) + + with pytest.raises(SystemExit): + main_mod._launch_tui(remote_control=True) + + assert captured["env"]["HERMES_TUI_REMOTE_BRIDGE"] == "1" + + def test_launch_tui_exit_code_42_relaunches_update(monkeypatch, main_mod): from unittest.mock import patch