- 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)
87 lines
2.8 KiB
Python
87 lines
2.8 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import inspect
|
|
import ssl
|
|
from collections.abc import Callable
|
|
|
|
from google.protobuf.message import Message
|
|
|
|
from proto_socket.communicator import ParserMap
|
|
from proto_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,
|
|
ssl_context: ssl.SSLContext | 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._ssl_context = ssl_context
|
|
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,
|
|
ssl=self._ssl_context,
|
|
)
|
|
|
|
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)
|