proto-socket/python/proto_socket/communicator.py

603 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:
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:
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