diff --git a/tui_gateway/compute_host.py b/tui_gateway/compute_host.py index 8778acf48e..52b5370339 100644 --- a/tui_gateway/compute_host.py +++ b/tui_gateway/compute_host.py @@ -40,21 +40,18 @@ class _HostTransport: return None -# Slice of ``ComputeHost.shutdown``'s budget held back for the post-drain finalize. -# ``HostSupervisor._terminate_pid`` SIGKILLs the host ``_SHUTDOWN_TIMEOUT_SECS`` (10s, -# same as ``shutdown``'s default ``wait``) after SIGTERM, so a drain allowed to consume -# the whole budget would leave the flush racing that kill and persist nothing at all. +# Slice of ``ComputeHost.shutdown``'s budget held back for the post-drain finalize: the +# supervisor SIGKILLs the host 10s (= default ``wait``) after SIGTERM, so a drain that ate +# the whole budget would leave the flush racing that kill and persist nothing. _FLUSH_RESERVE_SECS = 1.0 - # Fallback control.error text when a routed server method returns an error without a message. _CONTROL_FAILURES = { "session.save": "session save failed", "session.compress": "session compression failed"} class ComputeHost: - # frame ``type`` -> handler method name (resolved per call so instance - # monkeypatches of a handler still take effect). + # frame ``type`` -> handler method name (resolved per call so monkeypatches take effect). _FRAME_HANDLERS: dict[str, str] = { "turn.start": "_handle_turn_start", "interrupt": "_handle_interrupt", "respond": "_handle_respond", "reload_mcp": "_handle_reload_mcp", @@ -72,8 +69,7 @@ class ComputeHost: self._boot_id = uuid.uuid4().hex self._progress_counter = 0 self._progress_lock = threading.Lock() - # Future -> the ``sid`` whose turn it is running. ``shutdown`` needs to know - # *whose* turn is still live so it can leave those sessions unfinalized. + # Future -> the ``sid`` whose turn it runs; ``shutdown`` leaves live sids unfinalized. self._turn_futures: dict[concurrent.futures.Future, str] = {} self._turn_futures_lock = threading.Lock() self._transport = _HostTransport(self.emit) @@ -103,15 +99,11 @@ class ComputeHost: def shutdown(self, *, reason: str = "shutdown", wait: float = 10.0) -> None: """Drain in-flight turns, then finalize every session. - Order matters: ``_finalize_session`` is a one-shot latch, so finalizing before the - drain would spend it mid-turn, fire ``on_session_end(interrupted=True)`` on a running - session and release its lease. ``_FLUSH_RESERVE_SECS`` (at most half of ``wait``) is - withheld from the drain so the flush still runs when turns outlast the window. - Sessions still running at the deadline are skipped (``_executor.shutdown`` does not - join them): finalizing mid-turn would leave them un-finalizable with the lease - released; unfinalized keeps them recoverable. ``server._shutdown_sessions`` (atexit) - may re-finalize skipped sessions on SIGTERM / stdin_closed; ``os._exit`` (orphan) - bypasses atexit. + ``_finalize_session`` is a one-shot latch, so finalizing before the drain would spend + it mid-turn and release the lease. ``_FLUSH_RESERVE_SECS`` (at most half of ``wait``) + is withheld from the drain so the flush still runs when turns outlast the window. + Sessions still running at the deadline are skipped (unfinalized keeps them + recoverable; atexit ``server._shutdown_sessions`` may re-finalize them). """ self._closed.set() budget = max(0.0, wait) @@ -120,8 +112,7 @@ class ComputeHost: remaining = deadline - time.monotonic() if remaining <= 0 or not self._live_turns(): break - # Bounded by ``remaining``: a flat sleep would overshoot the deadline and - # eat the reserve it protects (all of it for small ``wait``). + # Bounded by ``remaining``: a flat sleep would eat the reserve it protects. time.sleep(min(0.05, remaining)) with self._turn_futures_lock: live_sids = {sid for f, sid in self._turn_futures.items() if sid and not f.done()} @@ -149,13 +140,12 @@ class ComputeHost: self.emit({ "type": "error", "request_id": frame.get("request_id"), "message": f"unknown frame type: {kind}"}) - return - getattr(self, handler)(frame) + else: + getattr(self, handler)(frame) def _handle_shutdown(self, frame: dict[str, Any]) -> None: self.emit({"type": "shutdown.ack", "request_id": frame.get("request_id")}) - # Explicit supervisor/test shutdown is a clean child-process close; - # SIGTERM and orphan paths are the durability flush paths. + # Explicit shutdown is a clean close; SIGTERM and orphan paths do the durability flush. self.close() def _track_turn_future(self, future: concurrent.futures.Future, sid: str) -> None: @@ -172,41 +162,44 @@ class ComputeHost: future = self._executor.submit(self._run_real_turn, dict(frame)) self._track_turn_future(future, str(frame.get("sid") or "")) - def _handle_interrupt(self, frame: dict[str, Any]) -> None: + def _guarded( + self, frame: dict[str, Any], error_kind: str, body: Callable, *, + on_error: Callable[[str], None] | None = None, **error_extra: Any) -> None: + """Run ``body(server, sid, request_id)``; any exception becomes an ``error_kind`` reply.""" sid = str(frame.get("sid") or "") request_id = frame.get("request_id") try: from tui_gateway import server + body(server, sid, request_id) + except Exception as exc: + if on_error is not None: + on_error(sid) + self._reply(error_kind, sid, request_id, **error_extra, message=str(exc)) + + def _handle_interrupt(self, frame: dict[str, Any]) -> None: + def body(server: Any, sid: str, request_id: Any) -> None: session = server._sessions.get(sid) if session is None: self._reply("interrupt.ack", sid, request_id, applied=False) return - # In the child, `_session_uses_compute_host()` is false, so the shared helper - # interrupts the local agent and releases this process's pending clarify - # Event; the parent only has a metadata mirror and cannot. + # In the child the shared helper interrupts the local agent and releases this + # process's pending clarify Event (the parent only has a metadata mirror). server._interrupt_session_turn(sid, session) self._reply("interrupt.ack", sid, request_id, applied=True, applied_ns=now_ns()) - except Exception as exc: - self._reply("interrupt.ack", sid, request_id, applied=False, message=str(exc)) + self._guarded(frame, "interrupt.ack", body, applied=False) def _handle_respond(self, frame: dict[str, Any]) -> None: """Resolve an interactive request in the host-owned pending registry.""" - sid = str(frame.get("sid") or "") - request_id = frame.get("request_id") - try: - from tui_gateway import server - if sid not in server._sessions: - self._reply("respond.error", sid, request_id, message="session not found") - return + def body(server: Any, sid: str, request_id: Any) -> None: params = frame.get("params") - if not isinstance(params, dict): - self._reply( - "respond.error", sid, request_id, message="response params must be an object") + error = ("session not found" if sid not in server._sessions + else None if isinstance(params, dict) else "response params must be an object") + if error: + self._reply("respond.error", sid, request_id, message=error) return response = server._methods["clarify.respond"](request_id, params) self._reply("respond.ack", sid, request_id, response=response) - except Exception as exc: - self._reply("respond.error", sid, request_id, message=str(exc)) + self._guarded(frame, "respond.error", body) def _run_real_turn(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") @@ -217,8 +210,8 @@ class ComputeHost: try: from tui_gateway import server session = self._ensure_server_session(server, frame) - text = frame.get("text") if "text" in frame else frame.get("prompt", "") - inflight = frame.get("text") if "text" in frame else frame.get("prompt") + text = frame["text"] if "text" in frame else frame.get("prompt", "") + inflight = frame["text"] if "text" in frame else frame.get("prompt") with session["history_lock"]: queued_gen = frame.get("queued_prompt_generation") current_gen = int(session.get("_queued_prompt_generation", 0)) @@ -247,7 +240,8 @@ class ComputeHost: meta = _history_meta(session) interrupted = bool(session.get("_turn_cancel_requested")) session_info = server._session_info(session.get("agent"), session) - self._bump_progress() + with self._progress_lock: + self._progress_counter += 1 self._reply( "turn.end", sid, request_id, **meta, interrupted=interrupted, ended_ns=now_ns(), session_info=session_info, session_info_emitted=True) @@ -263,25 +257,27 @@ class ComputeHost: def _ensure_server_session(self, server: Any, frame: dict[str, Any]) -> dict: sid = str(frame.get("sid") or "") - key = str(frame.get("session_key") or sid) session = server._sessions.get(sid) if session is not None: session["transport"] = self._transport if frame.get("cols") is not None: session["cols"] = int(frame.get("cols") or 80) - if frame.get("cwd"): - session["cwd"] = str(frame.get("cwd")) - if frame.get("profile_home"): - session["profile_home"] = str(frame.get("profile_home")) - if isinstance(frame.get("attached_images"), list): - session["attached_images"] = list(frame.get("attached_images") or []) - return session + for key in ("cwd", "profile_home"): + if frame.get(key): + session[key] = str(frame[key]) + else: + session = self._build_server_session(server, frame, sid) + if isinstance(frame.get("attached_images"), list): + session["attached_images"] = list(frame.get("attached_images") or []) + return session + + def _build_server_session(self, server: Any, frame: dict[str, Any], sid: str) -> dict: + """Build the agent under the frame's profile scope and register the session.""" + key = str(frame.get("session_key") or sid) history = frame.get("history") if isinstance(frame.get("history"), list) else [] profile_home = str(frame.get("profile_home") or "") - session_db = None + session_db = home_token = secret_token = None owns_db = False - home_token = None - secret_token = None try: if profile_home: from hermes_constants import set_hermes_home_override @@ -289,10 +285,8 @@ class ComputeHost: from hermes_state import get_shared_session_db home_token = set_hermes_home_override(profile_home) secret_token = set_secret_scope(build_profile_secret_scope(Path(profile_home))) - # DEDICATED handle — ours only until _make_agent succeeds; after that the - # agent (registered in server._sessions[sid] via _init_session or the - # fallback dict below) owns it. A RAISING _make_agent is the one path - # where nothing takes it, hence ``owns_db``. + # DEDICATED handle — ours only until _make_agent succeeds, then the agent owns + # it. A RAISING _make_agent is the one path where nothing takes it (``owns_db``). session_db = get_shared_session_db(Path(profile_home) / "state.db") owns_db = True agent = server._make_agent( @@ -327,9 +321,8 @@ class ComputeHost: finally: reset_transport(token) except Exception: - # If _init_session's side machinery (slash worker, approval notify) is - # unavailable, keep a minimal host-owned session rather than failing the - # turn after the expensive agent build succeeded. + # _init_session's side machinery (slash worker, approval notify) unavailable: keep a + # minimal host-owned session rather than failing after the expensive agent build. server._sessions[sid] = { "agent": agent, "session_key": key, "history": list(history), "history_lock": threading.Lock(), @@ -345,59 +338,44 @@ class ComputeHost: session = server._sessions[sid] session["transport"] = self._transport session["profile_home"] = profile_home or session.get("profile_home") - if isinstance(frame.get("attached_images"), list): - session["attached_images"] = list(frame.get("attached_images") or []) if frame.get("model_override") is not None: session["model_override"] = frame.get("model_override") return session def _handle_reload_mcp(self, frame: dict[str, Any]) -> None: - sid = str(frame.get("sid") or "") - request_id = frame.get("request_id") - try: - from tui_gateway import server + def body(server: Any, sid: str, request_id: Any) -> None: resp = server.handle_request({ "id": request_id, "method": "reload.mcp", "params": {"session_id": sid, "confirm": True}}) self._reply("reload_mcp.ack", sid, request_id, response=resp) - except Exception as exc: - self._reply("control.error", sid, request_id, message=str(exc)) + self._guarded(frame, "control.error", body) def _handle_control(self, frame: dict[str, Any]) -> None: - sid = str(frame.get("sid") or "") - request_id = frame.get("request_id") route_name = str(frame.get("route_name") or "") - def _error(message: str) -> None: - self._reply("control.error", sid, request_id, message=message) - try: - from tui_gateway import server + def body(server: Any, sid: str, request_id: Any) -> None: route = MUTATOR_ROUTE_TABLE.get(route_name) - if route is None: - _error(f"unclassified route: {route_name}") - return session = server._sessions.get(sid) - if session is None: - _error("session not found") - return - if route == "idle-gated" and session.get("running"): - _error("session busy") - return - if route_name == "reload.mcp": + error = (f"unclassified route: {route_name}" if route is None + else "session not found" if session is None + else "session busy" if route == "idle-gated" and session.get("running") + else None) + if error: + self._reply("control.error", sid, request_id, message=error) + elif route_name == "reload.mcp": self._handle_reload_mcp({**frame, "type": "reload_mcp"}) - return - ack = self._control_ack(server, frame, session) - if "error" in ack: - _error(ack["error"]) else: - self._reply("control.ack", sid, request_id, route_name=route_name, **ack) - except Exception as exc: + ack = self._control_ack(server, frame, session) + if "error" in ack: + self._reply("control.error", sid, request_id, message=ack["error"]) + else: + self._reply("control.ack", sid, request_id, route_name=route_name, **ack) + + def on_error(sid: str) -> None: if route_name in {"session.compress", "slash.compress"}: - # The compress mirror defers the context-engine boundary notification until - # the host commits. If anything raises between queueing and finalize (e.g. - # building the ack's session_info), discard the pending notification so it - # can't fire against a rejected boundary on a later compress. finalize is - # exactly-once, so this is a no-op if the mirror already emitted it. + # The compress mirror defers the context-engine boundary notification until the + # host commits; discard it so it can't fire against a rejected boundary later + # (finalize is exactly-once, so a no-op if the mirror already emitted it). with contextlib.suppress(Exception): from tui_gateway import server as _server from agent.conversation_compression import ( @@ -405,7 +383,7 @@ class ComputeHost: _agent = (_server._sessions.get(sid) or {}).get("agent") if _agent is not None: _finalize(_agent, committed=False) - _error(str(exc)) + self._guarded(frame, "control.error", body, on_error=on_error) def _control_ack(self, server: Any, frame: dict[str, Any], session: dict) -> dict: """control.ack payload for one classified route, or ``{"error": message}``.""" @@ -435,10 +413,6 @@ class ComputeHost: ack["session_info"] = server._session_info(session.get("agent"), session) return ack - def _bump_progress(self) -> None: - with self._progress_lock: - self._progress_counter += 1 - def _live_turns(self) -> list[concurrent.futures.Future]: with self._turn_futures_lock: return [f for f in self._turn_futures if not f.done()] diff --git a/tui_gateway/entry.py b/tui_gateway/entry.py index 33d83df476..209c48cbf9 100644 --- a/tui_gateway/entry.py +++ b/tui_gateway/entry.py @@ -1,9 +1,8 @@ import os import sys -# Stop a ``utils/`` (or ``proxy/``, ``ui/``) package in the launch directory from -# shadowing Hermes's own top-level modules. ``hermes_bootstrap`` lives at the repo -# root (its name can't collide with a user package), so importing it first is safe. +# Stop a ``utils/``-style package in the launch directory from shadowing Hermes's own +# top-level modules; ``hermes_bootstrap``'s name can't collide, so importing it first is safe. import hermes_bootstrap hermes_bootstrap.harden_import_path() @@ -29,15 +28,13 @@ logger = logging.getLogger(__name__) # Discovery thread spawned by THIS module; None when delegated to the shared owner in # hermes_cli.mcp_startup (current path). The wait/in-flight/join helpers consult both. _mcp_discovery_thread = None -# Set once MCP servers are found configured, so wait_for_mcp_discovery can re-invoke -# the idempotent spawn on later builds (retry-after-zero-connected) without a config -# re-probe — non-MCP sessions never pay the tools.mcp_tool import per build. +# Set once MCP servers are found configured so wait_for_mcp_discovery can re-invoke the +# idempotent spawn on later builds without a config re-probe. _mcp_discovery_enabled = False def _install_sidecar_publisher() -> None: - """Mirror every dispatcher emit to the dashboard sidebar via WS when - `HERMES_TUI_SIDECAR_URL` is set (best-effort: a dropped WS falls back to stdio-only).""" + """Mirror every dispatcher emit to the dashboard sidebar via WS when set (best-effort).""" url = os.environ.get("HERMES_TUI_SIDECAR_URL") if not url: return @@ -46,8 +43,7 @@ def _install_sidecar_publisher() -> None: # Grace for orderly shutdown before ``os._exit(0)`` so a worker wedged mid-flush can't -# strand the process; ``HERMES_TUI_GATEWAY_SHUTDOWN_GRACE_S`` overrides (a longer grace -# also means a longer wait on a real deadlock). +# strand the process; ``HERMES_TUI_GATEWAY_SHUTDOWN_GRACE_S`` overrides. _DEFAULT_SHUTDOWN_GRACE_S = 1.0 @@ -56,10 +52,6 @@ def _shutdown_grace_seconds() -> float: return value if value > 0 else _DEFAULT_SHUTDOWN_GRACE_S -def _stamp() -> str: - return time.strftime("%Y-%m-%d %H:%M:%S") - - def _mcp_startup_call(name: str, *args, default=None, log=None, **kwargs): """Call ``hermes_cli.mcp_startup.`` (lazy import); ``default`` on any failure, optionally logged as ``(level, message)``.""" @@ -88,10 +80,9 @@ def _append_crash_log(header: str, dump=None) -> None: def _log_signal(signum: int, frame) -> None: - """Capture WHICH thread and WHERE a termination signal hit us, then exit. - ``sys.exit(0)`` alone raced the worker pool (a thread holding ``_stdout_lock`` - mid-flush blocks interpreter shutdown), so: log all thread stacks, give the - configured grace to drain, then ``os._exit(0)``.""" + """Capture WHICH thread and WHERE a termination signal hit us, then exit. ``sys.exit(0)`` + alone raced the worker pool (a thread holding ``_stdout_lock`` mid-flush blocks interpreter + shutdown), so: log all thread stacks, give the configured grace to drain, then ``os._exit``.""" # SIGPIPE/SIGHUP don't exist on Windows — only look up attributes present. names = {int(sig): attr for attr in ("SIGPIPE", "SIGTERM", "SIGHUP", "SIGINT", "SIGBREAK") if (sig := getattr(signal, attr, None)) is not None} @@ -106,41 +97,35 @@ def _log_signal(signum: int, frame) -> None: f.write(f"\n--- thread {th.name} (id={tid}) ---\n") f.write("".join(traceback.format_stack(sys._current_frames().get(tid)))) - _append_crash_log(f"{name} received · {_stamp()}", _dump) + _append_crash_log(f"{name} received · {time.strftime('%Y-%m-%d %H:%M:%S')}", _dump) print(f"[gateway-signal] {name}", file=sys.stderr, flush=True) - - # ``os._exit`` skips atexit but breaks the mid-flush deadlock; the crash log - # + stderr line above are the forensic trail. + # ``os._exit`` skips atexit but breaks the mid-flush deadlock; the crash log is the trail. timer = threading.Timer(_shutdown_grace_seconds(), lambda: os._exit(0)) timer.daemon = True timer.start() - # atexit (_shutdown_sessions) can be blocked past the grace window by a worker - # holding the GIL/_stdout_lock; finalize explicitly so unpersisted messages reach - # state.db before the hard-exit timer fires. + # atexit (_shutdown_sessions) can be blocked past the grace window by a worker holding + # the GIL/_stdout_lock; finalize explicitly so unpersisted messages reach state.db first. with suppress(Exception): from tui_gateway.server import _shutdown_sessions _shutdown_sessions() - # Unwind the main thread so atexit + finalisers run inside the grace window; - # the daemon timer is the safety net if that unwind hangs. + # Unwind the main thread so atexit + finalisers run; the daemon timer is the safety net. sys.exit(0) def _install_signal(signame, handler): - """Install a signal handler if legal here: signal.signal() raises off the main - thread (Desktop build path: server._build imports entry from a worker), and - Windows lacks SIGPIPE/SIGHUP — both are skipped. Handlers are process-global.""" + """Install a signal handler if legal here: signal.signal() raises off the main thread + (Desktop build path imports entry from a worker) and Windows lacks SIGPIPE/SIGHUP.""" sig = getattr(signal, signame, None) if sig is None or threading.current_thread() is not threading.main_thread(): return - # Off the main thread despite the check, or handler rejected by the platform. - with suppress(ValueError, OSError, RuntimeError): + with suppress(ValueError, OSError, RuntimeError): # platform rejected the handler signal.signal(sig, handler) -# SIGPIPE: ignore, don't exit — SIG_DFL killed the process silently whenever a -# *background* thread (TTS, beep) wrote to a pipe the TUI had gone quiet on; ignoring -# lets write_json see BrokenPipeError and exit cleanly via _log_exit. Terminal signals -# route through _log_signal so kills/hangups are diagnosable (SIGBREAK = Windows SIGHUP). +# SIGPIPE: ignore, don't exit — SIG_DFL killed the process silently whenever a background +# thread wrote to a pipe the TUI had gone quiet on; ignoring lets write_json see +# BrokenPipeError and exit via _log_exit. Terminal signals route through _log_signal so +# kills/hangups are diagnosable (SIGBREAK = Windows SIGHUP). _install_signal("SIGPIPE", signal.SIG_IGN) _install_signal("SIGTERM", _log_signal) if hasattr(signal, "SIGHUP"): @@ -151,26 +136,23 @@ _install_signal("SIGINT", signal.SIG_IGN) def _log_exit(reason: str) -> None: - """Record why the gateway exits: every path collapses into a silent sys.exit(0), - and without this trail the TUI can't tell WHICH broken pipe triggered it.""" - _append_crash_log(f"gateway exit · {_stamp()} · reason={reason}") + """Record why the gateway exits (every path is a silent sys.exit(0) otherwise).""" + _append_crash_log(f"gateway exit · {time.strftime('%Y-%m-%d %H:%M:%S')} · reason={reason}") print(f"[gateway-exit] {reason}", file=sys.stderr, flush=True) def wait_for_mcp_discovery(timeout: "float | None" = None) -> None: - """Block until background MCP discovery finishes, up to the resolved bound - (``mcp_discovery_timeout`` from config; ``timeout`` overrides). The agent snapshots - its tool list ONCE at build time, so this bounded join lets already-spawning - servers land without re-introducing the startup hang.""" + """Block until background MCP discovery finishes, up to the resolved bound (config + ``mcp_discovery_timeout``; ``timeout`` overrides). The agent snapshots its tool list ONCE + at build time, so this bounded join lets already-spawning servers land.""" thread = _mcp_discovery_thread if thread is not None and thread.is_alive(): fallback = timeout if timeout is not None else 0.75 bound = _mcp_startup_call("_resolve_discovery_timeout", timeout, default=fallback) thread.join(timeout=bound) return - # Shared-owner path: re-invoke the idempotent spawn first so a zero-connected run - # gets its retry instead of latching the process MCP-less. Runs under the CALLER's - # profile context (agent build binds the session profile's HERMES_HOME first). + # Shared-owner path: re-invoke the idempotent spawn first so a zero-connected run gets + # its retry instead of latching the process MCP-less (runs under the CALLER's profile). if not _mcp_discovery_enabled: return _spawn_discovery(("debug", "TUI MCP discovery retry-spawn failed")) @@ -178,10 +160,8 @@ def wait_for_mcp_discovery(timeout: "float | None" = None) -> None: def mcp_discovery_in_flight() -> bool: - """True if ANY background MCP discovery thread is still running. Two owners by - surface (stdio thread here, ``hermes_cli.mcp_startup`` for desktop/dashboard); - the late-refresh scheduler calls this regardless of surface, so it MUST consult - both or slow MCP servers' tools never surface on desktop.""" + """True if ANY background MCP discovery thread is still running: the late-refresh + scheduler calls this regardless of surface, so it MUST consult both owners.""" thread = _mcp_discovery_thread if thread is not None and thread.is_alive(): return True @@ -189,9 +169,8 @@ def mcp_discovery_in_flight() -> bool: def join_mcp_discovery(timeout: float | None = None) -> bool: - """Join both discovery owners; True once neither is alive. Unlike - ``wait_for_mcp_discovery`` this accepts an unbounded wait (off-critical-path - late-refresh waiter); ``timeout`` bounds EACH join, entry thread first.""" + """Join both discovery owners; True once neither is alive. Accepts an unbounded wait + (off-critical-path late-refresh waiter); ``timeout`` bounds EACH join, entry thread first.""" entry_done = True thread = _mcp_discovery_thread if thread is not None: @@ -211,13 +190,10 @@ def _has_configured_mcp_servers() -> bool: def ensure_mcp_discovery_started() -> None: - """Start background MCP discovery for the current profile context, once. - ``main()`` calls this for stdio; WS/Desktop skip ``main()``, so - ``server._start_agent_build`` also calls it AFTER binding the session profile's - HERMES_HOME so discovery reads the SELECTED profile's ``mcp_servers``. MCP - registration is process-global: the FIRST profile to build an agent wins.""" + """Start background MCP discovery for the current profile context, once. ``main()`` calls + this for stdio; ``server._start_agent_build`` also calls it AFTER binding the session + profile's HERMES_HOME. MCP registration is process-global: the FIRST profile wins.""" global _mcp_discovery_enabled - if not _has_configured_mcp_servers(): return _mcp_discovery_enabled = True @@ -233,40 +209,30 @@ def _write_or_exit(payload: dict, reason: str) -> None: def main(): _install_sidecar_publisher() - # Heartbeat row lets the orphan sweep tell "live but idle" from "truly orphaned"; - # it must run BEFORE the sweep. The sweep is once-per-process and config-gated. + # The heartbeat row lets the orphan sweep tell "live but idle" from "truly orphaned", + # so it must start BEFORE the sweep. for start, what in ( - (server._start_backend_heartbeat_refresher, "backend heartbeat refresher start"), - (server._schedule_startup_orphan_sweep, "startup orphan sweep scheduling"), - ): + (server._start_backend_heartbeat_refresher, "backend heartbeat refresher start"), + (server._schedule_startup_orphan_sweep, "startup orphan sweep scheduling")): try: start() except Exception: logger.warning("%s failed", what, exc_info=True) - # Backgrounded so a dead MCP server (~7s of retries) can't freeze startup; - # _make_agent briefly joins it. The config gate keeps the MCP SDK import off - # the no-mcp_servers path. + # Backgrounded so a dead MCP server can't freeze startup; _make_agent briefly joins it. ensure_mcp_discovery_started() + # change_events: clients demote legacy polls; replay_epoch: WS restart detection. _write_or_exit({ - "jsonrpc": "2.0", - "method": "event", - "params": { - "type": "gateway.ready", - # change_events: clients demote legacy polls (see tui_gateway/ws.py). - # replay_epoch: WS restart detection; the stdio TUI ignores it. - "payload": { - "skin": resolve_skin(), "change_events": True, "replay_epoch": replay_epoch(), - }, - }, - }, "startup write failed (broken stdout pipe before first event)") + "jsonrpc": "2.0", "method": "event", + "params": {"type": "gateway.ready", "payload": { + "skin": resolve_skin(), "change_events": True, "replay_epoch": replay_epoch()}}}, + "startup write failed (broken stdout pipe before first event)") # Live-apply skins Hermes activates mid-conversation. server._ensure_skin_watcher() - # Warm the /model picker's provider-models cache in this idle window, else the - # first /model open blocks on serial /v1/models fetches. Fire-and-forget. + # Warm the /model picker's provider-models cache in this idle window (fire-and-forget). try: from hermes_cli.model_switch import prewarm_picker_cache_async prewarm_picker_cache_async() @@ -280,11 +246,9 @@ def main(): if not handle_spurious_eof(_recovery_times, _log_exit): break continue - line = raw.strip() if not line: continue - try: req = json.loads(line) except json.JSONDecodeError: diff --git a/tui_gateway/mcp_oauth_sessions.py b/tui_gateway/mcp_oauth_sessions.py index 4a1f1c8a16..650575fd0d 100644 --- a/tui_gateway/mcp_oauth_sessions.py +++ b/tui_gateway/mcp_oauth_sessions.py @@ -1,12 +1,9 @@ -"""Session-backed MCP OAuth flows for the gateway (mcp.servers.oauth.*). - -``start`` kicks off a background worker and returns ``{session_id, auth_url, flow}``; -``poll`` reports ``{status: pending|approved|error}`` until tokens land on disk. No OAuth -logic is reimplemented: ``hermes mcp login``'s probe under ``force_interactive_oauth`` -plus ``DashboardOAuthFlow`` as the bridge; the only new piece is a loopback listener -feeding ``deliver_callback``. Remote backends: the client hosts the listener, passes -``client_redirect_uri`` and relays via ``deliver_callback_flow`` (state check stays here). -""" +"""Session-backed MCP OAuth flows for the gateway (mcp.servers.oauth.*): ``start`` spawns a +worker and returns ``{session_id, auth_url, flow}``; ``poll`` reports ``{status}`` until tokens +land. Reuses ``hermes mcp login``'s probe under ``force_interactive_oauth`` plus +``DashboardOAuthFlow``; the only new piece is a loopback listener feeding ``deliver_callback``. +Remote backends host the listener (``client_redirect_uri``) and relay via +``deliver_callback_flow``.""" from __future__ import annotations @@ -23,18 +20,8 @@ from urllib.parse import parse_qs, urlparse _sessions: Dict[str, Dict[str, Any]] = {} _sessions_lock = threading.Lock() -# How long a completed/abandoned session lingers before GC (seconds). -_SESSION_TTL_SECONDS = 900 -# Cap concurrent in-flight flows so a runaway client can't exhaust ports/threads. -_MAX_PENDING = 12 - - -def _gc_sessions() -> None: - """Drop expired sessions. Called opportunistically on start.""" - cutoff = time.time() - _SESSION_TTL_SECONDS - with _sessions_lock: - for sid in [sid for sid, rec in _sessions.items() if rec["created_at"] < cutoff]: - _shutdown_listener(_sessions.pop(sid)) +_SESSION_TTL_SECONDS = 900 # completed/abandoned session lingers this long before GC +_MAX_PENDING = 12 # cap in-flight flows so a runaway client can't exhaust ports/threads def _shutdown_listener(rec: Dict[str, Any]) -> None: @@ -48,24 +35,20 @@ def _shutdown_listener(rec: Dict[str, Any]) -> None: def _validate_client_redirect_uri(uri: str) -> str: - """Accept only plain-http loopback URLs (RFC 8252 native-app rules) so the - gateway can't pin an attacker-controlled redirect into a DCR registration.""" + """Accept only plain-http loopback URLs (RFC 8252) so the gateway can't pin an + attacker-controlled redirect into a DCR registration.""" parsed = urlparse(str(uri or "").strip()) host = (parsed.hostname or "").lower() if (parsed.scheme != "http" or host not in ("127.0.0.1", "localhost", "::1") or not parsed.port or parsed.username is not None or parsed.password is not None): raise ValueError( - "client_redirect_uri must be a loopback http URL like " - "http://127.0.0.1:/callback" - ) + "client_redirect_uri must be a loopback http URL like http://127.0.0.1:/callback") return f"http://{'[' + host + ']' if ':' in host else host}:{parsed.port}{parsed.path or '/callback'}" def _start_loopback_listener(flow) -> "http.server.HTTPServer": """Bind a loopback callback listener feeding ``flow.deliver_callback``; returns the - HTTPServer already serving on a daemon thread. The caller pins ``flow.redirect_uri`` - from ``server_address`` BEFORE the worker starts (fixed at authorization).""" - + HTTPServer already serving on a daemon thread (caller pins ``flow.redirect_uri`` from it).""" class _Handler(http.server.BaseHTTPRequestHandler): def do_GET(self): # noqa: N802 — stdlib naming parsed = urlparse(self.path) @@ -74,11 +57,11 @@ def _start_loopback_listener(flow) -> "http.server.HTTPServer": self.end_headers() return qs = parse_qs(parsed.query) - code, state, error = ((qs.get(k) or [None])[0] for k in ("code", "state", "error")) body = b"

