moving logic to WebsocketServerInputTransport
This commit is contained in:
@@ -47,6 +47,7 @@ class WebsocketServerParams(TransportParams):
|
|||||||
class WebsocketServerCallbacks(BaseModel):
|
class WebsocketServerCallbacks(BaseModel):
|
||||||
on_client_connected: Callable[[websockets.WebSocketServerProtocol], Awaitable[None]]
|
on_client_connected: Callable[[websockets.WebSocketServerProtocol], Awaitable[None]]
|
||||||
on_client_disconnected: Callable[[websockets.WebSocketServerProtocol], Awaitable[None]]
|
on_client_disconnected: Callable[[websockets.WebSocketServerProtocol], Awaitable[None]]
|
||||||
|
on_session_timeout: Callable[[websockets.WebSocketServerProtocol], Awaitable[None]]
|
||||||
|
|
||||||
|
|
||||||
class WebsocketServerInputTransport(BaseInputTransport):
|
class WebsocketServerInputTransport(BaseInputTransport):
|
||||||
@@ -99,6 +100,10 @@ class WebsocketServerInputTransport(BaseInputTransport):
|
|||||||
# Notify
|
# Notify
|
||||||
await self._callbacks.on_client_connected(websocket)
|
await self._callbacks.on_client_connected(websocket)
|
||||||
|
|
||||||
|
# Create a task to monitor the websocket connection
|
||||||
|
if self._params.session_timeout:
|
||||||
|
self.get_event_loop().create_task(self._monitor_websocket(websocket))
|
||||||
|
|
||||||
# Handle incoming messages
|
# Handle incoming messages
|
||||||
async for message in websocket:
|
async for message in websocket:
|
||||||
frame = self._params.serializer.deserialize(message)
|
frame = self._params.serializer.deserialize(message)
|
||||||
@@ -125,6 +130,17 @@ class WebsocketServerInputTransport(BaseInputTransport):
|
|||||||
|
|
||||||
logger.info(f"Client {websocket.remote_address} disconnected")
|
logger.info(f"Client {websocket.remote_address} disconnected")
|
||||||
|
|
||||||
|
async def _monitor_websocket(self, websocket: websockets.WebSocketServerProtocol):
|
||||||
|
"""
|
||||||
|
Wait for self._params.session_timeout seconds, if the websocket is still open, trigger timeout event.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
await asyncio.sleep(self._params.session_timeout)
|
||||||
|
if not websocket.closed:
|
||||||
|
await self._callbacks.on_session_timeout(websocket)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
logger.info(f"Monitoring task cancelled for: {websocket.remote_address}")
|
||||||
|
|
||||||
|
|
||||||
class WebsocketServerOutputTransport(BaseOutputTransport):
|
class WebsocketServerOutputTransport(BaseOutputTransport):
|
||||||
def __init__(self, params: WebsocketServerParams, **kwargs):
|
def __init__(self, params: WebsocketServerParams, **kwargs):
|
||||||
@@ -209,6 +225,7 @@ class WebsocketServerTransport(BaseTransport):
|
|||||||
self._callbacks = WebsocketServerCallbacks(
|
self._callbacks = WebsocketServerCallbacks(
|
||||||
on_client_connected=self._on_client_connected,
|
on_client_connected=self._on_client_connected,
|
||||||
on_client_disconnected=self._on_client_disconnected,
|
on_client_disconnected=self._on_client_disconnected,
|
||||||
|
on_session_timeout=self._on_session_timeout
|
||||||
)
|
)
|
||||||
self._input: WebsocketServerInputTransport | None = None
|
self._input: WebsocketServerInputTransport | None = None
|
||||||
self._output: WebsocketServerOutputTransport | None = None
|
self._output: WebsocketServerOutputTransport | None = None
|
||||||
@@ -218,9 +235,7 @@ class WebsocketServerTransport(BaseTransport):
|
|||||||
# these handlers.
|
# these handlers.
|
||||||
self._register_event_handler("on_client_connected")
|
self._register_event_handler("on_client_connected")
|
||||||
self._register_event_handler("on_client_disconnected")
|
self._register_event_handler("on_client_disconnected")
|
||||||
|
self._register_event_handler("on_session_timeout")
|
||||||
if self._params.session_timeout:
|
|
||||||
self._register_event_handler("on_session_timeout")
|
|
||||||
|
|
||||||
def input(self) -> WebsocketServerInputTransport:
|
def input(self) -> WebsocketServerInputTransport:
|
||||||
if not self._input:
|
if not self._input:
|
||||||
@@ -234,22 +249,9 @@ class WebsocketServerTransport(BaseTransport):
|
|||||||
self._output = WebsocketServerOutputTransport(self._params, name=self._output_name)
|
self._output = WebsocketServerOutputTransport(self._params, name=self._output_name)
|
||||||
return self._output
|
return self._output
|
||||||
|
|
||||||
async def _monitor_websocket(self, websocket: websockets.WebSocketServerProtocol):
|
|
||||||
"""
|
|
||||||
Wait for self._params.session_timeout seconds, if the websocket is still open, trigger timeout event.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
await asyncio.sleep(self._params.session_timeout)
|
|
||||||
if not websocket.closed:
|
|
||||||
await self._call_event_handler("on_session_timeout", websocket)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
logger.info(f"Monitoring task cancelled for: {websocket.remote_address}")
|
|
||||||
|
|
||||||
async def _on_client_connected(self, websocket):
|
async def _on_client_connected(self, websocket):
|
||||||
if self._output:
|
if self._output:
|
||||||
await self._output.set_client_connection(websocket)
|
await self._output.set_client_connection(websocket)
|
||||||
if self._params.session_timeout:
|
|
||||||
self._loop.create_task(self._monitor_websocket(websocket))
|
|
||||||
await self._call_event_handler("on_client_connected", websocket)
|
await self._call_event_handler("on_client_connected", websocket)
|
||||||
else:
|
else:
|
||||||
logger.error("A WebsocketServerTransport output is missing in the pipeline")
|
logger.error("A WebsocketServerTransport output is missing in the pipeline")
|
||||||
@@ -261,3 +263,6 @@ class WebsocketServerTransport(BaseTransport):
|
|||||||
else:
|
else:
|
||||||
logger.error("A WebsocketServerTransport output is missing in the pipeline")
|
logger.error("A WebsocketServerTransport output is missing in the pipeline")
|
||||||
|
|
||||||
|
async def _on_session_timeout(self, websocket):
|
||||||
|
await self._call_event_handler("on_session_timeout", websocket)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user