From 2324dcd74435797ebb83127a5b2509cb6cc279c9 Mon Sep 17 00:00:00 2001 From: YueGuobin Date: Fri, 28 Aug 2026 23:35:33 +0800 Subject: [PATCH] fix: handle client disconnect in console WebSocket forwarding The compute-to-client forwarding loops in ws_console and vnc_console had no WebSocketDisconnect handling: when a client (WebUI, or an MCP-driven websocat session killed by timeout) disconnected while the compute was still streaming console output, the next send raised WebSocketDisconnect that leaked all the way up to uvicorn as an ERROR-level ASGI traceback. Catch it and log at info level, symmetric with the receive-side handlers. --- gns3server/api/routes/controller/nodes.py | 12 ++ tests/api/routes/controller/test_nodes.py | 188 +++++++++++++++++++++- 2 files changed, 199 insertions(+), 1 deletion(-) diff --git a/gns3server/api/routes/controller/nodes.py b/gns3server/api/routes/controller/nodes.py index 47f9e57aa..ed17a703e 100644 --- a/gns3server/api/routes/controller/nodes.py +++ b/gns3server/api/routes/controller/nodes.py @@ -736,6 +736,12 @@ async def ws_console( await websocket.send_bytes(msg.data) elif msg.type == aiohttp.WSMsgType.ERROR: break + except WebSocketDisconnect: + # the client disconnected while the compute was still streaming console output + log.info( + f"Client {websocket.client.host}:{websocket.client.port} has disconnected from controller" + f" console WebSocket" + ) except aiohttp.ClientError as e: log.error(f"Client error received when forwarding to compute console WebSocket: {e}") @@ -811,6 +817,12 @@ async def vnc_console( await websocket.send_bytes(msg.data) elif msg.type == aiohttp.WSMsgType.ERROR: break + except WebSocketDisconnect: + # the client disconnected while the compute was still streaming VNC console output + log.info( + f"Client {websocket.client.host}:{websocket.client.port} has disconnected from controller" + f" VNC console WebSocket" + ) except aiohttp.ClientError as e: log.error(f"Client error received when forwarding to compute VNC console WebSocket: {e}") diff --git a/tests/api/routes/controller/test_nodes.py b/tests/api/routes/controller/test_nodes.py index bded58728..02ca2c143 100644 --- a/tests/api/routes/controller/test_nodes.py +++ b/tests/api/routes/controller/test_nodes.py @@ -16,17 +16,27 @@ # along with this program. If not, see . +import aiohttp +import logging + import pytest -from fastapi import FastAPI, HTTPException, status +from types import SimpleNamespace +from typing import List, Optional + +from fastapi import FastAPI, HTTPException, WebSocketDisconnect, status from httpx import AsyncClient +from pydantic import SecretStr from unittest.mock import MagicMock from tests.utils import AsyncioMagicMock +from gns3server.config import Config from gns3server.controller.node import Node from gns3server.controller.project import Project from gns3server.controller.compute import Compute +from gns3server.utils.http_client import HTTPClient +from gns3server.api.routes.controller.nodes import ws_console, vnc_console pytestmark = pytest.mark.asyncio @@ -669,3 +679,179 @@ class TestNodeRoutes: # assert response.status_code == 201 # # compute.http_query.assert_called_with("POST", "/projects/{project_id}/files/project-files/vpcs/{node_id}/hello/nested".format(project_id=project.id, node_id=node.id), data=b'hello', timeout=None, raw=True) + + +class FakeComputeConsoleWebSocket: + """ + Stand-in for the aiohttp ClientWebSocketResponse returned when connecting + to the compute console WebSocket, yielding queued messages. + """ + + def __init__(self, messages: List[aiohttp.WSMessage]): + + self._messages = messages + self.closed = False + + async def __aenter__(self) -> "FakeComputeConsoleWebSocket": + + return self + + async def __aexit__(self, *exc_info) -> bool: + + self.closed = True + return False + + def __aiter__(self): + + return self._iterate_messages() + + async def _iterate_messages(self): + + for message in self._messages: + yield message + + async def close(self) -> None: + + self.closed = True + + async def send_str(self, data: str) -> None: + # client -> compute traffic, not exercised here + pass + + async def send_bytes(self, data: bytes) -> None: + pass + + +class FakeClientWebSocket: + """ + Stand-in for the starlette WebSocket facing the client. When fail_after is + set, raises WebSocketDisconnect once that many sends succeeded, mimicking + uvicorn raising ClientDisconnected when the client is gone mid-stream. + """ + + def __init__(self, fail_after: Optional[int] = None): + + self.url = SimpleNamespace(scheme="http") + self.client = SimpleNamespace(host="127.0.0.1", port=5000) + self.sent: List[tuple] = [] + self._fail_after = fail_after + + async def receive(self) -> dict: + + # the client is already gone from the receive side + return {"type": "websocket.disconnect"} + + async def receive_bytes(self) -> bytes: + + raise WebSocketDisconnect(code=1006) + + async def send_text(self, data: str) -> None: + + self._check_client_alive() + self.sent.append(("text", data)) + + async def send_bytes(self, data: bytes) -> None: + + self._check_client_alive() + self.sent.append(("bytes", data)) + + def _check_client_alive(self) -> None: + + if self._fail_after is not None and len(self.sent) >= self._fail_after: + raise WebSocketDisconnect(code=1006) + + +class TestNodeConsoleWebSocketRoutes: + """ + Exercise the console/VNC WebSocket forwarding handlers directly. + + The in-process ASGI WebSocket transport never raises on send once the + client disconnected (unlike uvicorn's real WebSocket protocol), so client + disconnects mid-stream are reproduced with a fake client WebSocket. + """ + + @pytest.fixture + def node(self, project: Project, compute: Compute) -> Node: + + compute.host = "127.0.0.1" + compute.port = 3080 + node = Node(project, compute, "test", node_type="vpcs") + project._nodes[node.id] = node + return node + + @pytest.fixture + def compute_credentials(self) -> None: + + server_config = Config.instance().settings.Server + server_config.compute_username = "admin" + server_config.compute_password = SecretStr("password") + + @staticmethod + def _forward_compute_ws(monkeypatch, messages: List[aiohttp.WSMessage]) -> FakeComputeConsoleWebSocket: + + compute_ws = FakeComputeConsoleWebSocket(messages) + monkeypatch.setattr( + HTTPClient, + "get_client", + classmethod(lambda cls: SimpleNamespace(ws_connect=lambda *args, **kwargs: compute_ws)) + ) + return compute_ws + + async def test_console_forwards_compute_output_to_client( + self, + compute_credentials, + node: Node, + monkeypatch + ) -> None: + + compute_ws = self._forward_compute_ws(monkeypatch, [ + aiohttp.WSMessage(aiohttp.WSMsgType.TEXT, "device output", None), + aiohttp.WSMessage(aiohttp.WSMsgType.BINARY, b"\x00\x01", None), + ]) + websocket = FakeClientWebSocket() + + await ws_console(websocket, current_user=MagicMock(), node=node) + + assert websocket.sent == [("text", "device output"), ("bytes", b"\x00\x01")] + assert compute_ws.closed + + async def test_console_client_disconnect_mid_stream( + self, + compute_credentials, + node: Node, + monkeypatch, + caplog + ) -> None: + # regression test: the client disconnects while the compute is still + # streaming console output, send must not leak WebSocketDisconnect + compute_ws = self._forward_compute_ws(monkeypatch, [ + aiohttp.WSMessage(aiohttp.WSMsgType.TEXT, "line 1", None), + aiohttp.WSMessage(aiohttp.WSMsgType.TEXT, "line 2", None), + aiohttp.WSMessage(aiohttp.WSMsgType.TEXT, "line 3", None), + ]) + websocket = FakeClientWebSocket(fail_after=1) + + with caplog.at_level(logging.INFO): + await ws_console(websocket, current_user=MagicMock(), node=node) + + assert websocket.sent == [("text", "line 1")] + assert compute_ws.closed + assert any("has disconnected from controller console WebSocket" in record.message for record in caplog.records) + + async def test_vnc_console_client_disconnect_mid_stream( + self, + compute_credentials, + node: Node, + monkeypatch + ) -> None: + # regression test: same as above for the VNC console forwarding loop + compute_ws = self._forward_compute_ws(monkeypatch, [ + aiohttp.WSMessage(aiohttp.WSMsgType.BINARY, b"\x01\x02", None), + aiohttp.WSMessage(aiohttp.WSMsgType.BINARY, b"\x03\x04", None), + ]) + websocket = FakeClientWebSocket(fail_after=0) + + await vnc_console(websocket, current_user=MagicMock(), node=node) + + assert websocket.sent == [] + assert compute_ws.closed