From 03c191076d1cc24b2d94e00501cabf2bab4edc5d Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:47:09 -0700 Subject: [PATCH] refactor(tui): second pass on tool_progress/ws/session_history/transport (group -25.4% LOC) - ws: _reply raises _SendFailed to end the read loop (one exit path instead of five 'if not await _reply(...): break' ladders); heartbeat/orphan-sweep starts loop-driven; inline streaming-frame check; frame bytes/log strings unchanged. - tool_progress: _verbose_text unifies args/result rendering; summary via _SUMMARY_COUNTERS table; walrus for optional payload fields. - session_history: _coerce_message_text list parts reuse _history_dict_text; dict.update() for inflight turn mutations; compact docstrings (all WHY kept). - transport: _raise_unless_peer_gone (void) classifier; docstrings compacted. Old-vs-new fuzz (2928 cases) identical; WIRE OK vs base2; _methods set identical. --- tui_gateway/session_history.py | 202 +++++++++----------------- tui_gateway/tool_progress.py | 199 ++++++++++---------------- tui_gateway/transport.py | 92 +++++------- tui_gateway/ws.py | 250 ++++++++++++--------------------- 4 files changed, 268 insertions(+), 475 deletions(-) diff --git a/tui_gateway/session_history.py b/tui_gateway/session_history.py index 2cd5637f73..4230051e5c 100644 --- a/tui_gateway/session_history.py +++ b/tui_gateway/session_history.py @@ -1,8 +1,5 @@ -"""Session history/message shaping: image-ref messages, content coercion, history->wire messages, -in-flight turn tracking and turn-failure detail. - -Bodies are rebound onto server.py's globals (method_ctx.bind_module) and reference them bare. -""" +"""Session history/message shaping: image-ref messages, content coercion, history->wire messages, in-flight +turn tracking and turn-failure detail. Bodies are rebound onto server.py's globals (method_ctx.bind_module).""" from __future__ import annotations @@ -12,14 +9,12 @@ from .method_ctx import bind_module def _active_image_routing_identity(agent: Any) -> tuple[str, str]: """Return the live provider/model, falling back before agent startup.""" from agent.auxiliary_client import _read_main_model, _read_main_provider - return (getattr(agent, "provider", "") or _read_main_provider(), getattr(agent, "model", "") or _read_main_model()) def _build_image_ref_message(user_text: str, image_paths: list[str]) -> str: - """Reference attached images by path so the agent analyzes them in-loop with ``vision_analyze``. - Pre-analyzing with the auxiliary vision model blocked submit 60-90s per photo and poisoned - auto-titles with the description.""" + """Reference attached images by path so the agent analyzes them in-loop with ``vision_analyze``: pre- + analyzing with the auxiliary vision model blocked submit 60-90s/photo and poisoned auto-titles.""" prefix = "\n\n".join( f"[The user attached an image: {p.name}]\n[Examine it with the vision_analyze tool using image_url: {p}]" for p in map(Path, image_paths) if p.exists() @@ -31,12 +26,10 @@ def _build_image_ref_message(user_text: str, image_paths: list[str]) -> str: def _build_persist_message_with_image_refs(user_text: str, image_paths: list[str]) -> str: - """Persisted form of the user's message: ``@image:`` directives (the desktop renders them - as images); ``_build_image_ref_message``'s ``image_url:`` hint is model-only, never persisted. - Caption first, directives last: session previews are the first 60 chars of the first user - message, so a leading directive would label the session with a truncated path.""" + """Persisted form of the user's message: ``@image:`` directives (the desktop renders them as + images); ``_build_image_ref_message``'s ``image_url:`` hint is model-only, never persisted. Caption + first, directives last: session previews are the first 60 chars of the first user message.""" from agent.context_references import format_reference_value - text = user_text or "" refs = "\n".join(f"@image:{format_reference_value(p)}" for p in image_paths if Path(p).exists()) if not refs: @@ -45,9 +38,9 @@ def _build_persist_message_with_image_refs(user_text: str, image_paths: list[str def _build_persist_user_message(user_text: str, image_paths: list[str], run_message: Any) -> Any: - """Shape the persisted user turn like the model payload: ``_flush_messages_to_session_db`` ignores - a plain-string override for a list (native-vision) payload, so swap only the text part for the - ``@image:`` form, keep image parts, and drop API-only text parts (barge-in note).""" + """Shape the persisted user turn like the model payload: ``_flush_messages_to_session_db`` ignores a + plain-string override for a list (native-vision) payload, so swap only the text part for the + ``@image:`` form, keep image parts, drop API-only text parts (barge-in note).""" persist_text = _build_persist_message_with_image_refs(user_text, image_paths) if not isinstance(run_message, list): return persist_text @@ -86,41 +79,24 @@ def _history_dict_text(content: dict, *, image_urls: bool) -> str: def _content_display_text(content: Any) -> str: if isinstance(content, list): - parts = (_content_display_text(part).strip() for part in content) - return "\n".join(text for text in parts if text) + return "\n".join(t for t in (_content_display_text(part).strip() for part in content) if t) if isinstance(content, dict): return _history_dict_text(content, image_urls=False) return "" if content is None else str(content) def _coerce_message_text(content: Any) -> str: - """Render ``message['content']`` (str, parts list, or one structured dict) as a plain string. - Image parts keep their URL inline so the desktop's ``extractEmbeddedImages`` and the resume payload - agree with the cached message (else the inline image flashed, then vanished); other structured - shapes become a bracketed placeholder so resume doesn't drop the message.""" + """Render ``message['content']`` (str, parts list, or one structured dict) as a plain string. Image parts + keep their URL inline so the desktop's ``extractEmbeddedImages`` and the resume payload agree with the + cached message (else the inline image flashed, then vanished); other shapes become a placeholder.""" if isinstance(content, list): chunks: list[str] = [] for part in content: - if isinstance(part, str): - chunks.append(part) - continue - if not isinstance(part, dict): - continue - text = part.get("text") - if isinstance(text, str): - chunks.append(text) - continue - kind = part.get("type") - if kind in _HISTORY_TEXT_KINDS: - t = part.get("text") or part.get("content") or "" - if t: - chunks.append(str(t)) - elif kind in _HISTORY_IMAGE_KINDS: - chunks.append(f"\n{_history_part_image_url(part) or '[image]'}") - elif kind in _HISTORY_AUDIO_KINDS: - chunks.append("\n[audio]") - elif kind: - chunks.append(f"\n[{kind}]") + if isinstance(part, str) or (isinstance(part, dict) and isinstance(part.get("text"), str)): + chunks.append(part if isinstance(part, str) else part["text"]) + elif isinstance(part, dict) and part.get("type"): + rendered = _history_dict_text(part, image_urls=True) + chunks.append(rendered if part["type"] in _HISTORY_TEXT_KINDS else f"\n{rendered}") return "".join(chunks) if isinstance(content, dict): return _history_dict_text(content, image_urls=True) @@ -134,44 +110,36 @@ def _history_text_only_part(part: dict) -> bool: def _is_text_only_busy_payload(content: Any) -> bool: """True when a busy submit carries only plain text, not attachments/media.""" - if isinstance(content, (str, int, float)): - return True if isinstance(content, list): return bool(content) and all( isinstance(part, str) or (isinstance(part, dict) and _history_text_only_part(part)) for part in content ) - return isinstance(content, dict) and _history_text_only_part(content) + return isinstance(content, (str, int, float)) or (isinstance(content, dict) and _history_text_only_part(content)) def _is_display_hidden_marker(role: str | None, text: str) -> bool: - """Gateway notices (model-switch, personality) persist as role=user ``[System: …]`` rows so strict - providers accept them mid-history; they must never render as a user bubble. Filtering in this one - projection hides them everywhere (raw marker stays in ``session["history"]``) and keeps them from - shifting the user-message ordinals the desktop reconciles against.""" + """Gateway notices (model-switch, personality) persist as role=user ``[System: …]`` rows so strict providers + accept them mid-history; they must never render as a user bubble. Filtering in this one projection hides + them everywhere (raw marker stays in ``session["history"]``) and keeps the desktop's user ordinals stable.""" return role == "user" and text.lstrip().startswith("[System:") def _skill_scaffold_projection(content_text: str) -> str: - """The invocation a slash-skill-expanded turn came from, else "" — every UI renders - ``/work fix the leak`` instead of the embedded skill body.""" + """The invocation a slash-skill-expanded turn came from, else "" — UIs render ``/work fix the leak``.""" return describe_skill_invocation(content_text, separator=" ") or "" def _expand_skill_invocation_for_replay(text: str, task_id: str) -> str: - """Inverse of :func:`_skill_scaffold_projection`: rewind/regenerate hands back the projected - invocation, and re-running it verbatim would drop the skill. Unchanged when not resolvable.""" + """Inverse of :func:`_skill_scaffold_projection`: rewind/regenerate hands back the projected invocation, + and re-running it verbatim would drop the skill. Unchanged when not resolvable.""" head, _, arg = (text or "").strip().partition(" ") if not head.startswith("/"): return text try: from agent.skill_commands import build_skill_invocation_message, resolve_skill_command_key - cmd_key = resolve_skill_command_key(head.lstrip("/")) - if cmd_key is None: - return text - return build_skill_invocation_message(cmd_key, arg.strip(), task_id=task_id) or text - except Exception: - # A skill that no longer resolves must not break the rewind. + return text if cmd_key is None else (build_skill_invocation_message(cmd_key, arg.strip(), task_id=task_id) or text) + except Exception: # a skill that no longer resolves must not break the rewind logger.debug("skill re-expansion failed for replay", exc_info=True) return text @@ -182,12 +150,9 @@ _AUTO_CONTINUE_NOTE_PREFIX = "[System note: Your previous turn was interrupted m def _legacy_display_kind(role: str, text: str) -> str | None: - """Infer the display type of a synthetic row persisted without one. New rows are typed at turn - start (``persist_user_display_kind``); this prefix sniff migrates untyped rows already on disk (a - turn killed mid-run never reached the stamp), which would otherwise paint as a user bubble.""" - if role == "user" and text.lstrip().startswith(_AUTO_CONTINUE_NOTE_PREFIX): - return "auto_continue" - return None + """Display type of a synthetic row persisted untyped: new rows are typed at turn start (``persist_user_display_kind``); + this prefix sniff migrates rows already on disk (a turn killed mid-run never reached the stamp).""" + return "auto_continue" if role == "user" and text.lstrip().startswith(_AUTO_CONTINUE_NOTE_PREFIX) else None _HISTORY_REASONING_KEYS = ("reasoning", "reasoning_content", "reasoning_details", "codex_reasoning_items") @@ -212,8 +177,7 @@ def _history_to_messages(history: list[dict]) -> list[dict]: continue if role == "assistant" and m.get("tool_calls"): for tc in m["tool_calls"]: - fn = tc.get("function", {}) - tc_id = tc.get("id", "") + fn, tc_id = tc.get("function", {}), tc.get("id", "") if tc_id and fn.get("name"): try: args = json.loads(fn.get("arguments", "{}")) @@ -223,15 +187,11 @@ def _history_to_messages(history: list[dict]) -> list[dict]: if not content_text.strip(): continue if role == "tool": - tc_id = m.get("tool_call_id", "") - tc_info = tool_call_args.get(tc_id) if tc_id else None - name = (tc_info[0] if tc_info else None) or m.get("tool_name") or "tool" - args = (tc_info[1] if tc_info else None) or {} - tool_msg = {"role": "tool", "name": name, "context": _tool_ctx(name, args)} + tc_name, tc_args = tool_call_args.get(m.get("tool_call_id") or "", (None, None)) + name = tc_name or m.get("tool_name") or "tool" + args = tc_args or {} # `context` is an 80-char preview; ship args so a full-call renderer isn't truncated. - if args: - tool_msg["args"] = args - messages.append(tool_msg) + messages.append({"role": "tool", "name": name, "context": _tool_ctx(name, args), **({"args": args} if args else {})}) continue # A reasoning-only assistant turn is kept so "Thinking…" still shows after resume/reload. has_reasoning = role == "assistant" and any(m.get(key) for key in _HISTORY_REASONING_KEYS) @@ -245,16 +205,12 @@ def _history_to_messages(history: list[dict]) -> list[dict]: # Durable row identity (_rows_to_conversation); reactions etc. address persisted messages by it. if m.get("_row_id") is not None: msg["row_id"] = m["_row_id"] - if role == "user": - invocation = _skill_scaffold_projection(content_text) - if invocation: - # The invocation, never the expanded body (rewind re-sends by ordinal). - msg["text"] = invocation - msg["display_kind"] = "skill_invocation" + # A user turn shows its skill invocation, never the expanded body (rewind re-sends by ordinal). + invocation = _skill_scaffold_projection(content_text) if role == "user" else "" + if invocation: + msg.update(text=invocation, display_kind="skill_invocation") if role == "assistant": - for key in _HISTORY_REASONING_KEYS: - if m.get(key) is not None: - msg[key] = m[key] + msg.update((key, m[key]) for key in _HISTORY_REASONING_KEYS if m.get(key) is not None) # Display-only timeline metadata (model switches, delegation events). display_kind = m.get("display_kind") or _legacy_display_kind(role, content_text) if display_kind: @@ -266,15 +222,11 @@ def _history_to_messages(history: list[dict]) -> list[dict]: def _coerce_seed_history(value: Any) -> list[dict]: - if not isinstance(value, list): - return [] history = [] - for item in value: + for item in value if isinstance(value, list) else (): if not isinstance(item, dict) or item.get("role") not in ("user", "assistant", "system"): continue - content = item.get("content") - if content is None: - content = item.get("text") + content = item.get("text") if item.get("content") is None else item.get("content") if isinstance(content, str) and content.strip(): history.append({"role": item["role"], "content": content}) return history @@ -286,9 +238,7 @@ def _inflight_text(value: Any) -> str: def _start_inflight_turn(session: dict, text: Any) -> None: now = time.time() - session["inflight_turn"] = { - "assistant": "", "started_at": now, "streaming": True, "updated_at": now, "user": _inflight_text(text), - } + session["inflight_turn"] = {"assistant": "", "started_at": now, "streaming": True, "updated_at": now, "user": _inflight_text(text)} def _append_inflight_delta(session: dict, delta: Any) -> None: @@ -298,9 +248,7 @@ def _append_inflight_delta(session: dict, delta: Any) -> None: turn = session.get("inflight_turn") if not isinstance(turn, dict): turn = {"assistant": "", "streaming": True, "user": ""} - turn["assistant"] = f"{turn.get('assistant') or ''}{text}" - turn["streaming"] = True - turn["updated_at"] = time.time() + turn.update(assistant=f"{turn.get('assistant') or ''}{text}", streaming=True, updated_at=time.time()) session["inflight_turn"] = turn @@ -311,10 +259,10 @@ def _record_inflight_correction(session: dict, text: Any) -> None: turn = session.get("inflight_turn") if not correction or not isinstance(turn, dict): return + # correction_offsets: arrival-order boundary (assistant chars already streamed) so resuming clients + # place the bubble between the output seen and the output redirected. turn = dict(turn) turn["corrections"] = [*(turn.get("corrections") or []), correction] - # Arrival-order boundary (assistant chars already streamed) so resuming clients place the bubble - # between the output seen and the output redirected. turn["correction_offsets"] = [*(turn.get("correction_offsets") or []), len(str(turn.get("assistant") or ""))] turn["updated_at"] = time.time() session["inflight_turn"] = turn @@ -325,27 +273,23 @@ def _clear_inflight_turn(session: dict) -> None: def _fail_inflight_turn(session: dict, error: Any, error_surface: Optional[dict] = None) -> None: - """Mark the in-flight turn terminal-error but keep it replayable: a failure's terminal frame can be - lost on WS disconnect and the turn may never have been committed, so the snapshot lets - ``session.resume`` replay prompt, partial text and error instead of stranding the client on a - spinner. Lives until the next turn starts or the session closes. Caller holds history_lock.""" + """Mark the in-flight turn terminal-error but keep it replayable: a failure's terminal frame can be lost on + WS disconnect and the turn may never have been committed, so the snapshot lets ``session.resume`` replay + prompt, partial text and error. Lives until the next turn starts or the session closes. Caller holds history_lock.""" message = str(error) if not isinstance(error, BaseException) else (str(error) or type(error).__name__) now = time.time() turn = session.get("inflight_turn") if not isinstance(turn, dict): turn = {"assistant": "", "user": "", "started_at": now} - turn["assistant"] = str(turn.get("assistant") or "") - turn["user"] = str(turn.get("user") or "") - turn["error"] = message or "turn failed" - turn["status"] = "error" - turn["recoverable"] = True - if error_surface: - # {layer, code, retryable} so a reconnect renders the same layered error card. + turn.update( + assistant=str(turn.get("assistant") or ""), user=str(turn.get("user") or ""), + error=message or "turn failed", status="error", recoverable=True, + ) + if error_surface: # {layer, code, retryable} so a reconnect renders the same layered error card turn["error_surface"] = dict(error_surface) else: turn.pop("error_surface", None) - turn["streaming"] = False - turn["updated_at"] = now + turn.update(streaming=False, updated_at=now) session["inflight_turn"] = turn @@ -357,10 +301,9 @@ _TURN_PROMPT_ECHO_MAX_PROMPT = 65536 def _strip_prompt_echo(message: str, prompt: Any) -> str: - """Blank runs of the submitted prompt that ``message`` quotes back: secret redaction is pattern-based - and a provider 4xx echoing the request carries private prose matching no pattern. Any run of - ``_TURN_PROMPT_ECHO_WINDOW``+ chars shared with the prompt (or its JSON-escaped form) becomes - ````. Shingle-set matching keeps it linear. Only verbatim echo is stopped — a floor.""" + """Blank runs of the submitted prompt that ``message`` quotes back: secret redaction is pattern-based and a + provider 4xx echoing the request carries private prose matching no pattern. Any ``_TURN_PROMPT_ECHO_WINDOW``+ + char run shared with the prompt (or its JSON-escaped form) becomes ````; shingles keep it linear.""" if not message or not prompt: return message needle = " ".join(str(prompt).split())[:_TURN_PROMPT_ECHO_MAX_PROMPT] @@ -368,15 +311,11 @@ def _strip_prompt_echo(message: str, prompt: Any) -> str: if len(needle) < window or len(message) < window: return message shingles = {needle[i:i + window] for i in range(len(needle) - window + 1)} - try: - escaped = json.dumps(needle)[1:-1] - except Exception: - escaped = "" - if escaped and escaped != needle: + escaped = json.dumps(needle)[1:-1] + if escaped != needle: shingles.update(escaped[i:i + window] for i in range(len(escaped) - window + 1)) out: list[str] = [] - i = 0 - n = len(message) + i, n = 0, len(message) while i <= n - window: if message[i:i + window] in shingles: j = i + window @@ -392,11 +331,9 @@ def _strip_prompt_echo(message: str, prompt: Any) -> str: def _turn_failure_detail(error: Any, reason: Any = None, prompt: Any = None) -> str: - """Why a turn failed, for the ``tui turn finished`` bookend: ``""`` when nothing to say, else a - fragment with its own leading space (distinguishes a provider 4xx from a budget wall or crashed - finalizer). Two content contracts: ``redact_sensitive_text`` removes credentials; - ``_strip_prompt_echo`` removes a 4xx body quoting ``prompt`` back. Invariant: this record may gain - failure classification and provider detail, never the user's own content.""" + """Why a turn failed, for the ``tui turn finished`` bookend: ``""`` when nothing to say, else a fragment with + its own leading space. ``redact_sensitive_text`` removes credentials; ``_strip_prompt_echo`` removes a 4xx + body quoting ``prompt`` back. This record may gain failure detail, never the user's own content.""" reason_text = str(reason or "").strip() message = str(error or "").strip() if isinstance(error, BaseException): @@ -405,7 +342,6 @@ def _turn_failure_detail(error: Any, reason: Any = None, prompt: Any = None) -> return "" try: from agent.redact import redact_sensitive_text - message = redact_sensitive_text(message, force=True) except Exception: message = "" # never fail open @@ -414,12 +350,8 @@ def _turn_failure_detail(error: Any, reason: Any = None, prompt: Any = None) -> message = _strip_prompt_echo(message, prompt) if len(message) > _TURN_FAILURE_DETAIL_LIMIT: message = message[:_TURN_FAILURE_DETAIL_LIMIT] + "\u2026" - out = "" - if reason_text: - out += " failure_reason=%s" % " ".join(reason_text.split()) - if message: - out += " cause=%r" % message - return out + out = " failure_reason=%s" % " ".join(reason_text.split()) if reason_text else "" + return out + (" cause=%r" % message if message else "") def register(server) -> None: diff --git a/tui_gateway/tool_progress.py b/tui_gateway/tool_progress.py index 81457a49df..cf3153dce8 100644 --- a/tui_gateway/tool_progress.py +++ b/tui_gateway/tool_progress.py @@ -1,8 +1,5 @@ -"""Tool lifecycle callbacks (tool.start/complete/progress events), verbose-text capping/redaction, todo-state projection. - -Bodies are rebound onto server.py's globals at install time (method_ctx.bind_module), so they -reference server.py globals bare. -""" +"""Tool lifecycle callbacks (tool.start/complete/progress events), verbose-text capping/redaction, todo-state +projection. Bodies are rebound onto server.py's globals (method_ctx.bind_module) and reference them bare.""" from __future__ import annotations @@ -20,15 +17,8 @@ _TODO_TOOL_NAMES = ("todo_list", "todo") # legacy alias: pre-rename replays def _cap_tui_verbose_text(text: str) -> str: if len(text) <= _TUI_VERBOSE_TEXT_MAX_CHARS and text.count("\n") < _TUI_VERBOSE_TEXT_MAX_LINES: return text - idx = len(text) - start = 0 - for _ in range(_TUI_VERBOSE_TEXT_MAX_LINES): - idx = text.rfind("\n", 0, idx) - if idx < 0: - start = 0 - break - start = idx + 1 - line_start = start + # Start of the last MAX_LINES lines, then pull forward to the char budget (never mid-line). + line_start = len(text) - len("\n".join(text.split("\n")[-_TUI_VERBOSE_TEXT_MAX_LINES:])) start = max(line_start, len(text) - _TUI_VERBOSE_TEXT_MAX_CHARS) if start > line_start: next_break = text.find("\n", start) @@ -44,29 +34,31 @@ def _cap_tui_verbose_text(text: str) -> str: def _redact_tui_verbose_text(text: str) -> str: try: from agent.redact import redact_sensitive_text - redacted = redact_sensitive_text(str(text), force=True) except Exception: return "" return _cap_tui_verbose_text(redacted) -def _tool_args_text(args: dict) -> str: +def _verbose_text(render, fallback) -> str: + """Redacted+capped ``render()``; ``fallback()`` when rendering raises.""" try: - raw = json.dumps(args or {}, indent=2, ensure_ascii=False, default=str) + raw = render() except Exception: - raw = str(args or {}) + raw = fallback() return _redact_tui_verbose_text(raw) +def _tool_args_text(args: dict) -> str: + return _verbose_text(lambda: json.dumps(args or {}, indent=2, ensure_ascii=False, default=str), lambda: str(args or {})) + + def _tool_result_text(result: object) -> str: - try: + def render(): from agent.tool_dispatch_helpers import _multimodal_text_summary + return _multimodal_text_summary(result) - raw = _multimodal_text_summary(result) - except Exception: - raw = str(result) - return _redact_tui_verbose_text(raw) + return _verbose_text(render, lambda: str(result)) def _fmt_tool_duration(seconds: float | None) -> str: @@ -110,9 +102,7 @@ def _tool_summary(name: str, result: str, duration_s: float | None) -> str | Non return f"{warning}{suffix}" entry = _SUMMARY_COUNTERS.get(name) n = entry[0](data) if entry else None - if n is None: - return None - return f"{entry[1]} {n} {entry[2] if n == 1 else entry[3]}{suffix}" + return f"{entry[1]} {n} {entry[2] if n == 1 else entry[3]}{suffix}" if n is not None else None def _normalize_todo_state(value: object) -> dict | None: @@ -124,8 +114,8 @@ def _normalize_todo_state(value: object) -> dict | None: except (TypeError, ValueError): return None todos = list(value["todos"]) - # Unused TodoStore snapshot() is {todos: [], revision: 0}: attaching it on resume stamps a - # client watermark and blocks unversioned tool.start merges. Empty at revision >= 1 is a real clear. + # Unused TodoStore snapshot() is {todos: [], revision: 0}: attaching it on resume stamps a client + # watermark and blocks unversioned tool.start merges. Empty at revision >= 1 is a real clear. if not todos and revision == 0: return None return {"todos": todos, "revision": revision} @@ -133,10 +123,8 @@ def _normalize_todo_state(value: object) -> dict | None: def _cache_todo_state(session: dict, state: dict | None) -> None: """Keep the newest snapshot on the session (revision-monotonic).""" - if state is None: - return - cached = _normalize_todo_state(session.get("todo_state")) - if cached is None or state["revision"] >= cached["revision"]: + cached = _normalize_todo_state(session.get("todo_state")) if state is not None else None + if state is not None and (cached is None or state["revision"] >= cached["revision"]): session["todo_state"] = state @@ -166,14 +154,12 @@ def _attach_todo_state(payload: dict, session: dict) -> dict: def _todo_state_from_history(history) -> dict | None: - """Latest todo snapshot from a loaded transcript, for resume paths that answer before an AIAgent - (and its live TodoStore) exists: the newest tool result paired with an assistant ``todo`` call - IS the durable snapshot.""" + """Latest todo snapshot from a loaded transcript, for resume paths that answer before an AIAgent (and + its live TodoStore) exists: the newest tool result paired with an assistant ``todo`` call IS it.""" if not isinstance(history, list) or not history: return None try: from tools.todo_tool import MAX_TODO_RESULT_CHARS - todo_call_ids = { call.get("id") for msg in history if isinstance(msg, dict) @@ -203,31 +189,26 @@ def _on_tool_start(sid: str, tool_call_id: str, name: str, args: dict): if session is not None: with contextlib.suppress(Exception): from agent.display import capture_local_edit_snapshot - snapshot = capture_local_edit_snapshot(name, args) if snapshot is not None: session.setdefault("edit_snapshots", {})[tool_call_id] = snapshot session.setdefault("tool_started_at", {})[tool_call_id] = time.time() if _tool_progress_enabled(sid) or _tool_lifecycle_required_for_ui(name): payload: dict[str, object] = {"tool_id": tool_call_id, "name": name, "context": _tool_ctx(name, args)} - # Full args (not just the 80-char `context` preview) so the desktop's expanded tool row is - # complete while the tool runs. args.todos may be a partial merge — tool.complete is the truth. + # Full args (not just the 80-char `context` preview) so the desktop's expanded tool row is complete + # while the tool runs. args.todos may be a partial merge — tool.complete is the truth. if args: payload["args"] = args - if _session_verbose(sid): - args_text = _tool_args_text(args) - if args_text: - payload["args_text"] = args_text + if _session_verbose(sid) and (args_text := _tool_args_text(args)): + payload["args_text"] = args_text _emit("tool.start", sid, payload) def _on_tool_complete(sid: str, tool_call_id: str, name: str, args: dict, result: str): payload = {"tool_id": tool_call_id, "name": name, "args": args} session = _sessions.get(sid) - snapshot = started_at = None - if session is not None: - snapshot = session.setdefault("edit_snapshots", {}).pop(tool_call_id, None) - started_at = session.setdefault("tool_started_at", {}).pop(tool_call_id, None) + snapshot = session.setdefault("edit_snapshots", {}).pop(tool_call_id, None) if session is not None else None + started_at = session.setdefault("tool_started_at", {}).pop(tool_call_id, None) if session is not None else None duration_s = time.time() - started_at if started_at else None if duration_s is not None: payload["duration_s"] = duration_s @@ -238,85 +219,60 @@ def _on_tool_complete(sid: str, tool_call_id: str, name: str, args: dict, result summary = _tool_summary(name, result, duration_s) if summary: payload["summary"] = summary - if _session_verbose(sid): - result_text = _tool_result_text(result) - if result_text: - payload["result_text"] = result_text - todo_state = None - if name in _TODO_TOOL_NAMES: - todo_state = _normalize_todo_state(payload.get("result")) - if todo_state is not None: - payload.update(todo_state) - if session is not None: - _cache_todo_state(session, todo_state) + if _session_verbose(sid) and (result_text := _tool_result_text(result)): + payload["result_text"] = result_text + todo_state = _normalize_todo_state(payload.get("result")) if name in _TODO_TOOL_NAMES else None + if todo_state is not None: + payload.update(todo_state) + if session is not None: + _cache_todo_state(session, todo_state) with contextlib.suppress(Exception): from agent.display import render_edit_diff_with_delta - rendered: list[str] = [] if render_edit_diff_with_delta(name, result, function_args=args, snapshot=snapshot, print_fn=rendered.append): payload["inline_diff"] = "\n".join(rendered) - if ( - _tool_progress_enabled(sid) - or payload.get("inline_diff") - or _tool_lifecycle_required_for_ui(name) - or name in _TODO_TOOL_NAMES - ): + if (_tool_progress_enabled(sid) or payload.get("inline_diff") or _tool_lifecycle_required_for_ui(name) + or name in _TODO_TOOL_NAMES): _emit("tool.complete", sid, payload) - # Task state is application data, not tool-progress chrome: a dedicated full-snapshot event - # lets every client reconcile without parsing tool args. + # Task state is application data, not tool-progress chrome: a dedicated full-snapshot event lets + # every client reconcile without parsing tool args. if todo_state is not None: _emit("todo.updated", sid, todo_state) -# ── _on_tool_progress dispatch ───────────────────────────────────────────── -# Each handler takes (sid, name, preview, kw). `tool.started` is dropped on purpose: _on_tool_start -# already emits the authoritative tool.start with the stable id and args; an id-less duplicate row -# makes the desktop live view diverge from hydrated history. - +# ── _on_tool_progress dispatch: each handler takes (sid, name, preview, kw) ───────────────────── +# `tool.started` is dropped on purpose: _on_tool_start already emits the authoritative tool.start with +# the stable id and args; an id-less duplicate row makes the desktop live view diverge from history. def _progress_output_risk(sid, name, preview, kw): metadata = kw.get("risk_metadata") - if not isinstance(metadata, dict): - return - _emit("tool.output_risk", sid, { - "tool_id": str(kw.get("tool_call_id") or ""), "name": str(name), - "risk": str(metadata.get("risk") or "low"), - "findings": [str(item) for item in metadata.get("findings", [])], - "redacted": bool(metadata.get("redacted", False)), - }) + if isinstance(metadata, dict): + _emit("tool.output_risk", sid, { + "tool_id": str(kw.get("tool_call_id") or ""), "name": str(name), "risk": str(metadata.get("risk") or "low"), + "findings": [str(item) for item in metadata.get("findings", [])], "redacted": bool(metadata.get("redacted", False)), + }) def _progress_reasoning(sid, name, preview, kw): - payload: dict[str, object] = {"text": str(preview)} - if _session_verbose(sid): - payload["verbose"] = True - _emit("reasoning.available", sid, payload) + _emit("reasoning.available", sid, {"text": str(preview), **({"verbose": True} if _session_verbose(sid) else {})}) def _progress_moa_reference(sid, name, preview, kw): # MoA reference-model output, rendered as a labelled block before the aggregator's response. # `name` is the slot label, `preview` the text. ref_payload: dict[str, object] = {"label": str(name), "text": str(preview or "")} - if kw.get("moa_index") is not None: - ref_payload["index"] = kw.get("moa_index") - if kw.get("moa_count") is not None: - ref_payload["count"] = kw.get("moa_count") + for key, out in (("moa_index", "index"), ("moa_count", "count")): + if kw.get(key) is not None: + ref_payload[out] = kw[key] _emit("moa.reference", sid, ref_payload) -def _progress_moa_aggregating(sid, name, preview, kw): - _emit("moa.aggregating", sid, {"aggregator": str(name or "")}) - - def _progress_moa_progress(sid, name, preview, kw): # Drives the status-bar `MOA: 2/3 refs done`; both counters required for deterministic rendering. - refs_done = kw.get("moa_refs_done") - refs_total = kw.get("moa_refs_total") + refs_done, refs_total = kw.get("moa_refs_done"), kw.get("moa_refs_total") if refs_done is None or refs_total is None: return - _emit("moa.progress", sid, { - "label": str(name or ""), "refs_done": int(refs_done), "refs_total": int(refs_total), - }) + _emit("moa.progress", sid, {"label": str(name or ""), "refs_done": int(refs_done), "refs_total": int(refs_total)}) def _progress_moa_phase(sid, name, preview, kw): @@ -342,34 +298,29 @@ def _str_list(v): def _int_or_skip(v): - # Per-branch rollups tolerate junk from older emitters: unparsable -> field omitted. + """Per-branch token/api rollups tolerate junk from older emitters: unparsable -> field omitted.""" try: return int(v) except (TypeError, ValueError): return None -# Optional subagent.* payload fields in WIRE ORDER: (source key, present-when, coerce). Identity -# fields are all optional: older emitters omit them and the TUI spawn tree falls back to flat -# rendering. `tool_name`/`text` are fed from the positional name/preview. +# Optional subagent.* payload fields in WIRE ORDER: (source key, present-when, coerce). Identity fields +# are all optional: older emitters omit them and the TUI spawn tree falls back to flat rendering. +# `tool_name`/`text` are fed from the positional name/preview; `output_tail` is a list of dicts. _SUBAGENT_FIELDS = ( ("subagent_id", bool, str), ("parent_id", bool, str), ("child_session_id", bool, str), - ("delegation_id", bool, str), ("depth", _not_none, int), ("model", bool, str), - ("tool_count", _not_none, int), ("toolsets", bool, _str_list), - ("input_tokens", _not_none, _int_or_skip), ("output_tokens", _not_none, _int_or_skip), + ("delegation_id", bool, str), ("depth", _not_none, int), ("model", bool, str), ("tool_count", _not_none, int), + ("toolsets", bool, _str_list), ("input_tokens", _not_none, _int_or_skip), ("output_tokens", _not_none, _int_or_skip), ("reasoning_tokens", _not_none, _int_or_skip), ("api_calls", _not_none, _int_or_skip), - ("files_read", bool, _str_list), ("files_written", bool, _str_list), - ("output_tail", bool, list), # list of dicts + ("files_read", bool, _str_list), ("files_written", bool, _str_list), ("output_tail", bool, list), ("tool_name", bool, str), ("text", bool, str), ("status", bool, str), ("summary", bool, str), ("duration_seconds", _not_none, float), ) def _progress_subagent(sid, name, preview, kw, event_type): - payload = { - "goal": str(kw.get("goal") or ""), "task_count": int(kw.get("task_count") or 1), - "task_index": int(kw.get("task_index") or 0), - } + payload = {"goal": str(kw.get("goal") or ""), "task_count": int(kw.get("task_count") or 1), "task_index": int(kw.get("task_index") or 0)} source = {**kw, "tool_name": name, "text": preview} for key, present, coerce in _SUBAGENT_FIELDS: if present(source.get(key)): @@ -379,20 +330,19 @@ def _progress_subagent(sid, name, preview, kw, event_type): if preview and event_type == "subagent.tool": payload["tool_preview"] = str(preview) payload["text"] = str(preview) - # subagent.text is the child's per-token reply, relayed solely to feed a watch window's live - # mirror (keyed off the child sid); on the parent it's hundreds of ignored frames, so skip it. + # subagent.text is the child's per-token reply, relayed solely to feed a watch window's live mirror + # (keyed off the child sid); on the parent it's hundreds of ignored frames, so skip it. if event_type != "subagent.text": _emit(event_type, sid, payload) _mirror_subagent_to_child(event_type, payload) -# event_type -> (handler, requires) where `requires` names the arg that must be truthy for the -# row to be emitted at all ("name" / "preview" / None). +# event_type -> (handler, requires): `requires` names the arg that must be truthy for the row to be +# emitted at all ("name" / "preview" / None). _PROGRESS_HANDLERS = { - "tool.output_risk": (_progress_output_risk, "name"), - "reasoning.available": (_progress_reasoning, "preview"), + "tool.output_risk": (_progress_output_risk, "name"), "reasoning.available": (_progress_reasoning, "preview"), "moa.reference": (_progress_moa_reference, "name"), - "moa.aggregating": (_progress_moa_aggregating, None), + "moa.aggregating": (lambda sid, name, preview, kw: _emit("moa.aggregating", sid, {"aggregator": str(name or "")}), None), "moa.progress": (_progress_moa_progress, None), "moa.phase": (_progress_moa_phase, None), } @@ -401,18 +351,13 @@ def _on_tool_progress( sid: str, event_type: str, name: str | None = None, preview: str | None = None, _args: dict | None = None, **_kwargs, ): - if not _tool_progress_enabled(sid): - return - if event_type == "tool.started" and name: - return - entry = _PROGRESS_HANDLERS.get(event_type) - if entry is not None: - handler, requires = entry - if requires is None or {"name": name, "preview": preview}[requires]: - handler(sid, name, preview, _kwargs) + if not _tool_progress_enabled(sid) or (event_type == "tool.started" and name): return if event_type.startswith("subagent."): - _progress_subagent(sid, name, preview, _kwargs, event_type) + return _progress_subagent(sid, name, preview, _kwargs, event_type) + handler, requires = _PROGRESS_HANDLERS.get(event_type, (None, None)) + if handler is not None and (requires is None or {"name": name, "preview": preview}[requires]): + handler(sid, name, preview, _kwargs) def register(server) -> None: diff --git a/tui_gateway/transport.py b/tui_gateway/transport.py index dcb171de4c..7b67e937ff 100644 --- a/tui_gateway/transport.py +++ b/tui_gateway/transport.py @@ -1,15 +1,15 @@ """Transport abstraction for the tui_gateway JSON-RPC server. -A :class:`Transport` accepts a JSON-serialisable dict and forwards it to its peer, so the same -dispatcher runs over stdio (``tui_gateway.entry``) or WebSocket (``tui_gateway.ws``). The active -transport for the current request lives in a ``ContextVar`` so handlers dispatched onto the worker -pool route writes to the right peer. ``server.write_json`` works with nothing bound: it falls back -to the module-level :class:`StdioTransport`, which resolves ``_real_stdout`` lazily through a -callback so tests that monkey-patch ``server._real_stdout`` keep working. +A :class:`Transport` forwards a JSON-serialisable dict to its peer, so one dispatcher runs over stdio +(``tui_gateway.entry``) or WebSocket (``tui_gateway.ws``). The request's transport lives in a +``ContextVar`` so pool-dispatched handlers write to the right peer; with nothing bound +``server.write_json`` falls back to the module-level :class:`StdioTransport`, which resolves +``_real_stdout`` lazily so tests that monkey-patch it keep working. """ from __future__ import annotations +import contextlib import contextvars import errno import json @@ -27,10 +27,9 @@ _PEER_GONE_ERRNOS = frozenset({ logger = logging.getLogger(__name__) -# When true, StdioTransport skips ``stream.flush`` after writing: on a half-closed pipe (TUI Node -# parent quit while the gateway still emits) flush can block long enough to starve the worker pool. -# Python text stdout is fully buffered on a pipe, so this ONLY makes sense with ``-u`` / -# ``PYTHONUNBUFFERED=1``; otherwise frames accumulate and the TUI hangs waiting for ``gateway.ready``. +# When true, StdioTransport skips ``stream.flush`` after writing: on a half-closed pipe (TUI Node parent quit +# while the gateway still emits) flush can block long enough to starve the worker pool. Python text stdout is +# fully buffered on a pipe, so this ONLY makes sense with ``-u``/``PYTHONUNBUFFERED=1``; otherwise the TUI hangs. _DISABLE_FLUSH = (os.environ.get("HERMES_TUI_GATEWAY_NO_FLUSH", "") or "").strip().lower() in {"1", "true", "yes", "on"} @@ -51,47 +50,38 @@ _current_transport: contextvars.ContextVar[Optional[Transport]] = contextvars.Co def current_transport() -> Optional[Transport]: - """Return the transport bound for the current request, if any.""" return _current_transport.get() def bind_transport(transport: Optional[Transport]): - """Bind *transport* for the current context. Returns a token for :func:`reset_transport`.""" + """Bind *transport* for the current context; returns a token for :func:`reset_transport`.""" return _current_transport.set(transport) def reset_transport(token) -> None: - """Restore the transport binding captured by :func:`bind_transport`.""" _current_transport.reset(token) -def _peer_gone(exc: Exception, what: str) -> bool: - """True when *exc* from a stream write/flush means the peer is gone; re-raise anything else. - - ``False`` from :meth:`StdioTransport.write` is the dispatcher's "broken stdout pipe" signal - (``entry.py`` exits cleanly on it), so programming errors and real host I/O bugs (non-JSON-safe - payloads, UnicodeEncodeError from a misconfigured locale, ENOSPC, EACCES, ...) MUST re-raise so - the crash log records them instead of masquerading as a clean disconnect. Peer-gone: - ``BrokenPipeError``, ``ValueError("...closed file...")``, ``OSError`` with errno in - :data:`_PEER_GONE_ERRNOS`. - """ +def _raise_unless_peer_gone(exc: Exception, what: str) -> None: + """Return when *exc* from a stream write/flush means the peer is gone; re-raise anything else. + ``False`` from :meth:`StdioTransport.write` is the dispatcher's "broken stdout pipe" signal (``entry.py`` + exits cleanly on it), so programming errors and real host I/O bugs (UnicodeEncodeError from a misconfigured + locale, ENOSPC, EACCES, ...) MUST re-raise so the crash log records them instead of masquerading as a clean + disconnect. Peer-gone: BrokenPipeError, ValueError("...closed file..."), OSError errno in _PEER_GONE_ERRNOS.""" if isinstance(exc, BrokenPipeError): - return True + return if isinstance(exc, ValueError): if isinstance(exc, UnicodeEncodeError) or "closed file" not in str(exc): raise exc - return True - if isinstance(exc, OSError): - if exc.errno not in _PEER_GONE_ERRNOS: - raise exc - logger.debug("StdioTransport %s peer gone: %s", what, exc) - return True - raise exc + return + if not isinstance(exc, OSError) or exc.errno not in _PEER_GONE_ERRNOS: + raise exc + logger.debug("StdioTransport %s peer gone: %s", what, exc) class StdioTransport: - """Writes JSON frames to a stream (usually ``sys.stdout``), resolved via a callable so runtime - monkey-patches of the underlying stream keep working.""" + """Writes JSON frames to a stream (usually ``sys.stdout``) resolved via a callable, so runtime + monkey-patches of the stream keep working.""" __slots__ = ("_stream_getter", "_lock") @@ -100,26 +90,25 @@ class StdioTransport: self._lock = lock def write(self, obj: dict) -> bool: - """Return ``True`` on success, ``False`` ONLY when the peer is gone (see :func:`_peer_gone`).""" - # Serialization is OUTSIDE the lock so a large payload can't block other threads emitting - # their own frames. A non-JSON-safe payload is a programming error: re-raise. + """Return ``True`` on success, ``False`` ONLY when the peer is gone (see :func:`_raise_unless_peer_gone`).""" + # Serialization is OUTSIDE the lock so a large payload can't block other threads' frames. A + # non-JSON-safe payload is a programming error: re-raise. line = json.dumps(obj, ensure_ascii=False) + "\n" with self._lock: stream = self._stream_getter() try: stream.write(line) except Exception as e: - if _peer_gone(e, "write"): - return False - # A flush that *raises* with a peer-gone errno means the dispatcher should exit cleanly. - # A flush that *hangs* on a half-closed pipe holds the lock until it returns — see - # ``_DISABLE_FLUSH`` for the "skip flush entirely" escape hatch. + _raise_unless_peer_gone(e, "write") + return False + # A flush that *raises* peer-gone means the dispatcher should exit cleanly; one that *hangs* + # on a half-closed pipe holds the lock until it returns — ``_DISABLE_FLUSH`` skips it entirely. if not _DISABLE_FLUSH: try: stream.flush() except Exception as e: - if _peer_gone(e, "flush"): - return False + _raise_unless_peer_gone(e, "flush") + return False return True def close(self) -> None: @@ -127,12 +116,9 @@ class StdioTransport: class TeeTransport: - """Mirrors writes to one primary plus N best-effort secondaries. - - The primary's return value (and exceptions) determine the result — secondaries swallow failures - so a wedged sidecar never stalls the main IO path. Used by the PTY child so every dispatcher - emit lands on stdio (Ink) AND on a back-WS feeding the dashboard sidebar. - """ + """Mirrors writes to one primary plus N best-effort secondaries. The primary's return value (and + exceptions) determine the result; secondaries swallow failures so a wedged sidecar never stalls the + main IO path. Used by the PTY child: every emit lands on stdio (Ink) AND a back-WS for the dashboard.""" __slots__ = ("_primary", "_secondaries") @@ -144,10 +130,8 @@ class TeeTransport: # Primary first so a slow sidecar (WS publisher) never delays Ink/stdio. ok = self._primary.write(obj) for sec in self._secondaries: - try: + with contextlib.suppress(Exception): sec.write(obj) - except Exception: - pass return ok def close(self) -> None: @@ -155,7 +139,5 @@ class TeeTransport: self._primary.close() finally: for sec in self._secondaries: - try: + with contextlib.suppress(Exception): sec.close() - except Exception: - pass diff --git a/tui_gateway/ws.py b/tui_gateway/ws.py index a1d04a9ae4..5ba1bdb6ca 100644 --- a/tui_gateway/ws.py +++ b/tui_gateway/ws.py @@ -1,11 +1,7 @@ -"""WebSocket transport for the tui_gateway JSON-RPC server. - -Reuses :func:`tui_gateway.server.dispatch` verbatim so every RPC method, slash command, -approval/clarify/sudo flow and agent event flows through the same handlers whether the client is -Ink over stdio or an iOS/web client over WS. Wire protocol is identical to stdio: newline-delimited -JSON-RPC both ways; the server emits ``gateway.ready`` right after accept, then echoes -responses/events. Mount as ``@app.websocket("/api/ws") async def ws(ws): await handle_ws(ws)``. -""" +"""WebSocket transport for the tui_gateway JSON-RPC server: reuses :func:`tui_gateway.server.dispatch` +verbatim so every RPC, slash command, approval flow and agent event takes the same handlers as Ink over +stdio. Wire protocol is identical to stdio (newline-delimited JSON-RPC both ways; ``gateway.ready`` right +after accept). Mount as ``@app.websocket("/api/ws") async def ws(ws): await handle_ws(ws)``.""" from __future__ import annotations @@ -41,7 +37,6 @@ def _note_dashboard_client_activity(*, force: bool = False) -> None: _dashboard_client_touched_at = now try: from gateway.scale_to_zero import touch_dashboard_client_heartbeat - touch_dashboard_client_heartbeat() except Exception: # noqa: BLE001 - liveness garnish must never break the WS _log.debug("dashboard client heartbeat touch failed", exc_info=True) @@ -52,13 +47,12 @@ def _note_dashboard_client_activity(*, force: bool = False) -> None: _WS_WRITE_TIMEOUT_S = 10.0 _WS_LOG_PAYLOAD_PREVIEW = 240 -# Per-token streaming frames are coalesced: buffered and flushed as a batch on a short timer -# instead of waking the loop once per token (each wakeup competes with the agent turn for the GIL). -# Keep this set to genuinely high-frequency, display-only events — anything a client must see -# promptly (tool/approval/status/completion) is non-streaming and flushes the buffer ahead of -# itself, so ordering is preserved. +# Per-token streaming frames are coalesced: buffered and flushed as a batch on a short timer instead +# of waking the loop once per token (each wakeup competes with the agent turn for the GIL). Keep this +# set to genuinely high-frequency, display-only events — anything a client must see promptly +# (tool/approval/status/completion) is non-streaming and flushes the buffer ahead of itself, so +# ordering is preserved. _TOKEN_COALESCE_S: max buffer wait (~30 fps; imperceptible). _STREAMING_EVENT_TYPES = frozenset({"message.delta", "reasoning.delta", "thinking.delta"}) -# Max time a streamed token waits in the buffer (~30 fps; imperceptible). _TOKEN_COALESCE_S = 0.033 # starlette stays optional at import time; fall back to a generic sentinel. @@ -69,23 +63,17 @@ except ImportError: # pragma: no cover - starlette is a required install path class WSTransport: - """Per-connection WS transport. - - ``write`` is safe from any thread *other than* the loop thread owning the socket (pool workers - marshal onto the loop and block on the future). Called from the loop thread itself it would - deadlock, so we detect that and fire-and-forget; loop-thread callers that need completion use - ``write_async``. - """ + """Per-connection WS transport. ``write`` is safe from any thread *other than* the loop thread owning the + socket (pool workers marshal onto the loop and block on the future); from the loop thread itself it would + deadlock, so it detects that and fires-and-forgets. Loop-thread callers needing completion use ``write_async``.""" def __init__(self, ws: Any, loop: asyncio.AbstractEventLoop, *, peer: str = "unknown", auth_identity: dict | None = None) -> None: self._ws = ws self._loop = loop self._peer = peer - #: Server-verified identity from the WS-upgrade credential (dashboard ticket / internal - #: credential), stamped by ``web_server._ws_auth_reason``. None for legacy-token/stdio - #: transports. RPC params can never populate this: it is the only identity authority for - #: browser-controller registration. + #: Server-verified identity from the WS-upgrade credential, stamped by ``web_server._ws_auth_reason``; None + #: for legacy-token/stdio. RPC params can never populate it: sole identity authority for browser controllers. self.auth_identity = auth_identity self._closed = False # Token-coalescing buffer. The lock guards the buffer + "armed" flag against worker threads @@ -94,15 +82,9 @@ class WSTransport: self._pending_tokens: list[str] = [] self._token_flush_handle: asyncio.TimerHandle | None = None self._token_flush_armed = False - # Socket writes need an async boundary: several batches can be queued on the owning loop - # while it recovers from a stall. + # Socket writes need an async boundary: several batches can queue on the loop during a stall. self._send_lock = asyncio.Lock() - @staticmethod - def _is_streaming_frame(obj: dict) -> bool: - params = obj.get("params") if isinstance(obj, dict) else None - return isinstance(params, dict) and params.get("type") in _STREAMING_EVENT_TYPES - def write(self, obj: dict) -> bool: if self._closed: return False @@ -111,22 +93,20 @@ class WSTransport: on_loop = asyncio.get_running_loop() is self._loop except RuntimeError: on_loop = False - # Streamed token: buffer it and arm the flush timer; the worker returns immediately. # call_soon_threadsafe is safe from a worker or the loop. - if self._is_streaming_frame(obj): + params = obj.get("params") if isinstance(obj, dict) else None + if isinstance(params, dict) and params.get("type") in _STREAMING_EVENT_TYPES: with self._token_lock: self._pending_tokens.append(line) if not self._token_flush_armed: self._token_flush_armed = True self._loop.call_soon_threadsafe(self._arm_token_flush) return not self._closed - # Non-streaming frame: append behind any buffered tokens and flush the whole batch NOW so it - # can never overtake them. The send is scheduled INSIDE the lock so wire order matches - # buffer order even if the coalesce timer fires on the loop at the same moment. + # can never overtake them. The send is scheduled INSIDE the lock so wire order matches buffer + # order even if the coalesce timer fires on the loop at the same moment. from agent.async_utils import safe_schedule_threadsafe - with self._token_lock: self._pending_tokens.append(line) batch, self._pending_tokens = self._pending_tokens, [] @@ -137,33 +117,27 @@ class WSTransport: if fut is None: self._closed = True return False - try: fut.result(timeout=_WS_WRITE_TIMEOUT_S) return not self._closed except concurrent.futures.TimeoutError: # builtin TimeoutError on 3.11+ - # The loop is stalled (GIL-heavy turn, delegation), NOT the socket dead: the send is - # already scheduled and flushes once the loop breathes. Latching _closed here permanently - # silenced live windows after one slow write; _safe_send_many latches on a real error. - _log.warning( - "ws write slow (loop stalled >%ss) peer=%s — frame left in flight", - _WS_WRITE_TIMEOUT_S, self._peer, - ) + # The loop is stalled (GIL-heavy turn, delegation), NOT the socket dead: the send is already + # scheduled and flushes once the loop breathes. Latching _closed here permanently silenced + # live windows after one slow write; _safe_send_many latches on a real error. + _log.warning("ws write slow (loop stalled >%ss) peer=%s — frame left in flight", _WS_WRITE_TIMEOUT_S, self._peer) return not self._closed except Exception as exc: self._closed = True _log.warning("ws write failed peer=%s error_type=%s error=%s", self._peer, type(exc).__name__, exc) return False - def _arm_token_flush(self) -> None: - """Arm the coalesce timer. Runs on the loop thread.""" - if self._closed: - return - self._token_flush_handle = self._loop.call_later(_TOKEN_COALESCE_S, self._flush_tokens) + def _arm_token_flush(self) -> None: # loop thread + if not self._closed: + self._token_flush_handle = self._loop.call_later(_TOKEN_COALESCE_S, self._flush_tokens) def _flush_tokens(self) -> None: - """Timer callback (loop thread): send buffered tokens as one batch. Scheduled under the lock - so wire order is fixed relative to a concurrent ``write``.""" + """Timer callback (loop thread): send buffered tokens as one batch, scheduled under the lock so + wire order is fixed relative to a concurrent ``write``.""" with self._token_lock: self._token_flush_handle = None self._token_flush_armed = False @@ -172,8 +146,8 @@ class WSTransport: self._loop.create_task(self._safe_send_many(batch)) async def write_async(self, obj: dict) -> bool: - """Send from the owning loop; awaits until the frame is on the wire. Buffered tokens are - flushed ahead of it in the SAME batch so nothing slips between.""" + """Send from the owning loop; awaits until the frame is on the wire. Buffered tokens are flushed + ahead of it in the SAME batch so nothing slips between.""" if self._closed: return False with self._token_lock: @@ -193,17 +167,14 @@ class WSTransport: return await self._ws.send_text(line) except Exception as exc: - # Latch while holding the writer lock so queued batches observe the failure before - # touching the socket. + # Latch while holding the writer lock so queued batches observe the failure first. self._closed = True _log.warning("ws send failed peer=%s error_type=%s error=%s", self._peer, type(exc).__name__, exc) - def close(self) -> None: + def close(self) -> None: # loop thread (handle_ws finally), so the TimerHandle is safe self._closed = True - # Runs on the loop thread (handle_ws finally), so the TimerHandle is safe. - handle = self._token_flush_handle - if handle is not None: - handle.cancel() + if self._token_flush_handle is not None: + self._token_flush_handle.cancel() self._token_flush_handle = None @@ -212,19 +183,15 @@ def _ws_peer_label(ws: Any) -> str: client = getattr(ws, "client", None) if client is None: return "unknown" - host = getattr(client, "host", None) or "unknown" - port = getattr(client, "port", None) + host, port = getattr(client, "host", None) or "unknown", getattr(client, "port", None) return f"{host}:{port}" if port is not None else host def _disable_nagle(ws: Any) -> None: - """Disable Nagle + enable TCP keepalive on the raw socket (best-effort). - - Without TCP_NODELAY the kernel coalesces small per-token frames, so a burst after the model's - think-pause lands in one tick and no client-side smoothing can recover the cadence. Without - keepalive a silently-dropped client (SSH tunnel reset, sleep) leaves the leg half-open forever: - receive_text() blocks and the disconnect teardown (detach + orphan reap + resume replay) never runs. - """ + """Disable Nagle + enable TCP keepalive on the raw socket (best-effort). Without TCP_NODELAY the kernel + coalesces small per-token frames, so a burst after the model's think-pause lands in one tick and no + client-side smoothing can recover the cadence. Without keepalive a silently-dropped client (SSH tunnel + reset, sleep) leaves the leg half-open forever: receive_text() blocks and the disconnect teardown never runs.""" try: scope = getattr(ws, "scope", None) or {} transport = (scope.get("extensions") or {}).get("transport") or getattr(ws, "transport", None) @@ -242,49 +209,43 @@ def _disable_nagle(ws: Any) -> None: _log.debug("ws TCP_NODELAY skip: %s", exc) -def _error_frame(code: int, message: str, req_id: Any) -> dict: - return {"jsonrpc": "2.0", "error": {"code": code, "message": message}, "id": req_id} +class _SendFailed(Exception): + """Raised by handle_ws._reply when a reply could not be written: ends the read loop.""" async def handle_ws(ws: Any, *, auth_identity: dict | None = None, subprotocol: str | None = None) -> None: - """Run one WebSocket session. Wire-compatible with ``tui_gateway.entry``. - - *auth_identity* is the server-minted ``{user_id, provider}`` recorded at WS-upgrade auth; stored - as ``WSTransport.auth_identity``, the only identity authority for browser-controller registration. - Callers that omit it (harnesses, the embedded TUI child) get a ``None`` transport identity. - """ - peer = _ws_peer_label(ws) - transport: WSTransport | None = None + """Run one WebSocket session. Wire-compatible with ``tui_gateway.entry``. *auth_identity* is the server-minted + ``{user_id, provider}`` recorded at WS-upgrade auth, stored as ``WSTransport.auth_identity`` (the only identity + authority for browser-controller registration); callers that omit it (harnesses, embedded TUI child) get None.""" + peer, transport = _ws_peer_label(ws), None messages = parse_errors = dispatch_crashes = send_failures = 0 disconnect_reason = "not_connected" - async def _reply(frame: dict, reason: str, msg: str, *args: Any) -> bool: - """write_async; on failure record *reason* and log *msg*. False => break.""" + async def _reply(frame: dict, reason: str, msg: str, *args: Any) -> None: + """write_async; on failure record *reason*, log *msg* and end the read loop.""" nonlocal disconnect_reason, send_failures - if await transport.write_async(frame): - return True - disconnect_reason = reason - send_failures += 1 - _log.warning(msg, *args) - return False + if not await transport.write_async(frame): + disconnect_reason = reason + send_failures += 1 + _log.warning(msg, *args) + raise _SendFailed + + def _error(code: int, message: str, req_id: Any) -> dict: + return {"jsonrpc": "2.0", "error": {"code": code, "message": message}, "id": req_id} try: await (ws.accept(subprotocol=subprotocol) if subprotocol else ws.accept()) disconnect_reason = "connected" - # A client is attached from the moment the upgrade is accepted — mark it before the - # (possibly slow) ready/skin setup so scale-to-zero sees it. + # Mark the client attached before the (possibly slow) ready/skin setup so scale-to-zero sees it. _note_dashboard_client_activity(force=True) _disable_nagle(ws) _log.info("ws accepted peer=%s", peer) - transport = WSTransport(ws, asyncio.get_running_loop(), peer=peer, auth_identity=auth_identity) - - # resolve_skin() is synchronous I/O + CPU work; run it in the pool so the WS read loop stays - # free to drain the frontend's initial RPC burst. + # resolve_skin() is sync I/O + CPU; pooled so the read loop can drain the frontend's initial RPC burst. skin_payload = await asyncio.to_thread(server.resolve_skin) - # change_events: this backend broadcasts pet/cron/sessions.changed, so clients can demote - # legacy polls to backstops. replay_epoch lets reconnecting clients detect a backend restart - # and reset their per-session seq watermarks (event_replay). + # change_events: this backend broadcasts pet/cron/sessions.changed, so clients can demote legacy + # polls to backstops. replay_epoch lets reconnecting clients detect a backend restart and reset + # their per-session seq watermarks (event_replay). ready_ok = await transport.write_async({ "jsonrpc": "2.0", "method": "event", "params": {"type": "gateway.ready", "payload": { @@ -292,21 +253,21 @@ async def handle_ws(ws: Any, *, auth_identity: dict | None = None, subprotocol: }}, }) if ready_ok: - # Live-apply skins Hermes activates mid-conversation, and track this peer for - # session-less global broadcasts write_json can't route. + # Live-apply skins Hermes activates mid-conversation, and track this peer for session-less + # global broadcasts write_json can't route. server._ensure_skin_watcher() server.register_live_transport(transport) # Cross-backend liveness: a heartbeat row lets the startup orphan sweep tell "live but idle - # backend" from "truly orphaned". Idempotent and once-per-process, like the orphan sweep - # (the desktop app and web dashboard reach the agent via this sidecar, not entry.main()). - try: - server._start_backend_heartbeat_refresher() - except Exception: - _log.warning("backend heartbeat refresher start failed", exc_info=True) - try: - server._schedule_startup_orphan_sweep() - except Exception: - _log.warning("startup orphan sweep scheduling failed", exc_info=True) + # backend" from "truly orphaned". Idempotent and once-per-process, like the orphan sweep (the + # desktop app and web dashboard reach the agent via this sidecar, not entry.main()). + for start, what in ( + (server._start_backend_heartbeat_refresher, "backend heartbeat refresher start"), + (server._schedule_startup_orphan_sweep, "startup orphan sweep scheduling"), + ): + try: + start() + except Exception: + _log.warning("%s failed", what, exc_info=True) if not ready_ok: disconnect_reason = "ready_send_failed" send_failures += 1 @@ -324,38 +285,24 @@ async def handle_ws(ws: Any, *, auth_identity: dict | None = None, subprotocol: disconnect_reason = "receive_failed" _log.exception("ws receive failed peer=%s", peer) break - line = raw.strip() if not line: continue messages += 1 - try: req = json.loads(line) except json.JSONDecodeError as exc: parse_errors += 1 - _log.warning( - "ws parse error peer=%s index=%d error=%s payload=%r", - peer, messages, exc, line[:_WS_LOG_PAYLOAD_PREVIEW], - ) - if not await _reply( - _error_frame(-32700, "parse error", None), - "send_failed_after_parse_error", "ws parse-error reply send failed peer=%s", peer, - ): - break + _log.warning("ws parse error peer=%s index=%d error=%s payload=%r", peer, messages, exc, line[:_WS_LOG_PAYLOAD_PREVIEW]) + await _reply(_error(-32700, "parse error", None), "send_failed_after_parse_error", + "ws parse-error reply send failed peer=%s", peer) continue - req_id = req.get("id") if isinstance(req, dict) else None req_method = req.get("method") if isinstance(req, dict) else None - if req_method == "gateway.ping": - if not await _reply( - {"jsonrpc": "2.0", "result": {"ok": True}, "id": req_id}, - "send_failed_after_heartbeat", "ws heartbeat reply send failed peer=%s id=%s", peer, req_id, - ): - break + await _reply({"jsonrpc": "2.0", "result": {"ok": True}, "id": req_id}, "send_failed_after_heartbeat", + "ws heartbeat reply send failed peer=%s id=%s", peer, req_id) continue - # dispatch() may schedule long handlers on the pool; it returns None then and the worker # writes the response itself via transport.write (a separate thread, so that is the safe # path). Inline handlers return the response dict, written here from the loop. @@ -364,47 +311,35 @@ async def handle_ws(ws: Any, *, auth_identity: dict | None = None, subprotocol: except Exception: dispatch_crashes += 1 _log.exception("ws dispatch crash peer=%s id=%s method=%s", peer, req_id, req_method) - if not await _reply( - _error_frame(-32603, "internal error", req_id), - "send_failed_after_dispatch_crash", - "ws dispatch-crash reply send failed peer=%s id=%s method=%s", peer, req_id, req_method, - ): - break + await _reply(_error(-32603, "internal error", req_id), "send_failed_after_dispatch_crash", + "ws dispatch-crash reply send failed peer=%s id=%s method=%s", peer, req_id, req_method) continue - if resp is not None and not await _reply( - resp, "send_failed_after_response", - "ws response send failed peer=%s id=%s method=%s", peer, req_id, req_method, - ): - break + if resp is not None: + await _reply(resp, "send_failed_after_response", + "ws response send failed peer=%s id=%s method=%s", peer, req_id, req_method) + except _SendFailed: + pass finally: reaped_sessions = detached_sessions = 0 if transport is not None: server.unregister_live_transport(transport) - - # Owner-safely park browser controllers this transport registered (a reconnect with the - # same identity may deliver a terminal result for in-flight work; no new dispatch is - # admitted while offline). Offloaded: disconnect takes the controller's send_lock, which - # a worker-thread dispatch may hold while blocking on THIS loop to transmit - # (result(timeout=10)); inline would park the loop behind it. + # Owner-safely park browser controllers this transport registered (a same-identity reconnect may + # deliver a terminal result for in-flight work). Offloaded: disconnect takes the controller's + # send_lock, which a worker-thread dispatch may hold while blocking on THIS loop to transmit. try: from gateway.browser_control_broker import get_browser_control_broker - await asyncio.to_thread(get_browser_control_broker().disconnect_owner, transport) except Exception: _log.exception("ws browser-controller disconnect failed peer=%s", peer) - transport.close() - try: await asyncio.to_thread(server._release_wake_for_transport, transport) except Exception: _log.exception("ws wake-word teardown failed peer=%s", peer) - - # The single WS-disconnect teardown path: reap sessions this transport owned - # (close_on_disconnect sidecars) or detach the rest to the drop sentinel so later emits - # don't hit a closed socket; detached ones go to the grace-windowed orphan reaper (a quick - # resume cancels it). Offloaded: worker.close() blocks (terminate + waits) plus a sync DB - # write, which inline would freeze the loop for every other peer. + # The single WS-disconnect teardown path: reap sessions this transport owned (close_on_disconnect + # sidecars) or detach the rest to the drop sentinel so later emits don't hit a closed socket; detached + # ones go to the grace-windowed orphan reaper (a quick resume cancels it). Offloaded: worker.close() + # blocks (terminate + waits) plus a sync DB write, which inline would freeze the loop for every peer. try: reaped_sessions, detached_sessions = await asyncio.to_thread( server._close_sessions_for_transport, transport, end_reason="ws_disconnect" @@ -418,6 +353,5 @@ async def handle_ws(ws: Any, *, auth_identity: dict | None = None, subprotocol: _log.info( "ws closed peer=%s reason=%s messages=%d parse_errors=%d " "dispatch_crashes=%d send_failures=%d reaped_sessions=%d detached_sessions=%d", - peer, disconnect_reason, messages, parse_errors, - dispatch_crashes, send_failures, reaped_sessions, detached_sessions, + peer, disconnect_reason, messages, parse_errors, dispatch_crashes, send_failures, reaped_sessions, detached_sessions, )