from __future__ import annotations import asyncio import inspect from collections.abc import Awaitable, Callable from dataclasses import dataclass from typing import Any, Protocol, TypeAlias, TypeVar, runtime_checkable from google.protobuf.message import Message from proto_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 def _process_decode(type_name: str, data: bytes, parser_cls: type[Message] | None) -> tuple[str, Message | None, bool]: if parser_cls is not None: try: msg = parser_cls.FromString(data) return type_name, msg, True except Exception: return type_name, None, False return type_name, None, False def _init_process() -> None: import sys import os current_dir = os.path.dirname(os.path.abspath(__file__)) packets_dir = os.path.join(current_dir, "packets") if packets_dir not in sys.path: sys.path.insert(0, packets_dir) @dataclass class _InboundItem: type_name: str data: bytes nonce: int response_nonce: int parsed_message: Message | None = None parse_exception: Exception | None = None 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._is_closing = 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._inbound_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._receive_task: asyncio.Task[None] | None = None self._write_error_handler: Callable[[Exception], Any] | None = None self._gateway_type: str | None = None self._payload_size_threshold: int = 1024 self._next_recv_seq = 0 self._next_dispatch_seq = 1 self._reorder_buffer: dict[int, Any] = {} self._process_executor: Any = None self._inbound_semaphore: asyncio.Semaphore | None = None self._parser_cls_map: dict[str, type[Message]] = {} self._gateway_tasks: set[asyncio.Task[Any]] = set() def initialize( self, transport: Transport, parser_map: ParserMap, gateway_type: str | None = None, payload_size_threshold: int = 1024, ) -> None: import sys import os current_dir = os.path.dirname(os.path.abspath(__file__)) packets_dir = os.path.join(current_dir, "packets") if packets_dir not in sys.path: sys.path.insert(0, packets_dir) self._transport = transport self._parser_map = dict(parser_map) self._parser_map[type_name_of(HeartBeat)] = HeartBeat.FromString self._parser_cls_map = {} for k, parser in self._parser_map.items(): if hasattr(parser, "__self__") and isinstance(parser.__self__, type): self._parser_cls_map[k] = parser.__self__ self._handlers = {} self._req_handlers = {} self._pending_requests = {} self._write_queue = asyncio.Queue(maxsize=64) self._inbound_queue = asyncio.Queue(maxsize=64) self._inbound_semaphore = asyncio.Semaphore(65) self._closed = asyncio.Event() self._is_alive = True self._is_closing = False self._gateway_tasks = set() self._gateway_type = gateway_type self._payload_size_threshold = payload_size_threshold self._next_recv_seq = 0 self._next_dispatch_seq = 1 self._reorder_buffer = {} if self._gateway_type == "process": import concurrent.futures self._process_executor = concurrent.futures.ProcessPoolExecutor( max_workers=2, initializer=_init_process ) else: self._process_executor = None self._write_loop_task = asyncio.create_task(self._write_loop()) self._receive_task = asyncio.create_task(self._receive_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._is_closing = False self._closed.set() for task in list(self._gateway_tasks): if not task.done(): task.cancel() self._gateway_tasks.clear() for _, pending in list(self._pending_requests.values()): if not pending.done(): pending.set_exception(NotConnectedError()) self._pending_requests.clear() if self._receive_task is not None: self._receive_task.cancel() try: self._write_queue.put_nowait(_STOP) except asyncio.QueueFull: if self._write_loop_task is not None: self._write_loop_task.cancel() if self._process_executor is not None: self._process_executor.shutdown(wait=False) self._process_executor = None async def close(self) -> None: if not self._is_alive or self._is_closing: return self._is_closing = True if self._gateway_tasks: await asyncio.gather(*list(self._gateway_tasks), return_exceptions=True) await self._inbound_queue.join() self._is_alive = False self._is_closing = False self._closed.set() for _, pending in list(self._pending_requests.values()): if not pending.done(): pending.set_exception(NotConnectedError()) self._pending_requests.clear() if self._receive_task is not None: self._receive_task.cancel() try: self._write_queue.put_nowait(_STOP) except asyncio.QueueFull: if self._write_loop_task is not None: self._write_loop_task.cancel() if self._transport is not None: await self._transport.close() if self._process_executor is not None: self._process_executor.shutdown(wait=False) self._process_executor = None 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 not self._is_alive or self._is_closing: return if response_nonce > 0: self._handle_response(type_name, data, response_nonce) return self._next_recv_seq += 1 seq = self._next_recv_seq task = asyncio.create_task(self._process_and_enqueue(seq, type_name, data, nonce, response_nonce)) self._gateway_tasks.add(task) task.add_done_callback(self._gateway_tasks.discard) async def enqueue_inbound( self, type_name: str, data: bytes, nonce: int = 0, response_nonce: int = 0, ) -> None: if not self._is_alive or self._is_closing: return if response_nonce > 0: self._handle_response(type_name, data, response_nonce) return self._next_recv_seq += 1 seq = self._next_recv_seq task = asyncio.create_task(self._process_and_enqueue(seq, type_name, data, nonce, response_nonce)) self._gateway_tasks.add(task) task.add_done_callback(self._gateway_tasks.discard) await task async def _receive_loop(self) -> None: while True: try: item = await self._inbound_queue.get() except asyncio.CancelledError: return try: await self._dispatch_inbound(item) except Exception: pass finally: self._inbound_queue.task_done() async def _dispatch_inbound(self, item: _InboundItem) -> None: try: if item.response_nonce > 0: self._handle_response(item.type_name, item.data, item.response_nonce) return req_handler = self._req_handlers.get(item.type_name) listeners = list(self._handlers.get(item.type_name, [])) if item.parsed_message is not None: message = item.parsed_message elif item.parse_exception is not None: return else: try: message = self._parse(item.type_name, item.data) except Exception: return if req_handler is not None: await self._run_request_handler(req_handler, message, item.nonce) return if not listeners: return for listener in listeners: listener(message) finally: if self._inbound_semaphore is not None: self._inbound_semaphore.release() async def _process_and_enqueue( self, seq: int, type_name: str, data: bytes, nonce: int, response_nonce: int, ) -> None: if self._inbound_semaphore is not None: await self._inbound_semaphore.acquire() acquired = True try: use_gateway = False if self._gateway_type in ("process", "asyncio"): if len(data) >= self._payload_size_threshold: use_gateway = True if not use_gateway: try: message = self._parse(type_name, data) exc = None except Exception as e: message = None exc = e acquired = False await self._dispatch_gateway_result(seq, type_name, message, exc, nonce, response_nonce) else: if self._gateway_type == "process" and self._process_executor is not None: loop = asyncio.get_running_loop() parser_cls = self._parser_cls_map.get(type_name) try: _, decoded_msg, is_success = await loop.run_in_executor( self._process_executor, _process_decode, type_name, data, parser_cls, ) acquired = False if is_success: await self._dispatch_gateway_result(seq, type_name, decoded_msg, None, nonce, response_nonce) else: await self._dispatch_gateway_result(seq, type_name, None, ValueError("process decode failed"), nonce, response_nonce) except Exception as e: acquired = False await self._dispatch_gateway_result(seq, type_name, None, e, nonce, response_nonce) elif self._gateway_type == "asyncio": try: await asyncio.sleep(0) message = self._parse(type_name, data) acquired = False await self._dispatch_gateway_result(seq, type_name, message, None, nonce, response_nonce) except Exception as e: acquired = False await self._dispatch_gateway_result(seq, type_name, None, e, nonce, response_nonce) except Exception as e: if acquired: acquired = False await self._dispatch_gateway_result(seq, type_name, None, e, nonce, response_nonce) finally: if acquired and self._inbound_semaphore is not None: self._inbound_semaphore.release() async def _dispatch_gateway_result( self, seq: int, type_name: str, message: Message | None, exc: Exception | None, nonce: int, response_nonce: int, ) -> None: self._reorder_buffer[seq] = (type_name, message, exc, nonce, response_nonce) while self._next_dispatch_seq in self._reorder_buffer: curr_seq = self._next_dispatch_seq item = self._reorder_buffer.pop(curr_seq) t_name, msg, e, n, r_n = item inbound_item = _InboundItem( type_name=t_name, data=b"", nonce=n, response_nonce=r_n, ) inbound_item.parsed_message = msg inbound_item.parse_exception = e await self._inbound_queue.put(inbound_item) self._next_dispatch_seq += 1 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