diff --git a/rolo/websocket/request.py b/rolo/websocket/request.py index d46a216..bf1e0ab 100644 --- a/rolo/websocket/request.py +++ b/rolo/websocket/request.py @@ -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 @@ -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: diff --git a/tests/serving/test_twisted.py b/tests/serving/test_twisted.py index 554e49b..3740977 100644 --- a/tests/serving/test_twisted.py +++ b/tests/serving/test_twisted.py @@ -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 diff --git a/tests/websocket/test_websockets.py b/tests/websocket/test_websockets.py index 8f7eab6..b830e94 100644 --- a/tests/websocket/test_websockets.py +++ b/tests/websocket/test_websockets.py @@ -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."""