from __future__ import annotations import asyncio import inspect from collections.abc import Awaitable, Callable from typing import Any, Protocol, TypeAlias, TypeVar, runtime_checkable from google.protobuf.message import Message from toki_socket.packets.message_common_pb2 import HeartBeat, PacketBase T = TypeVar("T", bound=Message) ParserMap: TypeAlias = dict[str, Callable[[bytes], Message]] RequestHandler: TypeAlias = Callable[..., Message | Awaitable[Message | None] | None] _STOP = object() MAX_NONCE = 2_147_483_647 class NotConnectedError(RuntimeError): def __init__(self) -> None: super().__init__("not connected") @runtime_checkable class Transport(Protocol): async def write_packet(self, base: PacketBase) -> None: ... async def close(self) -> None: ... class Communicator: def __init__(self) -> None: self._nonce = 0 self._is_alive = False self._parser_map: ParserMap = {} self._handlers: dict[str, list[Callable[[Message], Any]]] = {} self._req_handlers: dict[str, RequestHandler] = {} self._pending_requests: dict[int, tuple[str, asyncio.Future[Message]]] = {} self._write_queue: asyncio.Queue[Any] = asyncio.Queue(maxsize=64) self._closed = asyncio.Event() self._transport: Transport | None = None self._write_loop_task: asyncio.Task[None] | None = None self._write_error_handler: Callable[[Exception], Any] | None = None def initialize(self, transport: Transport, parser_map: ParserMap) -> None: self._transport = transport self._parser_map = dict(parser_map) self._parser_map[type_name_of(HeartBeat)] = HeartBeat.FromString self._handlers = {} self._req_handlers = {} self._pending_requests = {} self._write_queue = asyncio.Queue(maxsize=64) self._closed = asyncio.Event() self._is_alive = True self._write_loop_task = asyncio.create_task(self._write_loop()) def is_alive(self) -> bool: return self._is_alive def set_write_error_handler(self, fn: Callable[[Exception], Any]) -> None: self._write_error_handler = fn def next_nonce(self) -> int: if self._nonce >= MAX_NONCE: self._nonce = 0 self._nonce += 1 return self._nonce def shutdown(self) -> None: if not self._is_alive and self._closed.is_set(): return self._is_alive = False self._closed.set() for _, pending in list(self._pending_requests.values()): if not pending.done(): pending.set_exception(NotConnectedError()) self._pending_requests.clear() try: self._write_queue.put_nowait(_STOP) except asyncio.QueueFull: if self._write_loop_task is not None: self._write_loop_task.cancel() async def close(self) -> None: self.shutdown() if self._transport is not None: await self._transport.close() async def wait_closed(self) -> None: await self._closed.wait() async def _write_loop(self) -> None: while True: item = await self._write_queue.get() if item is _STOP: return base, done = item try: if self._transport is None: raise NotConnectedError() await self._transport.write_packet(base) if not done.done(): done.set_result(None) except Exception as exc: if not done.done(): done.set_exception(exc) if self._write_error_handler is not None: self._write_error_handler(exc) async def queue_packet(self, base: PacketBase) -> None: if not self.is_alive(): raise NotConnectedError() loop = asyncio.get_running_loop() done: asyncio.Future[None] = loop.create_future() item = (base, done) try: self._write_queue.put_nowait(item) except asyncio.QueueFull: closed_task = asyncio.create_task(self._closed.wait()) put_task = asyncio.create_task(self._write_queue.put(item)) finished, pending = await asyncio.wait( {closed_task, put_task}, return_when=asyncio.FIRST_COMPLETED, ) for task in pending: task.cancel() if closed_task in finished: raise NotConnectedError() closed_task = asyncio.create_task(self._closed.wait()) finished, pending = await asyncio.wait( {done, closed_task}, return_when=asyncio.FIRST_COMPLETED, ) for task in pending: task.cancel() if closed_task in finished and not done.done(): raise NotConnectedError() await done async def send(self, message: Message) -> None: if not self.is_alive(): raise NotConnectedError() await self.queue_packet( PacketBase( typeName=type_name_of(message), nonce=self.next_nonce(), data=message.SerializeToString(), ) ) async def send_request( self, req: Message, res_type: type[T] | T, timeout: float = 30.0, ) -> T: if not self.is_alive(): raise NotConnectedError() request_nonce = self.next_nonce() expected_type_name = type_name_of(res_type) loop = asyncio.get_running_loop() pending: asyncio.Future[Message] = loop.create_future() self._pending_requests[request_nonce] = (expected_type_name, pending) try: await self.queue_packet( PacketBase( typeName=type_name_of(req), nonce=request_nonce, data=req.SerializeToString(), ) ) timeout = timeout if timeout > 0 else 30.0 return await asyncio.wait_for(pending, timeout) except TimeoutError as exc: raise TimeoutError(f"request timeout for nonce {request_nonce}") from exc finally: self._pending_requests.pop(request_nonce, None) def add_listener(self, type_name: str, fn: Callable[[Message], Any]) -> None: if type_name in self._req_handlers: raise ValueError(f"type {type_name} is already registered with add_request_listener") self._handlers.setdefault(type_name, []).append(fn) def remove_listeners(self, type_name: str) -> None: self._handlers.pop(type_name, None) def add_request_listener(self, type_name: str, fn: RequestHandler) -> None: if self._handlers.get(type_name): raise ValueError(f"type {type_name} is already registered with add_listener") if type_name in self._req_handlers: raise ValueError(f"type {type_name} is already registered with add_request_listener") self._req_handlers[type_name] = fn def on_received_data( self, type_name: str, data: bytes, nonce: int = 0, response_nonce: int = 0, ) -> None: if response_nonce > 0: self._handle_response(type_name, data, response_nonce) return req_handler = self._req_handlers.get(type_name) listeners = list(self._handlers.get(type_name, [])) if req_handler is not None: try: message = self._parse(type_name, data) except Exception: return asyncio.create_task(self._run_request_handler(req_handler, message, nonce)) return if not listeners: return try: message = self._parse(type_name, data) except Exception: return for listener in listeners: listener(message) async def _run_request_handler( self, handler: RequestHandler, message: Message, request_nonce: int, ) -> None: try: result = handler(message, request_nonce) if _accepts_nonce(handler) else handler(message) if inspect.isawaitable(result): result = await result if result is None or not self.is_alive(): return await self.queue_packet( PacketBase( typeName=type_name_of(result), nonce=self.next_nonce(), responseNonce=request_nonce, data=result.SerializeToString(), ) ) except Exception: return def _handle_response(self, type_name: str, data: bytes, response_nonce: int) -> None: item = self._pending_requests.pop(response_nonce, None) if item is None: return expected_type_name, pending = item if type_name != expected_type_name: pending.set_exception( ValueError( f"response type mismatch for nonce {response_nonce}: " f"expected {expected_type_name}, got {type_name}" ) ) return try: pending.set_result(self._parse(type_name, data)) except Exception as exc: pending.set_exception(exc) def _parse(self, type_name: str, data: bytes) -> Message: parser = self._parser_map.get(type_name) if parser is None: raise ValueError(f"protobuf parser is not registered for type {type_name}") return parser(data) def type_name_of(message_or_class: Message | type[Message]) -> str: descriptor = getattr(message_or_class, "DESCRIPTOR", None) if descriptor is None: descriptor = message_or_class.__class__.DESCRIPTOR return descriptor.name def _accepts_nonce(handler: RequestHandler) -> bool: try: signature = inspect.signature(handler) except (TypeError, ValueError): return False positional = [ param for param in signature.parameters.values() if param.kind in (inspect.Parameter.POSITIONAL_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD) ] return len(positional) >= 2