diff --git a/tests/tui_gateway/test_compute_host.py b/tests/tui_gateway/test_compute_host.py index fa0019722f..f93b571946 100644 --- a/tests/tui_gateway/test_compute_host.py +++ b/tests/tui_gateway/test_compute_host.py @@ -6,10 +6,6 @@ import sys import threading from pathlib import Path -import pytest - -from tui_gateway.compute_host import ComputeHost, HostSession - def _stdout_queue(proc: subprocess.Popen) -> queue.Queue[dict]: out: queue.Queue[dict] = queue.Queue() @@ -30,7 +26,7 @@ def _read_json_line(out: queue.Queue[dict], timeout: float = 2.0) -> dict: raise AssertionError("timed out waiting for compute host JSON") from exc -def test_compute_host_line_json_seed_turn_interrupt(): +def test_compute_host_line_json_hello_and_shutdown(): repo = Path(__file__).resolve().parents[2] env = dict(os.environ) env["PYTHONPATH"] = str(repo) + os.pathsep + env.get("PYTHONPATH", "") @@ -51,34 +47,11 @@ def test_compute_host_line_json_seed_turn_interrupt(): assert hello["type"] == "hello" assert hello["host_pid"] == proc.pid - proc.stdin.write(json.dumps({"type": "session.seed", "sid": "s1", "request_id": "seed"}) + "\n") + proc.stdin.write(json.dumps({"type": "bogus", "request_id": "b"}) + "\n") proc.stdin.flush() - assert _read_json_line(out)["type"] == "session.seeded" - - proc.stdin.write( - json.dumps( - { - "type": "turn.start", - "sid": "s1", - "request_id": "turn", - "prompt": "hello", - "delta_count": 3, - "delay_s": 0, - } - ) - + "\n" - ) - proc.stdin.flush() - - seen = [] - while True: - frame = _read_json_line(out) - seen.append(frame["type"]) - if frame["type"] == "turn.end": - assert frame["history_version"] == 1 - assert frame["message_count"] == 2 - break - assert seen.count("delta") == 3 + error = _read_json_line(out) + assert error["type"] == "error" + assert error["message"] == "unknown frame type: bogus" proc.stdin.write(json.dumps({"type": "shutdown", "request_id": "stop"}) + "\n") proc.stdin.flush() @@ -87,42 +60,3 @@ def test_compute_host_line_json_seed_turn_interrupt(): finally: if proc.poll() is None: proc.kill() - - -@pytest.mark.parametrize("kind", ["legacy", "hard-only", "dynamic-getattr"]) -def test_compute_host_interrupt_uses_explicit_stop_compatibility(kind): - calls = [] - - class _Legacy: - def interrupt(self): - calls.append("legacy") - - class _HardOnly: - def hard_interrupt(self): - calls.append("hard") - - class _Dynamic: - def interrupt(self): - calls.append("legacy") - - def __getattr__(self, name): - if name == "hard_interrupt": - return lambda: calls.append("fabricated-hard") - raise AttributeError(name) - - agent = { - "legacy": _Legacy(), - "hard-only": _HardOnly(), - "dynamic-getattr": _Dynamic(), - }[kind] - host = ComputeHost(heartbeat_secs=0) - host._sessions["s1"] = HostSession(sid="s1", agent=agent) - emitted = [] - host.emit = emitted.append - try: - host._handle_interrupt({"sid": "s1", "request_id": "stop"}) - finally: - host.close() - - assert calls == ["hard" if kind == "hard-only" else "legacy"] - assert emitted[-1]["applied"] is True diff --git a/tui_gateway/compute_host.py b/tui_gateway/compute_host.py index 8b5205809a..0f3d49cf40 100644 --- a/tui_gateway/compute_host.py +++ b/tui_gateway/compute_host.py @@ -1,8 +1,7 @@ """Persistent dashboard compute-host process. -Phase 0 used this module as a deterministic line-JSON spike. Phase 1 keeps the -same transport and turns it into the long-lived child that owns live AIAgent -objects when ``dashboard.turn_isolation`` is enabled. +The long-lived child that owns live AIAgent objects when ``dashboard.turn_isolation`` +is enabled; frames are line-JSON over stdin/stdout. """ from __future__ import annotations @@ -18,11 +17,9 @@ import sys import threading import time import uuid -from dataclasses import dataclass, field from pathlib import Path from typing import Any, Callable, Collection -from agent.interrupt_compat import request_hard_interrupt from tui_gateway.host_supervisor import MUTATOR_ROUTE_TABLE, _build_sha @@ -30,61 +27,6 @@ def now_ns() -> int: return time.perf_counter_ns() -@dataclass -class SpikeAgent: - """A deterministic AIAgent-shaped object for pipe/interrupt measurements.""" - - session_id: str - history: list[dict[str, str]] = field(default_factory=list) - _interrupt: threading.Event = field(default_factory=threading.Event) - - def clear_interrupt(self) -> None: - self._interrupt.clear() - - def interrupt(self, *, hard_cancel: bool = False) -> None: - self._interrupt.set() - - def run_conversation( - self, - prompt: str, - *, - conversation_history: list[dict[str, str]] | None = None, - stream_callback: Callable[[str], None] | None = None, - delta_count: int = 24, - delay_s: float = 0.001, - ) -> dict[str, Any]: - base_history = list(conversation_history if conversation_history is not None else self.history) - chunks: list[str] = [] - interrupted = False - for index in range(max(0, int(delta_count))): - if self._interrupt.is_set(): - interrupted = True - break - chunk = f"{self.session_id}:{prompt}:{index:04d} " - chunks.append(chunk) - if stream_callback is not None: - stream_callback(chunk) - if delay_s > 0: - time.sleep(delay_s) - if self._interrupt.is_set(): - interrupted = True - final = "".join(chunks) - if interrupted: - final += "[interrupted]" - messages = [*base_history, {"role": "user", "content": prompt}, {"role": "assistant", "content": final}] - self.history = messages - return {"final_response": final, "messages": messages, "interrupted": interrupted} - - -@dataclass -class HostSession: - sid: str - agent: SpikeAgent - history_version: int = 0 - running: bool = False - lock: threading.Lock = field(default_factory=threading.Lock) - - class _HostTransport: def __init__(self, emit: Callable[[dict[str, Any]], None]) -> None: self._emit = emit @@ -101,11 +43,10 @@ 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. +# ``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. _FLUSH_RESERVE_SECS = 1.0 @@ -113,48 +54,35 @@ class ComputeHost: # frame ``type`` -> handler method name (resolved per call so instance # monkeypatches of a handler still take effect). _FRAME_HANDLERS: dict[str, str] = { - "session.seed": "_handle_seed", - "turn.start": "_handle_turn_start", - "interrupt": "_handle_interrupt", - "respond": "_handle_respond", - "reload_mcp": "_handle_reload_mcp", - "control": "_handle_control", - "shutdown": "_handle_shutdown", - } + "turn.start": "_handle_turn_start", "interrupt": "_handle_interrupt", + "respond": "_handle_respond", "reload_mcp": "_handle_reload_mcp", + "control": "_handle_control", "shutdown": "_handle_shutdown"} def __init__( - self, - *, - stdout: Any = None, - max_workers: int | None = None, - heartbeat_secs: int | float | None = None, - ) -> None: + self, *, stdout: Any = None, max_workers: int | None = None, + heartbeat_secs: int | float | None = None) -> None: self._stdout = stdout or sys.stdout self._write_lock = threading.Lock() - self._sessions: dict[str, HostSession] = {} self._executor = concurrent.futures.ThreadPoolExecutor( - max_workers=max_workers or _default_workers(), - thread_name_prefix="compute-host-turn", - ) + max_workers=max_workers or _default_workers(), thread_name_prefix="compute-host-turn") self._closed = threading.Event() self._parent_pid = os.getppid() 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; a bare set cannot answer that. + # Future -> the ``sid`` whose turn it is running. ``shutdown`` needs to know + # *whose* turn is still live so it can leave those sessions unfinalized. self._turn_futures: dict[concurrent.futures.Future, str] = {} self._turn_futures_lock = threading.Lock() self._transport = _HostTransport(self.emit) self._heartbeat_secs = ( - float(heartbeat_secs) - if heartbeat_secs is not None - else float(os.environ.get("HERMES_COMPUTE_HOST_HEARTBEAT_SECS") or "15") - ) + float(heartbeat_secs) if heartbeat_secs is not None + else float(os.environ.get("HERMES_COMPUTE_HOST_HEARTBEAT_SECS") or "15")) if self._heartbeat_secs > 0: - threading.Thread(target=self._heartbeat_loop, name="compute-host-heartbeat", daemon=True).start() - threading.Thread(target=self._parent_guard_loop, name="compute-host-ppid-guard", daemon=True).start() + for target, name in ( + (self._heartbeat_loop, "compute-host-heartbeat"), + (self._parent_guard_loop, "compute-host-ppid-guard")): + threading.Thread(target=target, name=name, daemon=True).start() def emit(self, frame: dict[str, Any]) -> None: frame.setdefault("host_ns", now_ns()) @@ -162,6 +90,10 @@ class ComputeHost: with self._write_lock: print(data, file=self._stdout, flush=True) + def _reply(self, kind: str, sid: str, request_id: Any, **extra: Any) -> None: + """Emit a per-session frame keyed by the request it answers.""" + self.emit({"type": kind, "sid": sid, "request_id": request_id, **extra}) + def close(self) -> None: self._closed.set() self._executor.shutdown(wait=False, cancel_futures=True) @@ -169,29 +101,17 @@ 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 (sets - ``session["_finalized"]``), so finalizing before the drain would spend - the flush's single chance while turns were still producing output, - fire ``on_session_end(interrupted=True)`` against a running session and - release the active-session lease under a live turn. - - ``_FLUSH_RESERVE_SECS`` of the budget (never more than half, so a short - explicit ``wait`` still gets a real drain) is withheld from the drain so - the flush still runs when in-flight turns outlast the window; ``wait`` - itself is unchanged, so no added latency or kill-escalation exposure. - - Sessions whose turn is *still running* at the drain deadline are - excluded from the flush: finalizing one mid-turn (``_executor.shutdown`` - below does not join it) would leave it permanently un-finalizable with - its lease released — the very race the drain closes. Leaving them - unfinalized keeps them recoverable. - - NOTE: ``server._shutdown_sessions`` (atexit) runs after ``shutdown()`` - returns on the SIGTERM / stdin_closed paths and may re-finalize skipped - sessions still in ``server._sessions``; the orphan path (``os._exit``) - bypasses atexit. Pre-existing gap, not worsened by this ordering; a - follow-up could gate ``_shutdown_sessions`` on - ``not _finalized and not running``. + Order matters: ``_finalize_session`` is a one-shot latch, so finalizing before + the drain would spend the flush's single chance mid-turn, fire + ``on_session_end(interrupted=True)`` against a running session and release the + active-session lease under a live turn. ``_FLUSH_RESERVE_SECS`` (never more than + half of ``wait``) is withheld from the drain so the flush still runs when turns + outlast the window. Sessions whose turn is *still running* at the deadline are + excluded from the flush (``_executor.shutdown`` does not join them): finalizing + one mid-turn would leave it un-finalizable with its lease released, whereas + leaving it unfinalized keeps it recoverable. ``server._shutdown_sessions`` + (atexit) may still re-finalize skipped sessions on the SIGTERM / stdin_closed + paths; the orphan path (``os._exit``) bypasses atexit. """ self._closed.set() budget = max(0.0, wait) @@ -204,28 +124,23 @@ class ComputeHost: pending = [f for f in self._turn_futures if not f.done()] if not pending: break - # Bounded by ``remaining``: a flat 0.05s sleep would overshoot the - # deadline and eat the reserve it protects (all of it for small ``wait``). + # Bounded by ``remaining``: a flat sleep would overshoot the deadline and + # eat the reserve it protects (all of it for small ``wait``). time.sleep(min(0.05, remaining)) with self._turn_futures_lock: - live_sids = {sid for future, sid in self._turn_futures.items() if sid and not future.done()} + live_sids = {sid for f, sid in self._turn_futures.items() if sid and not f.done()} self.flush_all_sessions(reason=reason, skip_sids=live_sids) self._executor.shutdown(wait=False, cancel_futures=True) def flush_all_sessions( - self, - *, - reason: str = "shutdown", - skip_sids: Collection[str] | None = None, - ) -> None: - """Finalize every server session except ``skip_sids`` (sessions whose - turn is still live and must not spend their one-shot finalize).""" + self, *, reason: str = "shutdown", skip_sids: Collection[str] | None = None) -> None: + """Finalize every server session except ``skip_sids`` (turn still live).""" try: from tui_gateway import server except Exception: return skip = set(skip_sids or ()) - for sid, session in list(getattr(server, "_sessions", {}).items()): + for sid, session in list(server._sessions.items()): if sid in skip: continue with contextlib.suppress(Exception): @@ -235,7 +150,9 @@ class ComputeHost: kind = str(frame.get("type") or "") handler = self._FRAME_HANDLERS.get(kind) if handler is None: - self.emit({"type": "error", "request_id": frame.get("request_id"), "message": f"unknown frame type: {kind}"}) + self.emit({ + "type": "error", "request_id": frame.get("request_id"), + "message": f"unknown frame type: {kind}"}) return getattr(self, handler)(frame) @@ -246,22 +163,9 @@ class ComputeHost: self._closed.set() self._executor.shutdown(wait=False, cancel_futures=True) - # ── Phase-0 deterministic spike frames ───────────────────────────── - - def _handle_seed(self, frame: dict[str, Any]) -> None: - sid = str(frame.get("sid") or "") - if not sid: - self.emit({"type": "error", "request_id": frame.get("request_id"), "message": "sid required"}) - return - history = frame.get("history") - if not isinstance(history, list): - history = [] - self._sessions[sid] = HostSession(sid=sid, agent=SpikeAgent(sid, list(history))) - self.emit({"type": "session.seeded", "sid": sid, "request_id": frame.get("request_id")}) - def _track_turn_future(self, future: concurrent.futures.Future, sid: str) -> None: - """Register an in-flight turn against the session running it; the done - callback must pop under the lock or the mapping grows for the host's life.""" + """Register an in-flight turn against its session; the done callback must pop + under the lock or the mapping grows for the host's life.""" with self._turn_futures_lock: self._turn_futures[future] = sid future.add_done_callback(self._untrack_turn_future) @@ -271,49 +175,25 @@ class ComputeHost: self._turn_futures.pop(future, None) def _handle_turn_start(self, frame: dict[str, Any]) -> None: - sid = str(frame.get("sid") or "") - if sid in self._sessions: - self._handle_spike_turn_start(frame) - return future = self._executor.submit(self._run_real_turn, dict(frame)) - self._track_turn_future(future, sid) - - def _handle_spike_turn_start(self, frame: dict[str, Any]) -> None: - sid = str(frame.get("sid") or "") - session = self._sessions.get(sid) - if session is None: - self.emit({"type": "turn.error", "sid": sid, "request_id": frame.get("request_id"), "message": "unknown session"}) - return - with session.lock: - if session.running: - self.emit({"type": "turn.error", "sid": sid, "request_id": frame.get("request_id"), "message": "session busy"}) - return - session.running = True - future = self._executor.submit(self._run_spike_turn, session, dict(frame)) - self._track_turn_future(future, sid) + self._track_turn_future(future, str(frame.get("sid") or "")) def _handle_interrupt(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") request_id = frame.get("request_id") - spike = self._sessions.get(sid) - if spike is not None: - request_hard_interrupt(spike.agent) - self.emit({"type": "interrupt.ack", "sid": sid, "request_id": request_id, "applied": True, "applied_ns": now_ns()}) - return try: from tui_gateway import server - session = server._sessions.get(sid) if session is None: - self.emit({"type": "interrupt.ack", "sid": sid, "request_id": request_id, "applied": False}) + 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, `_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. server._interrupt_session_turn(sid, session) - self.emit({"type": "interrupt.ack", "sid": sid, "request_id": request_id, "applied": True, "applied_ns": now_ns()}) + self._reply("interrupt.ack", sid, request_id, applied=True, applied_ns=now_ns()) except Exception as exc: - self.emit({"type": "interrupt.ack", "sid": sid, "request_id": request_id, "applied": False, "message": str(exc)}) + self._reply("interrupt.ack", sid, request_id, applied=False, message=str(exc)) def _handle_respond(self, frame: dict[str, Any]) -> None: """Resolve an interactive request in the host-owned pending registry.""" @@ -321,90 +201,53 @@ class ComputeHost: request_id = frame.get("request_id") try: from tui_gateway import server - if sid not in server._sessions: - self.emit({"type": "respond.error", "sid": sid, "request_id": request_id, "message": "session not found"}) + self._reply("respond.error", sid, request_id, message="session not found") return params = frame.get("params") if not isinstance(params, dict): - self.emit({"type": "respond.error", "sid": sid, "request_id": request_id, "message": "response params must be an object"}) + self._reply( + "respond.error", sid, request_id, message="response params must be an object") return response = server._methods["clarify.respond"](request_id, params) - self.emit({"type": "respond.ack", "sid": sid, "request_id": request_id, "response": response}) + self._reply("respond.ack", sid, request_id, response=response) except Exception as exc: - self.emit({"type": "respond.error", "sid": sid, "request_id": request_id, "message": str(exc)}) - - def _run_spike_turn(self, session: HostSession, frame: dict[str, Any]) -> None: - request_id = frame.get("request_id") or uuid.uuid4().hex - prompt = str(frame.get("prompt") or frame.get("text") or "") - delta_count = _coerce(frame.get("delta_count", 24), int, 24) - delay_s = _coerce(frame.get("delay_s", 0.001), float, 0.001) - with session.lock: - history = list(session.agent.history) - session.agent.clear_interrupt() - self.emit({"type": "turn.started", "sid": session.sid, "request_id": request_id, "started_ns": now_ns()}) - - def stream(delta: str) -> None: - self._bump_progress() - self.emit({"type": "delta", "sid": session.sid, "request_id": request_id, "text": delta, "emitted_ns": now_ns()}) - - try: - result = session.agent.run_conversation( - prompt, - conversation_history=history, - stream_callback=stream, - delta_count=delta_count, - delay_s=delay_s, - ) - with session.lock: - session.history_version += 1 - session.running = False - history_version = session.history_version - self._bump_progress() - self.emit( - {"type": "turn.end", "sid": session.sid, "request_id": request_id, "history_version": history_version, - "message_count": len(result.get("messages") or []), "interrupted": bool(result.get("interrupted")), "ended_ns": now_ns()} - ) - except Exception as exc: # pragma: no cover - defensive host boundary - with session.lock: - session.running = False - self.emit({"type": "turn.error", "sid": session.sid, "request_id": request_id, "message": str(exc)}) - - # ── Real dashboard turn path ─────────────────────────────────────── + self._reply("respond.error", sid, request_id, message=str(exc)) def _run_real_turn(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") request_id = str(frame.get("request_id") or uuid.uuid4().hex) if not sid: - self.emit({"type": "turn.error", "sid": sid, "request_id": request_id, "message": "sid required"}) + self._reply("turn.error", sid, request_id, message="sid required") return 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", "") with session["history_lock"]: queued_gen = frame.get("queued_prompt_generation") - if queued_gen is not None and int(session.get("_queued_prompt_generation", 0)) != int(queued_gen): - self.emit({"type": "turn.end", "sid": sid, "request_id": request_id, "interrupted": True, "ended_ns": now_ns()}) + current_gen = int(session.get("_queued_prompt_generation", 0)) + if queued_gen is not None and current_gen != int(queued_gen): + self._reply("turn.end", sid, request_id, interrupted=True, ended_ns=now_ns()) return if session.get("running"): - self.emit({"type": "turn.error", "sid": sid, "request_id": request_id, "message": "session busy"}) + self._reply("turn.error", sid, request_id, message="session busy") return session["running"] = True session["_turn_cancel_requested"] = False session["last_active"] = time.time() - server._start_inflight_turn(session, frame.get("text") if "text" in frame else frame.get("prompt")) - self.emit({"type": "turn.started", "sid": sid, "request_id": request_id, "started_ns": now_ns()}) + server._start_inflight_turn( + session, frame.get("text") if "text" in frame else frame.get("prompt")) + self._reply("turn.started", sid, request_id, started_ns=now_ns()) with contextlib.suppress(Exception): server._ensure_session_db_row(session) with contextlib.suppress(Exception): import hermes_undo - hermes_undo.on_user_message_appended(session["session_key"]) with contextlib.suppress(Exception): server._persist_branch_seed(session) - server._run_prompt_submit(request_id, sid, session, text, display_kind=frame.get("display_kind") or None) + server._run_prompt_submit( + request_id, sid, session, text, display_kind=frame.get("display_kind") or None) run_thread = session.get("_run_thread") if run_thread is not None and hasattr(run_thread, "join"): run_thread.join() @@ -415,20 +258,19 @@ class ComputeHost: session_key = str(session.get("session_key") or "") session_info = server._session_info(session.get("agent"), session) self._bump_progress() - self.emit( - {"type": "turn.end", "sid": sid, "request_id": request_id, "history_version": history_version, "session_key": session_key, - "message_count": message_count, "interrupted": interrupted, "ended_ns": now_ns(), "session_info": session_info, "session_info_emitted": True} - ) + self._reply( + "turn.end", sid, request_id, history_version=history_version, + session_key=session_key, message_count=message_count, interrupted=interrupted, + ended_ns=now_ns(), session_info=session_info, session_info_emitted=True) except Exception as exc: with contextlib.suppress(Exception): from tui_gateway import server - session = server._sessions.get(sid) if session is not None: with session.get("history_lock", threading.Lock()): session["running"] = False server._clear_inflight_turn(session) - self.emit({"type": "turn.error", "sid": sid, "request_id": request_id, "reason": "exception", "message": str(exc)}) + self._reply("turn.error", sid, request_id, reason="exception", message=str(exc)) def _ensure_server_session(self, server: Any, frame: dict[str, Any]) -> dict: sid = str(frame.get("sid") or "") @@ -445,7 +287,6 @@ class ComputeHost: if isinstance(frame.get("attached_images"), list): session["attached_images"] = list(frame.get("attached_images") or []) return session - history = frame.get("history") if isinstance(frame.get("history"), list) else [] profile_home = str(frame.get("profile_home") or "") session_db = None @@ -457,85 +298,61 @@ class ComputeHost: from hermes_constants import set_hermes_home_override from agent.secret_scope import build_profile_secret_scope, set_secret_scope 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; 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``. session_db = get_shared_session_db(Path(profile_home) / "state.db") owns_db = True agent = server._make_agent( - sid, - key, - session_id=key, - model_override=frame.get("model_override"), + sid, key, session_id=key, model_override=frame.get("model_override"), reasoning_config_override=frame.get("reasoning_config_override"), service_tier_override=frame.get("service_tier_override"), platform_override=frame.get("source"), - context_cwd_is_launch_artifact=bool(frame.get("context_cwd_is_launch_artifact", False)), - session_db=session_db, - ) + context_cwd_is_launch_artifact=bool( + frame.get("context_cwd_is_launch_artifact", False)), + session_db=session_db) if server._transfer_db_to_agent(agent, session_db): owns_db = False finally: if owns_db and session_db is not None: with contextlib.suppress(Exception): from hermes_state import release_or_close - release_or_close(session_db) if home_token is not None: with contextlib.suppress(Exception): from hermes_constants import reset_hermes_home_override from agent.secret_scope import reset_secret_scope - reset_hermes_home_override(home_token) reset_secret_scope(secret_token) try: from tui_gateway.transport import bind_transport, reset_transport - token = bind_transport(self._transport) try: server._init_session( - sid, - key, - agent, - list(history), - cols=int(frame.get("cols") or 80), - cwd=str(frame.get("cwd") or "") or None, - session_db=session_db, - source=frame.get("source"), - ) + sid, key, agent, list(history), cols=int(frame.get("cols") or 80), + cwd=str(frame.get("cwd") or "") or None, session_db=session_db, + source=frame.get("source")) 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. + # unavailable, keep a minimal host-owned session rather than failing the + # turn after the expensive agent build succeeded. server._sessions[sid] = { - "agent": agent, - "session_key": key, - "history": list(history), + "agent": agent, "session_key": key, "history": list(history), "history_lock": threading.Lock(), - "history_version": int(frame.get("history_version") or 0), - "inflight_turn": None, - "created_at": time.time(), - "last_active": time.time(), - "running": False, - "attached_images": [], - "image_counter": 0, - "cwd": str(frame.get("cwd") or os.getcwd()), - "cols": int(frame.get("cols") or 80), - "slash_worker": None, - "show_reasoning": server._load_show_reasoning(), - "tool_progress_mode": server._load_tool_progress_mode(), - "edit_snapshots": {}, - "tool_started_at": {}, - "model_override": frame.get("model_override"), + "history_version": int(frame.get("history_version") or 0), "inflight_turn": None, + "created_at": time.time(), "last_active": time.time(), "running": False, + "attached_images": [], "image_counter": 0, + "cwd": str(frame.get("cwd") or os.getcwd()), "cols": int(frame.get("cols") or 80), + "slash_worker": None, "show_reasoning": server._load_show_reasoning(), + "tool_progress_mode": server._load_tool_progress_mode(), "edit_snapshots": {}, + "tool_started_at": {}, "model_override": frame.get("model_override"), "source": server._sanitize_client_source(frame.get("source")), - "transport": self._transport, - } + "transport": self._transport} session = server._sessions[sid] session["transport"] = self._transport session["profile_home"] = profile_home or session.get("profile_home") @@ -550,11 +367,12 @@ class ComputeHost: request_id = frame.get("request_id") try: from tui_gateway import server - - resp = server.handle_request({"id": request_id, "method": "reload.mcp", "params": {"session_id": sid, "confirm": True}}) - self.emit({"type": "reload_mcp.ack", "sid": sid, "request_id": request_id, "response": resp}) + 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.emit({"type": "control.error", "sid": sid, "request_id": request_id, "message": str(exc)}) + self._reply("control.error", sid, request_id, message=str(exc)) def _handle_control(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") @@ -562,10 +380,10 @@ class ComputeHost: route_name = str(frame.get("route_name") or "") def _error(message: str) -> None: - self.emit({"type": "control.error", "sid": sid, "request_id": request_id, "message": message}) + self._reply("control.error", sid, request_id, message=message) def _ack(**extra: Any) -> None: - self.emit({"type": "control.ack", "sid": sid, "request_id": request_id, "route_name": route_name, **extra}) + self._reply("control.ack", sid, request_id, route_name=route_name, **extra) def _call_method(name: str, params: dict[str, Any], failure: str) -> dict | None: """Run a server method; emit control.error and return None on error.""" @@ -575,9 +393,14 @@ class ComputeHost: return None return response + def _history_meta() -> dict[str, Any]: + """Ack metadata read under ``history_lock`` (caller holds it).""" + return { + "session_key": str(session.get("session_key") or ""), + "history_version": int(session.get("history_version", 0)), + "message_count": len(session.get("history") or [])} try: from tui_gateway import server - route = MUTATOR_ROUTE_TABLE.get(route_name) if route is None: _error(f"unclassified route: {route_name}") @@ -599,49 +422,38 @@ class ComputeHost: return if route_name == "session.compress": focus_topic = str(frame.get("command") or "").removeprefix("/compress").strip() - params = {"session_id": sid, **({"focus_topic": focus_topic} if focus_topic else {})} + params = {"session_id": sid} + if focus_topic: + params["focus_topic"] = focus_topic response = _call_method("session.compress", params, "session compression failed") if response is None: return with session["history_lock"]: - session_key = str(session.get("session_key") or "") - history_version = int(session.get("history_version", 0)) - message_count = len(session.get("history") or []) + meta = _history_meta() _ack( - result=response.get("result") or {}, - session_key=session_key, - history_version=history_version, - message_count=message_count, - session_info=server._session_info(session.get("agent"), session), - ) + result=response.get("result") or {}, **meta, + session_info=server._session_info(session.get("agent"), session)) return command = str(frame.get("command") or "") output = server._mirror_slash_side_effects(sid, session, command) if command else "" with session["history_lock"]: messages = server._history_to_messages(list(session.get("history") or [])) - history_version = int(session.get("history_version", 0)) - message_count = len(session.get("history") or []) - session_key = str(session.get("session_key") or "") + meta = _history_meta() _ack( - output=output, - session_key=session_key, - history_version=history_version, - message_count=message_count, - messages=messages, - session_info=server._session_info(session.get("agent"), session), - ) + output=output, session_key=meta["session_key"], + history_version=meta["history_version"], message_count=meta["message_count"], + messages=messages, session_info=server._session_info(session.get("agent"), session)) except Exception as exc: 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 or discarded it. + # 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. with contextlib.suppress(Exception): from tui_gateway import server as _server - from agent.conversation_compression import finalize_context_engine_compression_notification - + from agent.conversation_compression import ( + finalize_context_engine_compression_notification) _agent = (_server._sessions.get(sid) or {}).get("agent") if _agent is not None: finalize_context_engine_compression_notification(_agent, committed=False) @@ -657,7 +469,9 @@ class ComputeHost: active_turns = sum(1 for f in self._turn_futures if not f.done()) with self._progress_lock: counter = self._progress_counter - self.emit({"type": "hb", "active_turns": active_turns, "progress_counter": counter, "rss_mb": _rss_mb(os.getpid())}) + self.emit({ + "type": "hb", "active_turns": active_turns, "progress_counter": counter, + "rss_mb": _rss_mb(os.getpid())}) def _parent_guard_loop(self) -> None: while not self._closed.wait(1.0): @@ -668,23 +482,21 @@ class ComputeHost: os._exit(0) -def _coerce(value: Any, cast: Callable[[Any], Any], default: Any) -> Any: - try: - return cast(value) - except (TypeError, ValueError): - return default - - def _rss_mb(pid: int) -> float: try: - out = subprocess.check_output(["ps", "-o", "rss=", "-p", str(pid)], text=True, encoding="utf-8", errors="replace", stdin=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=2).strip() + out = subprocess.check_output( + ["ps", "-o", "rss=", "-p", str(pid)], text=True, encoding="utf-8", errors="replace", + stdin=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=2).strip() return int(out.splitlines()[-1].strip()) / 1024.0 if out else 0.0 except Exception: return 0.0 def _default_workers() -> int: - return _coerce(os.environ.get("HERMES_TUI_RPC_POOL_WORKERS") or "8", lambda v: max(2, int(v)), 8) + try: + return max(2, int(os.environ.get("HERMES_TUI_RPC_POOL_WORKERS") or "8")) + except (TypeError, ValueError): + return 8 def run_host(stdin: Any = None, stdout: Any = None) -> None: @@ -699,15 +511,13 @@ def run_host(stdin: Any = None, stdout: Any = None) -> None: shutting_down.set() host.shutdown(reason="sigterm") raise SystemExit(0) - with contextlib.suppress(Exception): signal.signal(signal.SIGTERM, _signal_handler) signal.signal(signal.SIGINT, _signal_handler) - - host.emit( - {"type": "hello", "host_pid": os.getpid(), "boot_id": host._boot_id, "build_sha": _build_sha(), - "cwd": os.getcwd(), "hermes_home": os.environ.get("HERMES_HOME", "")} - ) + host.emit({ + "type": "hello", "host_pid": os.getpid(), "boot_id": host._boot_id, + "build_sha": _build_sha(), "cwd": os.getcwd(), + "hermes_home": os.environ.get("HERMES_HOME", "")}) def _reader() -> None: for raw in stdin: @@ -726,7 +536,6 @@ def run_host(stdin: Any = None, stdout: Any = None) -> None: os._exit(0) if host._closed.is_set(): break - reader = threading.Thread(target=_reader, name="compute-host-control-reader", daemon=True) reader.start() try: diff --git a/tui_gateway/compute_host_bridge.py b/tui_gateway/compute_host_bridge.py index 696f7b4156..4754a14efd 100644 --- a/tui_gateway/compute_host_bridge.py +++ b/tui_gateway/compute_host_bridge.py @@ -1,4 +1,5 @@ -"""Compute-host (turn isolation) bridge: relay prompts/controls to the per-session child process and mirror its metadata/clarify/compress acks back into the session. +"""Compute-host (turn isolation) bridge: relay prompts/controls to the child process +and mirror its metadata/clarify/compress acks back into the session. Bodies are rebound onto server.py's globals at install time (see method_ctx.bind_module), so they reference server.py globals bare. @@ -6,6 +7,7 @@ method_ctx.bind_module), so they reference server.py globals bare. from __future__ import annotations +import contextlib import threading from .method_ctx import HandlerRegistry, bind_module @@ -28,8 +30,7 @@ def _inside_compute_host_child() -> bool: def _turn_isolation_enabled(cfg: dict | None = None) -> bool: if _inside_compute_host_child(): return False - isolation_cfg = cfg or _load_dashboard_process_isolation_config() - return bool(isolation_cfg.get("turn_isolation")) + return bool((cfg or _load_dashboard_process_isolation_config()).get("turn_isolation")) def _session_uses_compute_host(session: dict, cfg: dict | None = None) -> bool: @@ -38,8 +39,7 @@ def _session_uses_compute_host(session: dict, cfg: dict | None = None) -> bool: # Routes lazy sessions whose AIAgent was never built in-process; already-built # sessions keep the in-process path unless a prior isolated turn marked host ownership. return bool(session.get("_compute_host_active")) or ( - session.get("agent") is None and session.get("agent_ready") is not None - ) + session.get("agent") is None and session.get("agent_ready") is not None) def _get_compute_host_supervisor(cfg: dict | None = None): @@ -48,43 +48,34 @@ def _get_compute_host_supervisor(cfg: dict | None = None): with _compute_host_supervisor_lock: if _compute_host_supervisor is None: from tui_gateway.host_supervisor import HostSupervisor - _compute_host_supervisor = HostSupervisor( rpc_sink=_relay_compute_host_rpc, heartbeat_secs=int(isolation_cfg.get("compute_host_heartbeat_secs") or 15), - respawn_max=int(isolation_cfg.get("compute_host_respawn_max") or 3), - ) + respawn_max=int(isolation_cfg.get("compute_host_respawn_max") or 3)) return _compute_host_supervisor def _compute_host_turn_frame( rid: str, sid: str, session: dict, text: Any, image_paths: list[str] | None = None, - queued_prompt_generation: int | None = None, display_kind: str | None = None, -) -> dict: + queued_prompt_generation: int | None = None, display_kind: str | None = None) -> dict: with session["history_lock"]: history = list(session.get("history", [])) history_version = int(session.get("history_version", 0)) - attached_images = list(image_paths if image_paths is not None else session.get("attached_images", [])) + attached_images = list( + image_paths if image_paths is not None else session.get("attached_images", [])) return { - "type": "turn.start", - "sid": sid, - "request_id": rid, - "session_key": session.get("session_key") or sid, - "text": text, - **({"display_kind": display_kind} if display_kind else {}), - "history": history, - "history_version": history_version, - "cols": int(session.get("cols", 80) or 80), + "type": "turn.start", "sid": sid, "request_id": rid, + "session_key": session.get("session_key") or sid, "text": text, + **({"display_kind": display_kind} if display_kind else {}), "history": history, + "history_version": history_version, "cols": int(session.get("cols", 80) or 80), "cwd": _session_cwd(session), "context_cwd_is_launch_artifact": _context_cwd_is_launch_artifact(session), "profile_home": session.get("profile_home") or "", "model_override": session.get("model_override"), "reasoning_config_override": session.get("create_reasoning_override"), "service_tier_override": session.get("create_service_tier_override"), - "source": _session_source(session), - "attached_images": attached_images, - "queued_prompt_generation": queued_prompt_generation, - } + "source": _session_source(session), "attached_images": attached_images, + "queued_prompt_generation": queued_prompt_generation} def _metadata_mirror(session: dict | None) -> dict: @@ -105,13 +96,9 @@ def _compute_host_adopt_frame_meta(session: dict, frame: dict) -> None: if frame.get("session_key"): session["session_key"] = str(frame.get("session_key")) if frame.get("history_version") is not None: - try: + with contextlib.suppress(Exception): session["history_version"] = max( - int(session.get("history_version", 0)), - int(frame.get("history_version") or 0), - ) - except Exception: - pass + int(session.get("history_version", 0)), int(frame.get("history_version") or 0)) def _relay_compute_host_rpc(message: dict) -> bool: @@ -126,33 +113,39 @@ def _relay_compute_host_rpc(message: dict) -> bool: with session.get("history_lock", threading.Lock()): if kind == "clarify.request": session["_compute_host_pending_clarify"] = dict(payload) - else: - pending = session.get("_compute_host_pending_clarify") - if isinstance(pending, dict) and pending.get("request_id") == request_id: - session.pop("_compute_host_pending_clarify", None) + elif _pending_clarify_matches(session, request_id): + session.pop("_compute_host_pending_clarify", None) return write_json(message) +def _pending_clarify_matches(session: dict, request_id) -> bool: + """Whether ``session``'s mirrored pending clarify is ``request_id``. Caller holds + history_lock.""" + pending = session.get("_compute_host_pending_clarify") + return isinstance(pending, dict) and pending.get("request_id") == request_id + + def _compute_host_clarify_session(request_id: str) -> tuple[str, dict] | None: """Find the parent mirror for one host-owned clarify request.""" if not request_id: return None for sid, session in list(_sessions.items()): with session.get("history_lock", threading.Lock()): - pending = session.get("_compute_host_pending_clarify") - if isinstance(pending, dict) and pending.get("request_id") == request_id: + if _pending_clarify_matches(session, request_id): return sid, session return None -def _update_compute_host_clarify_snapshot(sid: str, session: dict, params: dict, result: dict) -> None: +def _update_compute_host_clarify_snapshot( + sid: str, session: dict, params: dict, result: dict) -> None: """Keep reconnect snapshots accurate while a batch clarify is answered.""" request_id = str(params.get("request_id") or "") with session.get("history_lock", threading.Lock()): - pending = session.get("_compute_host_pending_clarify") - if not isinstance(pending, dict) or pending.get("request_id") != request_id: + if not _pending_clarify_matches(session, request_id): return - if result.get("status") == "expired" or not result.get("remaining") and not params.get("question_id"): + pending = session["_compute_host_pending_clarify"] + expired = result.get("status") == "expired" + if expired or not result.get("remaining") and not params.get("question_id"): session.pop("_compute_host_pending_clarify", None) return question_id = str(params.get("question_id") or "") @@ -183,7 +176,9 @@ def _respond_compute_host_clarify(rid: str, params: dict) -> dict | None: return _err(rid, 5019, "compute-host clarify response returned an invalid response") if "error" in response: error = response["error"] if isinstance(response["error"], dict) else {} - return _err(rid, int(error.get("code") or 5000), str(error.get("message") or "clarify response failed")) + return _err( + rid, int(error.get("code") or 5000), + str(error.get("message") or "clarify response failed")) result = response.get("result") if not isinstance(result, dict): return _err(rid, 5019, "compute-host clarify response returned an invalid result") @@ -192,23 +187,18 @@ def _respond_compute_host_clarify(rid: str, params: dict) -> dict | None: def _apply_compute_host_metadata_mirror(session: dict, frame: dict | None) -> None: - """Mirror host-owned session metadata: while turn isolation is active the host is - the only writer of live agent/history state, and UI reads must not build a - second in-process agent.""" + """Mirror host-owned session metadata: under turn isolation the host is the only + writer of live agent/history state, and UI reads must not build a second agent.""" if not isinstance(frame, dict): return with session.get("history_lock", threading.Lock()): _compute_host_adopt_frame_meta(session, frame) if frame.get("message_count") is not None: - try: + with contextlib.suppress(Exception): session["_metadata_message_count"] = int(frame.get("message_count") or 0) - except Exception: - pass info = frame.get("session_info") if isinstance(info, dict): - mirror = dict(_metadata_mirror(session)) - mirror.update(info) - session["_metadata_mirror"] = mirror + session["_metadata_mirror"] = {**_metadata_mirror(session), **info} session["_metadata_mirror_updated_at"] = time.time() @@ -231,13 +221,11 @@ def _on_compute_host_turn_done(rid: str, sid: str, session: dict, frame: dict) - def _submit_prompt_to_compute_host( rid: str, sid: str, session: dict, text: Any, image_paths: list[str] | None = None, - queued_prompt_generation: int | None = None, display_kind: str | None = None, -) -> dict: + queued_prompt_generation: int | None = None, display_kind: str | None = None) -> dict: cfg = _load_dashboard_process_isolation_config() frame = _compute_host_turn_frame( rid, sid, session, text, image_paths=image_paths, - queued_prompt_generation=queued_prompt_generation, display_kind=display_kind, - ) + queued_prompt_generation=queued_prompt_generation, display_kind=display_kind) def _complete(done: dict) -> None: # submit_turn reports a synchronous pipe failure via the callback before @@ -246,7 +234,6 @@ def _submit_prompt_to_compute_host( if done.get("reason") == "send_failed": return _on_compute_host_turn_done(rid, sid, session, done) - try: _get_compute_host_supervisor(cfg).submit_turn(frame, on_complete=_complete) except Exception as exc: @@ -260,44 +247,43 @@ def _submit_prompt_to_compute_host( def _send_compute_host_control( sid: str, *, route_name: str, command: str = "", payload: dict | None = None, - wait: bool = True, timeout: float = 30.0, on_late_ack=None, -) -> dict: + wait: bool = True, timeout: float = 30.0, on_late_ack=None) -> dict: frame = dict(payload or {}) frame.setdefault("type", "control") frame.setdefault("command", command) return _get_compute_host_supervisor().control( - sid, route_name=route_name, payload=frame, wait=wait, timeout=timeout, on_late_ack=on_late_ack - ) + sid, route_name=route_name, payload=frame, wait=wait, timeout=timeout, + on_late_ack=on_late_ack) def _compute_host_compress_wait_seconds(cfg: dict | None = None) -> float: - """RPC wait budget for a compute-host compress control: the configured - ``compression.context_total_ceiling_seconds`` plus slack, capped below the - desktop's RPC timeout (a fixed waiter reported false timeouts while the host - kept working); anything slower lands via the late-ack path.""" + """RPC wait budget for a compute-host compress control: the configured compression + ceiling plus slack, capped below the desktop's RPC timeout (a fixed waiter reported + false timeouts while the host kept working); slower acks land via the late-ack path.""" from agent.conversation_compression import resolve_context_compression_timeouts - try: compression_cfg = (cfg if cfg is not None else _load_cfg()).get("compression", {}) except Exception: compression_cfg = {} - _idle, ceiling = resolve_context_compression_timeouts(compression_cfg if isinstance(compression_cfg, dict) else {}) + if not isinstance(compression_cfg, dict): + compression_cfg = {} + _idle, ceiling = resolve_context_compression_timeouts(compression_cfg) return float(min(max(ceiling + 30.0, 120.0), _COMPUTE_HOST_COMPRESS_WAIT_CAP_SECS)) def _announce_compute_host_compress_done(sid: str, session: dict, ack: dict) -> None: - """Mirror a compress ack and push the same ``session.info`` + ``compacted`` edges - the in-process /compress path emits, so a client whose RPC wait expired still - learns the transcript changed.""" + """Mirror a compress ack and push the ``session.info`` + ``compacted`` edges the + in-process /compress path emits, so a client whose RPC wait expired still learns.""" _apply_compute_host_metadata_mirror(session, ack) _emit("session.info", sid, _compute_host_session_info(session)) _status_update(sid, "compacted", "✓ Context compression complete") -def _adopt_late_compute_host_compress_ack(sid: str, session: dict, ack: dict, *, route_name: str) -> None: - """Adopt a compress ack that arrived after its RPC waiter answered ``pending``: - the only place the rotated session_key / history_version / mirror can land and - the client's only signal. A late ``control.error`` goes out via ``error``.""" +def _adopt_late_compute_host_compress_ack( + sid: str, session: dict, ack: dict, *, route_name: str) -> None: + """Adopt a compress ack that arrived after its RPC waiter answered ``pending``: the + only place the rotated session_key / history_version / mirror can land and the + client's only signal. A late ``control.error`` goes out via ``error``.""" with _sessions_lock: live = _sessions.get(sid) if live is not session: diff --git a/tui_gateway/host_supervisor.py b/tui_gateway/host_supervisor.py index 25b9f084b9..dddb4401b9 100644 --- a/tui_gateway/host_supervisor.py +++ b/tui_gateway/host_supervisor.py @@ -1,9 +1,8 @@ """Supervisor for the dashboard compute-host child process. -The dashboard process owns sockets and JSON-RPC dispatch. When -``dashboard.turn_isolation`` is enabled, agent turns move behind one persistent -``python -m tui_gateway.compute_host`` child so compute-heavy agent threads do -not contend with the serving process' event loop for the same GIL. +When ``dashboard.turn_isolation`` is enabled, agent turns move behind one persistent +``python -m tui_gateway.compute_host`` child so compute-heavy agent threads do not +contend with the serving process' event loop for the same GIL. """ from __future__ import annotations @@ -30,33 +29,25 @@ logger = logging.getLogger(__name__) _Thread = threading.Thread MUTATOR_ROUTE_TABLE: dict[str, str] = { - "prompt.submit": "turn-path", - "session.interrupt": "turn-path", - "reload.mcp": "run-concurrent", - "session.save": "run-concurrent", - "session.compress": "idle-gated", - "prompt.submit.truncate": "idle-gated", - "slash.model": "idle-gated", - "slash.personality": "idle-gated", - "slash.prompt": "idle-gated", - "slash.compress": "idle-gated", - "session.reset": "idle-gated", - "session.history.reload": "idle-gated", - "slash.retry": "idle-gated", -} + "prompt.submit": "turn-path", "session.interrupt": "turn-path", "reload.mcp": "run-concurrent", + "session.save": "run-concurrent", "session.compress": "idle-gated", + "prompt.submit.truncate": "idle-gated", "slash.model": "idle-gated", + "slash.personality": "idle-gated", "slash.prompt": "idle-gated", "slash.compress": "idle-gated", + "session.reset": "idle-gated", "session.history.reload": "idle-gated", + "slash.retry": "idle-gated"} _REGISTRY_NAME = "dashboard-compute-host.json" _RESPAWN_WINDOW_SECS = 300.0 _SHUTDOWN_TIMEOUT_SECS = 10.0 -# Late control-ack handlers: a compress that outlives its RPC waiter can run for -# the full compression ceiling plus a stall-fallback retry, so keep -# registrations well past that — but bounded. +# Late control-ack handlers: a compress that outlives its RPC waiter can run for the +# full compression ceiling plus a stall-fallback retry, so keep registrations well +# past that — but bounded. _LATE_CONTROL_TTL_SECS = 1800.0 _LATE_CONTROL_MAX = 64 # Host frames whose ``request_id`` resolves a pending/late control waiter. -_CONTROL_REPLY_TYPES = frozenset( - {"control.ack", "control.error", "respond.ack", "respond.error", "interrupt.ack", "reload_mcp.ack", "shutdown.ack"} -) +_CONTROL_REPLY_TYPES = frozenset({ + "control.ack", "control.error", "respond.ack", "respond.error", "interrupt.ack", + "reload_mcp.ack", "shutdown.ack"}) def append_log_record(path: str | Path, record: str) -> None: @@ -64,10 +55,9 @@ def append_log_record(path: str | Path, record: str) -> None: p = Path(path) p.parent.mkdir(parents=True, exist_ok=True) text = record if record.endswith("\n") else f"{record}\n" - data = text.encode("utf-8", errors="replace") fd = os.open(str(p), os.O_WRONLY | os.O_CREAT | os.O_APPEND, 0o600) try: - os.write(fd, data) + os.write(fd, text.encode("utf-8", errors="replace")) finally: os.close(fd) @@ -76,25 +66,20 @@ def _repo_root() -> Path: return Path(__file__).resolve().parents[1] -def _build_sha() -> str: - """Current checkout's HEAD sha, or ``"unknown"``. Shared with ``compute_host`` - so the hello handshake and the supervisor's expectation agree byte-for-byte.""" +def _check_output(argv: list[str], **kwargs: Any) -> str: + """Stripped stdout of a short subprocess, or ``""`` on any failure.""" try: return subprocess.check_output( - ["git", "rev-parse", "HEAD"], - cwd=str(_repo_root()), - text=True, - encoding="utf-8", - errors="replace", - stderr=subprocess.DEVNULL, - timeout=2, - ).strip() + argv, text=True, encoding="utf-8", errors="replace", stderr=subprocess.DEVNULL, + timeout=2, **kwargs).strip() except Exception: - return "unknown" + return "" -def _default_registry_path() -> Path: - return get_hermes_home() / "state" / _REGISTRY_NAME +def _build_sha() -> str: + """Current checkout's HEAD sha, or ``"unknown"``. Shared with ``compute_host`` so + the hello handshake and the supervisor's expectation agree byte-for-byte.""" + return _check_output(["git", "rev-parse", "HEAD"], cwd=str(_repo_root())) or "unknown" def _pid_alive(pid: int) -> bool: @@ -118,17 +103,7 @@ def _pid_command(pid: int) -> str: data = (Path("/proc") / str(pid) / "cmdline").read_bytes() if data: return data.replace(b"\x00", b" ").decode("utf-8", errors="replace") - try: - return subprocess.check_output( - ["ps", "-p", str(pid), "-o", "command="], - text=True, - encoding="utf-8", - errors="replace", - stderr=subprocess.DEVNULL, - timeout=2, - ).strip() - except Exception: - return "" + return _check_output(["ps", "-p", str(pid), "-o", "command="]) def is_compute_host_identity(pid: int) -> bool: @@ -139,34 +114,26 @@ class HostSupervisor: """Own one persistent compute-host child and relay its frames.""" def __init__( - self, - *, - registry_path: str | Path | None = None, - argv: list[str] | None = None, - cwd: str | Path | None = None, - env: dict[str, str] | None = None, - rpc_sink: Callable[[dict], None] | None = None, - respawn_max: int = 3, - heartbeat_secs: int = 15, - expected_build_sha: str | None = None, - expected_hermes_home: str | None = None, - autostart: bool = True, - ) -> None: - self.registry_path = Path(registry_path) if registry_path is not None else _default_registry_path() + self, *, registry_path: str | Path | None = None, argv: list[str] | None = None, + cwd: str | Path | None = None, env: dict[str, str] | None = None, + rpc_sink: Callable[[dict], None] | None = None, respawn_max: int = 3, + heartbeat_secs: int = 15, expected_build_sha: str | None = None, + expected_hermes_home: str | None = None, autostart: bool = True) -> None: + self.registry_path = ( + Path(registry_path) if registry_path is not None + else get_hermes_home() / "state" / _REGISTRY_NAME) self.argv = argv or [sys.executable, "-m", "tui_gateway.compute_host"] self.cwd = Path(cwd) if cwd is not None else _repo_root() self.env = env self.rpc_sink = rpc_sink or (lambda _obj: None) self.respawn_max = max(0, int(respawn_max)) self.heartbeat_secs = max(1, int(heartbeat_secs)) - self.expected_build_sha = expected_build_sha if expected_build_sha is not None else _build_sha() - self.expected_hermes_home = expected_hermes_home if expected_hermes_home is not None else str(get_hermes_home()) - + self.expected_build_sha = ( + expected_build_sha if expected_build_sha is not None else _build_sha()) + self.expected_hermes_home = ( + expected_hermes_home if expected_hermes_home is not None else str(get_hermes_home())) self._lock = threading.RLock() self._proc: subprocess.Popen[str] | None = None - self._stdout_thread: threading.Thread | None = None - self._stderr_thread: threading.Thread | None = None - self._wait_thread: threading.Thread | None = None self._hello_event = threading.Event() self._hello: dict[str, Any] = {} self._closing = False @@ -174,13 +141,12 @@ class HostSupervisor: self._restart_times: list[float] = [] self._pending_turns: dict[str, tuple[str, Callable[[dict], None] | None]] = {} self._pending_controls: dict[str, queue.Queue[dict]] = {} - # request_id -> (registered_at, handler) for control waiters that timed - # out while their host work still runs; without it the eventual - # control.ack matched no queue and was silently dropped. + # request_id -> (registered_at, handler) for control waiters that timed out + # while their host work still runs; without it the eventual control.ack + # matched no queue and was silently dropped. self._late_control_handlers: dict[str, tuple[float, Callable[[dict], None]]] = {} self._stderr_tail: list[str] = [] self._last_progress_counter = 0 - if autostart: self.start() @@ -229,7 +195,6 @@ class HostSupervisor: except Exception: self._remove_registry() return "invalid-registry" - try: pid = int(data.get("host_pid") or 0) except Exception: @@ -241,17 +206,12 @@ class HostSupervisor: # PID was reused by another process. Never signal it. self._remove_registry() return "pid-reuse-ignored" - self._terminate_pid(pid, timeout=_SHUTDOWN_TIMEOUT_SECS) self._remove_registry() return "terminated" def submit_turn( - self, - frame: dict[str, Any], - *, - on_complete: Callable[[dict], None] | None = None, - ) -> str: + self, frame: dict[str, Any], *, on_complete: Callable[[dict], None] | None = None) -> str: self.start() request_id = str(frame.get("request_id") or uuid.uuid4().hex) sid = str(frame.get("sid") or "") @@ -264,85 +224,81 @@ class HostSupervisor: with self._lock: self._pending_turns.pop(request_id, None) if on_complete is not None: - on_complete({"type": "turn.error", "sid": sid, "request_id": request_id, "reason": "send_failed", "message": str(exc)}) + on_complete({ + "type": "turn.error", "sid": sid, "request_id": request_id, + "reason": "send_failed", "message": str(exc)}) raise return request_id def interrupt(self, sid: str, *, request_id: str | None = None) -> None: self.start() - self._send_frame({"type": "interrupt", "sid": sid, "request_id": request_id or uuid.uuid4().hex}) + self._send_frame({ + "type": "interrupt", "sid": sid, "request_id": request_id or uuid.uuid4().hex}) + + def _await_reply(self, frame: dict[str, Any], request_id: str, timeout: float) -> dict: + """Send ``frame`` and block for the host reply carrying ``request_id``.""" + q: queue.Queue[dict] = queue.Queue(maxsize=1) + with self._lock: + self._pending_controls[request_id] = q + try: + self._send_frame(frame) + return q.get(timeout=timeout) + finally: + with self._lock: + self._pending_controls.pop(request_id, None) def respond(self, sid: str, params: dict[str, Any], *, timeout: float = 15.0) -> dict: """Deliver an interactive prompt response to the host that owns it.""" self.start() request_id = uuid.uuid4().hex - q: queue.Queue[dict] = queue.Queue(maxsize=1) - with self._lock: - self._pending_controls[request_id] = q - try: - self._send_frame({"type": "respond", "sid": sid, "request_id": request_id, "params": dict(params)}) - return q.get(timeout=timeout) - finally: - with self._lock: - self._pending_controls.pop(request_id, None) + frame = {"type": "respond", "sid": sid, "request_id": request_id, "params": dict(params)} + return self._await_reply(frame, request_id, timeout) def reload_mcp(self, sid: str, *, request_id: str | None = None) -> dict: return self.control( - sid, - route_name="reload.mcp", - payload={"type": "reload_mcp", "sid": sid, "request_id": request_id or uuid.uuid4().hex}, - wait=True, - ) + sid, route_name="reload.mcp", wait=True, + payload={ + "type": "reload_mcp", "sid": sid, "request_id": request_id or uuid.uuid4().hex}) def control( - self, - sid: str, - *, - route_name: str, - payload: dict[str, Any] | None = None, - wait: bool = True, - timeout: float = 30.0, - on_late_ack: Callable[[dict], None] | None = None, + self, sid: str, *, route_name: str, payload: dict[str, Any] | None = None, + wait: bool = True, timeout: float = 30.0, on_late_ack: Callable[[dict], None] | None = None, ) -> dict: """Send a control frame; with ``wait`` block up to ``timeout`` for its ack. ``on_late_ack`` (only with ``wait``) keeps the request adoptable after the - waiter gives up: the host's eventual ``control.ack``/``control.error``/ - ``error`` for this ``request_id`` fires the handler once instead of being - dropped. Bounded by ``_LATE_CONTROL_TTL_SECS`` / ``_LATE_CONTROL_MAX``. + waiter gives up: the host's eventual ``control.ack``/``control.error``/``error`` + for this ``request_id`` fires the handler once instead of being dropped. + Bounded by ``_LATE_CONTROL_TTL_SECS`` / ``_LATE_CONTROL_MAX``. """ if route_name not in MUTATOR_ROUTE_TABLE: raise ValueError(f"unclassified host mutator route: {route_name}") self.start() request_id = str((payload or {}).get("request_id") or uuid.uuid4().hex) - frame = {"type": "control", **(payload or {}), "sid": sid, "route_name": route_name, "request_id": request_id} - q: queue.Queue[dict] | None = None - if wait: - q = queue.Queue(maxsize=1) - with self._lock: - self._pending_controls[request_id] = q - self._send_frame(frame) - if not wait or q is None: + frame = { + "type": "control", **(payload or {}), "sid": sid, "route_name": route_name, + "request_id": request_id} + if not wait: + self._send_frame(frame) return {"status": "sent", "request_id": request_id} try: - return q.get(timeout=timeout) + return self._await_reply(frame, request_id, timeout) except queue.Empty: if on_late_ack is not None: self._register_late_control_handler(request_id, on_late_ack) raise - finally: - with self._lock: - self._pending_controls.pop(request_id, None) - def _register_late_control_handler(self, request_id: str, handler: Callable[[dict], None]) -> None: + def _register_late_control_handler( + self, request_id: str, handler: Callable[[dict], None]) -> None: now = time.monotonic() with self._lock: - for rid in [r for r, (at, _cb) in self._late_control_handlers.items() if now - at > _LATE_CONTROL_TTL_SECS]: - self._late_control_handlers.pop(rid, None) - while len(self._late_control_handlers) >= _LATE_CONTROL_MAX: - oldest = min(self._late_control_handlers, key=lambda rid: self._late_control_handlers[rid][0]) - self._late_control_handlers.pop(oldest, None) - self._late_control_handlers[request_id] = (now, handler) + handlers = self._late_control_handlers + expired = [r for r, (at, _cb) in handlers.items() if now - at > _LATE_CONTROL_TTL_SECS] + for rid in expired: + handlers.pop(rid, None) + while len(handlers) >= _LATE_CONTROL_MAX: + handlers.pop(min(handlers, key=lambda rid: handlers[rid][0]), None) + handlers[request_id] = (now, handler) def _deliver_control_frame(self, request_id: str, frame: dict[str, Any]) -> None: with self._lock: @@ -357,7 +313,8 @@ class HostSupervisor: try: late[1](frame) except Exception: - logger.exception("compute host late control ack handler failed (request_id=%s)", request_id) + logger.exception( + "compute host late control ack handler failed (request_id=%s)", request_id) def _spawn_locked(self, *, reason: str) -> None: if self._stopped_respawning: @@ -374,26 +331,17 @@ class HostSupervisor: if root not in env["PYTHONPATH"].split(os.pathsep): env["PYTHONPATH"] = root + os.pathsep + env["PYTHONPATH"] proc = subprocess.Popen( - self.argv, - cwd=str(self.cwd), - env=env, - stdin=subprocess.PIPE, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, - # Lossy UTF-8 decode: a locale-mismatched byte must not raise inside - # the drain threads and kill the supervisor. - encoding="utf-8", - errors="replace", - bufsize=1, - start_new_session=True, - ) + self.argv, cwd=str(self.cwd), env=env, stdin=subprocess.PIPE, stdout=subprocess.PIPE, + stderr=subprocess.PIPE, text=True, + # Lossy UTF-8 decode: a locale-mismatched byte must not raise inside the + # drain threads and kill the supervisor. + encoding="utf-8", errors="replace", bufsize=1, start_new_session=True) self._proc = proc - self._stdout_thread = _Thread(target=self._drain_stdout, args=(proc,), name="compute-host-stdout", daemon=True) - self._stderr_thread = _Thread(target=self._drain_stderr, args=(proc,), name="compute-host-stderr", daemon=True) - self._wait_thread = _Thread(target=self._wait_for_exit, args=(proc,), name="compute-host-wait", daemon=True) - for t in (self._stdout_thread, self._stderr_thread, self._wait_thread): - t.start() + for target, name in ( + (self._drain_stdout, "compute-host-stdout"), + (self._drain_stderr, "compute-host-stderr"), (self._wait_for_exit, "compute-host-wait"), + ): + _Thread(target=target, args=(proc,), name=name, daemon=True).start() if not self._hello_event.wait(timeout=10.0): self._terminate_process(proc) raise RuntimeError(f"compute host did not send hello; stderr={self._stderr_tail[-5:]}") @@ -407,21 +355,20 @@ class HostSupervisor: raise RuntimeError("compute host missing hello") got_home = str(hello.get("hermes_home") or "") if got_home and got_home != self.expected_hermes_home: - raise RuntimeError(f"compute host HERMES_HOME mismatch: {got_home} != {self.expected_hermes_home}") + raise RuntimeError( + f"compute host HERMES_HOME mismatch: {got_home} != {self.expected_hermes_home}") got_sha = str(hello.get("build_sha") or "") - if self.expected_build_sha != "unknown" and got_sha not in {"", "unknown", self.expected_build_sha}: - raise RuntimeError(f"compute host build mismatch: {got_sha} != {self.expected_build_sha}") + expected = self.expected_build_sha + if expected != "unknown" and got_sha not in {"", "unknown", expected}: + raise RuntimeError(f"compute host build mismatch: {got_sha} != {expected}") def _persist_registry(self) -> None: self.registry_path.parent.mkdir(parents=True, exist_ok=True) tmp = self.registry_path.with_suffix(self.registry_path.suffix + ".tmp") payload = { - "host_pid": self.pid, - "boot_id": self._hello.get("boot_id") or "", - "build_sha": self._hello.get("build_sha") or "", - "started_at": time.time(), - "argv": self.argv, - } + "host_pid": self.pid, "boot_id": self._hello.get("boot_id") or "", + "build_sha": self._hello.get("build_sha") or "", "started_at": time.time(), + "argv": self.argv} tmp.write_text(json.dumps(payload, sort_keys=True), encoding="utf-8") tmp.replace(self.registry_path) @@ -469,19 +416,16 @@ class HostSupervisor: # host frame ``type`` -> handler method name (see also _CONTROL_REPLY_TYPES). _HOST_FRAME_HANDLERS: dict[str, str] = { - "hello": "_on_hello", - "hb": "_on_heartbeat", - "rpc": "_on_rpc", - "turn.end": "_complete_turn", - "turn.error": "_complete_turn", - } + "hello": "_on_hello", "hb": "_on_heartbeat", "rpc": "_on_rpc", "turn.end": "_complete_turn", + "turn.error": "_complete_turn"} def _on_hello(self, frame: dict[str, Any]) -> None: self._hello = dict(frame) self._hello_event.set() def _on_heartbeat(self, frame: dict[str, Any]) -> None: - self._last_progress_counter = int(frame.get("progress_counter") or self._last_progress_counter) + self._last_progress_counter = int( + frame.get("progress_counter") or self._last_progress_counter) logger.debug("compute host heartbeat: %s", frame) def _on_rpc(self, frame: dict[str, Any]) -> None: @@ -516,16 +460,17 @@ class HostSupervisor: pending = self._pending_turns self._pending_turns = {} for request_id, (sid, cb) in pending.items(): - self.rpc_sink( - { - "jsonrpc": "2.0", - "method": "event", - "params": {"type": "error", "session_id": sid, "payload": {"message": message, "reason": reason}}, - } - ) + self.rpc_sink({ + "jsonrpc": "2.0", + "method": "event", + "params": { + "type": "error", "session_id": sid, + "payload": {"message": message, "reason": reason}}}) if cb is not None: try: - cb({"type": "turn.error", "sid": sid, "request_id": request_id, "reason": reason, "message": message}) + cb({ + "type": "turn.error", "sid": sid, "request_id": request_id, + "reason": reason, "message": message}) except Exception: logger.exception("compute host error callback failed") # A crashed host never emits the late acks timed-out control waiters still @@ -535,7 +480,9 @@ class HostSupervisor: self._late_control_handlers = {} for request_id, (_registered_at, handler) in late.items(): try: - handler({"type": "control.error", "request_id": request_id, "reason": reason, "message": message}) + handler({ + "type": "control.error", "request_id": request_id, "reason": reason, + "message": message}) except Exception: logger.exception("compute host late control error handler failed") @@ -544,7 +491,9 @@ class HostSupervisor: self._restart_times = [t for t in self._restart_times if now - t <= _RESPAWN_WINDOW_SECS] if len(self._restart_times) >= self.respawn_max: self._stopped_respawning = True - logger.error("compute host crash loop: max %s restarts per 5min reached; not respawning", self.respawn_max) + logger.error( + "compute host crash loop: max %s restarts per 5min reached; not respawning", + self.respawn_max) return self._restart_times.append(now) # Small bounded backoff; tests and first recovery stay quick. @@ -559,9 +508,7 @@ class HostSupervisor: self._spawn_locked(reason="crash") except Exception: logger.exception("compute host respawn failed") - _Thread(target=_respawn, name="compute-host-respawn", daemon=True).start() - _pid_matches_compute_host = staticmethod(is_compute_host_identity) def _terminate_pid(self, pid: int, *, timeout: float = _SHUTDOWN_TIMEOUT_SECS) -> None: @@ -599,9 +546,4 @@ class HostSupervisor: proc.wait(timeout=2) -__all__ = [ - "MUTATOR_ROUTE_TABLE", - "HostSupervisor", - "append_log_record", - "is_compute_host_identity", -] +__all__ = ["MUTATOR_ROUTE_TABLE", "HostSupervisor", "append_log_record", "is_compute_host_identity"] diff --git a/tui_gateway/hosted_room_service.py b/tui_gateway/hosted_room_service.py index 8a86cb88e9..7303058c45 100644 --- a/tui_gateway/hosted_room_service.py +++ b/tui_gateway/hosted_room_service.py @@ -17,25 +17,14 @@ from gateway import hosted_room_discussion as discussion from gateway import hosted_room_driver as driver from gateway import hosted_room_links from gateway import hosted_rooms -from gateway.hosted_room_policy_checkpoint import ( - HostedRoomPolicyCheckpoint, - PolicySnapshot, -) +from gateway.hosted_room_policy_checkpoint import HostedRoomPolicyCheckpoint, PolicySnapshot from gateway.hosted_room_peer import ( - GatewayRoomCatalog, - HostedMemberDispatch, - PROTOCOL_VERSION, - room_grant_needs_dispatch_refresh, -) + GatewayRoomCatalog, HostedMemberDispatch, PROTOCOL_VERSION, room_grant_needs_dispatch_refresh) from tui_gateway.hosted_room_driver import HostedRoomBinding, HostedRoomRuntime from tui_gateway.hosted_room_server_rpc import HostedRoomServerRPC from tui_gateway.hosted_room_peer_http import PeerRunsHTTPClient, PeerRunsHTTPError from tui_gateway.hosted_room_peer_transport import ( - HostedRoomPeerClient, - PeerHostedRoomTransport, - PeerMemberRoute, - build_member_dispatch, -) + HostedRoomPeerClient, PeerHostedRoomTransport, PeerMemberRoute, build_member_dispatch) _HOSTED_ROOM_IDLE_FALLBACK_SECONDS = 5.0 @@ -45,6 +34,7 @@ _HOSTED_ROOM_TERMINAL_GRACE_SECONDS = 30.0 _TERMINAL_STATUSES = ("deferred", "settled", "failed", "cancelled") _LIVE_STATUSES = ("queued", "running", "stopping") _STOPPABLE_STATUSES = ("queued", "running", "indeterminate", "deferred", "stopping") +_RETRYABLE_STATUSES = ("indeterminate", "deferred") def _hosted_room_turn_timeout_seconds() -> float: @@ -60,22 +50,26 @@ def _hosted_room_turn_timeout_seconds() -> float: def _grant_revoke_is_terminal(exc: PeerRunsHTTPError) -> bool: """Return whether the peer proves the scoped grant is already unusable.""" return exc.status_code in {401, 403} and exc.error_code in { - "invalid_room_grant", - "room_reauthorization_required", - } + "invalid_room_grant", "room_reauthorization_required"} + + +def _hook(obj: Any, name: str): + """Optional callable attribute of a duck-typed peer client, or None.""" + value = getattr(obj, name, None) + return value if callable(value) else None + + +def _authority(room: Mapping[str, Any]) -> tuple[str, int]: + return str(room["authority_gateway_id"]), int(room["authority_epoch"]) class HostedRoomService: """Own the hosted Discussion policy and its transport-free worker.""" def __init__( - self, - server: ModuleType, - *, - db_path: Path | str | None = None, + self, server: ModuleType, *, db_path: Path | str | None = None, peer_routes: Mapping[tuple[str, str], PeerMemberRoute] | None = None, - peer_clients: Mapping[Any, HostedRoomPeerClient] | None = None, - ) -> None: + peer_clients: Mapping[Any, HostedRoomPeerClient] | None = None) -> None: self.server = server self.db_path = Path(db_path or hosted_rooms.default_db_path()) hosted_rooms.prune_disbanded_rooms(self.db_path) @@ -101,18 +95,13 @@ class HostedRoomService: if client is not None: self.peer_clients[key] = client self.runtime = HostedRoomRuntime( - db_path=self.db_path, - rooms=self.bindings, - rpc=self.rpc, - transport_resolver=self._resolve_member_transport, - turn_lock=self._turn_lock, - prepare_room=self.prepare_room, - publish_terminal=self.publish_terminal, + db_path=self.db_path, rooms=self.bindings, rpc=self.rpc, + transport_resolver=self._resolve_member_transport, turn_lock=self._turn_lock, + prepare_room=self.prepare_room, publish_terminal=self.publish_terminal, pending_action=self._set_pending_action, poll_interval_seconds=_HOSTED_ROOM_IDLE_FALLBACK_SECONDS, active_poll_interval_seconds=_HOSTED_ROOM_ACTIVE_POLL_SECONDS, - turn_timeout_seconds=_hosted_room_turn_timeout_seconds(), - ) + turn_timeout_seconds=_hosted_room_turn_timeout_seconds()) def _load_stored_links(self) -> None: """Rehydrate persisted peer routes; collect per-link errors into one string.""" @@ -125,20 +114,14 @@ class HostedRoomService: continue self.peer_routes[key] = PeerMemberRoute( home_install_id=hosted_rooms.local_authority_gateway_id(), - member_id=stored.member_id, - target_install_id=stored.catalog.installation_id, + member_id=stored.member_id, target_install_id=stored.catalog.installation_id, target_profile=stored.target_profile, capability_digest=stored.catalog.catalog_digest, execution_policy_digest=stored.catalog.execution_policy.policy_digest, - cancellation_scope_id=stored.cancellation_scope_id, - trace_id=stored.trace_id, - grant=stored.grant, - ) + cancellation_scope_id=stored.cancellation_scope_id, trace_id=stored.trace_id, + grant=stored.grant) self.peer_clients[key] = PeerRunsHTTPClient( - base_url=stored.target_url, - api_key="", - receipt_db_path=self.db_path, - ) + base_url=stored.target_url, api_key="", receipt_db_path=self.db_path) self._peer_route_status[key] = stored.status if errors: self._link_load_error = ",".join(errors) @@ -158,26 +141,24 @@ class HostedRoomService: local_gateway_id = hosted_rooms.local_authority_gateway_id() return tuple( HostedRoomBinding( - room_id=str(room["room_id"]), - gateway_id=str(room["authority_gateway_id"]), - authority_epoch=int(room["authority_epoch"]), - ) + room_id=str(room["room_id"]), gateway_id=str(room["authority_gateway_id"]), + authority_epoch=int(room["authority_epoch"])) for room in hosted_rooms.list_rooms(self.db_path) - if str(room["authority_gateway_id"]) == local_gateway_id - ) + if str(room["authority_gateway_id"]) == local_gateway_id) + + def _room(self, room_id: str) -> dict[str, Any]: + return hosted_rooms.room_state(self.db_path, room_id=room_id) def _owned_room(self, room_id: str) -> dict[str, Any]: - room = hosted_rooms.room_state(self.db_path, room_id=room_id) + room = self._room(room_id) if str(room["authority_gateway_id"]) != hosted_rooms.local_authority_gateway_id(): raise hosted_rooms.AuthorityConflictError( - "This Group Chat is managed by another gateway." - ) + "This Group Chat is managed by another gateway.") return room @contextlib.contextmanager def _turn_lock(self, profile: str) -> Iterator[None]: from tools.bot_relay import acquire_turn_lock - with acquire_turn_lock(self.root, profile): yield @@ -194,45 +175,37 @@ class HostedRoomService: for status in statuses: yield from driver.list_tasks(self.db_path, room_id=room_id, status=status) + def _save_link( + self, *, room_id: str, member_id: str, target_url: str, target_profile: str, grant: str, + catalog: GatewayRoomCatalog, cancellation_scope_id: str, trace_id: str) -> None: + hosted_room_links.save_room_link( + self.db_path, + hosted_room_links.make_stored_link( + room_id=room_id, member_id=member_id, target_url=target_url, + target_profile=target_profile, grant=grant, catalog=catalog, + cancellation_scope_id=cancellation_scope_id, trace_id=trace_id)) + def register_peer_route( - self, - *, - room_id: str, - member_id: str, - route: PeerMemberRoute, - client: HostedRoomPeerClient, - target_url: str | None = None, - catalog: GatewayRoomCatalog | None = None, - ) -> None: + self, *, room_id: str, member_id: str, route: PeerMemberRoute, + client: HostedRoomPeerClient, target_url: str | None = None, + catalog: GatewayRoomCatalog | None = None) -> None: """Register one verified route and optionally persist its scoped grant.""" - bind_store = getattr(client, "bind_receipt_store", None) - if callable(bind_store): + bind_store = _hook(client, "bind_receipt_store") + if bind_store is not None: bind_store(self.db_path) if catalog is not None: if not route.execution_policy_digest: route = replace( - route, - execution_policy_digest=catalog.execution_policy.policy_digest, - ) + route, execution_policy_digest=catalog.execution_policy.policy_digest) if ( route.capability_digest != catalog.catalog_digest - or route.execution_policy_digest != catalog.execution_policy.policy_digest - ): + or route.execution_policy_digest != catalog.execution_policy.policy_digest): raise ValueError("peer route does not match its target catalog") if target_url is not None: - hosted_room_links.save_room_link( - self.db_path, - hosted_room_links.make_stored_link( - room_id=room_id, - member_id=member_id, - target_url=target_url, - target_profile=route.target_profile, - grant=route.grant, - catalog=catalog, - cancellation_scope_id=route.cancellation_scope_id, - trace_id=route.trace_id, - ), - ) + self._save_link( + room_id=room_id, member_id=member_id, target_url=target_url, + target_profile=route.target_profile, grant=route.grant, catalog=catalog, + cancellation_scope_id=route.cancellation_scope_id, trace_id=route.trace_id) # Persistence is the publication boundary: a failed disk write must never # leave a process-local route that disappears after restart. with self._policy_lock: @@ -250,15 +223,14 @@ class HostedRoomService: with self._policy_lock: routes = [(key, route) for key, route in self.peer_routes.items() if key[0] == room_id] for key, route in routes: - revoke = getattr(self.peer_clients.get(key), "revoke_grant", None) - if not callable(revoke): + revoke = _hook(self.peer_clients.get(key), "revoke_grant") + if revoke is None: raise RuntimeError("peer room grant cannot be revoked safely") try: revoke(grant=route.grant) except PeerRunsHTTPError as exc: if not _grant_revoke_is_terminal(exc): raise - hosted_rooms.delete_room_link_records(self.db_path, room_id=room_id) with self._policy_lock: for key, _route in routes: @@ -267,11 +239,7 @@ class HostedRoomService: self.peer_clients.pop(key, None) return len(routes) - def _resolve_member_transport( - self, - binding: HostedRoomBinding, - task: Mapping[str, Any], - ): + def _resolve_member_transport(self, binding: HostedRoomBinding, task: Mapping[str, Any]): payload = task.get("payload", {}) member_id = str(payload.get("target_member_id") or payload.get("target_profile") or "") route = self.peer_routes.get((binding.room_id, member_id)) @@ -284,76 +252,55 @@ class HostedRoomService: raise RuntimeError("peer room client is unavailable") identity = task.get("identity") execution_generation = int(task.get("execution_generation") or 0) - bind_observation = getattr(client, "bind_observation", None) + bind_observation = _hook(client, "bind_observation") if ( - callable(bind_observation) + bind_observation is not None and isinstance(identity, driver.TaskIdentity) - and execution_generation > 0 - ): + and execution_generation > 0): bind_observation(task_id=identity.task_id, execution_generation=execution_generation) def set_status(status: str): return lambda: self._set_route_status(binding.room_id, member_id, status) - tracked_client = _RouteStatusPeerClient( client, on_ready=set_status("ready"), on_reauthorization=set_status("needs_reauthorization"), on_unavailable=set_status("unavailable"), on_refreshed=lambda grant, catalog=None: self._rotate_route_grant( - binding.room_id, member_id, grant, catalog - ), - ) + binding.room_id, member_id, grant, catalog)) self._recover_peer_admission(binding, task, route, tracked_client) return PeerHostedRoomTransport( - binding=binding, - route=route, - client=tracked_client, + binding=binding, route=route, client=tracked_client, source_event_seq=int(payload.get("source_event_seq") or 0), - task_id=getattr(identity, "task_id", None), - execution_generation=execution_generation, - ) + task_id=getattr(identity, "task_id", None), execution_generation=execution_generation) def _recover_peer_admission( - self, - binding: HostedRoomBinding, - task: Mapping[str, Any], - route: PeerMemberRoute, - client: Any, - ) -> None: + self, binding: HostedRoomBinding, task: Mapping[str, Any], route: PeerMemberRoute, + client: Any) -> None: """Rediscover an admitted peer run without advancing its generation.""" - recover = getattr(client, "recover_dispatch", None) + recover = _hook(client, "recover_dispatch") identity = task.get("identity") payload = task.get("payload") execution_generation = int(task.get("execution_generation") or 0) if ( - not callable(recover) + recover is None or not isinstance(identity, driver.TaskIdentity) or not isinstance(payload, Mapping) or execution_generation < 1 - or task.get("status") not in {"running", "indeterminate", "stopping"} - ): + or task.get("status") not in {"running", "indeterminate", "stopping"}): return prompt = payload.get("prompt") source_event_seq = int(payload.get("source_event_seq") or 0) if not isinstance(prompt, str) or source_event_seq < 1 or not route.trace_id: raise RuntimeError("peer room admission identity is unavailable for recovery") dispatch = build_member_dispatch( - binding=binding, - route=route, - room_id=identity.room_id, - task_id=identity.task_id, - target_profile=route.target_profile, - execution_generation=execution_generation, - source_event_seq=source_event_seq, - prompt=prompt, - trace_id=route.trace_id, - ) + binding=binding, route=route, room_id=identity.room_id, task_id=identity.task_id, + target_profile=route.target_profile, execution_generation=execution_generation, + source_event_seq=source_event_seq, prompt=prompt, trace_id=route.trace_id) recover(dispatch=dispatch.as_mapping(), grant=route.grant) def _member_is_peer(self, room_id: str, member_id: str) -> bool: - room = hosted_rooms.room_state(self.db_path, room_id=room_id) - for member in room.get("members") or []: + for member in self._room(room_id).get("members") or []: if not isinstance(member, Mapping): continue if str(member.get("member_id") or member.get("profile") or "") != member_id: @@ -369,15 +316,10 @@ class HostedRoomService: return self._peer_route_status[key] = status hosted_room_links.mark_room_link_status( - self.db_path, room_id=room_id, member_id=member_id, status=status - ) + self.db_path, room_id=room_id, member_id=member_id, status=status) def _set_pending_action( - self, - room_id: str, - member_id: str, - action: Mapping[str, Any] | None, - ) -> None: + self, room_id: str, member_id: str, action: Mapping[str, Any] | None) -> None: key = (room_id, member_id) with self._policy_lock: if action is None: @@ -386,11 +328,7 @@ class HostedRoomService: self._pending_actions[key] = {**action, "member_id": member_id} def _rotate_route_grant( - self, - room_id: str, - member_id: str, - grant: str, - catalog: GatewayRoomCatalog | None = None, + self, room_id: str, member_id: str, grant: str, catalog: GatewayRoomCatalog | None = None ) -> None: """Persist a target-refreshed scoped grant before publishing it live.""" key = (room_id, member_id) @@ -399,12 +337,9 @@ class HostedRoomService: raise RuntimeError("peer room route is unavailable") stored = next( ( - link - for link in hosted_room_links.load_room_links(self.db_path) - if (link.room_id, link.member_id) == key - ), - None, - ) + link for link in hosted_room_links.load_room_links(self.db_path) + if (link.room_id, link.member_id) == key), + None) if stored is None: raise RuntimeError("peer room route cannot be renewed before persistence") digests = {} @@ -415,30 +350,18 @@ class HostedRoomService: or PROTOCOL_VERSION not in catalog.protocol_versions or "direct" not in catalog.link_modes or not catalog.text - or catalog.execution_policy.policy_digest != route.execution_policy_digest - ): + or catalog.execution_policy.policy_digest != route.execution_policy_digest): self._set_route_status(room_id, member_id, "needs_reauthorization") raise RuntimeError( - "peer room execution policy changed; reauthorization is required" - ) + "peer room execution policy changed; reauthorization is required") digests = { "capability_digest": catalog.catalog_digest, - "execution_policy_digest": catalog.execution_policy.policy_digest, - } + "execution_policy_digest": catalog.execution_policy.policy_digest} rotated_route = replace(route, grant=grant, **digests) - hosted_room_links.save_room_link( - self.db_path, - hosted_room_links.make_stored_link( - room_id=room_id, - member_id=member_id, - target_url=stored.target_url, - target_profile=stored.target_profile, - grant=grant, - catalog=catalog or stored.catalog, - cancellation_scope_id=stored.cancellation_scope_id, - trace_id=stored.trace_id, - ), - ) + self._save_link( + room_id=room_id, member_id=member_id, target_url=stored.target_url, + target_profile=stored.target_profile, grant=grant, catalog=catalog or stored.catalog, + cancellation_scope_id=stored.cancellation_scope_id, trace_id=stored.trace_id) with self._policy_lock: self.peer_routes[key] = rotated_route self._peer_route_status[key] = "ready" @@ -448,8 +371,7 @@ class HostedRoomService: rows = [ {"room_id": key[0], "member_id": key[1], "status": status} for key, status in self._peer_route_status.items() - if room_id is None or key[0] == room_id - ] + if room_id is None or key[0] == room_id] return sorted(rows, key=lambda row: (row["room_id"], row["member_id"])) def _events(self, room_id: str) -> list[dict[str, Any]]: @@ -457,11 +379,7 @@ class HostedRoomService: cursor = 0 while True: page = hosted_rooms.read_events( - self.db_path, - room_id=room_id, - since_seq=cursor, - limit=hosted_rooms.MAX_LOG_LIMIT, - ) + self.db_path, room_id=room_id, since_seq=cursor, limit=hosted_rooms.MAX_LOG_LIMIT) rows = page.get("events") if isinstance(rows, list): events.extend(row for row in rows if isinstance(row, dict)) @@ -478,9 +396,7 @@ class HostedRoomService: def _policy_snapshot(self, room: Mapping[str, Any]) -> PolicySnapshot: return self.policy_checkpoint.snapshot( - room_id=str(room["room_id"]), - latest_seq=int(room["latest_seq"]), - ) + room_id=str(room["room_id"]), latest_seq=int(room["latest_seq"])) def _publish_terminal_tasks(self, room: Mapping[str, Any]) -> bool: changed = False @@ -490,102 +406,73 @@ class HostedRoomService: for task in driver.list_tasks(self.db_path, room_id=room_id, status=status): execution_generation = int(task["execution_generation"]) if self.policy_checkpoint.publication_exists( - room_id=room_id, - task_id=task["identity"].task_id, - status=status, - execution_generation=execution_generation, - ): + room_id=room_id, task_id=task["identity"].task_id, status=status, + execution_generation=execution_generation): continue task_events = self.policy_checkpoint.events_for_task( - room_id=room_id, - source_event_seq=int(task["payload"]["source_event_seq"]), - ) + room_id=room_id, source_event_seq=int(task["payload"]["source_event_seq"])) plan = discussion.reconstruct_task_plan( - room, task_events, task, local_profiles=local_profiles - ) + room, task_events, task, local_profiles=local_profiles) publication = discussion.plan_publication( - room, - task_events, - plan, - status=status, - result=task.get("result"), + room, task_events, plan, status=status, result=task.get("result"), execution_generation=execution_generation if status == "deferred" else None, - local_profiles=local_profiles, - ) + local_profiles=local_profiles) self._append_plan(room_id, publication) changed = True return changed def _append_room_status( - self, - room: Mapping[str, Any], - decision: discussion.DiscussionDecision, - ) -> None: + self, room: Mapping[str, Any], decision: discussion.DiscussionDecision) -> None: if decision.discussion_event_id is None: return + gateway_id, epoch = _authority(room) hosted_rooms.append_event( self.db_path, room_id=str(room["room_id"]), event_id=f"dactivity:{decision.discussion_event_id}:{decision.reason}", kind="room.activity", - actor={"kind": "gateway", "id": str(room["authority_gateway_id"])}, + actor={"kind": "gateway", "id": gateway_id}, payload={ - "status": decision.status, - "reason_code": decision.reason, + "status": decision.status, "reason_code": decision.reason, "thread_id": decision.thread_id, - "discussion_event_id": decision.discussion_event_id, - }, - authority_gateway_id=str(room["authority_gateway_id"]), - authority_epoch=int(room["authority_epoch"]), - ) + "discussion_event_id": decision.discussion_event_id}, + authority_gateway_id=gateway_id, + authority_epoch=epoch) def prepare_room(self, binding: HostedRoomBinding) -> None: with self._policy_lock: - room = hosted_rooms.room_state(self.db_path, room_id=binding.room_id) + room = self._room(binding.room_id) snapshot = self._policy_snapshot(room) if self._publish_terminal_tasks(room): - room = hosted_rooms.room_state(self.db_path, room_id=binding.room_id) + room = self._room(binding.room_id) snapshot = self._policy_snapshot(room) self.policy_checkpoint.compact_completed(room_id=binding.room_id) driver.prune_published_terminal_tasks( - self.db_path, room_id=binding.room_id, clock=self.runtime.clock - ) + self.db_path, room_id=binding.room_id, clock=self.runtime.clock) if any(True for _ in self._list_tasks(binding.room_id, _LIVE_STATUSES)): return decision = discussion.plan_next_task( - room, - list(snapshot.events), - local_profiles=self.local_profiles(), - initial_watermarks=snapshot.watermarks, - ) + room, list(snapshot.events), local_profiles=self.local_profiles(), + initial_watermarks=snapshot.watermarks) if decision.status == "task" and decision.task is not None: driver.admit_task( - self.db_path, - decision.task.identity, - payload=decision.task.payload, - clock=time.time, - ) - # A stop can race the policy read from another process. Re-read - # after admission and cancel before the runtime can execute a - # task whose source event is now behind the room stop fence. - fresh_room = hosted_rooms.room_state(self.db_path, room_id=binding.room_id) - stopped_through_seq = self._policy_snapshot(fresh_room).stopped_through_seq + self.db_path, decision.task.identity, payload=decision.task.payload, + clock=time.time) + # A stop can race the policy read from another process. Re-read after + # admission and cancel before the runtime can execute a task whose + # source event is now behind the room stop fence. + stopped_through_seq = self._policy_snapshot( + self._room(binding.room_id) + ).stopped_through_seq if ( decision.source_event_seq is not None - and decision.source_event_seq < stopped_through_seq - ): + and decision.source_event_seq < stopped_through_seq): self.runtime.cancel( - decision.task.identity, - cancel_id=f"stop-fence:{stopped_through_seq}", - ) + decision.task.identity, cancel_id=f"stop-fence:{stopped_through_seq}") elif decision.status in {"settled", "bounded"}: self._append_room_status(room, decision) - def publish_terminal( - self, - binding: HostedRoomBinding, - _task: Mapping[str, Any], - ) -> None: + def publish_terminal(self, binding: HostedRoomBinding, _task: Mapping[str, Any]) -> None: self.prepare_room(binding) self.runtime.wakeup() @@ -597,32 +484,21 @@ class HostedRoomService: name=name, members=[ { - "member_id": member.member_id, - "profile": member.profile, - "handle": member.handle, - "target": dict(member.target or {}), - **({"display_name": member.display_name} if member.display_name else {}), - } - for member in normalized - ], - authority_gateway_id=hosted_rooms.local_authority_gateway_id(), - ) + "member_id": member.member_id, "profile": member.profile, + "handle": member.handle, "target": dict(member.target or {}), + **({"display_name": member.display_name} if member.display_name else {})} + for member in normalized], + authority_gateway_id=hosted_rooms.local_authority_gateway_id()) self.runtime.wakeup() return room def send(self, *, room_id: str, event_id: str, payload: Any) -> dict[str, Any]: normalized = discussion.validate_user_payload(payload) - room = self._owned_room(room_id) + gateway_id, epoch = _authority(self._owned_room(room_id)) event = hosted_rooms.append_event( - self.db_path, - room_id=room_id, - event_id=event_id, - kind="message.user", - actor={"kind": "user", "id": "desktop"}, - payload=normalized, - authority_gateway_id=str(room["authority_gateway_id"]), - authority_epoch=int(room["authority_epoch"]), - ) + self.db_path, room_id=room_id, event_id=event_id, kind="message.user", + actor={"kind": "user", "id": "desktop"}, payload=normalized, + authority_gateway_id=gateway_id, authority_epoch=epoch) binding = next((b for b in self.bindings() if b.room_id == room_id), None) if binding is None: raise hosted_rooms.RoomNotFoundError("hosted room not found") @@ -631,39 +507,27 @@ class HostedRoomService: return event def stop_room( - self, - room_id: str, - *, - cancel_id: str, - require_acknowledged: bool = False, - ) -> int: - room = self._owned_room(room_id) + self, room_id: str, *, cancel_id: str, require_acknowledged: bool = False) -> int: + gateway_id, epoch = _authority(self._owned_room(room_id)) hosted_rooms.request_room_stop( - self.db_path, - room_id=room_id, - cancel_id=cancel_id, - expected_gateway_id=str(room["authority_gateway_id"]), - expected_epoch=int(room["authority_epoch"]), - ) + self.db_path, room_id=room_id, cancel_id=cancel_id, expected_gateway_id=gateway_id, + expected_epoch=epoch) cancelled = 0 pending = 0 with self._policy_lock: tasks = { (task["identity"].room_id, task["identity"].task_id): task - for task in self._list_tasks(room_id, _STOPPABLE_STATUSES) - } + for task in self._list_tasks(room_id, _STOPPABLE_STATUSES)} for task in tasks.values(): task_cancel_id = ( - str(task.get("cancel_id") or "") if task.get("status") == "stopping" else "" - ) - result = self.runtime.cancel(task["identity"], cancel_id=task_cancel_id or cancel_id) + str(task.get("cancel_id") or "") if task.get("status") == "stopping" else "") + result = self.runtime.cancel( + task["identity"], cancel_id=task_cancel_id or cancel_id) cancelled += 1 if result["status"] == "stopping": pending += 1 if require_acknowledged and pending: - raise RuntimeError( - "room work is still stopping; retry deletion after Stop completes" - ) + raise RuntimeError("room work is still stopping; retry deletion after Stop completes") self.runtime.wakeup() return cancelled @@ -671,26 +535,16 @@ class HostedRoomService: """Retry one uncertain or deferred task only after explicit user action.""" task = next( ( - candidate - for candidate in self._list_tasks(room_id, ("indeterminate", "deferred")) - if candidate["identity"].task_id == task_id - ), - None, - ) + candidate for candidate in self._list_tasks(room_id, _RETRYABLE_STATUSES) + if candidate["identity"].task_id == task_id), + None) if task is None: raise driver.InvalidTaskTransitionError("no retryable room task matches task_id") return self.runtime.retry_indeterminate(task["identity"]) def approve_room_task( - self, - room_id: str, - *, - member_id: str, - task_id: str, - execution_generation: int, - choice: str, - request_id: str | None = None, - ) -> Mapping[str, Any]: + self, room_id: str, *, member_id: str, task_id: str, execution_generation: int, + choice: str, request_id: str | None = None) -> Mapping[str, Any]: """Resolve one exact local or peer approval and wake room observation.""" key = (room_id, member_id) route = self.peer_routes.get(key) @@ -704,29 +558,22 @@ class HostedRoomService: pending is not None and str(pending.get("request_id") or "") == requested_approval_id and pending.get("task_id") == task_id - and int(pending.get("execution_generation") or 0) == execution_generation - ) - + and int(pending.get("execution_generation") or 0) == execution_generation) if not requested_approval_id or not matches(action): raise RuntimeError("room approval is no longer pending") if choice not in {"once", "deny"}: raise RuntimeError("room approval choice must be once or deny") - approve = getattr(client, "approve_receipt", None) - if route is not None and callable(approve): + approve = _hook(client, "approve_receipt") + if route is not None and approve is not None: result = approve( - task_id=task_id, - execution_generation=execution_generation, - request_id=requested_approval_id, - choice=choice, - grant=route.grant, - ) + task_id=task_id, execution_generation=execution_generation, + request_id=requested_approval_id, choice=choice, grant=route.grant) else: session_id = str(action.get("session_id") or "") if not session_id: raise RuntimeError("local room approval identity is unavailable") result = self.rpc.approve( - session_id=session_id, request_id=requested_approval_id, choice=choice - ) + session_id=session_id, request_id=requested_approval_id, choice=choice) if result is None: raise RuntimeError("room approval target is unavailable") with self._policy_lock: @@ -746,39 +593,28 @@ class HostedRoomService: pending_actions = [ {"kind": "retry", "task_id": task["identity"].task_id} for task in tasks - if task["status"] in {"indeterminate", "deferred"} - ] + if task["status"] in _RETRYABLE_STATUSES] with self._policy_lock: pending_actions.extend( dict(action) for (action_room_id, _member_id), action in self._pending_actions.items() - if action_room_id == room_id - ) + if action_room_id == room_id) return { "running": runtime["running"], "working": bool( - counts.get("running") or counts.get("queued") or counts.get("stopping") - ), + counts.get("running") or counts.get("queued") or counts.get("stopping")), "blocked": room_id in runtime["blocked_rooms"] or bool(counts.get("indeterminate") or counts.get("stopping")), "counts": dict(counts), "pending_actions": pending_actions, - "peer_routes": self._route_statuses(room_id), - } + "peer_routes": self._route_statuses(room_id)} class _RouteStatusPeerClient: """Classify scoped-auth failures without exposing route credentials.""" def __init__( - self, - client, - *, - on_ready, - on_reauthorization, - on_unavailable, - on_refreshed, - ) -> None: + self, client, *, on_ready, on_reauthorization, on_unavailable, on_refreshed) -> None: self._client = client self._on_ready = on_ready self._on_reauthorization = on_reauthorization @@ -788,28 +624,25 @@ class _RouteStatusPeerClient: def _refresh_grant(self, kwargs: dict) -> dict: """Rotate an expiring grant before dispatch; return the kwargs to send. - Refresh failures only escalate to reauthorization when the peer says so - or the grant is already past its hard expiry; otherwise the original - grant is tried as-is. A refreshed catalog whose digests drift from the - dispatch is a policy change and is refused before any dispatch. + Refresh failures only escalate to reauthorization when the peer says so or the + grant is already past its hard expiry; otherwise the original grant is tried + as-is. A refreshed catalog whose digests drift from the dispatch is a policy + change and is refused before any dispatch. """ grant = kwargs["grant"] if not room_grant_needs_dispatch_refresh(grant): return kwargs checked = HostedMemberDispatch.from_mapping(kwargs["dispatch"]) - refresh = getattr(self._client, "refresh_grant", None) - if not callable(refresh): + refresh = _hook(self._client, "refresh_grant") + if refresh is None: return kwargs try: refreshed = refresh( - grant=grant, - capability_digest=checked.capability_digest, - execution_policy_digest=checked.execution_policy_digest, - ) + grant=grant, capability_digest=checked.capability_digest, + execution_policy_digest=checked.execution_policy_digest) except Exception as exc: if getattr(exc, "needs_reauthorization", False) or ( - room_grant_needs_dispatch_refresh(grant, leeway_seconds=0) - ): + room_grant_needs_dispatch_refresh(grant, leeway_seconds=0)): self._on_reauthorization() raise return kwargs @@ -821,14 +654,17 @@ class _RouteStatusPeerClient: refreshed_catalog = GatewayRoomCatalog.from_mapping(refreshed.get("catalog")) drift = None if refreshed_catalog.execution_policy.policy_digest != checked.execution_policy_digest: - drift = ("peer room execution policy needs reauthorization", "room_execution_policy_changed") + drift = ( + "peer room execution policy needs reauthorization", + "room_execution_policy_changed") elif refreshed_catalog.catalog_digest != checked.capability_digest: - drift = ("peer room capabilities need reauthorization", "room_capability_catalog_changed") + drift = ( + "peer room capabilities need reauthorization", "room_capability_catalog_changed" + ) if drift is not None: self._on_reauthorization() raise PeerRunsHTTPError( - drift[0], status_code=403, error_code=drift[1], not_admitted=True - ) + drift[0], status_code=403, error_code=drift[1], not_admitted=True) self._on_refreshed(replacement, refreshed_catalog) return {**kwargs, "grant": replacement} @@ -851,5 +687,4 @@ class _RouteStatusPeerClient: if name != "prepare": self._on_ready() return result - return tracked diff --git a/tui_gateway/methods_groups.py b/tui_gateway/methods_groups.py index e83f26a761..e286266174 100644 --- a/tui_gateway/methods_groups.py +++ b/tui_gateway/methods_groups.py @@ -4,11 +4,10 @@ These methods expose durable room identity, replay, and the process-owned same-gateway Discussion driver. ``groups.capabilities`` keeps that boundary machine-readable so older clients stay on the renderer-owned room path. -Handlers are rebound onto server.py's globals at install (see method_ctx.py), -so bodies see only server globals plus the names methods_bot_relay.register -publishes; module-private helpers reach them through keyword defaults. -``_room_method`` wraps each handler with the shared service-lookup / -error-code envelope so the bodies hold only the room logic. +Handlers are rebound onto server.py's globals at install (see method_ctx.py), so +bodies see only server globals plus the names methods_bot_relay.register publishes; +module-private helpers reach them through keyword defaults. ``_room_method`` wraps +each handler with the shared service-lookup / error-code envelope. """ from .method_ctx import HandlerRegistry @@ -21,25 +20,10 @@ method = _registry.method #: Wire order of ``groups.capabilities.methods``; every one runs on the RPC pool. _METHODS = ( - "groups.capabilities", - "groups.list", - "groups.create", - "groups.state", - "groups.send", - "groups.rename", - "groups.log", - "groups.disband", - "groups.replicate", - "groups.replica_state", - "groups.promote", - "groups.demote", - "groups.stop", - "groups.retry", - "groups.approve", - "groups.peer.invite", - "groups.peer.revoke", - "groups.peer.register", -) + "groups.capabilities", "groups.list", "groups.create", "groups.state", "groups.send", + "groups.rename", "groups.log", "groups.disband", "groups.replicate", "groups.replica_state", + "groups.promote", "groups.demote", "groups.stop", "groups.retry", "groups.approve", + "groups.peer.invite", "groups.peer.revoke", "groups.peer.register") LONG_HANDLERS = frozenset(_METHODS) _service_lock = threading.Lock() @@ -47,9 +31,7 @@ _run_store_lock = threading.Lock() _bound_server = None _service = None -_WORKER_UNAVAILABLE = ( - "Group Chat worker is unavailable. Restart the Hermes gateway and try again." -) +_WORKER_UNAVAILABLE = "Group Chat worker is unavailable. Restart the Hermes gateway and try again." _DRIVER_UNAVAILABLE = "hosted room driver is unavailable" @@ -67,7 +49,6 @@ def start_hosted_room_service(): return None from gateway.hosted_rooms import default_db_path from tui_gateway.hosted_room_service import HostedRoomService - db_path = default_db_path() with _service_lock: if _service is not None and _service.db_path != db_path: @@ -128,7 +109,6 @@ def _requested_profile(params: dict) -> str: def _api_server_key(profile: str | None = None) -> str: 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 "" @@ -137,7 +117,6 @@ def _api_server_key(profile: str | None = None) -> str: return str(build_profile_secret_scope(home).get("API_SERVER_KEY") or "").strip() try: from agent.secret_scope import get_secret - scoped = (get_secret("API_SERVER_KEY", "") or "").strip() if scoped: return scoped @@ -150,7 +129,6 @@ def _profile_execution_policy(profile: str) -> dict: """Resolve execution policy under the exact multiplexed profile home.""" from gateway.hosted_room_execution_policy import execution_policy_mapping 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) @@ -172,12 +150,10 @@ def _room_link_run_storage_durable() -> bool: 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; a separately enabled API adapter uses the same file. + # 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. from gateway.platforms.api_server import RunIdempotencyStore - with _run_store_lock: store = getattr(_bound_server, "_run_idempotency_store", None) if store is None: @@ -189,35 +165,27 @@ def _room_link_run_storage_durable() -> bool: def _local_catalog(installation_id: str, profile: str, execution_policy: dict) -> dict: """Advertise this gateway's direct-only, text-only RoomLink catalog.""" from gateway.hosted_room_peer import PROTOCOL_VERSION, local_catalog_mapping - return local_catalog_mapping( - installation_id=installation_id, - protocol_versions=(PROTOCOL_VERSION,), - link_modes=("direct",), - text=True, - attachments=False, - target_profile=profile, - execution_policy=execution_policy, - ) + installation_id=installation_id, protocol_versions=(PROTOCOL_VERSION,), + link_modes=("direct",), text=True, attachments=False, target_profile=profile, + execution_policy=execution_policy) + + +def _grant_expiry(claims: dict) -> float: + return float(claims.get("status_expires_at", claims["expires_at"])) def _room_method( - name: str, - *, - code: int, - room_code: int | None = None, - replica_only: bool = False, - with_reason: bool = True, - service_code: int | None = None, - service_message: str = _DRIVER_UNAVAILABLE, -): + name: str, *, code: int, room_code: int | None = None, replica_only: bool = False, + with_reason: bool = True, service_code: int | None = None, + service_message: str = _DRIVER_UNAVAILABLE, db: bool = False): """Register ``fn`` under ``name`` with the shared hosted-room error envelope. - ``service_code`` set: the live service is required and passed as a third - argument; when absent the handler fails with that code. ``room_code`` - maps ``HostedRoomError`` (or only ``ReplicaError`` when ``replica_only``) - to a 4xxx client error, attaching ``{"reason"}`` data when ``with_reason``; - any other exception maps to ``code``. + ``service_code`` set: the live service is required and passed as a third argument; + when absent the handler fails with that code. ``db``: the default room db path is + passed as the next argument. ``room_code`` maps ``HostedRoomError`` (or only + ``ReplicaError`` when ``replica_only``) to a 4xxx client error, attaching + ``{"reason"}`` data when ``with_reason``; any other exception maps to ``code``. """ def dec(fn): @@ -228,25 +196,25 @@ def _room_method( if service is None: return _err(rid, service_code, service_message) args += (service,) + if db: + from gateway.hosted_rooms import default_db_path + args += (default_db_path(),) try: return fn(*args) except Exception as exc: if room_code is not None: from gateway.hosted_rooms import HostedRoomError - klass = HostedRoomError if replica_only: from gateway.hosted_room_replicas import ReplicaError - klass = ReplicaError if isinstance(exc, klass): reason = getattr(exc, "reason", None) if with_reason else None - return _err(rid, room_code, str(exc), {"reason": reason} if reason else None) + data = {"reason": reason} if reason else None + return _err(rid, room_code, str(exc), data) return _err(rid, code, str(exc)) - handler.__doc__ = fn.__doc__ return method(name)(handler) - return dec @@ -254,73 +222,48 @@ def _room_method( def _(rid, params: dict, _catalog=_local_catalog, _methods=_METHODS) -> dict: """Describe the hosted-room protocol implemented by this gateway.""" from gateway.hosted_rooms import MAX_LOG_LIMIT, PROTOCOL_VERSION, local_authority_gateway_id - service = get_hosted_room_service() driver_ready = bool(service and service.runtime.status()["running"]) try: from gateway.hosted_room_peer import gateway_room_grant_secret - profile = _requested_profile(params) if not _room_link_run_storage_durable(): raise ValueError("durable run idempotency storage is required") gateway_room_grant_secret() - catalog = _catalog(local_authority_gateway_id(), profile, _profile_execution_policy(profile)) + 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" - ), - } - return _ok( - rid, - { - "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, - "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, - }, - ) + "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, + "persistent_process": bool(room_link.get("catalog", {}).get("persistent_process", False)), + "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}) -@_room_method("groups.peer.invite", code=4120) -def _(rid, params: dict, _catalog=_local_catalog) -> dict: +@_room_method("groups.peer.invite", code=4120, db=True) +def _(rid, params: dict, db_path, _catalog=_local_catalog, _expiry=_grant_expiry) -> dict: """Mint one target-issued room/profile grant for a prospective home.""" from gateway.hosted_room_peer import ( - decode_room_grant, - gateway_room_grant_secret, - issue_room_grant, - ) - from gateway import hosted_rooms - + decode_room_grant, gateway_room_grant_secret, issue_room_grant) + from gateway.hosted_rooms import local_authority_gateway_id, reserve_peer_room if not _room_link_run_storage_durable(): raise ValueError("durable run idempotency storage is required") - installation_id = hosted_rooms.local_authority_gateway_id() + installation_id = local_authority_gateway_id() profile = _requested_profile(params) ttl = float(params.get("ttl_seconds", 3600)) if not 60 <= ttl <= 24 * 60 * 60: @@ -328,56 +271,35 @@ def _(rid, params: dict, _catalog=_local_catalog) -> dict: grant_secret = gateway_room_grant_secret() execution_policy = _profile_execution_policy(profile) token = issue_room_grant( - grant_secret, - grant_id=str(params.get("grant_id") or f"grant-{os.urandom(16).hex()}"), + grant_secret, grant_id=str(params.get("grant_id") or f"grant-{os.urandom(16).hex()}"), room_id=str(params.get("room_id") or ""), home_install_id=str(params.get("home_install_id") or ""), authority_gateway_id=str(params.get("authority_gateway_id") or ""), authority_epoch=int(params.get("authority_epoch") or 0), - member_id=str(params.get("member_id") or ""), - target_install_id=installation_id, - target_profile=profile, - execution_policy_digest=execution_policy["policy_digest"], - ttl_seconds=ttl, - ) + member_id=str(params.get("member_id") or ""), target_install_id=installation_id, + target_profile=profile, execution_policy_digest=execution_policy["policy_digest"], + ttl_seconds=ttl) claims = decode_room_grant(grant_secret, token, permission="status") - hosted_rooms.reserve_peer_room( - hosted_rooms.default_db_path(), - claims=claims, - expires_at=float(claims.get("status_expires_at", claims["expires_at"])), - ) + reserve_peer_room(db_path, claims=claims, expires_at=_expiry(claims)) catalog = _catalog(installation_id, profile, execution_policy) - return _ok( - rid, - { - "grant": token, - "target_profile": profile, - "catalog": catalog, - "endpoint": catalog["endpoint"], - }, - ) + return _ok(rid, { + "grant": token, "target_profile": profile, "catalog": catalog, + "endpoint": catalog["endpoint"]}) -@_room_method("groups.peer.revoke", code=4122) -def _(rid, params: dict) -> dict: +@_room_method("groups.peer.revoke", code=4122, db=True) +def _(rid, params: dict, db_path, _expiry=_grant_expiry) -> dict: """Revoke one target-issued grant using its exact profile scope.""" - from gateway import hosted_rooms from gateway.hosted_room_peer import decode_room_grant, gateway_room_grant_secret - + from gateway.hosted_rooms import local_authority_gateway_id, revoke_room_grant_scope profile = _requested_profile(params) claims = decode_room_grant( - gateway_room_grant_secret(), str(params.get("grant") or ""), permission="status" - ) + gateway_room_grant_secret(), str(params.get("grant") or ""), permission="status") if ( claims["target_profile"] != profile - or claims["target_install_id"] != hosted_rooms.local_authority_gateway_id() - ): + or claims["target_install_id"] != local_authority_gateway_id()): raise ValueError("room grant target does not match this profile") - hosted_rooms.revoke_room_grant_scope( - hosted_rooms.default_db_path(), - claims=claims, - expires_at=float(claims.get("status_expires_at", claims["expires_at"])), - ) + revoke_room_grant_scope(db_path, claims=claims, expires_at=_expiry(claims)) return _ok(rid, {"revoked": True}) @@ -385,20 +307,14 @@ def _(rid, params: dict) -> dict: def _(rid, params: dict, service) -> dict: """Register and probe one scoped target route on the room home.""" from gateway.hosted_room_peer import ( - GatewayRoomCatalog, - PROTOCOL_VERSION as ROOM_LINK_PROTOCOL_VERSION, - validate_room_link_url, - ) + GatewayRoomCatalog, PROTOCOL_VERSION as ROOM_LINK_PROTOCOL_VERSION, validate_room_link_url) from gateway.hosted_rooms import local_authority_gateway_id, room_state from tui_gateway.hosted_room_peer_http import PeerRunsHTTPClient from tui_gateway.hosted_room_peer_transport import PeerMemberRoute - target_url, transport_security = validate_room_link_url(params.get("target_url")) catalog = GatewayRoomCatalog.from_mapping(params.get("catalog")) if ROOM_LINK_PROTOCOL_VERSION not in catalog.protocol_versions: - raise ValueError( - f"target does not support RoomLink protocol v{ROOM_LINK_PROTOCOL_VERSION}" - ) + raise ValueError(f"target does not support RoomLink protocol v{ROOM_LINK_PROTOCOL_VERSION}") if "direct" not in catalog.link_modes: raise ValueError("target does not support a direct RoomLink") target_profile = str(params.get("target_profile") or "") @@ -410,8 +326,7 @@ def _(rid, params: dict, service) -> dict: 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 - ): + 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 "") @@ -423,8 +338,7 @@ def _(rid, params: dict, service) -> dict: or probe.get("authority_gateway_id") != home_room.get("authority_gateway_id") or int(probe.get("authority_epoch") or 0) != int(home_room.get("authority_epoch") or 0) or probe.get("member_id") != member_id - or probe.get("target_profile") != target_profile - ): + or probe.get("target_profile") != target_profile): raise ValueError("room grant scope does not match this route") route = PeerMemberRoute( home_install_id=home_install_id, @@ -434,53 +348,33 @@ def _(rid, params: dict, service) -> dict: capability_digest=catalog.catalog_digest, execution_policy_digest=catalog.execution_policy.policy_digest, cancellation_scope_id=str( - params.get("cancellation_scope_id") or f"cancel-{params.get('room_id') or ''}" - ), + params.get("cancellation_scope_id") or f"cancel-{params.get('room_id') or ''}"), trace_id=str(params.get("trace_id") or f"trace-{os.urandom(16).hex()}"), - grant=grant, - ) + grant=grant) service.register_peer_route( - room_id=room_id, - member_id=member_id, - route=route, - client=client, - target_url=target_url, - catalog=catalog, - ) - return _ok( - rid, - { - "registered": True, - "mode": "direct", - "transport_security": transport_security, - "target_install_id": catalog.installation_id, - "target_profile": target_profile, - }, - ) + room_id=room_id, member_id=member_id, route=route, client=client, target_url=target_url, + catalog=catalog) + return _ok(rid, { + "registered": True, "mode": "direct", "transport_security": transport_security, + "target_install_id": catalog.installation_id, "target_profile": target_profile}) -@_room_method("groups.list", code=5110) -def _(rid, params: dict) -> dict: +@_room_method("groups.list", code=5110, db=True) +def _(rid, params: dict, db_path) -> dict: """List rooms hosted by this gateway.""" - from gateway.hosted_rooms import MAX_ROOM_LIST_LIMIT, default_db_path, list_rooms - + from gateway.hosted_rooms import MAX_ROOM_LIST_LIMIT, list_rooms limit = params.get("limit", MAX_ROOM_LIST_LIMIT) offset = params.get("offset", 0) rooms = list_rooms( - default_db_path(), - include_disbanded=params.get("include_disbanded") is True, - limit=limit, - offset=offset, - ) - return _ok( - rid, - {"rooms": rooms, "next_offset": offset + limit if len(rooms) == limit else None}, - ) + db_path, include_disbanded=params.get("include_disbanded") is True, limit=limit, + offset=offset) + next_offset = offset + limit if len(rooms) == limit else None + return _ok(rid, {"rooms": rooms, "next_offset": next_offset}) @_room_method( - "groups.create", code=5111, room_code=4110, service_code=4123, service_message=_WORKER_UNAVAILABLE -) + "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. @@ -488,23 +382,17 @@ def _(rid, params: dict, service) -> dict: derived from this gateway's stable install identity, never from the client. """ room = service.create_room( - room_id=params.get("room_id"), - name=params.get("name"), - members=params.get("members"), - ) + room_id=params.get("room_id"), name=params.get("name"), members=params.get("members")) return _ok(rid, {"room": room}) -@_room_method("groups.state", code=5115, room_code=4114) -def _(rid, params: dict) -> dict: +@_room_method("groups.state", code=5115, room_code=4114, db=True) +def _(rid, params: dict, db_path) -> dict: """Return one hosted room's replay cursor and fenced authority state.""" - from gateway.hosted_rooms import default_db_path, room_state - + from gateway.hosted_rooms import room_state room = room_state( - default_db_path(), - room_id=params.get("room_id"), - include_disbanded=params.get("include_disbanded") is True, - ) + db_path, room_id=params.get("room_id"), + include_disbanded=params.get("include_disbanded") is True) service = get_hosted_room_service() result = {"room": room} if service is not None and room.get("disbanded_at") is None: @@ -523,80 +411,54 @@ def _(rid, params: dict, service) -> dict: method; the actor is server-owned rather than trusted from params. """ from gateway.hosted_rooms import user_event_id - client_event_id = params.get("event_id") event = service.send( - room_id=params.get("room_id"), - event_id=user_event_id(client_event_id), - payload=params.get("payload"), - ) - return _ok( - rid, - { - "event": event, - "client_event_id": client_event_id, - "accepted": True, - "driver_started": True, - }, - ) + room_id=params.get("room_id"), event_id=user_event_id(client_event_id), + payload=params.get("payload")) + return _ok(rid, { + "event": event, "client_event_id": client_event_id, "accepted": True, "driver_started": True + }) -@_room_method("groups.rename", code=5117, room_code=4117) -def _(rid, params: dict) -> dict: +@_room_method("groups.rename", code=5117, room_code=4117, db=True) +def _(rid, params: dict, db_path) -> dict: """Rename one hosted room atomically with its replay event.""" - from gateway.hosted_rooms import default_db_path, rename_room - + from gateway.hosted_rooms import rename_room renamed = rename_room( - default_db_path(), - room_id=params.get("room_id"), - event_id=params.get("event_id"), - name=params.get("name"), - ) + db_path, room_id=params.get("room_id"), event_id=params.get("event_id"), + name=params.get("name")) return _ok(rid, {"room": renamed}) @_room_method( - "groups.disband", code=5114, room_code=4113, service_code=4123, service_message=_WORKER_UNAVAILABLE -) + "groups.disband", code=5114, room_code=4113, service_code=4123, + service_message=_WORKER_UNAVAILABLE) def _(rid, params: dict, service) -> dict: """Permanently tombstone a hosted room id.""" from gateway.hosted_rooms import ( - AuthorityConflictError, - RoomHistoryExpiredError, - disband_room, - local_authority_gateway_id, - room_state, - ) - + AuthorityConflictError, RoomHistoryExpiredError, disband_room, local_authority_gateway_id, + room_state) room_id = str(params.get("room_id") or "") def disband_with_state(state: dict | None = None) -> dict: local_gateway_id = local_authority_gateway_id() if state is not None and str(state["authority_gateway_id"]) != local_gateway_id: raise AuthorityConflictError("This Group Chat is managed by another gateway.") - return _ok( - rid, - { - "tombstone": disband_room( - service.db_path, - room_id=params.get("room_id"), - expected_gateway_id=str(local_gateway_id), - expected_epoch=int(state["authority_epoch"] if state is not None else 1), - ) - }, - ) - + tombstone = disband_room( + service.db_path, room_id=params.get("room_id"), + expected_gateway_id=str(local_gateway_id), + expected_epoch=int(state["authority_epoch"] if state is not None else 1)) + return _ok(rid, {"tombstone": tombstone}) try: - existing = room_state(service.db_path, room_id=params.get("room_id"), include_disbanded=True) + existing = room_state( + service.db_path, room_id=params.get("room_id"), include_disbanded=True) except RoomHistoryExpiredError: return disband_with_state() if existing.get("disbanded_at") is not None: return disband_with_state(existing) service.stop_room( - room_id, - cancel_id=str(params.get("cancel_id") or "room-disbanded"), - require_acknowledged=True, - ) + room_id, cancel_id=str(params.get("cancel_id") or "room-disbanded"), + require_acknowledged=True) service.revoke_room_routes(room_id) return disband_with_state(existing) @@ -605,9 +467,7 @@ def _(rid, params: dict, service) -> dict: def _(rid, params: dict, service) -> dict: """Durably cancel queued or running work for one hosted room.""" count = service.stop_room( - str(params.get("room_id") or ""), - cancel_id=str(params.get("cancel_id") or "desktop-stop"), - ) + str(params.get("room_id") or ""), cancel_id=str(params.get("cancel_id") or "desktop-stop")) return _ok(rid, {"cancelled": count}) @@ -615,13 +475,10 @@ def _(rid, params: dict, service) -> dict: def _(rid, params: dict, service) -> dict: """Resolve one exact approval requested by a local or peer room member.""" result = service.approve_room_task( - str(params.get("room_id") or ""), - member_id=str(params.get("member_id") or ""), + str(params.get("room_id") or ""), member_id=str(params.get("member_id") or ""), task_id=str(params.get("task_id") or ""), execution_generation=int(params.get("execution_generation") or 0), - choice=str(params.get("choice") or ""), - request_id=str(params.get("request_id") or ""), - ) + choice=str(params.get("choice") or ""), request_id=str(params.get("request_id") or "")) return _ok(rid, {"approved": True, "result": result}) @@ -629,41 +486,33 @@ def _(rid, params: dict, service) -> dict: def _(rid, params: dict, service) -> dict: """Retry one indeterminate room task after explicit user confirmation.""" task = service.retry_room_task( - str(params.get("room_id") or ""), - task_id=str(params.get("task_id") or ""), - ) + str(params.get("room_id") or ""), task_id=str(params.get("task_id") or "")) if not isinstance(task, dict): task = {} identity = task.get("identity") receipt = { **{ field: str(getattr(identity, field, "") or "") - for field in ("room_id", "task_id", "thread_id", "turn_id") - }, + for field 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), - } + "cancel_generation": int(task.get("cancel_generation") or 0)} return _ok(rid, {"retried": True, "task": receipt}) -@_room_method("groups.log", code=5113, room_code=4112) -def _(rid, params: dict) -> dict: +@_room_method("groups.log", code=5113, room_code=4112, db=True) +def _(rid, params: dict, db_path) -> dict: """Return a monotonic room-log delta after ``since_seq``.""" - from gateway.hosted_rooms import default_db_path, read_events - + from gateway.hosted_rooms import read_events delta = read_events( - default_db_path(), - room_id=params.get("room_id"), - since_seq=params.get("since_seq", 0), - limit=params.get("limit", 100), - include_disbanded=params.get("include_disbanded") is True, - ) + db_path, room_id=params.get("room_id"), since_seq=params.get("since_seq", 0), + limit=params.get("limit", 100), include_disbanded=params.get("include_disbanded") is True) return _ok(rid, delta) -@_room_method("groups.replicate", code=5116, room_code=4116, replica_only=True, with_reason=False) -def _(rid, params: dict) -> dict: +@_room_method( + "groups.replicate", code=5116, room_code=4116, replica_only=True, with_reason=False, db=True) +def _(rid, params: dict, db_path) -> dict: """Persist one authority-stamped replay page into the local replica store. ``page`` is the verbatim ``groups.log`` result read from the room's @@ -671,66 +520,49 @@ def _(rid, params: dict) -> dict: authority-epoch regressions. """ from gateway.hosted_room_replicas import ingest_page - from gateway.hosted_rooms import default_db_path - result = ingest_page( - default_db_path(), - room_id=params.get("room_id"), - room_name=params.get("room_name"), - members=params.get("members"), - page=params.get("page"), - ) + db_path, room_id=params.get("room_id"), room_name=params.get("room_name"), + members=params.get("members"), page=params.get("page")) return _ok(rid, result) @_room_method( - "groups.replica_state", code=5117, room_code=4117, replica_only=True, with_reason=False + "groups.replica_state", code=5117, room_code=4117, replica_only=True, with_reason=False, db=True ) -def _(rid, params: dict) -> dict: +def _(rid, params: dict, db_path) -> dict: """Report the local replica's coverage and authority lineage.""" from gateway.hosted_room_replicas import replica_state - from gateway.hosted_rooms import default_db_path - - return _ok(rid, replica_state(default_db_path(), room_id=params.get("room_id"))) + return _ok(rid, replica_state(db_path, room_id=params.get("room_id"))) -@_room_method("groups.promote", code=5118, room_code=4118, with_reason=False) -def _(rid, params: dict) -> dict: +@_room_method("groups.promote", code=5118, room_code=4118, with_reason=False, db=True) +def _(rid, params: dict, db_path) -> dict: """Continue a replicated room on THIS gateway at ``epoch + 1``. Requires ``confirm: true`` — the caller asserts the previous authority can no longer commit (explicit user action; a lease/quorum driver later). """ from gateway.hosted_room_replicas import promote_replica - from gateway.hosted_rooms import default_db_path - if params.get("confirm") is not True: return _err( - rid, - 4118, + rid, 4118, "promotion requires confirm=true acknowledging the previous " - "authority can no longer commit", - ) + "authority can no longer commit") result = promote_replica( - default_db_path(), - room_id=params.get("room_id"), - reason=params.get("reason", "authority-unreachable"), + db_path, room_id=params.get("room_id"), reason=params.get("reason", "authority-unreachable") ) return _ok(rid, result) -@_room_method("groups.demote", code=5119, room_code=4119, replica_only=True, with_reason=False) -def _(rid, params: dict) -> dict: +@_room_method( + "groups.demote", code=5119, room_code=4119, replica_only=True, with_reason=False, db=True) +def _(rid, params: dict, db_path) -> dict: """Fence this gateway's stale room authority against a proven newer epoch.""" from gateway.hosted_room_replicas import demote_room - from gateway.hosted_rooms import default_db_path - result = demote_room( - default_db_path(), - room_id=params.get("room_id"), + db_path, room_id=params.get("room_id"), observed_gateway_id=params.get("observed_gateway_id"), - observed_epoch=params.get("observed_epoch"), - ) + observed_epoch=params.get("observed_epoch")) return _ok(rid, result)