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
71 changes: 68 additions & 3 deletions rolo/serving/twisted.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from typing import Iterator, Sequence, Tuple, Union

from twisted.internet import reactor
from twisted.internet.interfaces import IPushProducer
from twisted.internet.protocol import Protocol
from twisted.internet.threads import blockingCallFromThread
from twisted.protocols.policies import ProtocolWrapper
Expand Down Expand Up @@ -290,7 +291,9 @@ def render(self, request: Request):
return super().render(request)

def _processWebsocket(self, request: Request):
channel = WebSocketChannel(request)
channel = WebSocketChannel(request, self.original._reactor)
# lets the channel know when the transport's send buffer is full
request.registerProducer(channel, True)
if isinstance(request.channel.transport, ProtocolWrapper):
request.transport.wrappedProtocol = channel
else:
Expand All @@ -304,6 +307,7 @@ def _processWebsocket(self, request: Request):
self.original._threadpool.callInThread(self.websocketListener, environment)


@implementer(IPushProducer)
class WebSocketChannel(Protocol):
"""
Websocket protocol implementation over twisted. Note this is a ``twisted.internet.Protocol``, not a
Expand All @@ -312,11 +316,25 @@ class WebSocketChannel(Protocol):

eventQueue: Queue[events.Event]

def __init__(self, request: Request):
closeTimeout: float = 5
"""Seconds to wait for the client's close frame after the server sent its own, before terminating the TCP
connection anyway. The timeout starts once the transport's send buffer is no longer full, so a slow client
has the time to read the messages sent before the close frame."""

closeAbortTimeout: float = 30
"""Seconds after the server sent its close frame, after which the TCP connection is aborted if it is still
open, discarding any data still buffered. This bounds the close of a client that stopped reading."""

def __init__(self, request: Request, reactor=reactor):
self.request = request
self.reactor = reactor
self.wsproto = WSConnection(ConnectionType.SERVER)
self.eventQueue = Queue()
self.upgraded = False
self._closeTimeoutCall = None
self._closeAbortCall = None
self._transportPaused = False
self._closeTimeoutPending = False

@property
def closed(self):
Expand All @@ -332,6 +350,8 @@ def initiateUpgrade(self):
self.close()

def connectionLost(self, reason):
if self._closeAbortCall and self._closeAbortCall.active():
self._closeAbortCall.cancel()
self.close()

def dataReceived(self, data: bytes) -> None:
Expand All @@ -340,6 +360,9 @@ def dataReceived(self, data: bytes) -> None:
if isinstance(event, events.Ping):
self.wsSend(events.Pong(event.payload))
continue
if self.wsproto.state == ConnectionState.LOCAL_CLOSING:
# the server closed the websocket already, the listener doesn't consume any more frames
continue
# TODO: filter other event types that are not expected by WebSocketAdapter
# queue the event before ``close()`` queues its poison pill, so the consumer sees the
# client's close code and reason
Expand Down Expand Up @@ -387,14 +410,56 @@ def wsReject(
self.close()

def wsClose(self, code: int = 1000, reason: t.Optional[str] = None):
if self.request.finished or self.wsproto.state == ConnectionState.LOCAL_CLOSING:
return
if self.wsproto.state != ConnectionState.OPEN:
self.close()
return
try:
self.wsSend(events.CloseConnection(code, reason))
finally:
except BaseException:
self.close()
raise
# the server terminates the TCP connection once it has both sent and received a close frame (RFC 6455
# section 5.5.1), so ``dataReceived`` terminates it when the client echoes the close frame. terminating it
# right away would stop reading the client's frames, and closing a socket with unread data makes the kernel
# reset the connection, which discards the frames the client has not read yet (RFC 6455 section 1.4).
# a client that doesn't echo the close frame gets the TCP connection terminated after ``closeTimeout``.
if self._transportPaused:
self._closeTimeoutPending = True
else:
self._startCloseTimeout()
self._closeAbortCall = self.reactor.callLater(self.closeAbortTimeout, self._abort)
# special internal poison pill, the websocket is closed for the listener already
self.eventQueue.put_nowait(events.CloseConnection(None))

def _abort(self):
self.close()
self.request.transport.abortConnection()

def _startCloseTimeout(self):
self._closeTimeoutPending = False
self._closeTimeoutCall = self.reactor.callLater(self.closeTimeout, self.close)

def pauseProducing(self):
self._transportPaused = True

def resumeProducing(self):
self._transportPaused = False
if self._closeTimeoutPending:
self._startCloseTimeout()

def stopProducing(self):
pass

def close(self):
self._closeTimeoutPending = False
if self._closeTimeoutCall and self._closeTimeoutCall.active():
self._closeTimeoutCall.cancel()
if self.request.finished:
return
if getattr(self.request, "producer", None) is self:
self.request.unregisterProducer()
if self.upgraded:
# 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
Expand Down
127 changes: 125 additions & 2 deletions tests/websocket/test_websockets.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,20 @@
import json
import socket
import struct
import threading
import time
from queue import Queue

import pytest
import websocket
from _pytest.fixtures import SubRequest
from twisted.internet.threads import blockingCallFromThread
from twisted.python import threadable
from werkzeug.datastructures import Headers

from rolo import Response, Router
from rolo.serving.twisted import WebSocketChannel
from rolo.testing.pytest import poll_condition
from rolo.websocket.request import (
WebSocketDisconnectedError,
WebSocketProtocolError,
Expand Down Expand Up @@ -201,8 +205,8 @@ def app(request: WebSocketRequest):


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."""
"""When the server closes the websocket, the client has to receive a proper close frame. Once the client
echoed it, the server has to terminate the TCP connection."""

@WebSocketRequest.listener
def app(request: WebSocketRequest):
Expand All @@ -215,6 +219,32 @@ def app(request: WebSocketRequest):
client.connect(server.url.replace("http://", "ws://"))
assert client.recv() == "hello"

frame = client.recv_frame()
assert frame.opcode == websocket.ABNF.OPCODE_CLOSE
client.send_close(status=struct.unpack("!H", frame.data[:2])[0])

client.sock.settimeout(5)
assert client.sock.recv(1) == b"", "expected the server to terminate the TCP connection"


def test_close_handshake_server_initiated_client_does_not_echo(
serve_twisted_websocket_listener, monkeypatch
):
"""When the client never echoes the server's close frame, the server has to terminate the TCP connection
after a timeout. Only tested with twisted, since the timeout is specific to its channel."""
monkeypatch.setattr(WebSocketChannel, "closeTimeout", 0.5)

@WebSocketRequest.listener
def app(request: WebSocketRequest):
with request.accept() as ws:
ws.send("hello")

server = serve_twisted_websocket_listener(app)

client = websocket.WebSocket()
client.connect(server.url.replace("http://", "ws://"))
assert client.recv() == "hello"

frame = client.recv_frame()
assert frame.opcode == websocket.ABNF.OPCODE_CLOSE

Expand Down Expand Up @@ -355,3 +385,96 @@ def app(request: WebSocketRequest):
assert client.recv() == f"{i:08d}{payload}"
client.send("done")
client.close()


def test_server_close_while_client_is_sending(serve_twisted_websocket_listener):
"""When the server closes the websocket while the client is still sending, the server has to wait for the
client's close frame before terminating the TCP connection (RFC 6455 section 5.5.1). Terminating it right
away leaves the client's frames unread on the server, which makes the kernel reset the connection, and the
client loses the messages it has not read yet (RFC 6455 section 1.4). Only tested with twisted, since
hypercorn terminates the TCP connection right away as well."""
messages = 50_000

@WebSocketRequest.listener
def app(request: WebSocketRequest):
with request.accept() as ws:
assert ws.receive() == "start"
for _ in range(messages):
ws.send("x" * 32)

server = serve_twisted_websocket_listener(app)

for _ in range(3):
client = websocket.WebSocket()
client.connect(server.url.replace("http://", "ws://"), timeout=5)
stop = threading.Event()

def _ping(_client: websocket.WebSocket, _stop: threading.Event):
while not _stop.is_set():
try:
_client.ping(b"ping")
except websocket.WebSocketException:
return
time.sleep(0.001)

pinger = threading.Thread(target=_ping, args=(client, stop), daemon=True)
pinger.start()
client.send("start")
try:
for _ in range(messages):
assert client.recv_data(control_frame=False)[1] == b"x" * 32
finally:
stop.set()
pinger.join()

# the server may still answer pings it received before it sent its close frame
opcode, _ = client.recv_data(control_frame=True)
while opcode == websocket.ABNF.OPCODE_PONG:
opcode, _ = client.recv_data(control_frame=True)
assert opcode == websocket.ABNF.OPCODE_CLOSE
client.close()


def test_server_close_client_stops_reading(
twisted_reactor, serve_twisted_websocket_listener, monkeypatch
):
"""When the client stops reading after the server closed the websocket, the send buffer never drains, so
neither the client's close frame nor the start of the close timeout ever comes. The server has to abort the
TCP connection after ``closeAbortTimeout``. Only tested with twisted, since the timeout is specific to its
channel."""
monkeypatch.setattr(WebSocketChannel, "closeTimeout", 0.2)
monkeypatch.setattr(WebSocketChannel, "closeAbortTimeout", 0.5)
channels = Queue()

@WebSocketRequest.listener
def app(request: WebSocketRequest):
with request.accept() as ws:
channel = ws.socket.channel
channels.put(channel)
# small socket buffers, so the send buffer stays full while the client doesn't read
blockingCallFromThread(
twisted_reactor,
channel.request.transport.socket.setsockopt,
socket.SOL_SOCKET,
socket.SO_SNDBUF,
4096,
)
ws.send(b"x" * (8 * 1024 * 1024))

server = serve_twisted_websocket_listener(app)

client = websocket.WebSocket()
client.connect(
server.url.replace("http://", "ws://"),
timeout=5,
sockopt=[(socket.SOL_SOCKET, socket.SO_RCVBUF, 4096)],
)
try:
channel = channels.get(timeout=5)
transport = channel.request.transport
assert poll_condition(
lambda: blockingCallFromThread(twisted_reactor, lambda: transport.disconnected),
timeout=5,
), "expected the server to abort the TCP connection"
finally:
client.shutdown()
Loading