124 lines
4.2 KiB
Python
124 lines
4.2 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from typing import Any
|
|
|
|
import websockets
|
|
from websockets.exceptions import ConnectionClosed
|
|
|
|
from yuxi.channel.transport.protocol import MessageHandler, TransportState
|
|
from yuxi.utils.logging_config import logger
|
|
|
|
|
|
class WebSocketTransport:
|
|
def __init__(
|
|
self,
|
|
url: str,
|
|
headers: dict | None = None,
|
|
heartbeat_interval: float = 30.0,
|
|
heartbeat_message: str | bytes = b"ping",
|
|
):
|
|
self.url = url
|
|
self.headers = headers or {}
|
|
self.heartbeat_interval = heartbeat_interval
|
|
self.heartbeat_message = heartbeat_message
|
|
self._state = TransportState.DISCONNECTED
|
|
self._ws: websockets.WebSocketClientProtocol | None = None
|
|
self._message_handler: MessageHandler | None = None
|
|
self._tasks: set[asyncio.Task] = set()
|
|
self._closed = True
|
|
self._lock = asyncio.Lock()
|
|
|
|
@property
|
|
def state(self) -> TransportState:
|
|
return self._state
|
|
|
|
def _spawn_task(self, coro: Any) -> asyncio.Task[Any]:
|
|
task = asyncio.create_task(coro)
|
|
self._tasks.add(task)
|
|
task.add_done_callback(self._tasks.discard)
|
|
return task
|
|
|
|
async def start(self) -> None:
|
|
async with self._lock:
|
|
self._tasks = {t for t in self._tasks if not t.done()}
|
|
if self._state in (TransportState.CONNECTED, TransportState.CONNECTING):
|
|
return
|
|
self._closed = False
|
|
self._state = TransportState.CONNECTING
|
|
try:
|
|
self._ws = await websockets.connect(self.url, additional_headers=self.headers)
|
|
self._state = TransportState.CONNECTED
|
|
self._spawn_task(self._receive_loop())
|
|
if self.heartbeat_interval > 0:
|
|
self._spawn_task(self._heartbeat_loop())
|
|
except Exception:
|
|
self._state = TransportState.DISCONNECTED
|
|
self._closed = True
|
|
raise
|
|
|
|
async def stop(self) -> None:
|
|
async with self._lock:
|
|
self._closed = True
|
|
self._state = TransportState.STOPPED
|
|
tasks = list(self._tasks)
|
|
self._tasks.clear()
|
|
for task in tasks:
|
|
task.cancel()
|
|
if tasks:
|
|
await asyncio.gather(*tasks, return_exceptions=True)
|
|
ws = self._ws
|
|
self._ws = None
|
|
if ws:
|
|
await ws.close()
|
|
|
|
def on_message(self, handler: MessageHandler) -> None:
|
|
self._message_handler = handler
|
|
|
|
async def send(self, data: bytes | str) -> None:
|
|
if self._state != TransportState.CONNECTED or self._ws is None:
|
|
raise RuntimeError("WebSocket is not connected")
|
|
await self._ws.send(data)
|
|
|
|
async def _receive_loop(self) -> None:
|
|
while not self._closed and self._ws is not None:
|
|
try:
|
|
message = await self._ws.recv()
|
|
except ConnectionClosed:
|
|
self._state = TransportState.DISCONNECTED
|
|
break
|
|
except asyncio.CancelledError:
|
|
break
|
|
except Exception:
|
|
logger.exception("WebSocket receive error")
|
|
continue
|
|
|
|
if self._message_handler is None:
|
|
continue
|
|
|
|
payload = message.encode() if isinstance(message, str) else message
|
|
try:
|
|
await self._message_handler(payload)
|
|
except Exception:
|
|
logger.exception("WebSocket message handler error")
|
|
|
|
async def _heartbeat_loop(self) -> None:
|
|
while not self._closed and self._ws is not None:
|
|
try:
|
|
await asyncio.sleep(self.heartbeat_interval)
|
|
except asyncio.CancelledError:
|
|
break
|
|
|
|
if self._closed or self._ws is None:
|
|
break
|
|
|
|
try:
|
|
await self._ws.send(self.heartbeat_message)
|
|
except ConnectionClosed:
|
|
self._state = TransportState.DISCONNECTED
|
|
break
|
|
except Exception:
|
|
logger.warning("WebSocket heartbeat failed")
|
|
self._state = TransportState.DISCONNECTED
|
|
break
|