From ead47d48ed2b5fbe90266a4102638fc40e973c4e Mon Sep 17 00:00:00 2001 From: Benjamin Simon Date: Wed, 30 Sep 2026 17:14:30 +0200 Subject: [PATCH] fix twisted websocket close after a client-initiated close Since the client's close event is queued before the reactor echoes it and finishes the request, the application can close the websocket while wsproto is already CLOSED but the request is not finished yet, which raised a LocalProtocolError. wsSend is now a no-op once a close was sent. Co-Authored-By: Claude Opus 5.5 --- rolo/serving/twisted.py | 4 +- tests/serving/test_twisted.py | 89 ++++++++++++++++++++++++++++++ tests/websocket/test_websockets.py | 42 ++++++++++++++ 3 files changed, 134 insertions(+), 1 deletion(-) diff --git a/rolo/serving/twisted.py b/rolo/serving/twisted.py index 699b0ae..1b38aac 100644 --- a/rolo/serving/twisted.py +++ b/rolo/serving/twisted.py @@ -18,7 +18,7 @@ from twisted.web.server import Request as TwistedRequest from twisted.web.wsgi import WSGIResource, _WSGIResponse, _wsgiString from werkzeug.datastructures import Headers -from wsproto import ConnectionType, WSConnection, events +from wsproto import ConnectionState, ConnectionType, WSConnection, events from zope.interface import implementer from rolo.gateway import Gateway @@ -352,6 +352,8 @@ def wsSend(self, event: events.Event): request = self.request if request.finished: return + if self.wsproto.state in (ConnectionState.LOCAL_CLOSING, ConnectionState.CLOSED): + return data = self.wsproto.send(event) if isinstance(event, events.AcceptConnection): self.upgraded = True diff --git a/tests/serving/test_twisted.py b/tests/serving/test_twisted.py index 3740977..832e1cf 100644 --- a/tests/serving/test_twisted.py +++ b/tests/serving/test_twisted.py @@ -2,12 +2,19 @@ import io import json +import pytest import requests +from twisted.web.http_headers import Headers as TwistedHeaders +from wsproto import ConnectionType, WSConnection, events +from wsproto.connection import ConnectionState from rolo import Request, Router, route from rolo.dispatcher import handler_dispatcher from rolo.gateway import Gateway from rolo.gateway.handlers import RouterHandler +from rolo.serving.twisted import TwistedWebSocketAdapter, WebSocketChannel +from rolo.websocket.adapter import CreateConnection, TextMessage +from rolo.websocket.request import WebSocketDisconnectedError def test_large_file_upload(serve_twisted_gateway): @@ -52,3 +59,85 @@ def hello(request: Request): response_data = json.loads(response.read()) assert response_data["path"] == "/hello" assert response_data["raw_uri"].startswith("http") + + +class _FakeTransport: + def __init__(self): + self.written = [] + + def write(self, data: bytes): + self.written.append(data) + + def loseConnection(self): + pass + + +class _UnfinishedRequest: + """A twisted request stand-in whose ``finish()`` does not mark it finished, which freezes the channel in + the window where the closing handshake is complete, but the reactor has not finished the request yet. + """ + + def __init__(self, headers: list[tuple[bytes, bytes]]): + self.requestHeaders = TwistedHeaders() + for k, v in headers: + self.requestHeaders.addRawHeader(k, v) + self.path = b"/" + self.transport = _FakeTransport() + self.finished = False + self.startedWriting = 0 + + def finish(self): + pass + + +def _accepted_websocket() -> tuple[TwistedWebSocketAdapter, WSConnection, _FakeTransport]: + client = WSConnection(ConnectionType.CLIENT) + request_head = client.send(events.Request(host="localhost", target="/")).split(b"\r\n\r\n")[0] + headers = [ + tuple(part.strip() for part in line.split(b":", 1)) + for line in request_head.split(b"\r\n")[1:] + ] + request = _UnfinishedRequest(headers) + channel = WebSocketChannel(request) + channel.initiateUpgrade() + adapter = TwistedWebSocketAdapter(channel) + assert isinstance(adapter.receive(timeout=1), CreateConnection) + adapter.accept() + + client.receive_data(b"".join(request.transport.written)) + assert isinstance(next(client.events()), events.AcceptConnection) + request.transport.written.clear() + return adapter, client, request.transport + + +def test_websocket_close_after_client_close_before_request_finished(): + adapter, client, transport = _accepted_websocket() + + adapter.channel.dataReceived(client.send(events.CloseConnection(4001, "bye"))) + assert adapter.channel.wsproto.state == ConnectionState.CLOSED + assert not adapter.channel.closed + echoed = b"".join(transport.written) + client.receive_data(echoed) + assert next(client.events()) == events.CloseConnection(4001, "bye") + + with pytest.raises(WebSocketDisconnectedError) as e: + adapter.receive(timeout=1) + assert (e.value.code, e.value.reason) == (4001, "bye") + + adapter.close(1000) + adapter.close(1000) + + assert b"".join(transport.written) == echoed + + +def test_websocket_close_twice_before_request_finished(): + adapter, client, transport = _accepted_websocket() + + adapter.close(1000) + assert adapter.channel.wsproto.state == ConnectionState.LOCAL_CLOSING + sent = b"".join(transport.written) + + adapter.close(1000) + adapter.send(TextMessage("too late")) + + assert b"".join(transport.written) == sent diff --git a/tests/websocket/test_websockets.py b/tests/websocket/test_websockets.py index b830e94..3e59534 100644 --- a/tests/websocket/test_websockets.py +++ b/tests/websocket/test_websockets.py @@ -6,9 +6,11 @@ import pytest import websocket from _pytest.fixtures import SubRequest +from twisted.python import threadable from werkzeug.datastructures import Headers from rolo import Response, Router +from rolo.serving.twisted import WebSocketChannel from rolo.websocket.request import ( WebSocketDisconnectedError, WebSocketProtocolError, @@ -158,6 +160,46 @@ def app(request: WebSocketRequest): assert closes.get(timeout=3) == (4001, "test reason") +def test_server_close_after_client_close(serve_twisted_websocket_listener, monkeypatch): + """The server application closing the websocket after the client closed it must not fail, even when + it runs before the reactor has finished the request of the completed closing handshake. The reactor is + held in that window until the application has closed. Only tested with twisted, since the window is + specific to its channel.""" + app_closed = threading.Event() + results = Queue() + + original_close = WebSocketChannel.close + + def close(self): + if threadable.isInIOThread(): + app_closed.wait(timeout=3) + original_close(self) + + monkeypatch.setattr(WebSocketChannel, "close", close) + + @WebSocketRequest.listener + def app(request: WebSocketRequest): + try: + with request.accept() as ws: + with pytest.raises(WebSocketDisconnectedError): + ws.receive() + ws.close() + except Exception as e: + results.put(e) + else: + results.put("ok") + finally: + app_closed.set() + + server = serve_twisted_websocket_listener(app) + + client = websocket.WebSocket() + client.connect(server.url.replace("http://", "ws://")) + client.close(status=4001, reason=b"test reason") + + assert results.get(timeout=5) == "ok" + + def test_close_handshake_server_initiated(serve_websocket_listener): """When the server closes the websocket, the client has to receive a proper close frame, followed by the termination of the TCP connection."""