608 lines
22 KiB
Python
608 lines
22 KiB
Python
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: ...
|
|
|
|
|
|
def _short_type_name(type_name: str) -> str:
|
|
last_dot = type_name.rfind(".")
|
|
if last_dot < 0 or last_dot == len(type_name) - 1:
|
|
return type_name
|
|
return type_name[last_dot + 1 :]
|
|
|
|
|
|
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()
|
|
self._canonical_name_map: dict[str, str] = {}
|
|
|
|
def _canonicalize(self, type_name: str) -> str:
|
|
if not hasattr(self, "_canonical_name_map") or not self._canonical_name_map:
|
|
return type_name
|
|
if type_name in self._canonical_name_map:
|
|
return self._canonical_name_map[type_name]
|
|
short = _short_type_name(type_name)
|
|
if short in self._canonical_name_map:
|
|
return self._canonical_name_map[short]
|
|
return type_name
|
|
|
|
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
|
|
|
|
# Build canonical_name_map and check collision
|
|
self._canonical_name_map = {}
|
|
short_to_full = {}
|
|
|
|
temp_parsers = dict(parser_map)
|
|
temp_parsers[type_name_of(HeartBeat)] = HeartBeat.FromString
|
|
|
|
for full_name in temp_parsers:
|
|
short_name = _short_type_name(full_name)
|
|
existing = short_to_full.get(short_name)
|
|
if existing is not None and existing != full_name:
|
|
raise ValueError(f"Duplicate alias mapping: {short_name} maps to both {existing} and {full_name}")
|
|
short_to_full[short_name] = full_name
|
|
|
|
for full_name in temp_parsers:
|
|
self._canonical_name_map[full_name] = full_name
|
|
short_name = _short_type_name(full_name)
|
|
self._canonical_name_map[short_name] = full_name
|
|
|
|
self._parser_map = {}
|
|
for k, v in temp_parsers.items():
|
|
self._parser_map[self._canonicalize(k)] = v
|
|
|
|
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:
|
|
attempts = 0
|
|
while attempts < MAX_NONCE:
|
|
if self._nonce >= MAX_NONCE:
|
|
self._nonce = 0
|
|
self._nonce += 1
|
|
if self._nonce not in self._pending_requests:
|
|
return self._nonce
|
|
attempts += 1
|
|
raise RuntimeError("All nonces are currently pending")
|
|
|
|
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:
|
|
canonical = self._canonicalize(type_name)
|
|
if canonical in self._req_handlers:
|
|
raise ValueError(f"type {type_name} is already registered with add_request_listener")
|
|
self._handlers.setdefault(canonical, []).append(fn)
|
|
|
|
def remove_listeners(self, type_name: str) -> None:
|
|
canonical = self._canonicalize(type_name)
|
|
self._handlers.pop(canonical, None)
|
|
|
|
def add_request_listener(self, type_name: str, fn: RequestHandler) -> None:
|
|
canonical = self._canonicalize(type_name)
|
|
if self._handlers.get(canonical):
|
|
raise ValueError(f"type {type_name} is already registered with add_listener")
|
|
if canonical in self._req_handlers:
|
|
raise ValueError(f"type {type_name} is already registered with add_request_listener")
|
|
self._req_handlers[canonical] = 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
|
|
|
|
canonical = self._canonicalize(item.type_name)
|
|
req_handler = self._req_handlers.get(canonical)
|
|
listeners = list(self._handlers.get(canonical, []))
|
|
|
|
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()
|
|
canonical = self._canonicalize(type_name)
|
|
parser_cls = self._parser_cls_map.get(canonical)
|
|
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
|
|
canonical_expected = self._canonicalize(expected_type_name)
|
|
canonical_received = self._canonicalize(type_name)
|
|
if canonical_received != canonical_expected:
|
|
pending.set_exception(
|
|
ValueError(
|
|
f"response type mismatch for nonce {response_nonce}: "
|
|
f"expected {canonical_expected}, 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:
|
|
canonical = self._canonicalize(type_name)
|
|
parser = self._parser_map.get(canonical)
|
|
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.full_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
|