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.
This commit is contained in:
Teknium
2026-09-02 22:47:09 -07:00
parent d41d546d0e
commit 03c191076d
4 changed files with 268 additions and 475 deletions

View File

@@ -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:<path>`` 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:<path>`` 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
``<prompt>``. 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 ``<prompt>``; 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 = "<unredactable>" # 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:

View File

@@ -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:

View File

@@ -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

View File

@@ -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,
)