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."""