diff --git a/rolo/serving/twisted.py b/rolo/serving/twisted.py index 821afc8..4de0c6a 100644 --- a/rolo/serving/twisted.py +++ b/rolo/serving/twisted.py @@ -335,6 +335,7 @@ def __init__(self, request: Request, reactor=reactor): self._closeAbortCall = None self._transportPaused = False self._closeTimeoutPending = False + self._messageParts: list[str | bytes] = [] @property def closed(self): @@ -363,6 +364,15 @@ def dataReceived(self, data: bytes) -> None: if self.wsproto.state == ConnectionState.LOCAL_CLOSING: # the server closed the websocket already, the listener doesn't consume any more frames continue + if isinstance(event, events.Message): + # wsproto emits the data of a message as it arrives, frame by frame and in chunks of a + # frame, while the listener receives complete messages + self._messageParts.append(event.data) + if not event.message_finished: + continue + data = event.data[:0].join(self._messageParts) + self._messageParts = [] + event = type(event)(data=data) # 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 diff --git a/tests/websocket/test_websockets.py b/tests/websocket/test_websockets.py index f0527c9..b0fd491 100644 --- a/tests/websocket/test_websockets.py +++ b/tests/websocket/test_websockets.py @@ -320,6 +320,55 @@ def echo_headers(request: WebSocketRequest): assert received.get(timeout=5) == b"bar" +def test_receive_large_message(serve_websocket_listener): + """A frame bigger than a single socket read must still be received as one message.""" + received = Queue() + + @WebSocketRequest.listener + def app(request: WebSocketRequest): + with request.accept() as ws: + received.put(ws.receive()) + received.put(ws.receive()) + + server = serve_websocket_listener(app) + + client = websocket.WebSocket() + client.connect(server.url.replace("http://", "ws://")) + client.send("x" * 1024 * 1024) + client.send_binary(b"y" * 1024 * 1024) + + assert received.get(timeout=5) == "x" * 1024 * 1024 + assert received.get(timeout=5) == b"y" * 1024 * 1024 + client.close() + + +def test_receive_fragmented_message(serve_websocket_listener): + """A message sent as several frames (RFC 6455 section 5.4) must be received as one message.""" + received = Queue() + + @WebSocketRequest.listener + def app(request: WebSocketRequest): + with request.accept() as ws: + received.put(ws.receive()) + received.put(ws.receive()) + + server = serve_websocket_listener(app) + + client = websocket.WebSocket() + client.connect(server.url.replace("http://", "ws://")) + client.send_frame(websocket.ABNF.create_frame("foo", websocket.ABNF.OPCODE_TEXT, fin=0)) + client.send_frame(websocket.ABNF.create_frame("bar", websocket.ABNF.OPCODE_CONT, fin=0)) + # control frames can be sent in the middle of a fragmented message + client.ping("ping") + client.send_frame(websocket.ABNF.create_frame("baz", websocket.ABNF.OPCODE_CONT, fin=1)) + client.send_frame(websocket.ABNF.create_frame(b"foo", websocket.ABNF.OPCODE_BINARY, fin=0)) + client.send_frame(websocket.ABNF.create_frame(b"bar", websocket.ABNF.OPCODE_CONT, fin=1)) + + assert received.get(timeout=5) == "foobarbaz" + assert received.get(timeout=5) == b"foobar" + client.close() + + def test_send_non_confirming_data(serve_websocket_listener): match = Queue()