From 59004a62356f3a4697ab0fe8ad5086d2b405e2a6 Mon Sep 17 00:00:00 2001 From: kshitijk4poor <82637225+kshitijk4poor@users.noreply.github.com> Date: Thu, 24 Sep 2026 21:54:57 +0530 Subject: [PATCH] fix(tui): socket-close only on fanout overflow; plain WSTransport.close() back to main WSTransport.close() scheduled ws.close(code=1011) on every call, so handle_ws's normal teardown reported 1011 before its own close. Move the off-loop socket close into a one-shot WSTransport.abort() that the fanout overflow path calls; close() is byte-identical to main again. Rename _close_stalled_socket -> _close_socket(code, reason) with accurate log text, share an _on_loop() helper with write(), and make the overflow test's slow peer a real WSTransport whose socket close must be awaited with 1011. Co-authored-by: KoNit-K --- tests/tui_gateway/test_multi_client_fanout.py | 38 +++++++----- tui_gateway/transport.py | 8 ++- tui_gateway/ws.py | 59 ++++++++++--------- 3 files changed, 61 insertions(+), 44 deletions(-) diff --git a/tests/tui_gateway/test_multi_client_fanout.py b/tests/tui_gateway/test_multi_client_fanout.py index 523879d062..54b24c6919 100644 --- a/tests/tui_gateway/test_multi_client_fanout.py +++ b/tests/tui_gateway/test_multi_client_fanout.py @@ -46,9 +46,9 @@ def _await_frame_count(transport, count, timeout=2.0): def _overflow_slow_peer(fan, healthy, slow): - """Emit until the slow mailbox overflows; pace healthy one receipt per emit.""" + """Emit non-streaming frames (a WS peer blocks on each) until the slow mailbox overflows.""" for n in range(FanoutTransport._MAX_PENDING_FRAMES + 64): - frame = {"params": {"type": "message.delta", "n": n}} + frame = {"params": {"type": "tool.progress", "n": n}} assert fan.write(frame) _await_frame_count(healthy, n + 1) if not fan.contains(slow): @@ -93,15 +93,18 @@ class PipeClient: class _SocketWS: - """ASGI-ws stand-in for SocketClient: send goes to the socketpair, close is a no-op.""" - def __init__(self, client): - self.client = client + """ASGI-ws stand-in: send goes to the socketpair (or never completes when client is None); + close records its code.""" + def __init__(self, client=None): + self.client, self.close_codes = client, [] async def send_text(self, payload): + if self.client is None: + await asyncio.Event().wait() await self.client.send_text(payload) async def close(self, code=1000): - return None + self.close_codes.append(code) class SocketClient(WSTransport): @@ -295,18 +298,23 @@ def test_backpressure_never_blocks_later_frames_or_other_subscribers(slow_first, assert not fan.write({"after": "close"}) - def test_overflow_closes_only_the_slow_peer_and_healthy_keeps_streaming(): - class BoomOnClose(RecordingTransport): - def close(self): - super().close() - raise RuntimeError("overflow close exploded") - + loop = asyncio.new_event_loop() + loop_thread = threading.Thread(target=loop.run_forever, daemon=True) + loop_thread.start() + stalled_ws = _SocketWS() # send_text never completes: the real slow WS peer + slow = WSTransport(stalled_ws, loop, peer="slow") healthy = RecordingTransport() - slow = BoomOnClose(delay=30.0) fan = FanoutTransport(healthy, slow) try: last_n = _overflow_slow_peer(fan, healthy, slow) + deadline = time.monotonic() + 2.0 + while not stalled_ws.close_codes and time.monotonic() < deadline: + time.sleep(0.001) + assert stalled_ws.close_codes == [1011] # the overflow itself aborted the socket + slow.abort() # a second overflow signal must not schedule a second socket close + time.sleep(0.05) + assert stalled_ws.close_codes == [1011] assert slow.closed is True assert healthy.closed is False assert fan.contains(healthy) @@ -315,8 +323,10 @@ def test_overflow_closes_only_the_slow_peer_and_healthy_keeps_streaming(): _await_frame_count(healthy, last_n + 2) assert healthy.frames[-1] == after finally: - slow.release() fan.close() + loop.call_soon_threadsafe(loop.stop) + loop_thread.join(timeout=2) + loop.close() def test_fanout_close_and_detach_leave_peer_sockets_open(): diff --git a/tui_gateway/transport.py b/tui_gateway/transport.py index f5bf36bdcd..689f561429 100644 --- a/tui_gateway/transport.py +++ b/tui_gateway/transport.py @@ -239,10 +239,12 @@ class FanoutTransport: return def _signal_overflow_detach(self, transport: Transport) -> None: - # Outside the fanout lock: close() may re-enter contains/detach, and a - # WS close must not stall the emit turn or other subscribers. + # Outside the fanout lock: abort()/close() may re-enter contains/detach, and + # a WS close must not stall the emit turn or other subscribers. WSTransport + # aborts (1011 socket close, off-loop safe); other transports just close. try: - transport.close() + abort = getattr(transport, "abort", None) + (abort or transport.close)() except Exception: logger.debug("fanout overflow close failed; membership already dropped", exc_info=True) diff --git a/tui_gateway/ws.py b/tui_gateway/ws.py index 93663350bd..f21dfb24ab 100644 --- a/tui_gateway/ws.py +++ b/tui_gateway/ws.py @@ -107,15 +107,19 @@ class WSTransport: self._token_flush_armed = False # Socket writes need an async boundary: several batches can queue on the loop during a stall. self._send_lock = asyncio.Lock() + self._abort_requested = False + + def _on_loop(self) -> bool: + try: + return asyncio.get_running_loop() is self._loop + except RuntimeError: + return False def write(self, obj: dict) -> bool: if self._closed: return False line = serialize_frame(obj, self._peer, _log) - try: - on_loop = asyncio.get_running_loop() is self._loop - except RuntimeError: - on_loop = False + on_loop = self._on_loop() # Streamed token: buffer it and arm the flush timer; the worker returns immediately. # call_soon_threadsafe is safe from a worker or the loop. params = obj.get("params") if isinstance(obj, dict) else None @@ -202,7 +206,7 @@ class WSTransport: self._closed = True _log.warning("ws send deadline exceeded (socket stalled, loop responsive) peer=%s deadline=%ss — closing", self._peer, _WS_SEND_DEADLINE_S) - self._loop.create_task(self._close_stalled_socket()) + self._loop.create_task(self._close_socket(1011, "send deadline")) return except UnicodeEncodeError as exc: # A single illegal UTF-8 frame (lone surrogate) must not tear down the socket. @@ -214,37 +218,38 @@ class WSTransport: _log.warning("ws send failed peer=%s error_type=%s error=%s", self._peer, type(exc).__name__, exc) return - def close(self) -> None: - # Latch first so heartbeats/writes fail immediately. Fanout overflow may - # call this off-loop; TimerHandle.cancel and ws.close belong on the loop. + def close(self) -> None: # loop thread (handle_ws finally), so the TimerHandle is safe self._closed = True - - def _finish_close() -> None: # loop thread - handle = self._token_flush_handle + if self._token_flush_handle is not None: + self._token_flush_handle.cancel() self._token_flush_handle = None - if handle is not None: - handle.cancel() - if self._ws is not None: - self._loop.create_task(self._close_stalled_socket()) - try: - on_loop = asyncio.get_running_loop() is self._loop - except RuntimeError: - on_loop = False - if on_loop: - _finish_close() + def abort(self) -> None: + """Close from any thread and drop the socket with 1011 so the client reconnects and replays + (fanout overflow). One-shot: N mirrored sessions overflowing on this socket schedule one close.""" + self._closed = True + with self._token_lock: + if self._abort_requested: + return + self._abort_requested = True + if self._on_loop(): + self._finish_abort() return # A loop that already shut down has nothing left to cancel or close. with contextlib.suppress(RuntimeError): - self._loop.call_soon_threadsafe(_finish_close) + self._loop.call_soon_threadsafe(self._finish_abort) - async def _close_stalled_socket(self) -> None: - """Close the peer socket after a send deadline so ``handle_ws``'s ``receive_text`` unblocks and its - disconnect teardown runs. The server library bounds this (websockets ``close_timeout`` → abort).""" + def _finish_abort(self) -> None: # loop thread + self.close() + self._loop.create_task(self._close_socket(1011, "fanout overflow")) + + async def _close_socket(self, code: int, reason: str) -> None: + """Close the peer socket so ``handle_ws``'s ``receive_text`` unblocks and its disconnect teardown + runs. The server library bounds this (websockets ``close_timeout`` → abort).""" try: - await self._ws.close(code=1011) + await self._ws.close(code=code) except Exception as exc: # noqa: BLE001 - the peer is already gone; teardown is what matters - _log.debug("ws close after send deadline failed peer=%s error=%s", self._peer, exc) + _log.debug("ws close after %s failed peer=%s error=%s", reason, self._peer, exc) def _ws_peer_label(ws: Any) -> str: