proto-socket/python/proto_socket/communicator.py
toki a9480c5afb feat: implement inbound queue ordering across all language implementations
- 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
2026-06-02 05:36:52 +09:00

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