Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand Down
46 changes: 40 additions & 6 deletions rolo/serving/twisted.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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. 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

def __init__(self, channel: WebSocketChannel):
def __init__(self, channel: WebSocketChannel, reactor=reactor):
self.channel = channel
self.reactor = reactor

def _mustScheduleInReactor(self) -> bool:
# without a running reactor, there's no reactor thread to race with
return 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:
Expand All @@ -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__}")

Expand All @@ -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,
Expand All @@ -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)

Expand Down
99 changes: 98 additions & 1 deletion tests/serving/test_twisted.py
Original file line number Diff line number Diff line change
@@ -1,18 +1,37 @@
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

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

Expand Down Expand Up @@ -141,3 +160,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()
27 changes: 27 additions & 0 deletions tests/websocket/test_websockets.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading