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
4 changes: 3 additions & 1 deletion rolo/serving/twisted.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
from twisted.web.server import Request as TwistedRequest
from twisted.web.wsgi import WSGIResource, _WSGIResponse, _wsgiString
from werkzeug.datastructures import Headers
from wsproto import ConnectionType, WSConnection, events
from wsproto import ConnectionState, ConnectionType, WSConnection, events
from zope.interface import implementer

from rolo.gateway import Gateway
Expand Down Expand Up @@ -352,6 +352,8 @@ def wsSend(self, event: events.Event):
request = self.request
if request.finished:
return
if self.wsproto.state in (ConnectionState.LOCAL_CLOSING, ConnectionState.CLOSED):
return
data = self.wsproto.send(event)
if isinstance(event, events.AcceptConnection):
self.upgraded = True
Expand Down
89 changes: 89 additions & 0 deletions tests/serving/test_twisted.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,19 @@
import io
import json

import pytest
import requests
from twisted.web.http_headers import Headers as TwistedHeaders
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.websocket.adapter import CreateConnection, TextMessage
from rolo.websocket.request import WebSocketDisconnectedError


def test_large_file_upload(serve_twisted_gateway):
Expand Down Expand Up @@ -52,3 +59,85 @@ def hello(request: Request):
response_data = json.loads(response.read())
assert response_data["path"] == "/hello"
assert response_data["raw_uri"].startswith("http")


class _FakeTransport:
def __init__(self):
self.written = []

def write(self, data: bytes):
self.written.append(data)

def loseConnection(self):
pass


class _UnfinishedRequest:
"""A twisted request stand-in whose ``finish()`` does not mark it finished, which freezes the channel in
the window where the closing handshake is complete, but the reactor has not finished the request yet.
"""

def __init__(self, headers: list[tuple[bytes, bytes]]):
self.requestHeaders = TwistedHeaders()
for k, v in headers:
self.requestHeaders.addRawHeader(k, v)
self.path = b"/"
self.transport = _FakeTransport()
self.finished = False
self.startedWriting = 0

def finish(self):
pass


def _accepted_websocket() -> tuple[TwistedWebSocketAdapter, WSConnection, _FakeTransport]:
client = WSConnection(ConnectionType.CLIENT)
request_head = client.send(events.Request(host="localhost", target="/")).split(b"\r\n\r\n")[0]
headers = [
tuple(part.strip() for part in line.split(b":", 1))
for line in request_head.split(b"\r\n")[1:]
]
request = _UnfinishedRequest(headers)
channel = WebSocketChannel(request)
channel.initiateUpgrade()
adapter = TwistedWebSocketAdapter(channel)
assert isinstance(adapter.receive(timeout=1), CreateConnection)
adapter.accept()

client.receive_data(b"".join(request.transport.written))
assert isinstance(next(client.events()), events.AcceptConnection)
request.transport.written.clear()
return adapter, client, request.transport


def test_websocket_close_after_client_close_before_request_finished():
adapter, client, transport = _accepted_websocket()

adapter.channel.dataReceived(client.send(events.CloseConnection(4001, "bye")))
assert adapter.channel.wsproto.state == ConnectionState.CLOSED
assert not adapter.channel.closed
echoed = b"".join(transport.written)
client.receive_data(echoed)
assert next(client.events()) == events.CloseConnection(4001, "bye")

with pytest.raises(WebSocketDisconnectedError) as e:
adapter.receive(timeout=1)
assert (e.value.code, e.value.reason) == (4001, "bye")

adapter.close(1000)
adapter.close(1000)

assert b"".join(transport.written) == echoed


def test_websocket_close_twice_before_request_finished():
adapter, client, transport = _accepted_websocket()

adapter.close(1000)
assert adapter.channel.wsproto.state == ConnectionState.LOCAL_CLOSING
sent = b"".join(transport.written)

adapter.close(1000)
adapter.send(TextMessage("too late"))

assert b"".join(transport.written) == sent
42 changes: 42 additions & 0 deletions tests/websocket/test_websockets.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,11 @@
import pytest
import websocket
from _pytest.fixtures import SubRequest
from twisted.python import threadable
from werkzeug.datastructures import Headers

from rolo import Response, Router
from rolo.serving.twisted import WebSocketChannel
from rolo.websocket.request import (
WebSocketDisconnectedError,
WebSocketProtocolError,
Expand Down Expand Up @@ -158,6 +160,46 @@ def app(request: WebSocketRequest):
assert closes.get(timeout=3) == (4001, "test reason")


def test_server_close_after_client_close(serve_twisted_websocket_listener, monkeypatch):
"""The server application closing the websocket after the client closed it must not fail, even when
it runs before the reactor has finished the request of the completed closing handshake. The reactor is
held in that window until the application has closed. Only tested with twisted, since the window is
specific to its channel."""
app_closed = threading.Event()
results = Queue()

original_close = WebSocketChannel.close

def close(self):
if threadable.isInIOThread():
app_closed.wait(timeout=3)
original_close(self)

monkeypatch.setattr(WebSocketChannel, "close", close)

@WebSocketRequest.listener
def app(request: WebSocketRequest):
try:
with request.accept() as ws:
with pytest.raises(WebSocketDisconnectedError):
ws.receive()
ws.close()
except Exception as e:
results.put(e)
else:
results.put("ok")
finally:
app_closed.set()

server = serve_twisted_websocket_listener(app)

client = websocket.WebSocket()
client.connect(server.url.replace("http://", "ws://"))
client.close(status=4001, reason=b"test reason")

assert results.get(timeout=5) == "ok"


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."""
Expand Down
Loading