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
8 changes: 8 additions & 0 deletions rolo/asgi.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import io
import logging
import math
import time
import typing as t
from asyncio import AbstractEventLoop
from concurrent.futures import Executor
Expand Down Expand Up @@ -355,6 +356,7 @@ def __init__(
self._receive = receive
self._send = send
self._loop = loop
self._last_received_at = time.monotonic()

async def asgi_send_async(self, event: "_WebsocketResponse"):
await self._send(event)
Expand Down Expand Up @@ -401,6 +403,7 @@ def receive(self, timeout: float = None) -> rolows.CreateConnection | rolows.Mes

if event["type"] == "websocket.receive":
event: "WebsocketReceiveEvent"
self._last_received_at = time.monotonic()
text = event.get("text")
if text is not None:
return rolows.TextMessage(text)
Expand All @@ -418,6 +421,11 @@ def receive(self, timeout: float = None) -> rolows.CreateConnection | rolows.Mes
event: "WebsocketDisconnectEvent"
raise WebSocketDisconnectedError(event["code"], event.get("reason"))

@property
def last_received_at(self) -> float:
# ASGI doesn't pass control frames on, and messages only count once the listener consumes them
return self._last_received_at

def send(self, event: rolows.Message, timeout: float = None):
if isinstance(event, rolows.TextMessage):
asgi_event = {
Expand Down
7 changes: 7 additions & 0 deletions rolo/serving/twisted.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
Bindings to serve rolo through Twisted.
"""
import logging
import time
import typing as t
from io import BytesIO
from queue import Empty, Queue
Expand Down Expand Up @@ -336,6 +337,7 @@ def __init__(self, request: Request, reactor=reactor):
self._transportPaused = False
self._closeTimeoutPending = False
self._messageParts: list[str | bytes] = []
self.lastReceivedAt = time.monotonic()

@property
def closed(self):
Expand All @@ -356,6 +358,7 @@ def connectionLost(self, reason):
self.close()

def dataReceived(self, data: bytes) -> None:
self.lastReceivedAt = time.monotonic()
self.wsproto.receive_data(data)
for event in self.wsproto.events():
if isinstance(event, events.Ping):
Expand Down Expand Up @@ -539,6 +542,10 @@ def send(self, event: rolows.Message, timeout: float = None):
else:
raise TypeError(f"Unexpected event type {event.__class__.__name__}")

@property
def last_received_at(self) -> float:
return self.channel.lastReceivedAt

def reject(
self,
status_code: int,
Expand Down
9 changes: 9 additions & 0 deletions rolo/websocket/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,15 @@ def close(self, code: int = 1001, reason: str = None, timeout: float = None):
"""
raise NotImplementedError

@property
def last_received_at(self) -> float:
"""
The ``time.monotonic()`` time at which data was last received from the client, or at which the
connection was created if nothing was received yet. Servers that handle the control frames
themselves count them as well when they can see them, so a client that only pings is active.
"""
raise NotImplementedError


class WebSocketListener(t.Protocol):
"""
Expand Down
8 changes: 8 additions & 0 deletions rolo/websocket/request.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,14 @@ def close(self, code: int = 1000, reason: t.Optional[str] = None, timeout: float
"""
self.socket.close(code, reason, timeout)

@property
def last_received_at(self) -> float:
"""
The ``time.monotonic()`` time at which data was last received from the client, or at which the
connection was created if nothing was received yet. See ``WebSocketAdapter.last_received_at``.
"""
return self.socket.last_received_at


class WebSocketRequest(_SansIORequest):
"""
Expand Down
50 changes: 50 additions & 0 deletions tests/websocket/test_websockets.py
Original file line number Diff line number Diff line change
Expand Up @@ -320,6 +320,56 @@ def echo_headers(request: WebSocketRequest):
assert received.get(timeout=5) == b"bar"


def test_last_received_at_message(serve_websocket_listener):
received_at = Queue()

@WebSocketRequest.listener
def app(request: WebSocketRequest):
with request.accept() as ws:
received_at.put(ws.last_received_at)
ws.receive()
received_at.put(ws.last_received_at)
ws.send("done")

server = serve_websocket_listener(app)

client = websocket.WebSocket()
client.connect(server.url.replace("http://", "ws://"))
connected_at = received_at.get(timeout=5)
sent_at = time.monotonic()
client.send("foobar")
assert client.recv() == "done"

assert connected_at < sent_at <= received_at.get(timeout=5)
client.close()


def test_last_received_at_ping(serve_twisted_websocket_listener):
# ASGI servers answer pings without passing them on to the application
activity = Queue()

@WebSocketRequest.listener
def app(request: WebSocketRequest):
with request.accept() as ws:
connected_at = ws.last_received_at
activity.put(connected_at)
poll_condition(lambda: ws.last_received_at > connected_at, timeout=5, interval=0.01)
activity.put(ws.last_received_at)
ws.receive()

server = serve_twisted_websocket_listener(app)

client = websocket.WebSocket()
client.connect(server.url.replace("http://", "ws://"))
connected_at = activity.get(timeout=5)
pinged_at = time.monotonic()
client.ping("ping")

assert connected_at < pinged_at <= activity.get(timeout=5)
client.send("done")
client.close()


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()
Expand Down
Loading