from __future__ import annotations import asyncio import ssl import struct from contextlib import suppress from google.protobuf.message import DecodeError from proto_socket.base_client import BaseClient from proto_socket.communicator import ParserMap, type_name_of from proto_socket.packets.message_common_pb2 import HeartBeat, PacketBase MAX_PACKET_SIZE = 64 << 20 class TcpClient(BaseClient): def __init__( self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, interval_sec: int, wait_sec: int, parser_map: ParserMap, ) -> None: super().__init__(interval_sec, wait_sec, self._do_close_impl) self._reader = reader self._writer = writer self._init_base() self.communicator.initialize(self, parser_map) self.communicator.set_write_error_handler(lambda _: asyncio.create_task(self._on_disconnected())) self.communicator.add_listener( type_name_of(HeartBeat), lambda _: asyncio.create_task(self._on_heartbeat()), ) self._read_task = asyncio.create_task(self._read_loop()) asyncio.create_task(self._send_heartbeat()) async def write_packet(self, base: PacketBase) -> None: data = base.SerializeToString() if len(data) > MAX_PACKET_SIZE: raise ValueError("packet exceeds maximum size") self._writer.write(struct.pack(">I", len(data)) + data) await self._writer.drain() async def _do_close_impl(self) -> None: self._writer.close() with suppress(Exception): await self._writer.wait_closed() async def _read_loop(self) -> None: while self.communicator.is_alive(): try: header = await self._reader.readexactly(4) length = struct.unpack(">I", header)[0] if length == 0: continue if length > MAX_PACKET_SIZE: await self._on_disconnected() return data = await self._reader.readexactly(length) base = PacketBase() base.ParseFromString(data) self.communicator.on_received_data( base.typeName, base.data, base.nonce, base.responseNonce, ) asyncio.create_task(self._send_heartbeat()) except (asyncio.IncompleteReadError, ConnectionError, OSError, DecodeError): await self._on_disconnected() return async def connect_tcp( host: str, port: int, interval_sec: int, wait_sec: int, parser_map: ParserMap, ) -> TcpClient: reader, writer = await asyncio.open_connection(host, port) return TcpClient(reader, writer, interval_sec, wait_sec, parser_map) async def connect_tcp_tls( host: str, port: int, ssl_context: ssl.SSLContext, interval_sec: int, wait_sec: int, parser_map: ParserMap, ) -> TcpClient: reader, writer = await asyncio.open_connection(host, port, ssl=ssl_context) return TcpClient(reader, writer, interval_sec, wait_sec, parser_map)