Authorization received

You can close this tab and return to Hermes.

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

OAuth callback rejected

The callback was invalid or already used.

" status = 400 @@ -129,9 +112,9 @@ def _probe_with_rollback( raise -def _worker(session_id: str, hermes_home: str, server_name: str, cfg: dict, reconnect_live: bool) -> None: - """Drive the interactive MCP OAuth probe under the shared dashboard bridge (same - HERMES_HOME + secret-scope + force_interactive_oauth + dashboard_oauth_flow wrapping +def _worker( + session_id: str, hermes_home: str, server_name: str, cfg: dict, reconnect_live: bool) -> None: + """Drive the interactive MCP OAuth probe under the shared dashboard bridge (same wrapping as ``web_server._run_dashboard_mcp_oauth``), keyed to our session record.""" from hermes_constants import reset_hermes_home_override, set_hermes_home_override rec = _sessions.get(session_id) @@ -166,23 +149,18 @@ def _worker(session_id: str, hermes_home: str, server_name: str, cfg: dict, reco def start_flow( - hermes_home: str, - server_name: str, - cfg: dict, - *, - reconnect_live: bool = False, - url_timeout: float = 30.0, - client_redirect_uri: Optional[str] = None) -> Dict[str, Any]: - """Begin an MCP OAuth flow and return ``{session_id, auth_url, flow}``; blocks up - to ``url_timeout`` for the authorization URL. With ``client_redirect_uri`` (remote - backend; invalid values raise ``ValueError``) no gateway-side listener is bound and - the client relays ``code``/``state`` via ``deliver_callback_flow``.""" + hermes_home: str, server_name: str, cfg: dict, *, reconnect_live: bool = False, + url_timeout: float = 30.0, client_redirect_uri: Optional[str] = None) -> Dict[str, Any]: + """Begin an MCP OAuth flow and return ``{session_id, auth_url, flow}``; blocks up to + ``url_timeout`` for the authorization URL. With ``client_redirect_uri`` (invalid values + raise ``ValueError``) no gateway-side listener is bound.""" from tools.mcp_dashboard_oauth import DashboardOAuthFlow if client_redirect_uri is not None: client_redirect_uri = _validate_client_redirect_uri(client_redirect_uri) - - _gc_sessions() - + cutoff = time.time() - _SESSION_TTL_SECONDS # opportunistic GC of expired sessions + with _sessions_lock: + for sid in [sid for sid, rec in _sessions.items() if rec["created_at"] < cutoff]: + _shutdown_listener(_sessions.pop(sid)) with _sessions_lock: active = [r for r in _sessions.values() if not r["flow"].worker_done] if len(active) >= _MAX_PENDING: @@ -199,18 +177,14 @@ def start_flow( httpd = None if client_redirect_uri else _start_loopback_listener(flow) flow.redirect_uri = ( client_redirect_uri or f"http://127.0.0.1:{httpd.server_address[1]}/callback") - rec = { "session_id": session_id, "server_name": server_name, "hermes_home": hermes_home, - "flow": flow, "httpd": httpd, "created_at": time.time(), - } + "flow": flow, "httpd": httpd, "created_at": time.time()} with _sessions_lock: _sessions[session_id] = rec - threading.Thread( target=_worker, args=(session_id, hermes_home, server_name, dict(cfg), reconnect_live), daemon=True, name=f"mcp-oauth-{server_name}").start() - try: auth_url = None # wait_for_authorization_url is async; run its wait synchronously. @@ -220,7 +194,8 @@ def start_flow( if auth_url := snap.get("authorization_url"): break if snap.get("status") == "error": - raise RuntimeError(snap.get("error") or "MCP OAuth flow failed before authorization") + raise RuntimeError( + snap.get("error") or "MCP OAuth flow failed before authorization") time.sleep(0.1) if not auth_url: raise TimeoutError("Timed out waiting for MCP authorization URL") @@ -228,9 +203,7 @@ def start_flow( flow.mark_error("Timed out waiting for MCP authorization URL") _shutdown_listener(rec) raise - - # ``flow`` mirrors the provider-OAuth discriminator: open a URL then poll - # (no user_code to type, unlike device_code). + # ``flow`` mirrors the provider-OAuth discriminator: open a URL then poll (no user_code). return {"session_id": session_id, "auth_url": auth_url, "flow": "pkce"} @@ -246,21 +219,19 @@ def _lookup(session_id: str, server_name: str) -> "tuple[Dict[str, Any] | None, def poll_flow(session_id: str, server_name: str) -> Dict[str, Any]: - """Poll a session → ``{status, error_message?, auth_url?, tools?}``; ``status`` - is ``pending`` | ``approved`` | ``error`` (the bridge's ``authorization_required`` - maps to ``pending`` — the client only needs to know whether to keep waiting).""" + """Poll a session → ``{status, error_message?, auth_url?, tools?}``; ``status`` is + ``pending`` | ``approved`` | ``error`` (the bridge's ``authorization_required`` maps to + ``pending``).""" rec, err = _lookup(session_id, server_name) if rec is None: return {"status": "error", "error_message": err} - flow = rec["flow"] snap = flow.snapshot() raw = snap.get("status") status = raw if raw in ("approved", "error") else "pending" out: Dict[str, Any] = { "session_id": session_id, "status": status, "error_message": snap.get("error"), - "auth_url": snap.get("authorization_url"), - } + "auth_url": snap.get("authorization_url")} if status == "approved": out["tools"] = list(getattr(flow, "tools", []) or []) return out @@ -268,12 +239,10 @@ def poll_flow(session_id: str, server_name: str) -> Dict[str, Any]: def deliver_callback_flow( session_id: str, server_name: str, *, code: Optional[str], state: Optional[str], - error: Optional[str] = None, -) -> Dict[str, Any]: - """Relay a client-captured OAuth redirect into a session's flow (remote-backend - companion to ``start_flow(client_redirect_uri=...)``). Security is unchanged: - ``DashboardOAuthFlow.deliver_callback`` verifies ``state`` (constant-time) and - rejects replays. Returns ``{ok: true}`` or ``{ok: false, error_message}``.""" + error: Optional[str] = None) -> Dict[str, Any]: + """Relay a client-captured OAuth redirect into a session's flow (remote-backend companion + to ``start_flow(client_redirect_uri=...)``); ``deliver_callback`` still verifies ``state`` + and rejects replays. Returns ``{ok: true}`` or ``{ok: false, error_message}``.""" rec, err = _lookup(session_id, server_name) if rec is None: return {"ok": False, "error_message": err} diff --git a/tui_gateway/methods_groups.py b/tui_gateway/methods_groups.py index a0480aedfe..44a5c9b12b 100644 --- a/tui_gateway/methods_groups.py +++ b/tui_gateway/methods_groups.py @@ -1,9 +1,7 @@ """Hosted-room JSON-RPC contract: durable room identity, replay, and the process-owned -same-gateway Discussion driver. ``groups.capabilities`` keeps that boundary -machine-readable so older clients stay on the renderer-owned room path. +same-gateway Discussion driver; ``groups.capabilities`` keeps that boundary machine-readable. -Handlers are rebound onto server.py's globals at install (method_ctx.py), so bodies see -only server globals plus what methods_bot_relay.register publishes; module-private +Handlers are rebound onto server.py's globals at install (method_ctx.py); module-private helpers reach them through keyword defaults. ``_room_method`` is the shared envelope.""" from .method_ctx import HandlerRegistry @@ -91,6 +89,14 @@ def _current_profile() -> str: return str(_bound_server._current_profile_name() or "").strip() +def _foreign_profile_home(profile: str): + """Home of a routed profile other than the process's own, or ``ValueError``.""" + home = _bound_server._profile_home(profile) + if home is None: + raise ValueError(f"profile '{profile}' is unavailable") + return home + + def _requested_profile(params: dict) -> str: requested = str(params.get("profile") or "").strip() if not requested: @@ -99,19 +105,18 @@ def _requested_profile(params: dict) -> str: raise ValueError("profile routing is unavailable") if requested == _current_profile(): return requested - if _bound_server._profile_home(requested) is None: - raise ValueError(f"profile '{requested}' is unavailable") + _foreign_profile_home(requested) return str(_bound_server._response_profile_name(requested) or requested) def _api_server_key(profile: str | None = None) -> str: + # Published onto the server by methods_bot_relay.register (an explicit routed profile is + # authoritative: never borrow the process profile's key on a multiplexed gateway). if profile and _bound_server is not None and profile != _current_profile(): from agent.secret_scope import build_profile_secret_scope home = _bound_server._profile_home(profile) if home is None: return "" - # An explicit routed profile is authoritative. Never borrow the - # process/default profile's API key on a multiplexed gateway. return str(build_profile_secret_scope(home).get("API_SERVER_KEY") or "").strip() scoped = "" with contextlib.suppress(Exception): @@ -126,10 +131,7 @@ def _profile_execution_policy(profile: str) -> dict: from hermes_constants import reset_hermes_home_override, set_hermes_home_override token = None if _bound_server is not None and profile not in {_current_profile(), _profile_name()}: - home = _bound_server._profile_home(profile) - if home is None: - raise ValueError(f"profile '{profile}' is unavailable") - token = set_hermes_home_override(str(home)) + token = set_hermes_home_override(str(_foreign_profile_home(profile))) try: return execution_policy_mapping(target_profile=profile) finally: @@ -140,14 +142,12 @@ def _profile_execution_policy(profile: str) -> dict: def _room_link_run_storage_durable() -> bool: """Return whether peer-run replay survives this gateway process.""" if _bound_server is None: - # Direct method-contract tests and embedded callers without a bound API - # server do not expose peer-run transport; production always binds first. + # Embedded callers without a bound server expose no peer-run transport. return True store = getattr(_bound_server, "_run_idempotency_store", None) if store is None: - # The dashboard/TUI process owns groups.* but does not construct the API adapter - # that owns this store. Open the same shared SQLite-backed store lazily so - # capability negotiation reflects the real /v1/runs replay boundary. + # This process does not construct the API adapter that owns the store; open the + # same shared SQLite store lazily so negotiation reflects the real replay boundary. from gateway.platforms.api_server import RunIdempotencyStore with _run_store_lock: store = getattr(_bound_server, "_run_idempotency_store", None) @@ -169,6 +169,10 @@ def _grant_expiry(claims: dict) -> float: return float(claims.get("status_expires_at", claims["expires_at"])) +def _include_disbanded(params: dict) -> bool: + return params.get("include_disbanded") is True + + def _room_error_class(replica_only: bool) -> type: if replica_only: from gateway.hosted_room_replicas import ReplicaError @@ -182,12 +186,10 @@ def _room_method( with_reason: bool = True, service_code: int | None = None, service_message: str = _DRIVER_UNAVAILABLE, db: bool = False): """Register ``fn`` under ``name`` with the shared hosted-room error envelope. - - ``service_code``: the live service is required (else fail with that code) and passed - as a third argument; ``db``: the default room db path follows. ``room_code`` maps - ``HostedRoomError`` (only ``ReplicaError`` when ``replica_only``) to a 4xxx client - error with ``{"reason"}`` data when ``with_reason``; anything else maps to ``code``. - """ + ``service_code``: the live service is required (else that error) and passed as a third + argument; ``db``: the default room db path follows. ``room_code`` maps ``HostedRoomError`` + (only ``ReplicaError`` when ``replica_only``) to a client error with ``{"reason"}`` data + when ``with_reason``; anything else maps to ``code``.""" error_class = _room_error_class # closure cell: handlers run under server.py globals def dec(fn): @@ -228,26 +230,21 @@ def _(rid, params: dict, _catalog=_local_catalog, _methods=_METHODS) -> dict: policy = _profile_execution_policy(profile) catalog = _catalog(local_authority_gateway_id(), profile, policy) room_link = { - "enabled": True, "profile": profile, "catalog": catalog, "endpoint": catalog["endpoint"] - } + "enabled": True, "profile": profile, "catalog": catalog, + "endpoint": catalog["endpoint"]} except Exception: - room_link = { - "enabled": False, - "reason": ( - "durable_run_storage_required" if not _room_link_run_storage_durable() - else "gateway_roomlink_secret_unavailable")} + room_link = {"enabled": False, "reason": ( + "durable_run_storage_required" if not _room_link_run_storage_durable() + else "gateway_roomlink_secret_unavailable")} return _ok(rid, { - "protocol_version": PROTOCOL_VERSION, - "driver": driver_ready, + "protocol_version": PROTOCOL_VERSION, "driver": driver_ready, "persistent_process": bool(room_link.get("catalog", {}).get("persistent_process", False)), - "authority_gateway_id": local_authority_gateway_id(), - "room_link": room_link, + "authority_gateway_id": local_authority_gateway_id(), "room_link": room_link, "features": [ "authority_epoch", "coordinator_fencing", "room_identity", "monotonic_log", "idempotent_send", "replayable_disband", "typed_events", "actor_identity", "log_replication", "authority_takeover"], - "methods": list(_methods), - "max_log_limit": MAX_LOG_LIMIT}) + "methods": list(_methods), "max_log_limit": MAX_LOG_LIMIT}) @_room_method("groups.peer.invite", code=4120, db=True) @@ -290,9 +287,8 @@ def _(rid, params: dict, db_path, _expiry=_grant_expiry) -> dict: profile = _requested_profile(params) claims = decode_room_grant( gateway_room_grant_secret(), str(params.get("grant") or ""), permission="status") - if ( - claims["target_profile"] != profile - or claims["target_install_id"] != local_authority_gateway_id()): + if (claims["target_profile"] != profile + or claims["target_install_id"] != local_authority_gateway_id()): raise ValueError("room grant target does not match this profile") revoke_room_grant_scope(db_path, claims=claims, expires_at=_expiry(claims)) return _ok(rid, {"revoked": True}) @@ -316,13 +312,9 @@ def _(rid, params: dict, service) -> dict: grant = str(params.get("grant") or "") client = PeerRunsHTTPClient(base_url=target_url, api_key="", receipt_db_path=service.db_path) probe = client.probe(grant=grant) - live_catalog = GatewayRoomCatalog.from_mapping(probe.get("catalog")) - if live_catalog != catalog: + # Frozen dataclass equality: an equal live catalog already passed the checks above. + if GatewayRoomCatalog.from_mapping(probe.get("catalog")) != catalog: raise ValueError("target capability catalog changed during setup") - if ( - ROOM_LINK_PROTOCOL_VERSION not in live_catalog.protocol_versions - or "direct" not in live_catalog.link_modes): - raise ValueError("target RoomLink capability is incompatible") room_id = str(params.get("room_id") or "") member_id = str(params.get("member_id") or "") home_install_id = local_authority_gateway_id() @@ -331,9 +323,9 @@ def _(rid, params: dict, service) -> dict: "room_id": room_id, "home_install_id": home_install_id, "authority_gateway_id": home_room.get("authority_gateway_id"), "member_id": member_id, "target_profile": target_profile} - if ( - any(probe.get(key) != value for key, value in expected_scope.items()) - or int(probe.get("authority_epoch") or 0) != int(home_room.get("authority_epoch") or 0)): + if (any(probe.get(k) != v for k, v in expected_scope.items()) + or int(probe.get("authority_epoch") or 0) + != int(home_room.get("authority_epoch") or 0)): raise ValueError("room grant scope does not match this route") route = PeerMemberRoute( home_install_id=home_install_id, member_id=member_id, @@ -368,8 +360,7 @@ def _(rid, params: dict, db_path) -> dict: "groups.create", code=5111, room_code=4110, service_code=4123, service_message=_WORKER_UNAVAILABLE) def _(rid, params: dict, service) -> dict: - """Create a hosted room idempotently; authority comes from this gateway's stable - install identity, never from the client.""" + """Create a hosted room idempotently; authority is this gateway's stable install identity.""" room = service.create_room( room_id=params.get("room_id"), name=params.get("name"), members=params.get("members")) return _ok(rid, {"room": room}) @@ -390,19 +381,18 @@ def _(rid, params: dict, db_path) -> dict: @_room_method( - "groups.send", code=5112, room_code=4111, service_code=4123, service_message=_WORKER_UNAVAILABLE -) + "groups.send", code=5112, room_code=4111, service_code=4123, + service_message=_WORKER_UNAVAILABLE) def _(rid, params: dict, service) -> dict: - """Append one typed event idempotently. Only inert ``message.user`` events are - accepted from clients; the actor is server-owned rather than trusted from params.""" + """Append one typed event idempotently (inert ``message.user`` only; actor is server-owned).""" from gateway.hosted_rooms import user_event_id client_event_id = params.get("event_id") event = service.send( room_id=params.get("room_id"), event_id=user_event_id(client_event_id), payload=params.get("payload")) return _ok(rid, { - "event": event, "client_event_id": client_event_id, "accepted": True, "driver_started": True - }) + "event": event, "client_event_id": client_event_id, "accepted": True, + "driver_started": True}) @_room_method( @@ -466,9 +456,8 @@ def _(rid, params: dict, service) -> dict: task = {} identity = task.get("identity") receipt = { - **{ - field: str(getattr(identity, field, "") or "") - for field in ("room_id", "task_id", "thread_id", "turn_id")}, + **{f: str(getattr(identity, f, "") or "") + for f in ("room_id", "task_id", "thread_id", "turn_id")}, "status": str(task.get("status") or ""), "execution_generation": int(task.get("execution_generation") or 0), "cancel_generation": int(task.get("cancel_generation") or 0)} @@ -478,30 +467,21 @@ def _(rid, params: dict, service) -> dict: def _passthrough( name: str, module: str, fn_name: str, doc: str, *, code: int, room_code: int, params: tuple, replica_only: bool = False, wrap: str | None = None) -> None: - """Register a method whose result is ``module.fn(db_path, **params)`` verbatim (or under - key ``wrap``). ``params`` items are ``key`` (-> ``params.get(key)``) or - ``(key, extractor(params))``.""" - + """Register a method whose result is ``module.fn(db_path, **params)`` verbatim (or under key + ``wrap``). ``params`` items are ``key`` (-> ``params.get(key)``) or ``(key, extractor)``.""" @_room_method( name, code=code, room_code=room_code, replica_only=replica_only, with_reason=not replica_only, db=True) def handler(rid, params_in: dict, db_path, _import=importlib.import_module) -> dict: - fn = getattr(_import(module), fn_name) - kwargs = {} - for spec in params: - if isinstance(spec, str): - kwargs[spec] = params_in.get(spec) - else: - kwargs[spec[0]] = spec[1](params_in) - result = fn(db_path, **kwargs) + kwargs = { + (spec if isinstance(spec, str) else spec[0]): + (params_in.get(spec) if isinstance(spec, str) else spec[1](params_in)) + for spec in params} + result = getattr(_import(module), fn_name)(db_path, **kwargs) return _ok(rid, {wrap: result} if wrap else result) handler.__doc__ = doc -def _include_disbanded(params: dict) -> bool: - return params.get("include_disbanded") is True - - _passthrough( "groups.rename", "gateway.hosted_rooms", "rename_room", """Rename one hosted room atomically with its replay event.""", @@ -531,10 +511,8 @@ def _(rid, params: dict, db_path) -> dict: true`` — the caller asserts the previous authority can no longer commit.""" from gateway.hosted_room_replicas import promote_replica if params.get("confirm") is not True: - return _err( - rid, 4118, - "promotion requires confirm=true acknowledging the previous " - "authority can no longer commit") + return _err(rid, 4118, "promotion requires confirm=true acknowledging the previous " + "authority can no longer commit") reason = params.get("reason", "authority-unreachable") return _ok(rid, promote_replica(db_path, room_id=params.get("room_id"), reason=reason)) diff --git a/tui_gateway/methods_projects.py b/tui_gateway/methods_projects.py index 6a6d7cac0d..0023eb4d5c 100644 --- a/tui_gateway/methods_projects.py +++ b/tui_gateway/methods_projects.py @@ -1,8 +1,5 @@ """Projects RPC surface: per-profile multi-folder workspaces, repo discovery, sidebar tree. - -Bodies are rebound onto server.py's globals at install time (see -method_ctx.bind_module), so they reference server.py globals bare. -""" +Bodies are rebound onto server.py's globals at install (method_ctx.bind_module).""" from __future__ import annotations @@ -12,10 +9,8 @@ _registry = HandlerRegistry() method = _registry.method -# JSON-RPC error codes for the projects surface. -_E_PROJECTS = 5061 # generic failure -_E_NO_PROJECT = 5062 # id resolved to nothing -_E_PROJECT_ARG = 5063 # invalid argument (e.g. bad name/slug) +# JSON-RPC error codes: generic failure / id resolved to nothing / invalid argument. +_E_PROJECTS, _E_NO_PROJECT, _E_PROJECT_ARG = 5061, 5062, 5063 class _NoProject(Exception): @@ -30,13 +25,8 @@ def _projects_payload(conn) -> dict: def _projects_method(name: str): - """Register a projects RPC, injecting (pdb, conn) and unifying error mapping. - - Binds ``params['profile']`` (via ``@_profile_scoped``) so app-global remote - mode reads that profile's ``projects.db``. Missing id maps to 5062, bad args - to 5063, everything else to 5061. - """ - + """Register a projects RPC, injecting (pdb, conn) and unifying error mapping; profile-scoped + so app-global remote mode reads that profile's ``projects.db``.""" def decorator(fn): @method(name) @_registry.profile_scoped @@ -67,20 +57,9 @@ def _pick(params: dict, *keys: str) -> dict: return {k: params.get(k) for k in keys} -# Per-project mutators: (rpc suffix, pdb function, takes params['path'], extra kwargs). -# Each resolves ``params['id']`` (5062 when missing), mutates, and answers with the -# refreshed project. -_PROJECT_MUTATORS = ( - ("update", "update_project", False, - lambda p: _pick(p, "name", "description", "icon", "color", "board_slug")), - ("add_folder", "add_folder", True, - lambda p: {"label": p.get("label"), "is_primary": bool(p.get("is_primary"))}), - ("remove_folder", "remove_folder", True, lambda p: {}), - ("set_primary", "set_primary", True, lambda p: {}), -) - - def _register_project_mutator(suffix: str, fn_name: str, takes_path: bool, kwargs_of) -> None: + """``projects.``: resolve ``params['id']`` (5062 when missing), call + ``pdb.(conn, id[, path], **kwargs_of(params))``, answer with the refreshed project.""" @_projects_method(f"projects.{suffix}") def _(rid, params, pdb, conn) -> dict: proj = _require_project(pdb, conn, params) @@ -89,9 +68,14 @@ def _register_project_mutator(suffix: str, fn_name: str, takes_path: bool, kwarg return _ok(rid, {"project": pdb.get_project(conn, proj.id).to_dict()}) -for _spec in _PROJECT_MUTATORS: - _register_project_mutator(*_spec) -del _spec +_register_project_mutator( + "update", "update_project", False, + lambda p: _pick(p, "name", "description", "icon", "color", "board_slug")) +_register_project_mutator( + "add_folder", "add_folder", True, + lambda p: {"label": p.get("label"), "is_primary": bool(p.get("is_primary"))}) +_register_project_mutator("remove_folder", "remove_folder", True, lambda p: {}) +_register_project_mutator("set_primary", "set_primary", True, lambda p: {}) @_projects_method("projects.list") @@ -136,25 +120,26 @@ def _(rid, params, pdb, conn) -> dict: @_projects_method("projects.for_cwd") def _(rid, params, pdb, conn) -> dict: - cwd = _completion_cwd({"cwd": str(params.get("cwd") or "").strip()} if params.get("cwd") else {}) + cwd = _completion_cwd( + {"cwd": str(params.get("cwd") or "").strip()} if params.get("cwd") else {}) proj = pdb.project_for_path(conn, cwd) - return _ok(rid, {"project": proj.to_dict() if proj else None, "cwd": cwd, "branch": _git_branch_for_cwd(cwd)}) + return _ok(rid, { + "project": proj.to_dict() if proj else None, "cwd": cwd, + "branch": _git_branch_for_cwd(cwd)}) def _non_workspace_dirs() -> set[str]: - """Never-a-workspace dirs: ``/``, the user's home, and the dir homes live in. Both - POSIX spellings are excluded on every host (macOS ships an empty ``/home`` autofs - stub; containers/remote shells hand back Linux paths) — promoting one mints a - catch-all project and ``/home`` renders as a second "home" row beside Home.""" + """Never-a-workspace dirs: ``/``, the user's home, the dir homes live in, plus both POSIX + spellings on every host (remote shells hand back Linux paths; promoting one mints a + catch-all project).""" home = os.path.realpath(os.path.expanduser("~")) candidates = (os.sep, home, os.path.dirname(home), "/home", "/Users") return {os.path.normcase(os.path.realpath(path)) for path in candidates if path} def _is_repo_junk(root: str) -> bool: - """A git root never auto-surfaced as a project: a non-workspace dir or anything - under HERMES_HOME (config/sessions/skills). User-created projects pointing - there are still honored.""" + """A git root never auto-surfaced as a project: a non-workspace dir or anything under + HERMES_HOME. User-created projects pointing there are still honored.""" if not root: return True from hermes_constants import get_hermes_home @@ -167,9 +152,8 @@ def _is_repo_junk(root: str) -> bool: def _is_session_cwd_junk(cwd: str) -> bool: - """A non-git cwd that stays in flat Recents rather than auto-grouping. Unlike git - roots, a selected DESCENDANT of HERMES_HOME may be an intentional prose/data - workspace, so only HERMES_HOME itself and ``_non_workspace_dirs`` are excluded.""" + """A non-git cwd that stays in flat Recents. A DESCENDANT of HERMES_HOME may be an + intentional prose/data workspace, so only HERMES_HOME itself is excluded here.""" if not cwd: return True from hermes_constants import get_hermes_home @@ -194,25 +178,19 @@ def _repo_discovery_policy(raw: dict | None = None) -> dict: if not isinstance(values, list): return list(defaults[long]) return [v.strip() for v in values if isinstance(v, str) and v.strip()] - enabled = _get("enabled", "repo_scan_enabled") return { "enabled": enabled if isinstance(enabled, bool) else defaults["repo_scan_enabled"], "roots": _paths("roots", "repo_scan_roots"), - "exclude_paths": _paths("exclude_paths", "repo_scan_exclude_paths"), - } + "exclude_paths": _paths("exclude_paths", "repo_scan_exclude_paths")} def _repo_discovery_policy_key(policy: dict) -> str: def _paths(values: list[str]) -> list[str]: - normalized = set() home = os.path.expanduser("~") - for value in values: - expanded = os.path.expanduser(value) - if not os.path.isabs(expanded): - expanded = os.path.join(home, expanded) - normalized.add(os.path.normcase(os.path.abspath(expanded))) - return sorted(normalized) + return sorted({ + os.path.normcase(os.path.abspath(os.path.join(home, os.path.expanduser(v)))) + for v in values}) canonical = { "enabled": bool(policy["enabled"]), "roots": _paths(policy["roots"]), "exclude_paths": _paths(policy["exclude_paths"])} @@ -226,11 +204,10 @@ def _repo_discovery_policy_is_default(policy: dict) -> bool: def _scan_discovered_repos_remote(conn, policy: dict) -> bool: - """Backend-side disk scan of the policy roots into the discovery cache (the desktop's - native scan only sees the local filesystem). Best-effort: failures log and leave - the cache untouched. Returns True only when the scan is authoritative (every root - walked to completion, cap not hit) — only then is the cache write ``replace=True``; - a partial/errored scan must MERGE, never wipe, or a failed refresh blanks the sidebar.""" + """Backend-side disk scan of the policy roots into the discovery cache. Best-effort: + failures log and leave the cache untouched. True only when the scan is authoritative + (every root walked to completion, cap not hit) — only then is the cache write + ``replace=True``; a partial/errored scan must MERGE, or a failed refresh blanks the sidebar.""" from hermes_cli import projects_db as pdb roots = policy.get("roots") or [] excludes = policy.get("exclude_paths") or [] @@ -239,11 +216,11 @@ def _scan_discovered_repos_remote(conn, policy: dict) -> bool: authoritative = True def _is_excluded(path: str) -> bool: - return any(path == ex or path.startswith(ex.rstrip("/\\") + os.sep) for ex in excludes if ex) + return any( + path == ex or path.startswith(ex.rstrip("/\\") + os.sep) for ex in excludes if ex) for root in roots: if not os.path.isdir(root): - # `os.walk` on a missing root yields nothing instead of raising; an unmounted - # volume would look like an empty scan and let the replace wipe its cache. + # `os.walk` on a missing root yields nothing; an unmounted volume must not wipe. authoritative = False logger.debug("discover_repos scan root missing, skipping: %s", root) continue @@ -264,8 +241,7 @@ def _scan_discovered_repos_remote(conn, policy: dict) -> bool: except Exception: authoritative = False logger.debug("discover_repos scan failed for root %s", root, exc_info=True) - if len(pairs) >= 500: - # Cap hit: the walk didn't cover the full roots -> not authoritative. + if len(pairs) >= 500: # cap hit: the walk didn't cover the full roots authoritative = False break if pairs: @@ -280,18 +256,16 @@ def _scan_discovered_repos_remote(conn, policy: dict) -> bool: def _discover_repos_payload( db, *, conn=None, backfill: bool = True, include_cached: bool = True) -> list[dict]: - """Merge filesystem-scanned repos (cached; may have zero sessions) with - session-derived roots, junk-filtered, with session totals. ``conn`` reuses an open - projects.db connection; ``backfill`` persists resolved roots onto session rows — - kept OFF the per-turn tree path and done only on explicit refresh.""" + """Merge cached filesystem-scanned repos with session-derived roots, junk-filtered, with + session totals. ``backfill`` persists resolved roots onto session rows — kept OFF the + per-turn tree path and done only on explicit refresh.""" repos: dict[str, dict] = {} def _agg(root: str) -> dict: - return repos.setdefault(root, {"root": root, "label": "", "sessions": 0, "last_active": 0.0}) - + return repos.setdefault( + root, {"root": root, "label": "", "sessions": 0, "last_active": 0.0}) cwd_rows = list(db.distinct_session_cwds()) - # Warm the per-cwd git probes in parallel so a cold first paint doesn't - # serialize one subprocess per distinct cwd before this loop reads the cache. + # Parallel-warm the per-cwd git probes so a cold first paint doesn't serialize them. git_probe.warm_roots(str(r.get("cwd") or "") for r in cwd_rows) cwd_to_root: dict[str, str] = {} for row in cwd_rows: @@ -311,8 +285,7 @@ def _discover_repos_payload( except Exception: logger.debug("failed to backfill repo roots", exc_info=True) if include_cached: - # `last_seen` is scan time, not user activity — never fold it into - # `last_active` (made every scanned repo "just now"). + # `last_seen` is scan time, not user activity — never fold it into `last_active`. try: from hermes_cli import projects_db as pdb with (contextlib.nullcontext(conn) if conn is not None else pdb.connect_closing()) as c: @@ -330,27 +303,23 @@ def _discover_repos_payload( return out -# Not user conversations (cron has its own section; kanban runs are read on -# the board). Subagent/compression children are dropped by include_children=False. +# Not user conversations; subagent/compression children are dropped by include_children=False. _PROJECT_TREE_EXCLUDED_SOURCES = ["cron", "kanban"] def _project_tree_row(r: dict) -> dict: - """Project a SessionDB row to the minimal shape the sidebar renders: the - grouping fields (cwd/git_branch/git_repo_root) + everything ``SidebarSessionRow`` - reads (parent_session_id for the └─ connector, cost for Show → cost), minus the - heavy columns.""" + """Project a SessionDB row to the minimal shape the sidebar renders (grouping fields + + what ``SidebarSessionRow`` reads), minus the heavy columns.""" row = {k: r.get(k) for k in ( "id", "_lineage_root_id", "_lineage_ids", "parent_session_id", "title", "preview")} row.update( started_at=r.get("started_at") or 0, ended_at=r.get("ended_at"), last_active=r.get("last_active") or r.get("started_at") or 0, - source=r.get("source"), archived=bool(r.get("archived"))) - row.update({k: r.get(k) or 0 for k in ( - "message_count", "tool_call_count", "input_tokens", "output_tokens")}) - row.update({k: r.get(k) for k in ("actual_cost_usd", "estimated_cost_usd", "model")}) - row["is_active"] = False - row.update({k: r.get(k) for k in ("cwd", "git_branch", "git_repo_root")}) + source=r.get("source"), archived=bool(r.get("archived")), + **{k: r.get(k) or 0 for k in ( + "message_count", "tool_call_count", "input_tokens", "output_tokens")}, + **{k: r.get(k) for k in ("actual_cost_usd", "estimated_cost_usd", "model")}, + is_active=False, **{k: r.get(k) for k in ("cwd", "git_branch", "git_repo_root")}) return row @@ -358,17 +327,15 @@ def _project_tree_inputs( db, session_limit: int, *, include_discovered: bool ) -> tuple[list[dict], list[dict], list[dict], str | None]: """Gather (sessions, projects, discovered_repos, active_id) for build_tree. - ``include_discovered`` is the zero-session-repo overview tier; drill-in skips it, - avoiding the distinct-cwd scan + git probes on that per-turn path.""" - # compact_rows: `_project_tree_row` drops the system-prompt blob; selecting it - # only to discard it costs tens of MB of B-tree reads per build on a big DB. + ``include_discovered`` is the zero-session-repo overview tier; drill-in skips it (and + the distinct-cwd scan + git probes) on that per-turn path.""" + # compact_rows: selecting the system-prompt blob only to drop it costs tens of MB of reads. rows = db.list_sessions_rich( limit=session_limit, offset=0, order_by_last_active=True, min_message_count=1, include_children=False, exclude_sources=_PROJECT_TREE_EXCLUDED_SOURCES, include_archived=False, compact_rows=True) sessions = [_project_tree_row(r) for r in rows] - # Parallel-warm the git cache so build_tree's resolver reads it instead of - # cold-probing each cwd in sequence (matters on the drill-in path). + # Parallel-warm the git cache so build_tree's resolver doesn't cold-probe each cwd in turn. git_probe.warm_roots(s["cwd"] for s in sessions if s.get("cwd")) from hermes_cli import projects_db as pdb policy = _repo_discovery_policy() @@ -387,19 +354,15 @@ def _project_tree_inputs( return sessions, projects, discovered, active_id -# Per-build memo for `_dir_exists_cached`; cleared at the top of every -# `_build_project_tree` so a dir created/deleted between refreshes is seen. +# Per-build memo for `_dir_exists_cached`; cleared by every `_build_project_tree`. _DIR_EXISTS_CACHE: dict[str, bool] = {} def _dir_exists_cached(path: str) -> bool: - """``os.path.isdir`` memoized per build — ``build_tree`` asks per SESSION, not - per distinct path, so hundreds of sessions in a few dirs would otherwise fire - hundreds of redundant stats per sidebar open.""" + """``os.path.isdir`` memoized per build — ``build_tree`` asks per SESSION, not per path.""" hit = _DIR_EXISTS_CACHE.get(path) if hit is None: - hit = os.path.isdir(path) - _DIR_EXISTS_CACHE[path] = hit + hit = _DIR_EXISTS_CACHE[path] = os.path.isdir(path) return hit @@ -411,8 +374,7 @@ def _build_project_tree( _DIR_EXISTS_CACHE.clear() sessions, projects, discovered, active_id = _project_tree_inputs( db, session_limit, include_discovered=include_discovered) - # build_tree also resolves every declared project folder and discovered repo - # root — not session cwds, so warm them too or they probe git one at a time. + # build_tree also resolves declared project folders and discovered roots — warm them too. git_probe.warm_roots( [str(f.get("path") or "") for p in projects for f in (p.get("folders") or [])] + [str(r.get("root") or "") for r in discovered]) diff --git a/tui_gateway/project_tree.py b/tui_gateway/project_tree.py index 31c6dcaffb..126b99f539 100644 --- a/tui_gateway/project_tree.py +++ b/tui_gateway/project_tree.py @@ -12,11 +12,10 @@ from __future__ import annotations import re from typing import Any, Callable, Optional -# cwd -> ``{"repo_root", "worktree_root"}`` (COMMON main root shared across worktrees / -# this cwd's own checkout root); ``None`` when not in git or unprobeable (remote backend). +# cwd -> ``{"repo_root", "worktree_root"}`` (COMMON main root / this cwd's checkout root); +# ``None`` when not in git or unprobeable (remote backend). Resolve = Callable[[str], Optional[dict]] -# "does this directory still exist?"; defaults to always-True so callers that can't -# stat (remote backends) don't hide a project living on the other host. +# "does this directory still exist?"; always-True default keeps remote-host projects visible. Exists = Callable[[str], bool] # Only KANBAN-TASK worktrees (`/.worktrees/t_`, the id kanban_db mints) @@ -25,23 +24,21 @@ _KANBAN_DIR_RE = re.compile(r"^(.*[/\\]\.worktrees)[/\\]t_[0-9a-f]+[/\\]?$") _TRUNK_BRANCHES = {"main", "master", "trunk", "develop"} DEFAULT_BRANCH_LABEL = "main" -# Synthetic bucket for every session no project claimed (no cwd, bare home, HERMES -# state, deleted workspace); the id/flag name what it MEANS since membership keys off them. +# Synthetic bucket for every session no project claimed (no cwd, bare home, deleted workspace). NO_PROJECT_ID = "__no_project__" NO_PROJECT_LABEL = "Home" -# Sibling probes when recovering a deleted worktree's parent repo; each miss is a git -# invocation and real suffixes are one or two segments. +# Sibling probes when recovering a deleted worktree's parent repo (each miss is a git call). _MAX_SIBLING_PROBES = 4 def stamp_profile(projects: list[dict], profile: str) -> None: - """Stamp every session row with the request-scope profile (authoritative even - for legacy rows whose ``profile_name`` is NULL) for cross-profile routing.""" + """Stamp every session row with the request-scope profile (authoritative even for legacy + rows whose ``profile_name`` is NULL) for cross-profile routing.""" for project in projects: lanes = [g for repo in project.get("repos") or [] for g in repo.get("groups") or []] - lane_rows = [s for g in lanes for s in g.get("sessions") or []] - for session in (project.get("previewSessions") or []) + lane_rows: + for session in (project.get("previewSessions") or []) + [ + s for g in lanes for s in g.get("sessions") or []]: session["profile"] = profile @@ -59,15 +56,14 @@ def _segments(path: str) -> list[str]: def _is_windows_path(path: str) -> bool: - # Drive-letter (`C:\…`), UNC (`\\srv`, `//srv`), or any backslash-rooted path - # (`\wsl.localhost\…`, `\Users\…`). A single leading `/` stays POSIX. + # Drive-letter, UNC (`\\srv`, `//srv`) or backslash-rooted; a single leading `/` stays POSIX. value = (path or "").strip() return bool(re.match(r"^[A-Za-z]:[/\\]", value)) or value.startswith(("\\", "//")) def _comparison_segments(path: str) -> list[str]: - """Segments for identity comparison: Windows paths casefold (even when running - on POSIX); display paths and emitted IDs keep their spelling.""" + """Segments for identity comparison: Windows paths casefold (even on POSIX); display + paths and emitted IDs keep their spelling.""" segs = _segments(path) return [s.casefold() for s in segs] if _is_windows_path(path) else segs @@ -78,8 +74,7 @@ def _path_key(path: str) -> str: def _lane_key(path_or_lane: str) -> str: - """Canonicalize only the path portion of a lane id; branch labels stay - byte-preserved so equivalent Windows spellings don't fork lanes.""" + """Canonicalize only the path portion of a lane id (branch labels stay byte-preserved).""" marker = next((m for m in ("::branch::", "::kanban") if m in path_or_lane), None) if marker is None: return _path_key(path_or_lane) @@ -98,25 +93,15 @@ def kanban_worktree_dir(path: str) -> Optional[str]: return m.group(1) if m else None -def _with_base_name(path: str, name: str) -> str: - return re.sub(r"[^/\\]+$", name, (path or "").rstrip("/\\")) - - -def _parent_dir(path: str) -> str: - """The containing directory of ``path`` (``""`` once the root is passed).""" - return _with_base_name(path, "").rstrip("/\\") +def _with_base_name(path: str, name: str = "") -> str: + """Swap the last segment for ``name`` (``""`` -> the parent dir, ``""`` past the root).""" + return re.sub(r"[^/\\]+$", name, (path or "").rstrip("/\\")).rstrip("" if name else "/\\") def _field(row: dict, key: str) -> str: return (row.get(key) or "").strip() -def _branch_label(branch: str) -> str: - # An unrecorded branch folds into the one trunk lane so a repo never shows two - # "main" lanes (recorded "main" + the empty-branch bucket). - return (branch or "").strip() or DEFAULT_BRANCH_LABEL - - def _session_time(session: dict) -> float: return float(session.get("last_active") or session.get("started_at") or 0) @@ -131,12 +116,12 @@ def _placement( return { "repo_key": repo_root, "repo_label": base_name(repo_root) or repo_root, "lane_key": lane_key, "lane_label": lane_label, "lane_path": lane_path, - "is_main": is_main, "is_kanban": is_kanban, - } + "is_main": is_main, "is_kanban": is_kanban} def _trunk_placement(repo_root: str, branch: str) -> dict: - b = _branch_label(branch) + # An unrecorded branch folds into the trunk lane so a repo never shows two "main" lanes. + b = (branch or "").strip() or DEFAULT_BRANCH_LABEL return _placement(repo_root, _branch_lane_id(repo_root, b), b, repo_root, True, False) @@ -145,13 +130,9 @@ def _kanban_placement(repo_root: str, kanban_dir: str) -> dict: def _probe_sibling_worktree(cwd: str, resolve: Resolve) -> str: - """The parent repo root of a deleted ``-`` worktree, else ``""``. - - A deleted dir can't be probed, so trim one ``-`` at a time off its name - and return the first sibling that resolves. The cwd is often a SUBDIR of the dead - worktree (``-/apps/desktop``), so the trim runs on each ANCESTOR, - deepest first. Probes are bounded in total (each is a git invocation). - """ + """The parent repo root of a deleted ``-`` worktree, else ``""``: trim one + ``-`` at a time off each ancestor's name (the cwd is often a SUBDIR of the dead + worktree), deepest first, returning the first sibling that resolves; probes are bounded.""" probes = 0 path = (cwd or "").rstrip("/\\") while path and probes < _MAX_SIBLING_PROBES: @@ -163,7 +144,7 @@ def _probe_sibling_worktree(cwd: str, resolve: Resolve) -> str: info = resolve(_with_base_name(path, "-".join(parts[:i]))) if info and info.get("repo_root"): return (info["repo_root"] or "").strip() - path = _parent_dir(path) + path = _with_base_name(path) return "" @@ -174,14 +155,15 @@ def _place_by_heuristic(path: str) -> Optional[dict]: return None kanban_dir = kanban_worktree_dir(path) if kanban_dir: - return _kanban_placement(_parent_dir(kanban_dir), kanban_dir) + return _kanban_placement(_with_base_name(kanban_dir), kanban_dir) m = re.match(r"^(.+)-wt-(.+)$", base) if m: return _placement(_with_base_name(path, m.group(1)), path, m.group(2), path, False, False) return _placement(path, _branch_lane_id(path, DEFAULT_BRANCH_LABEL), base, path, True, False) -def _place(cwd: str, branch: str, resolve: Optional[Resolve], persisted_root: str) -> Optional[dict]: +def _place( + cwd: str, branch: str, resolve: Optional[Resolve], persisted_root: str) -> Optional[dict]: info = resolve(cwd) if resolve else None if info and info.get("repo_root") and info.get("worktree_root"): repo_root, worktree_root = info["repo_root"], info["worktree_root"] @@ -193,16 +175,14 @@ def _place(cwd: str, branch: str, resolve: Optional[Resolve], persisted_root: st label = base_name(worktree_root) or worktree_root return _placement(repo_root, worktree_root, label, worktree_root, False, False) - # No live probe: trust the backend-persisted root (split main by the recorded - # branch). Kanban tasks still collapse by path shape. + # No live probe: trust the persisted root; kanban tasks still collapse by path shape. if persisted_root: kanban_dir = kanban_worktree_dir(cwd) if kanban_dir: return _kanban_placement(persisted_root, kanban_dir) return _trunk_placement(persisted_root, branch) - # Unresolvable cwd: a deleted ``-`` worktree still belongs to its - # parent; absorb it into the trunk lane rather than stranding a dead-path lane. + # Unresolvable cwd: a deleted ``-`` worktree still belongs to its parent. sibling_root = _probe_sibling_worktree(cwd, resolve) if resolve else "" if sibling_root: return _trunk_placement(sibling_root, branch) @@ -228,8 +208,7 @@ def _session_repo_root(session: dict, resolve: Optional[Resolve]) -> str: def _lane_sort_key(group: dict) -> tuple: - # Trunk pins to the top; the kanban aggregate sinks to the bottom; the rest - # (branches + linked worktrees) sort by most-recent activity, then label. + # Trunk pins to the top, the kanban aggregate to the bottom; the rest by recency, then label. is_trunk = bool(group.get("isMain")) and group["label"].lower() in _TRUNK_BRANCHES return (0 if is_trunk else 1, 1 if group.get("isKanban") else 0, -_last_active(group.get("sessions") or []), group["label"].lower()) @@ -240,7 +219,6 @@ def _disambiguate_labels(items: list[dict]) -> None: by_label: dict[str, list[dict]] = {} for item in items: by_label.setdefault(item["label"], []).append(item) - for bucket in by_label.values(): pathed = [g for g in bucket if g.get("path")] if len(pathed) < 2: @@ -277,7 +255,6 @@ def _build_repos(sessions: list[dict], resolve: Optional[Resolve], hydrate: bool group["sessions"] = [] lanes[lane_identity] = (group, placement) lanes[lane_identity][0]["sessions"].append(session) - repos: dict[str, dict] = {} for group, placement in lanes.values(): group["sessions"].sort(key=_session_time, reverse=True) @@ -285,13 +262,11 @@ def _build_repos(sessions: list[dict], resolve: Optional[Resolve], hydrate: bool repo = repos.setdefault(_path_key(repo_key), _repo_node(repo_key, placement["repo_label"])) repo["groups"].append(group) repo["sessionCount"] += len(group["sessions"]) - repo_list = list(repos.values()) for repo in repo_list: repo["groups"] = sorted(repo["groups"], key=_lane_sort_key) _disambiguate_labels(repo["groups"]) - # Drop per-lane rows only AFTER sorting: _lane_sort_key derives recency - # from them. Counts were captured above, so the overview payload stays slim. + # Drop per-lane rows only AFTER sorting (_lane_sort_key reads them); counts stay. if not hydrate: for group in repo["groups"]: group["sessions"] = [] @@ -299,11 +274,10 @@ def _build_repos(sessions: list[dict], resolve: Optional[Resolve], hydrate: bool return repo_list -def _seed_folder_repos(repos: list[dict], folders: list[dict], resolve: Optional[Resolve]) -> list[dict]: - """Ensure every declared project folder shows as a repo, even with 0 sessions: - otherwise the desktop's entered-project view renders blank and the optimistic - live-session overlay has no lane for a fresh session until a full refresh. - Folders already covered by a session-derived repo (same git root) are untouched.""" +def _seed_folder_repos( + repos: list[dict], folders: list[dict], resolve: Optional[Resolve]) -> list[dict]: + """Ensure every declared project folder shows as a repo, even with 0 sessions (else the + entered-project view renders blank); folders covered by a session-derived repo are untouched.""" seen = {_path_key(v) for repo in repos for v in (repo.get("id"), repo.get("path")) if v} seeded = list(repos) for folder in folders or []: @@ -323,8 +297,7 @@ def _seed_folder_repos(repos: list[dict], folders: list[dict], resolve: Optional class _FolderIndex: - """Normalized folder path -> (owning project, depth): a session is matched by - walking its cwd's ancestors instead of scanning every project x folder.""" + """Normalized folder path -> (owning project, depth); matched by walking cwd ancestors.""" def __init__(self, projects: list[dict]) -> None: self._by_path: dict[str, tuple[dict, int]] = {} @@ -345,7 +318,8 @@ class _FolderIndex: return None, -1 -def _project_for_session(session: dict, index: _FolderIndex, resolve: Optional[Resolve]) -> Optional[dict]: +def _project_for_session( + session: dict, index: _FolderIndex, resolve: Optional[Resolve]) -> Optional[dict]: cwd = _field(session, "cwd") if not cwd: return None @@ -355,73 +329,49 @@ def _project_for_session(session: dict, index: _FolderIndex, resolve: Optional[R return max((index.match(t) for t in candidates), key=lambda hit: hit[1])[0] -def _session_cost(session: dict) -> float: - """A session's spend, billed if the provider reported it, else estimated.""" - for key in ("actual_cost_usd", "estimated_cost_usd"): - if session.get(key): - return float(session[key]) - return 0.0 - - def _project_node( pid: str, label: str, path: Optional[str], repos: list[dict], session_count: int, last_active: float, preview_sessions: list[dict], sessions: Optional[list[dict]] = None, - **flags: Any, -) -> dict: - """``flags`` overrides ``color`` / ``icon`` / ``isAuto`` / ``isNoProject`` (key order is - fixed by the defaults below — the renderer's wire shape).""" + **flags: Any) -> dict: + """``flags`` overrides ``color``/``icon``/``isAuto``/``isNoProject``; key order = wire shape.""" + rows = sessions or [] node = { "id": pid, "label": label, "path": path, "color": None, "icon": None, "isAuto": False, "isNoProject": False, "sessionCount": session_count, "lastActive": last_active, - # Totals over the same sessions `sessionCount` counts, so a project header - # adds up to what its rows show. - "totalTokens": sum((s.get("input_tokens") or 0) + (s.get("output_tokens") or 0) for s in sessions or []), - "totalCostUsd": sum(_session_cost(s) for s in sessions or []), - "repos": repos, "previewSessions": preview_sessions, - } + # Totals over the same sessions `sessionCount` counts (billed cost, else estimated). + "totalTokens": sum( + (s.get("input_tokens") or 0) + (s.get("output_tokens") or 0) for s in rows), + "totalCostUsd": sum( + float(s.get("actual_cost_usd") or s.get("estimated_cost_usd") or 0) for s in rows), + "repos": repos, "previewSessions": preview_sessions} node.update(flags) return node def _auto_buckets( unowned: list[dict], resolve: Optional[Resolve], junk: Callable, junk_cwd: Callable, - exists: Callable, -) -> tuple[dict[str, dict], list[dict]]: - """Group leftover sessions by auto-project root; the rest go to the Home bucket. - Prefer the common git root, then the session cwd for non-git workspaces (the - pre-Projects desktop grouped every cwd; dropping that flattens them into Recents).""" + exists: Callable) -> tuple[dict[str, dict], list[dict]]: + """Group leftover sessions by auto-project root (common git root, else the session cwd + for non-git workspaces); the rest go to the Home bucket.""" by_auto_root: dict[str, dict] = {} homeless: list[dict] = [] - - def _add_auto(root: str, session: dict) -> None: - key = _path_key(root) - if not key: - homeless.append(session) - return - by_auto_root.setdefault(key, {"root": root, "sessions": []})["sessions"].append(session) - for session in unowned: root = _session_repo_root(session, resolve) if root: - # A real git root uses the stricter repo policy; never reinterpret a - # filtered internal repo as a cwd-only project. A root no longer on - # disk is a stale persisted value and must not resurrect as a project. - if not junk(root) and exists(root): - _add_auto(root, session) - else: - homeless.append(session) - continue - cwd = _field(session, "cwd") - if not cwd or junk_cwd(cwd): - homeless.append(session) - continue - placement = _place_session(session, resolve) - # A placement that only echoes back an unresolvable cwd is the path-only - # heuristic guessing. If that dir is also gone from disk, promoting it - # mints a phantom project that can only be dismissed by hand -> Home. - if placement and exists(placement["repo_key"]): - _add_auto(placement["repo_key"], session) + # Stricter repo policy for real git roots; a root gone from disk is stale and + # must not resurrect as a project (never reinterpret it as a cwd-only project). + if junk(root) or not exists(root): + root = "" + elif (cwd := _field(session, "cwd")) and not junk_cwd(cwd): + # A path-only heuristic placement whose dir is gone from disk would mint a phantom + # project that can only be dismissed by hand -> Home. + placement = _place_session(session, resolve) + if placement and exists(placement["repo_key"]): + root = placement["repo_key"] + key = _path_key(root) if root else "" + if key: + by_auto_root.setdefault(key, {"root": root, "sessions": []})["sessions"].append(session) else: homeless.append(session) return by_auto_root, homeless @@ -431,35 +381,26 @@ def _home_project(homeless: list[dict], hydrate: bool, previews: list[dict]) -> """The synthetic Home bucket: no folder => no repo/lane structure, one lane carries the rows.""" lane = { "id": NO_PROJECT_ID, "label": NO_PROJECT_LABEL, "path": None, "isMain": False, - "isKanban": False, "sessions": homeless if hydrate else [], - } + "isKanban": False, "sessions": homeless if hydrate else []} home_repo = { "id": NO_PROJECT_ID, "label": NO_PROJECT_LABEL, "path": None, "groups": [lane], - "sessionCount": len(homeless), - } + "sessionCount": len(homeless)} return _project_node( NO_PROJECT_ID, NO_PROJECT_LABEL, None, [home_repo], len(homeless), _last_active(homeless), previews, homeless, isNoProject=True) def build_tree( - projects: list[dict], - sessions: list[dict], - discovered_repos: list[dict], - resolve: Optional[Resolve] = None, - *, - preview_limit: int = 3, - hydrate: bool = False, + projects: list[dict], sessions: list[dict], discovered_repos: list[dict], + resolve: Optional[Resolve] = None, *, preview_limit: int = 3, hydrate: bool = False, is_junk_root: Optional[Callable[[str], bool]] = None, - is_junk_cwd: Optional[Callable[[str], bool]] = None, - exists: Optional[Exists] = None) -> dict: + is_junk_cwd: Optional[Callable[[str], bool]] = None, exists: Optional[Exists] = None) -> dict: """Build the authoritative project tree -> ``{"projects", "scoped_session_ids"}``. - ``is_junk_root`` flags git roots that must never become an AUTO project (bare home, - HERMES_HOME); ``is_junk_cwd`` is the narrower policy for non-git folders; explicit - projects are honored regardless. ``exists`` keeps a DELETED workspace from becoming - a phantom AUTO project (omit on remote backends). ``hydrate`` False (overview) empties - lane ``sessions`` but keeps counts + ``preview_limit`` ``previewSessions``. + ``is_junk_root`` flags git roots that must never become an AUTO project; ``is_junk_cwd`` + is the narrower non-git policy (explicit projects are honored regardless); ``exists`` + keeps a DELETED workspace from becoming a phantom AUTO project (omit on remote backends). + ``hydrate`` False empties lane ``sessions`` but keeps counts + ``previewSessions``. """ active_projects = [p for p in projects if not p.get("archived")] _junk = is_junk_root or (lambda _root: False) @@ -482,7 +423,6 @@ def build_tree( def _scope(project_sessions: list[dict]) -> None: scoped_ids.extend(s["id"] for s in project_sessions if s.get("id")) - # Tier 1: explicit, user-created projects (always shown, even with 0 sessions). for project in active_projects: psessions = by_project.get(project["id"], []) @@ -513,8 +453,7 @@ def build_tree( repo_node["sessionCount"], _last_active(auto_sessions), _previews(auto_sessions), auto_sessions, isAuto=True)) - # Tier 3: repos discovered from full history / disk scan with no loaded - # sessions, folded to their common root and not owned by an explicit project. + # Tier 3: discovered repos with no loaded sessions, folded to their common root. for repo in discovered_repos or []: raw_root = _field(repo, "root") if not raw_root: @@ -530,12 +469,10 @@ def build_tree( root, label, root, [_repo_node(root, label)], int(repo.get("sessions") or 0), float(repo.get("last_active") or 0), [], isAuto=True)) - # Auto projects are labelled by repo basename, which can collide; grow path - # prefixes so each is distinct. Explicit projects keep their user-chosen names. + # Auto-project basename labels can collide; explicit projects keep their user-chosen names. _disambiguate_labels([p for p in result if p.get("isAuto")]) - # Tier 0: everything above could not place, so the grouped view loses no - # session. Leads the list; omitted entirely when empty. + # Tier 0: whatever the tiers above could not place. Leads the list; omitted when empty. if homeless: homeless.sort(key=_session_time, reverse=True) _scope(homeless)