diff --git a/tests/test_tui_gateway_ws.py b/tests/test_tui_gateway_ws.py index 096f5a4c2c..7dbebed65c 100644 --- a/tests/test_tui_gateway_ws.py +++ b/tests/test_tui_gateway_ws.py @@ -153,6 +153,50 @@ def test_ws_starts_mcp_discovery_before_ready(monkeypatch): assert events == ["accept", "ready_after_0"] +def test_ws_ready_advertises_heartbeat_and_ping_is_inline(monkeypatch): + sent = [] + inbound = iter( + [ + json.dumps( + { + "jsonrpc": "2.0", + "id": "heartbeat-1", + "method": "gateway.ping", + "params": {}, + } + ) + ] + ) + monkeypatch.setattr(server, "_WS_ORPHAN_REAP_GRACE_S", 0) + + class FakeWS: + async def accept(self): + pass + + async def send_text(self, line): + sent.append(json.loads(line)) + + async def receive_text(self): + try: + return next(inbound) + except StopIteration: + raise ws_mod._WebSocketDisconnect() + + async def close(self): + pass + + asyncio.run(ws_mod.handle_ws(FakeWS())) + + ready = sent[0]["params"] + assert ready["type"] == "gateway.ready" + assert ready["payload"]["heartbeat"] is True + assert sent[1] == { + "jsonrpc": "2.0", + "result": {"ok": True}, + "id": "heartbeat-1", + } + + def test_ws_transport_serializes_concurrent_sends(): active_sends = 0 max_active_sends = 0 diff --git a/tui_gateway/ws.py b/tui_gateway/ws.py index dc1367d38f..be324ae55f 100644 --- a/tui_gateway/ws.py +++ b/tui_gateway/ws.py @@ -29,6 +29,7 @@ import json import logging import socket import threading +import time from typing import Any from tui_gateway import server @@ -103,6 +104,7 @@ class WSTransport: #: browser-controller registration. self.auth_identity = auth_identity self._closed = False + self._last_inbound_at = time.monotonic() # Token-coalescing buffer (CF-2). Streamed token frames land here and a # short timer flushes the batch. The lock guards the buffer + the # "armed" flag against the worker threads that call write(); the timer @@ -116,6 +118,17 @@ class WSTransport: # the owning loop while it recovers from a stall. self._send_lock = asyncio.Lock() + @property + def closed(self) -> bool: + return self._closed + + @property + def last_inbound_at(self) -> float: + return self._last_inbound_at + + def mark_inbound(self) -> None: + self._last_inbound_at = time.monotonic() + @staticmethod def _is_streaming_frame(obj: dict) -> bool: """True for high-frequency per-token frames eligible for coalescing.""" @@ -361,7 +374,11 @@ async def handle_ws( # change_events: this backend broadcasts pet.changed / # cron.changed / sessions.changed, so clients can demote # their legacy polls to slow backstops. - "payload": {"skin": skin_payload, "change_events": True}, + "payload": { + "skin": skin_payload, + "change_events": True, + "heartbeat": True, + }, }, } ) @@ -404,6 +421,7 @@ async def handle_ws( line = raw.strip() if not line: continue + transport.mark_inbound() messages += 1 try: @@ -438,6 +456,22 @@ async def handle_ws( # response dict, which we write here from the loop. 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": + ok = await transport.write_async( + { + "jsonrpc": "2.0", + "result": {"ok": True}, + "id": req_id, + } + ) + if not ok: + disconnect_reason = "send_failed_after_heartbeat" + send_failures += 1 + _log.warning("ws heartbeat reply send failed peer=%s id=%s", peer, req_id) + break + continue + try: resp = await asyncio.to_thread(server.dispatch, req, transport) except Exception: