Initial Commit
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
import aiohttp
|
||||
|
||||
from .auth import IngressAuth
|
||||
from .client import IngressClient, T
|
||||
|
||||
|
||||
def create_client(ctx: T, server_url: str, websession: aiohttp.ClientSession) -> IngressClient[T]:
|
||||
return IngressClient(ctx, IngressAuth(server_url, websession))
|
||||
@@ -0,0 +1,22 @@
|
||||
import aiohttp
|
||||
|
||||
|
||||
class IngressAuth:
|
||||
def __init__(self, server_url: str, websession: aiohttp.ClientSession) -> None:
|
||||
self._server_url = server_url
|
||||
self._websession = websession
|
||||
|
||||
@property
|
||||
def ws_url(self):
|
||||
return f"ws{self._server_url[4:]}/ws"
|
||||
|
||||
@property
|
||||
def websession(self):
|
||||
return self._websession
|
||||
|
||||
@property
|
||||
def access_token(self):
|
||||
return ""
|
||||
|
||||
async def refresh(self):
|
||||
pass
|
||||
@@ -0,0 +1,239 @@
|
||||
import asyncio
|
||||
from typing import Any, Generic, TypeVar, TYPE_CHECKING, cast
|
||||
|
||||
from .connection import WebsocketConnection
|
||||
from .exceptions import (
|
||||
InvalidMessage,
|
||||
InvalidAuth,
|
||||
ConnectionLost,
|
||||
ClientException,
|
||||
ResultException,
|
||||
)
|
||||
from .helpers import LOGGER, json_loads, json_dumps, to_dict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable, Awaitable, AsyncGenerator
|
||||
|
||||
from .auth import IngressAuth
|
||||
|
||||
type ClientEventListener[T] = Callable[[T, "IngressClient", Any], Awaitable[None] | None]
|
||||
type BinaryHandler[T] = Callable[[T, "IngressClient", bytes], Awaitable[None] | None]
|
||||
type EventHandler[T] = Callable[[T, "IngressClient", dict[str, Any]], Awaitable[None] | None]
|
||||
type EventSubscriber[T] = tuple[EventHandler[T], tuple[str] | None, tuple[str] | None]
|
||||
|
||||
|
||||
MSG_TYPE_AUTH = "auth"
|
||||
MSG_TYPE_AUTH_OK = "auth_ok"
|
||||
MSG_TYPE_AUTH_INVALID = "auth_invalid"
|
||||
MSG_TYPE_RESULT = "result"
|
||||
MSG_TYPE_EVENT = "event"
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class IngressClient(Generic[T]):
|
||||
"""Manage an Ingress server remotely."""
|
||||
|
||||
def __init__(self, ctx: T, auth: "IngressAuth"):
|
||||
self._ctx = ctx
|
||||
self.auth = auth
|
||||
self._conn = WebsocketConnection(auth.ws_url, auth.websession)
|
||||
self._event_listeners: dict[str, list[ClientEventListener[T]]] = {}
|
||||
self._binary_handlers: dict[int, BinaryHandler[T]] = {}
|
||||
self._subscribers: list[EventSubscriber[T]] = []
|
||||
self._result_futures: dict[int | None, asyncio.Future[Any]] = {}
|
||||
self._receive_task: asyncio.Task | None = None
|
||||
|
||||
def on_client_event(
|
||||
self, event_type: str, callback: "ClientEventListener[T]"
|
||||
) -> "Callable[[], None]":
|
||||
listeners = self._event_listeners.setdefault(event_type, [])
|
||||
listeners.append(callback)
|
||||
return lambda: listeners.remove(callback)
|
||||
|
||||
def _fire_event(self, event_type: str, event_data: Any = None) -> None:
|
||||
async def fire_event():
|
||||
for callback in self._event_listeners.get(event_type, []):
|
||||
if asyncio.iscoroutinefunction(callback):
|
||||
await callback(self._ctx, self, event_data)
|
||||
else:
|
||||
callback(self._ctx, self, event_data)
|
||||
|
||||
self._loop.create_task(fire_event())
|
||||
|
||||
def register_binary(self, id: int, handler: "BinaryHandler[T]"):
|
||||
if handler:
|
||||
self._binary_handlers[id] = handler
|
||||
else:
|
||||
self._binary_handlers.pop(id, None)
|
||||
|
||||
def subscribe(
|
||||
self,
|
||||
cb_func: "EventHandler[T]",
|
||||
ev_type: str | tuple[str] | None = None,
|
||||
ev_id: str | tuple[str] | None = None,
|
||||
) -> "Callable[[], None]":
|
||||
"""Add callback to event listeners. Returns function to remove the listener."""
|
||||
subscriber: EventSubscriber[T] = (
|
||||
cb_func,
|
||||
((ev_type,) if isinstance(ev_type, str) else ev_type),
|
||||
((ev_id,) if isinstance(ev_id, str) else ev_id),
|
||||
)
|
||||
self._subscribers.append(subscriber)
|
||||
return lambda: self._subscribers.remove(subscriber)
|
||||
|
||||
def _handle_event(self, msg: dict[str, Any]) -> None:
|
||||
async def handle_event():
|
||||
ev_type, ev_id = msg.get("ev_type"), msg.get("ev_id")
|
||||
for cb_func, ev_types, ev_ids in self._subscribers:
|
||||
if ev_types is not None and ev_type not in ev_types:
|
||||
continue
|
||||
if ev_ids is not None and ev_id not in ev_ids:
|
||||
continue
|
||||
if asyncio.iscoroutinefunction(cb_func):
|
||||
await cb_func(self._ctx, self, msg)
|
||||
else:
|
||||
cb_func(self._ctx, self, msg)
|
||||
|
||||
self._loop.create_task(handle_event())
|
||||
|
||||
async def receive_message(self) -> "AsyncGenerator[dict[str, Any]]":
|
||||
"""Receive the next message from the server."""
|
||||
try:
|
||||
while True:
|
||||
if not (data := await self._conn.receive_message()):
|
||||
continue
|
||||
if not isinstance(data, bytes):
|
||||
break
|
||||
if not (handler := self._binary_handlers.get(data[0])):
|
||||
if data[0] in b"[{":
|
||||
break
|
||||
raise InvalidMessage(f"Received invalid binary: {data[0]}")
|
||||
handler(self._ctx, self, data[1:])
|
||||
msg = json_loads(data)
|
||||
except (TypeError, ValueError) as err:
|
||||
raise InvalidMessage(f"Received invalid json: {err}") from err
|
||||
|
||||
for msg in msg if isinstance(msg, list) else [msg]:
|
||||
if not isinstance(msg, dict) or not msg.get("type"):
|
||||
LOGGER.warning("Received invalid msg: %s", msg)
|
||||
continue
|
||||
LOGGER.debug("Received message: %s", msg)
|
||||
yield cast(dict[str, Any], msg)
|
||||
|
||||
async def send_message(self, obj: Any) -> None:
|
||||
"""Send a message to the server."""
|
||||
msg = json_dumps(obj)
|
||||
LOGGER.debug("Publishing message: %s", msg)
|
||||
await self._conn.send_message(msg)
|
||||
|
||||
async def _connect(self) -> bool:
|
||||
"""Connect to the server."""
|
||||
if not (await self._conn.connect()):
|
||||
return False
|
||||
server_info = None
|
||||
|
||||
try:
|
||||
await self.send_message({"type": MSG_TYPE_AUTH, "token": self.auth.access_token})
|
||||
while not self._stop_called and server_info is None:
|
||||
async for msg in self.receive_message():
|
||||
msg_type = msg["type"]
|
||||
if msg_type == MSG_TYPE_AUTH_OK:
|
||||
server_info = msg
|
||||
break
|
||||
elif msg_type == MSG_TYPE_AUTH_INVALID:
|
||||
raise InvalidAuth
|
||||
finally:
|
||||
if server_info is None:
|
||||
await self._conn.disconnect()
|
||||
|
||||
self._binary_handlers.clear()
|
||||
self._msg_id = 0
|
||||
return True
|
||||
|
||||
async def reconnect(self) -> None:
|
||||
for future in self._result_futures.values():
|
||||
future.set_exception(ConnectionLost)
|
||||
self._result_futures.clear()
|
||||
self._fire_event("disconnected")
|
||||
|
||||
while not self._stop_called:
|
||||
await self._conn.disconnect()
|
||||
try:
|
||||
await self._connect()
|
||||
self._fire_event("ready")
|
||||
break
|
||||
except InvalidAuth as err:
|
||||
raise
|
||||
except ClientException:
|
||||
await asyncio.sleep(5)
|
||||
|
||||
async def connect(self) -> bool:
|
||||
async def receive_task():
|
||||
client_event = None
|
||||
try:
|
||||
while not self._stop_called:
|
||||
try:
|
||||
async for msg in self.receive_message():
|
||||
await self._handle_message(msg)
|
||||
except InvalidMessage as err:
|
||||
LOGGER.warning(err)
|
||||
except ConnectionLost:
|
||||
await self.reconnect()
|
||||
except InvalidAuth as err:
|
||||
client_event = ("reconnect-error", err)
|
||||
finally:
|
||||
await self._conn.disconnect()
|
||||
self._receive_task = None
|
||||
if client_event:
|
||||
self._fire_event(*client_event)
|
||||
|
||||
if self._receive_task is not None:
|
||||
return False
|
||||
self._stop_called = False
|
||||
await self._connect()
|
||||
self._loop = asyncio.get_running_loop()
|
||||
self._receive_task = self._loop.create_task(receive_task())
|
||||
self._fire_event("ready")
|
||||
return True
|
||||
|
||||
def disconnect(self):
|
||||
"""Disconnect the client."""
|
||||
self._stop_called = True
|
||||
for future in self._result_futures.values():
|
||||
future.cancel()
|
||||
self._result_futures.clear()
|
||||
|
||||
async def send_command(
|
||||
self, cmd: Any, wait: bool = True, return_exceptions: bool = False
|
||||
) -> Any | list[Any] | None:
|
||||
is_list = isinstance(cmd, (list, tuple))
|
||||
msgs = [to_dict(i) for i in cmd] if is_list else [to_dict(cmd)]
|
||||
futures: list[asyncio.Future[Any]] = []
|
||||
for msg in msgs:
|
||||
if msg.get("type") in (MSG_TYPE_RESULT, MSG_TYPE_EVENT):
|
||||
continue
|
||||
self._msg_id += 1
|
||||
msg["id"] = self._msg_id
|
||||
if wait:
|
||||
futures.append(self._loop.create_future())
|
||||
self._result_futures[msg["id"]] = futures[-1]
|
||||
|
||||
await self.send_message(msgs[0] if len(msgs) == 1 else msgs)
|
||||
if wait:
|
||||
result = await asyncio.gather(*futures, return_exceptions=return_exceptions)
|
||||
return result if is_list else result[0]
|
||||
|
||||
async def _handle_message(self, msg: dict[str, Any]):
|
||||
msg_type = msg["type"]
|
||||
if msg_type == MSG_TYPE_RESULT:
|
||||
future = self._result_futures.pop(msg.get("id"), None)
|
||||
if future is None:
|
||||
pass
|
||||
elif not msg.get("fail"):
|
||||
future.set_result(msg.get("result"))
|
||||
else:
|
||||
err = msg.get("error", {})
|
||||
future.set_exception(ResultException(err.get("code"), err.get("msg")))
|
||||
elif msg_type == MSG_TYPE_EVENT:
|
||||
self._handle_event(msg)
|
||||
@@ -0,0 +1,70 @@
|
||||
import aiohttp
|
||||
from aiohttp import WSMsgType
|
||||
from typing import cast
|
||||
|
||||
from .exceptions import CannotConnect, ConnectionLost, InvalidMessage
|
||||
from .helpers import LOGGER
|
||||
|
||||
|
||||
class WebsocketConnection:
|
||||
"""Websocket connection to server."""
|
||||
|
||||
def __init__(self, server_url: str, websession: aiohttp.ClientSession):
|
||||
self._server_url = server_url
|
||||
self._websession = websession
|
||||
self._client: aiohttp.ClientWebSocketResponse | None = None
|
||||
|
||||
@property
|
||||
def connected(self) -> bool:
|
||||
"""Return if we're currently connected."""
|
||||
return self._client is not None and not self._client.closed
|
||||
|
||||
async def connect(self) -> bool:
|
||||
"""Connect to the websocket server."""
|
||||
if self.connected:
|
||||
return False
|
||||
|
||||
LOGGER.debug("Trying to connect")
|
||||
try:
|
||||
self._client = await self._websession.ws_connect(
|
||||
self._server_url, heartbeat=55, compress=15, max_msg_size=0
|
||||
)
|
||||
except (aiohttp.WSServerHandshakeError, aiohttp.ClientError) as err:
|
||||
raise CannotConnect(err) from err
|
||||
return True
|
||||
|
||||
async def disconnect(self) -> bool:
|
||||
"""Disconnect the client."""
|
||||
LOGGER.debug("Closing client connection")
|
||||
if self._client is None:
|
||||
return False
|
||||
await self._client.close()
|
||||
self._client = None
|
||||
return True
|
||||
|
||||
async def receive_message(self) -> str | bytes:
|
||||
"""Receive the next message from the server (or raise on error)."""
|
||||
assert self._client
|
||||
ws_msg = await self._client.receive()
|
||||
|
||||
if ws_msg.type == WSMsgType.TEXT:
|
||||
return cast(str, ws_msg.data)
|
||||
elif ws_msg.type == WSMsgType.BINARY:
|
||||
return cast(bytes, ws_msg.data)
|
||||
elif ws_msg.type in (WSMsgType.CLOSE, WSMsgType.CLOSING, WSMsgType.CLOSED):
|
||||
raise ConnectionLost
|
||||
elif ws_msg.type == WSMsgType.ERROR:
|
||||
raise ConnectionLost(ws_msg.data)
|
||||
else:
|
||||
raise InvalidMessage(f"Received unknown type message: {ws_msg.type}")
|
||||
|
||||
async def send_message(self, msg: str | bytes, binary: bool = False) -> None:
|
||||
"""Send a message to the server."""
|
||||
if not self.connected:
|
||||
raise ConnectionLost
|
||||
assert self._client
|
||||
|
||||
await self._client.send_frame(
|
||||
msg.encode() if isinstance(msg, str) else msg,
|
||||
WSMsgType.BINARY if binary else WSMsgType.TEXT,
|
||||
)
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Client-specific exceptions."""
|
||||
|
||||
|
||||
class ClientException(Exception):
|
||||
"""Generic exception."""
|
||||
|
||||
def __init__(self, error: str | Exception | None = None):
|
||||
if error is not None:
|
||||
super().__init__(error if isinstance(error, str) else str(error))
|
||||
|
||||
|
||||
class CannotConnect(ClientException):
|
||||
"""Exception raised when failed to connect the client."""
|
||||
|
||||
|
||||
class InvalidAuth(ClientException):
|
||||
"""Exception raised when authenticate failed."""
|
||||
|
||||
|
||||
class ConnectionLost(ClientException):
|
||||
"""Exception raised when the connection is lost."""
|
||||
|
||||
|
||||
class InvalidMessage(ClientException):
|
||||
"""Exception raised when an invalid message is received."""
|
||||
|
||||
|
||||
class ResultException(Exception):
|
||||
"""Result exception."""
|
||||
|
||||
def __init__(self, code: str, msg: str) -> None:
|
||||
self.code = code
|
||||
self.msg = msg
|
||||
@@ -0,0 +1,53 @@
|
||||
from dataclasses import is_dataclass, asdict
|
||||
import logging
|
||||
import orjson
|
||||
from typing import Any
|
||||
|
||||
|
||||
# logger
|
||||
LOGGER = logging.getLogger(__package__)
|
||||
|
||||
|
||||
# json helpers
|
||||
def omit_none(obj: Any) -> Any:
|
||||
"""Omit dict none fields."""
|
||||
return _omit_none(obj) if isinstance(obj, dict) else obj
|
||||
|
||||
|
||||
def _omit_none(obj: dict):
|
||||
for k in [k for k, v in obj.items() if v is None]:
|
||||
del obj[k]
|
||||
for v in obj.values():
|
||||
if isinstance(v, dict):
|
||||
_omit_none(v)
|
||||
return obj
|
||||
|
||||
|
||||
def to_dict(obj: Any) -> dict:
|
||||
"""Convert obj(msg) to dict."""
|
||||
if isinstance(obj, dict):
|
||||
pass
|
||||
elif callable(to_dict := getattr(obj, "to_dict", None)):
|
||||
obj = to_dict()
|
||||
elif is_dataclass(obj) and not isinstance(obj, type):
|
||||
obj = asdict(obj)
|
||||
else:
|
||||
raise TypeError(f"Type can not to_dict: {type(obj).__name__}")
|
||||
return _omit_none(obj)
|
||||
|
||||
|
||||
def _json_dumps_default(obj: Any) -> Any:
|
||||
"""orjson.dumps default handler."""
|
||||
if callable(to_dict := getattr(obj, "to_dict", None)):
|
||||
return omit_none(to_dict())
|
||||
if is_dataclass(obj) and not isinstance(obj, type):
|
||||
return _omit_none(asdict(obj))
|
||||
raise TypeError(f"Type is not JSON serializable: {type(obj).__name__}")
|
||||
|
||||
|
||||
json_loads = orjson.loads
|
||||
|
||||
|
||||
def json_dumps(obj: Any, option: int = 0) -> bytes:
|
||||
option |= orjson.OPT_PASSTHROUGH_DATACLASS
|
||||
return orjson.dumps(omit_none(obj), default=_json_dumps_default, option=option)
|
||||
Reference in New Issue
Block a user