From 3830e0c5c1ebe497eb2fe768a0410d316c753412 Mon Sep 17 00:00:00 2001 From: Benjamin Simon Date: Wed, 30 Sep 2026 19:25:36 +0200 Subject: [PATCH 1/2] schedule twisted websocket writes onto the reactor on TLS connections The websocket listener runs in a threadpool thread, but wrote to the transport directly. On TLS connections, this races with the reactor processing incoming TLS records on the same OpenSSL connection, which breaks the connection (`OpenSSL.SSL.Error: []`) or corrupts frames. On TLS connections, `send` is now scheduled onto the reactor without waiting for it, and `accept`, `reject` and `close` wait for their completion. Plain connections keep writing directly. Co-Authored-By: Claude Opus 5.5 --- pyproject.toml | 2 +- rolo/serving/twisted.py | 46 +++++++++++++-- tests/serving/test_twisted.py | 102 +++++++++++++++++++++++++++++++++- 3 files changed, 142 insertions(+), 8 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 140b8f8..4e153d7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,7 +45,7 @@ dev = [ "websocket-client>=1.7.0", "coverage[toml]>=5.0.0", "coveralls>=3.3", - "twisted>=24", + "twisted[tls]>=24", "ruff==0.1.0" ] docs = [ diff --git a/rolo/serving/twisted.py b/rolo/serving/twisted.py index 1b38aac..7cafb3e 100644 --- a/rolo/serving/twisted.py +++ b/rolo/serving/twisted.py @@ -9,7 +9,9 @@ from twisted.internet import reactor from twisted.internet.protocol import Protocol +from twisted.internet.threads import blockingCallFromThread from twisted.protocols.policies import ProtocolWrapper +from twisted.python import threadable from twisted.python.components import proxyForInterface from twisted.web.http import HTTPChannel, _GenericHTTPChannelProtocol, urlparse from twisted.web.http_headers import Headers as TwistedHeaders @@ -297,7 +299,7 @@ def _processWebsocket(self, request: Request): channel.initiateUpgrade() environment = to_websocket_environment(request) - environment["rolo.websocket"] = TwistedWebSocketAdapter(channel) + environment["rolo.websocket"] = TwistedWebSocketAdapter(channel, self.original._reactor) # WSGIResource also dispatches requests through the threadpool self.original._threadpool.callInThread(self.websocketListener, environment) @@ -404,10 +406,36 @@ def close(self): class TwistedWebSocketAdapter(rolows.WebSocketAdapter): + """ + Adapter between the ``WebSocketChannel``, which lives in the reactor thread, and the ``WebSocketListener``, + which runs in a threadpool thread. On TLS connections, every operation that touches the channel's connection + state or transport is scheduled onto the reactor thread, since writing to the TLS connection from the listener + thread races with the reactor processing incoming TLS records, which corrupts the connection. ``send`` is + scheduled without waiting for the write, the other operations wait for their completion. Plain connections + keep writing directly from the listener thread, which avoids the cost of the thread handoff. + """ + channel: WebSocketChannel - def __init__(self, channel: WebSocketChannel): + def __init__(self, channel: WebSocketChannel, reactor=reactor): self.channel = channel + self.reactor = reactor + self._isTLS = channel.request.isSecure() + + def _mustScheduleInReactor(self) -> bool: + # without a running reactor, there's no reactor thread to race with + return self._isTLS and self.reactor.running and not threadable.isInIOThread() + + def _callInReactor(self, f: t.Callable, *args): + if self._mustScheduleInReactor(): + return blockingCallFromThread(self.reactor, f, *args) + return f(*args) + + def _sendInReactor(self, event: events.Event): + if self._mustScheduleInReactor(): + self.reactor.callFromThread(self.channel.wsSend, event) + else: + self.channel.wsSend(event) def receive(self, timeout: float = None) -> rolows.CreateConnection | rolows.Message: try: @@ -430,9 +458,9 @@ def receive(self, timeout: float = None) -> rolows.CreateConnection | rolows.Mes def send(self, event: rolows.Message, timeout: float = None): if isinstance(event, rolows.TextMessage): - self.channel.wsSend(events.TextMessage(event.data)) + self._sendInReactor(events.TextMessage(event.data)) elif isinstance(event, rolows.BytesMessage): - self.channel.wsSend(events.BytesMessage(event.data)) + self._sendInReactor(events.BytesMessage(event.data)) else: raise TypeError(f"Unexpected event type {event.__class__.__name__}") @@ -443,7 +471,10 @@ def reject( body: t.Iterable[bytes] = None, timeout: float = None, ): - self.channel.wsReject(status_code, headers, body) + # consume the body here, so the reactor thread doesn't run arbitrary (possibly blocking) iterators + self._callInReactor( + self.channel.wsReject, status_code, headers, list(body) if body else None + ) def accept( self, @@ -461,9 +492,12 @@ def accept( # TODO: extensions event = events.AcceptConnection(subprotocol, extensions=[], extra_headers=raw_headers) - self.channel.wsSend(event) + self._callInReactor(self.channel.wsSend, event) def close(self, code: int = 1001, reason: str = None, timeout: float = None): + self._callInReactor(self._close, code, reason) + + def _close(self, code: int, reason: str | None): if not self.channel.closed: self.channel.wsClose(code, reason) diff --git a/tests/serving/test_twisted.py b/tests/serving/test_twisted.py index 832e1cf..da64241 100644 --- a/tests/serving/test_twisted.py +++ b/tests/serving/test_twisted.py @@ -1,10 +1,20 @@ +import datetime import http.client import io import json +import ssl as stdlib_ssl import pytest import requests +import websocket +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import ec +from cryptography.x509.oid import NameOID +from twisted.internet import ssl +from twisted.protocols.tls import TLSMemoryBIOFactory from twisted.web.http_headers import Headers as TwistedHeaders +from twisted.web.server import Site from wsproto import ConnectionType, WSConnection, events from wsproto.connection import ConnectionState @@ -12,7 +22,16 @@ 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.serving.twisted import ( + HeaderPreservingHTTPChannel, + HeaderPreservingWSGIResource, + TwistedRequestAdapter, + TwistedWebSocketAdapter, + WebSocketChannel, + WebsocketResourceDecorator, +) +from rolo.testing.pytest import _ServerInfo, get_random_tcp_port, wait_server_is_up +from rolo.websocket import WebSocketListener, WebSocketRequest from rolo.websocket.adapter import CreateConnection, TextMessage from rolo.websocket.request import WebSocketDisconnectedError @@ -89,6 +108,9 @@ def __init__(self, headers: list[tuple[bytes, bytes]]): def finish(self): pass + def isSecure(self): + return False + def _accepted_websocket() -> tuple[TwistedWebSocketAdapter, WSConnection, _FakeTransport]: client = WSConnection(ConnectionType.CLIENT) @@ -141,3 +163,81 @@ def test_websocket_close_twice_before_request_finished(): adapter.send(TextMessage("too late")) assert b"".join(transport.written) == sent + + +@pytest.fixture(scope="module") +def self_signed_cert() -> tuple[bytes, bytes]: + """A self-signed certificate and private key for ``localhost``, both PEM encoded.""" + key = ec.generate_private_key(ec.SECP256R1()) + name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "localhost")]) + now = datetime.datetime.now(datetime.timezone.utc) + cert = ( + x509.CertificateBuilder() + .subject_name(name) + .issuer_name(name) + .public_key(key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(days=1)) + .not_valid_after(now + datetime.timedelta(days=1)) + .add_extension(x509.SubjectAlternativeName([x509.DNSName("localhost")]), critical=False) + .sign(key, hashes.SHA256()) + ) + key_pem = key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + return cert.public_bytes(serialization.Encoding.PEM), key_pem + + +@pytest.fixture +def serve_twisted_tls_websocket_listener(twisted_reactor, self_signed_cert): + """Like ``serve_twisted_websocket_listener``, but serves the listener over TLS.""" + ports = [] + + def _create(websocket_listener: WebSocketListener): + site = Site( + WebsocketResourceDecorator( + original=HeaderPreservingWSGIResource( + twisted_reactor, twisted_reactor.getThreadPool(), None + ), + websocketListener=websocket_listener, + ), + requestFactory=TwistedRequestAdapter, + ) + site.protocol = HeaderPreservingHTTPChannel.protocol_factory + + cert_pem, key_pem = self_signed_cert + certificate = ssl.PrivateCertificate.loadPEM(cert_pem + key_pem) + factory = TLSMemoryBIOFactory(certificate.options(), False, site) + + port = get_random_tcp_port() + ports.append(twisted_reactor.listenTCP(port, factory)) + srv = _ServerInfo("localhost", port, f"wss://localhost:{port}") + assert wait_server_is_up(srv), f"gave up waiting for {srv}" + return srv + + yield _create + + for _port in ports: + _port.stopListening() + + +def test_websocket_tls_send_and_close_from_listener_thread(serve_twisted_tls_websocket_listener): + """The listener runs in a threadpool thread, but twisted transports, and the TLS connection state + in particular, may only be touched from the reactor thread. Writing to them directly from the + listener thread races with the reactor processing the client's TLS records, which breaks the + connection before the client reads the upgrade response.""" + + @WebSocketRequest.listener + def app(request: WebSocketRequest): + with request.accept() as ws: + ws.send("hello") + + server = serve_twisted_tls_websocket_listener(app) + + for _ in range(200): + client = websocket.WebSocket(sslopt={"cert_reqs": stdlib_ssl.CERT_NONE}) + client.connect(server.url, timeout=5) + assert client.recv() == "hello" + client.close() From 9390e0e92497f04140c22031a65e1a774617a360 Mon Sep 17 00:00:00 2001 From: Benjamin Simon Date: Wed, 30 Sep 2026 23:24:06 +0200 Subject: [PATCH 2/2] schedule twisted websocket writes onto the reactor on plain connections too Writing to the transport from the listener thread also breaks plain connections: it races with the reactor flushing the same transport, which loses and reorders data once the messages exceed the socket buffers. Co-Authored-By: Claude Opus 5.5 --- rolo/serving/twisted.py | 14 +++++++------- tests/serving/test_twisted.py | 3 --- tests/websocket/test_websockets.py | 27 +++++++++++++++++++++++++++ 3 files changed, 34 insertions(+), 10 deletions(-) diff --git a/rolo/serving/twisted.py b/rolo/serving/twisted.py index 7cafb3e..4864d07 100644 --- a/rolo/serving/twisted.py +++ b/rolo/serving/twisted.py @@ -408,11 +408,12 @@ def close(self): class TwistedWebSocketAdapter(rolows.WebSocketAdapter): """ Adapter between the ``WebSocketChannel``, which lives in the reactor thread, and the ``WebSocketListener``, - which runs in a threadpool thread. On TLS connections, every operation that touches the channel's connection - state or transport is scheduled onto the reactor thread, since writing to the TLS connection from the listener - thread races with the reactor processing incoming TLS records, which corrupts the connection. ``send`` is - scheduled without waiting for the write, the other operations wait for their completion. Plain connections - keep writing directly from the listener thread, which avoids the cost of the thread handoff. + which runs in a threadpool thread. Twisted is not thread-safe, so every operation that touches the channel's + connection state or transport is scheduled onto the reactor thread. Writing from the listener thread directly + races with the reactor writing to the same transport, which loses or reorders data, and on TLS connections also + with the reactor processing incoming TLS records, which breaks the connection. ``send`` is scheduled without + waiting for the write (the reactor runs scheduled calls in order), the other operations wait for their + completion. """ channel: WebSocketChannel @@ -420,11 +421,10 @@ class TwistedWebSocketAdapter(rolows.WebSocketAdapter): def __init__(self, channel: WebSocketChannel, reactor=reactor): self.channel = channel self.reactor = reactor - self._isTLS = channel.request.isSecure() def _mustScheduleInReactor(self) -> bool: # without a running reactor, there's no reactor thread to race with - return self._isTLS and self.reactor.running and not threadable.isInIOThread() + return self.reactor.running and not threadable.isInIOThread() def _callInReactor(self, f: t.Callable, *args): if self._mustScheduleInReactor(): diff --git a/tests/serving/test_twisted.py b/tests/serving/test_twisted.py index da64241..1823a83 100644 --- a/tests/serving/test_twisted.py +++ b/tests/serving/test_twisted.py @@ -108,9 +108,6 @@ def __init__(self, headers: list[tuple[bytes, bytes]]): def finish(self): pass - def isSecure(self): - return False - def _accepted_websocket() -> tuple[TwistedWebSocketAdapter, WSConnection, _FakeTransport]: client = WSConnection(ConnectionType.CLIENT) diff --git a/tests/websocket/test_websockets.py b/tests/websocket/test_websockets.py index 3e59534..14f0962 100644 --- a/tests/websocket/test_websockets.py +++ b/tests/websocket/test_websockets.py @@ -328,3 +328,30 @@ def _handler(request: WebSocketRequest, request_args: dict): assert client.recv() == "foo" assert client.recv() == "id=bar" assert "CasedHeader" in json.loads(client.recv()) + + +def test_send_many_messages(serve_websocket_listener): + """Sending more data than the socket buffers hold must deliver every message intact and in order. The + twisted listener runs in a threadpool thread, and writing to the transport directly from there races + with the reactor flushing the same transport, which loses and reorders data.""" + messages = 10_000 + payload = "x" * 1024 + + @WebSocketRequest.listener + def app(request: WebSocketRequest): + with request.accept() as ws: + assert ws.receive() == "start" + for i in range(messages): + ws.send(f"{i:08d}{payload}") + assert ws.receive() == "done" + + server = serve_websocket_listener(app) + + for _ in range(3): + client = websocket.WebSocket() + client.connect(server.url.replace("http://", "ws://"), timeout=5) + client.send("start") + for i in range(messages): + assert client.recv() == f"{i:08d}{payload}" + client.send("done") + client.close()