from __future__ import annotations import asyncio import inspect from collections.abc import Callable from google.protobuf.message import Message from toki_socket.communicator import ParserMap from toki_socket.tcp_client import TcpClient class TcpServer: def __init__( self, host: str, port: int, new_client: Callable[[asyncio.StreamReader, asyncio.StreamWriter], TcpClient] | None = None, interval_sec: int = 0, wait_sec: int = 0, parser_map: ParserMap | None = None, ) -> None: self._host = host self._port = port self._new_client = new_client self._interval_sec = interval_sec self._wait_sec = wait_sec self._parser_map = parser_map self._clients: list[TcpClient] = [] self._server: asyncio.Server | None = None self.on_client_connected: Callable[[TcpClient], object] = lambda _: None @property def port(self) -> int: if self._server is None or not self._server.sockets: return self._port return int(self._server.sockets[0].getsockname()[1]) def clients(self) -> list[TcpClient]: return list(self._clients) async def start(self) -> None: self._server = await asyncio.start_server(self._handle_connection, self._host, self._port) async def stop(self) -> None: if self._server is not None: self._server.close() await self._server.wait_closed() self._server = None clients = list(self._clients) self._clients.clear() for client in clients: await client.close() async def broadcast(self, message: Message) -> None: for client in self.clients(): await client.communicator.send(message) async def _handle_connection( self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, ) -> None: if self._new_client is not None: client = self._new_client(reader, writer) else: if self._parser_map is None: raise ValueError("parser_map is required when new_client is not provided") client = TcpClient(reader, writer, self._interval_sec, self._wait_sec, self._parser_map) self._clients.append(client) client.add_disconnect_listener(lambda c: self._remove_client(c)) result = self.on_client_connected(client) if inspect.isawaitable(result): await result def _remove_client(self, client: TcpClient) -> None: if client in self._clients: self._clients.remove(client)