diff --git a/rolo/serving/twisted.py b/rolo/serving/twisted.py index 5f19a1a..031b1c9 100644 --- a/rolo/serving/twisted.py +++ b/rolo/serving/twisted.py @@ -477,8 +477,12 @@ def close(self): # the 101 upgrade response was written raw to the transport, so ``Request.finish()`` # must not write its own (never started) HTTP response into the websocket stream self.request.startedWriting = 1 + # the HTTP channel is registered as producer of its transport, and a TLS transport defers its shutdown until + # no producer is registered. ``HTTPChannel.loseConnection`` unregisters it first, and ``finish`` detaches + # the channel from the request. + channel = self.request.channel self.request.finish() - self.request.transport.loseConnection() + channel.loseConnection() # special internal poison pill self.eventQueue.put_nowait(events.CloseConnection(None)) diff --git a/tests/serving/test_twisted.py b/tests/serving/test_twisted.py index 1823a83..da486ea 100644 --- a/tests/serving/test_twisted.py +++ b/tests/serving/test_twisted.py @@ -3,6 +3,8 @@ import io import json import ssl as stdlib_ssl +import struct +import threading import pytest import requests @@ -102,6 +104,8 @@ def __init__(self, headers: list[tuple[bytes, bytes]]): self.requestHeaders.addRawHeader(k, v) self.path = b"/" self.transport = _FakeTransport() + # the channel is only used to lose the connection, which the fake transport does as well + self.channel = self.transport self.finished = False self.startedWriting = 0 @@ -238,3 +242,31 @@ def app(request: WebSocketRequest): client.connect(server.url, timeout=5) assert client.recv() == "hello" client.close() + + +def test_websocket_tls_close_handshake_client_initiated(serve_twisted_tls_websocket_listener): + """After the closing handshake, the server terminates the connection over TLS too. The HTTP channel is + registered as producer of the TLS transport, which defers its shutdown until no producer is registered, + so terminating the TLS transport directly never closes the connection.""" + disconnected = threading.Event() + + @WebSocketRequest.listener + def app(request: WebSocketRequest): + with request.accept() as ws: + with pytest.raises(WebSocketDisconnectedError): + ws.receive() + disconnected.set() + + server = serve_twisted_tls_websocket_listener(app) + + client = websocket.WebSocket(sslopt={"cert_reqs": stdlib_ssl.CERT_NONE}) + client.connect(server.url, timeout=5) + client.send_close(websocket.STATUS_NORMAL) + + frame = client.recv_frame() + assert frame.opcode == websocket.ABNF.OPCODE_CLOSE + assert struct.unpack("!H", frame.data[:2])[0] == websocket.STATUS_NORMAL + + client.sock.settimeout(5) + assert client.sock.recv(1) == b"", "expected the server to terminate the TLS connection" + assert disconnected.wait(timeout=3)