diff --git a/rolo/serving/twisted.py b/rolo/serving/twisted.py index 4864d07..821afc8 100644 --- a/rolo/serving/twisted.py +++ b/rolo/serving/twisted.py @@ -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 @@ -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: @@ -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 @@ -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): @@ -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: @@ -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 @@ -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 diff --git a/tests/websocket/test_websockets.py b/tests/websocket/test_websockets.py index 14f0962..f0527c9 100644 --- a/tests/websocket/test_websockets.py +++ b/tests/websocket/test_websockets.py @@ -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, @@ -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): @@ -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 @@ -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()