from __future__ import annotations import asyncio import ssl from pathlib import Path import pytest from toki_socket.communicator import type_name_of from toki_socket.packets.message_common_pb2 import TestData as ProtoTestData from toki_socket.tcp_client import TcpClient, connect_tcp, connect_tcp_tls from toki_socket.tcp_server import TcpServer def _make_server_ssl_context() -> ssl.SSLContext: ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) certs_dir = Path(__file__).parent / "certs" ctx.load_cert_chain(certs_dir / "server.crt", certs_dir / "server.key") return ctx def _make_client_ssl_context() -> ssl.SSLContext: ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) certs_dir = Path(__file__).parent / "certs" ctx.load_verify_locations(certs_dir / "server.crt") ctx.check_hostname = False return ctx async def _wait_for_clients(server: TcpServer, count: int = 1, timeout: float = 1.0) -> None: loop = asyncio.get_running_loop() deadline = loop.time() + timeout while loop.time() < deadline: if len(server.clients()) >= count: return await asyncio.sleep(0.01) raise TimeoutError("server did not see expected client count") def parser_map(): return {type_name_of(ProtoTestData): ProtoTestData.FromString} @pytest.mark.asyncio async def test_tcp_send_receive(): received = asyncio.get_running_loop().create_future() server = TcpServer( "127.0.0.1", 0, lambda reader, writer: TcpClient(reader, writer, 0, 0, parser_map()), ) server.on_client_connected = lambda client: client.communicator.add_listener( type_name_of(ProtoTestData), lambda message: received.set_result(message) if not received.done() else None, ) await server.start() client = await connect_tcp("127.0.0.1", server.port, 0, 0, parser_map()) await client.communicator.send(ProtoTestData(index=1, message="hello")) result = await asyncio.wait_for(received, 2) assert result.index == 1 assert result.message == "hello" await client.close() await server.stop() @pytest.mark.asyncio async def test_tcp_request_response(): server = TcpServer( "127.0.0.1", 0, lambda reader, writer: TcpClient(reader, writer, 0, 0, parser_map()), ) server.on_client_connected = lambda client: client.communicator.add_request_listener( type_name_of(ProtoTestData), lambda req: ProtoTestData(index=req.index * 2, message="echo: " + req.message), ) await server.start() client = await connect_tcp("127.0.0.1", server.port, 0, 0, parser_map()) result = await client.communicator.send_request( ProtoTestData(index=21, message="single request"), ProtoTestData, timeout=2, ) assert result.index == 42 assert result.message == "echo: single request" await client.close() await server.stop() @pytest.mark.asyncio async def test_tcp_tls_send_receive(): received = asyncio.get_running_loop().create_future() server = TcpServer( "127.0.0.1", 0, lambda reader, writer: TcpClient(reader, writer, 0, 0, parser_map()), ssl_context=_make_server_ssl_context(), ) server.on_client_connected = lambda client: client.communicator.add_listener( type_name_of(ProtoTestData), lambda message: received.set_result(message) if not received.done() else None, ) await server.start() client = await connect_tcp_tls( "127.0.0.1", server.port, _make_client_ssl_context(), 0, 0, parser_map(), ) await client.communicator.send(ProtoTestData(index=3, message="hello over tls")) result = await asyncio.wait_for(received, 2) assert result.index == 3 assert result.message == "hello over tls" await client.close() await server.stop() @pytest.mark.asyncio async def test_tcp_tls_request_response(): server = TcpServer( "127.0.0.1", 0, lambda reader, writer: TcpClient(reader, writer, 0, 0, parser_map()), ssl_context=_make_server_ssl_context(), ) server.on_client_connected = lambda client: client.communicator.add_request_listener( type_name_of(ProtoTestData), lambda req: ProtoTestData(index=req.index * 2, message="tls echo: " + req.message), ) await server.start() client = await connect_tcp_tls( "127.0.0.1", server.port, _make_client_ssl_context(), 0, 0, parser_map(), ) result = await client.communicator.send_request( ProtoTestData(index=24, message="secure request"), ProtoTestData, timeout=2, ) assert result.index == 48 assert result.message == "tls echo: secure request" await client.close() await server.stop() @pytest.mark.asyncio async def test_tcp_server_push(): pushed = asyncio.get_running_loop().create_future() server = TcpServer( "127.0.0.1", 0, lambda reader, writer: TcpClient(reader, writer, 0, 0, parser_map()), ) await server.start() client = await connect_tcp("127.0.0.1", server.port, 0, 0, parser_map()) client.communicator.add_listener( type_name_of(ProtoTestData), lambda msg: pushed.set_result(msg) if not pushed.done() else None, ) await _wait_for_clients(server) await server.broadcast(ProtoTestData(index=200, message="push from python server")) result = await asyncio.wait_for(pushed, 2) assert result.index == 200 assert result.message == "push from python server" await client.close() await server.stop() @pytest.mark.asyncio async def test_tcp_concurrent_requests(): server = TcpServer( "127.0.0.1", 0, lambda reader, writer: TcpClient(reader, writer, 0, 0, parser_map()), ) server.on_client_connected = lambda client: client.communicator.add_request_listener( type_name_of(ProtoTestData), lambda req: ProtoTestData(index=req.index * 2, message="echo: " + req.message), ) await server.start() client = await connect_tcp("127.0.0.1", server.port, 0, 0, parser_map()) async def do_request(i: int) -> None: index = 30 + i message = f"request {i}" result = await client.communicator.send_request( ProtoTestData(index=index, message=message), ProtoTestData, timeout=2, ) assert result.index == index * 2 assert result.message == f"echo: {message}" await asyncio.gather(*[do_request(i) for i in range(5)]) await client.close() await server.stop()