82 lines
2.6 KiB
Python
82 lines
2.6 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from collections.abc import Awaitable, Callable
|
|
|
|
from yuxi.channel.transport.protocol import MessageHandler, TransportState
|
|
from yuxi.utils.logging_config import logger
|
|
|
|
PollFunc = Callable[[], Awaitable[list[bytes]]]
|
|
|
|
|
|
class PollingTransport:
|
|
def __init__(self, poll_func: PollFunc, interval: float = 5.0):
|
|
self.poll_func = poll_func
|
|
self.interval = interval
|
|
self._state = TransportState.DISCONNECTED
|
|
self._message_handler: MessageHandler | None = None
|
|
self._task: asyncio.Task | None = None
|
|
self._closed = True
|
|
self._lock = asyncio.Lock()
|
|
|
|
@property
|
|
def state(self) -> TransportState:
|
|
return self._state
|
|
|
|
async def start(self) -> None:
|
|
async with self._lock:
|
|
if self._state == TransportState.CONNECTED:
|
|
return
|
|
self._closed = False
|
|
self._state = TransportState.CONNECTED
|
|
self._task = asyncio.create_task(self._poll_loop())
|
|
|
|
async def stop(self) -> None:
|
|
async with self._lock:
|
|
self._closed = True
|
|
self._state = TransportState.STOPPED
|
|
task = self._task
|
|
self._task = None
|
|
if task:
|
|
task.cancel()
|
|
try:
|
|
await task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
|
|
def on_message(self, handler: MessageHandler) -> None:
|
|
self._message_handler = handler
|
|
|
|
async def send(self, data: bytes | str) -> None:
|
|
raise NotImplementedError("Polling transport does not support direct send")
|
|
|
|
async def _poll_loop(self) -> None:
|
|
consecutive_failures = 0
|
|
while not self._closed:
|
|
try:
|
|
messages = await self.poll_func()
|
|
consecutive_failures = 0
|
|
except asyncio.CancelledError:
|
|
break
|
|
except Exception:
|
|
logger.exception("Polling error")
|
|
consecutive_failures += 1
|
|
messages = []
|
|
if consecutive_failures >= 3:
|
|
self._state = TransportState.DISCONNECTED
|
|
break
|
|
|
|
if self._message_handler is not None:
|
|
for message in messages:
|
|
try:
|
|
await self._message_handler(message)
|
|
except Exception:
|
|
logger.exception("Polling message handler error")
|
|
|
|
try:
|
|
await asyncio.sleep(self.interval)
|
|
except asyncio.CancelledError:
|
|
break
|
|
if self._closed:
|
|
break
|