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
18 changes: 16 additions & 2 deletions rolo/websocket/request.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,10 +27,16 @@ class WebSocket:

request: "WebSocketRequest"
socket: WebSocketAdapter
close_code: t.Optional[int]
"""The close code of the first disconnect seen by ``receive``, or ``None`` if not disconnected yet."""
close_reason: t.Optional[str]
"""The close reason of the first disconnect seen by ``receive``, if the client sent one."""

def __init__(self, request: "WebSocketRequest", socket: WebSocketAdapter):
self.request = request
self.socket = socket
self.close_code = None
self.close_reason = None

def __enter__(self):
return self
Expand Down Expand Up @@ -70,11 +76,19 @@ def receive(self) -> str | bytes:
Receive the next data package from the websocket. Will be string or byte data and set the
underlying binary for the frame automatically.

:raise WebSocketDisconnectedError: if the websocket was closed in the meantime
:raise WebSocketDisconnectedError: if the websocket was closed in the meantime. The close code
and reason are also stored in ``close_code`` and ``close_reason``.
:raise WebSocketProtocolError: error in the interaction between the app and the webserver
:return: the next data package from the websocket
"""
event = self.socket.receive()
try:
event = self.socket.receive()
except WebSocketDisconnectedError as e:
# keep the first disconnect, later calls may only see the server's internal close event
if self.close_code is None:
self.close_code = e.code
self.close_reason = e.reason
raise
if isinstance(event, Message):
data = event.data
if data is None:
Expand Down
9 changes: 0 additions & 9 deletions tests/serving/test_twisted.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,12 +52,3 @@ def hello(request: Request):
response_data = json.loads(response.read())
assert response_data["path"] == "/hello"
assert response_data["raw_uri"].startswith("http")


def test_twisted_gateway_import_and_wsproto_dependency():
import wsproto

from rolo.serving.twisted import TwistedGateway

assert hasattr(wsproto, "WSConnection")
assert TwistedGateway is not None
24 changes: 24 additions & 0 deletions tests/websocket/test_websockets.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,30 @@ def app(request: WebSocketRequest):
assert error.reason == "test reason"


def test_close_code_and_reason_after_iter(serve_twisted_websocket_listener):
"""The iterator ends silently on disconnect, but the client's close code and reason are kept on the
``WebSocket``. Only tested with twisted, see ``test_close_code_and_reason_client_initiated``."""
closes = Queue()

@WebSocketRequest.listener
def app(request: WebSocketRequest):
with request.accept() as ws:
assert ws.close_code is None
assert ws.close_reason is None
for _ in iter(ws):
pass
closes.put((ws.close_code, ws.close_reason))

server = serve_twisted_websocket_listener(app)

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

assert closes.get(timeout=3) == (4001, "test reason")


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