- Add gateway contract for inbound queue ordering (03+02_gateway_contract) - Update receive thread model across Dart, Go, Kotlin, Python, TypeScript - Implement queue ordering logic in BaseClient and Communicator - Add comprehensive tests for queue ordering in all languages - Update documentation (PORTING_GUIDE, PROTOCOL, README) - Archive task completion logs for completed subtasks
370 lines
13 KiB
Python
370 lines
13 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
|
|
|
|
|
|
@dataclass
|
|
class _InboundItem:
|
|
type_name: str
|
|
data: bytes
|
|
nonce: int
|
|
response_nonce: int
|
|
|
|
|
|
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
|
|
|
|
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._inbound_queue = asyncio.Queue(maxsize=64)
|
|
self._closed = asyncio.Event()
|
|
self._is_alive = True
|
|
self._is_closing = False
|
|
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 _, 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()
|
|
|
|
async def close(self) -> None:
|
|
if not self._is_alive or self._is_closing:
|
|
return
|
|
self._is_closing = 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()
|
|
|
|
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
|
|
try:
|
|
self._inbound_queue.put_nowait(_InboundItem(type_name, data, nonce, response_nonce))
|
|
except asyncio.QueueFull:
|
|
asyncio.create_task(self.enqueue_inbound(type_name, data, nonce, response_nonce))
|
|
|
|
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
|
|
await self._inbound_queue.put(_InboundItem(type_name, data, nonce, response_nonce))
|
|
|
|
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:
|
|
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 req_handler is not None:
|
|
try:
|
|
message = self._parse(item.type_name, item.data)
|
|
except Exception:
|
|
return
|
|
await self._run_request_handler(req_handler, message, item.nonce)
|
|
return
|
|
|
|
if not listeners:
|
|
return
|
|
try:
|
|
message = self._parse(item.type_name, item.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
|