proto-socket/python/proto_socket/tcp_client.py
toki d0754a353a refactor: rename toki_socket to proto_socket across all languages
- Rename package/module from toki_socket to proto_socket in Dart, Kotlin, Python
- Update crosstest implementations to use renamed packages
- Add new proto_socket skill, deprecate add-toki-socket-crosstest-language skill
- Update domain rules for all languages
- Update documentation (README, PORTING_GUIDE, PROTOCOL, VERSIONING)
2026-05-02 07:19:12 +09:00

96 lines
3.1 KiB
Python

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)