from __future__ import annotations import asyncio import pytest from toki_socket.communicator import type_name_of from toki_socket.packets.message_common_pb2 import TestData as ProtoTestData from toki_socket.ws_client import WsClient, connect_ws from toki_socket.ws_server import WsServer def parser_map(): return {type_name_of(ProtoTestData): ProtoTestData.FromString} async def _wait_for_clients(server: WsServer, 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") @pytest.mark.asyncio async def test_ws_send_receive(): received = asyncio.get_running_loop().create_future() server = WsServer( "127.0.0.1", 0, "/", lambda ws: WsClient(ws, 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_ws("127.0.0.1", server.port, "/", 0, 0, parser_map()) await client.communicator.send(ProtoTestData(index=1, message="hello ws")) result = await asyncio.wait_for(received, 2) assert result.index == 1 assert result.message == "hello ws" await client.close() await server.stop() @pytest.mark.asyncio async def test_ws_server_push(): pushed = asyncio.get_running_loop().create_future() server = WsServer( "127.0.0.1", 0, "/", lambda ws: WsClient(ws, 0, 0, parser_map()), ) await server.start() client = await connect_ws("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 ws server")) result = await asyncio.wait_for(pushed, 2) assert result.index == 200 assert result.message == "push from python ws server" await client.close() await server.stop() @pytest.mark.asyncio async def test_ws_request_response(): server = WsServer( "127.0.0.1", 0, "/", lambda ws: WsClient(ws, 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_ws("127.0.0.1", server.port, "/", 0, 0, parser_map()) result = await client.communicator.send_request( ProtoTestData(index=21, message="ws request"), ProtoTestData, timeout=2, ) assert result.index == 42 assert result.message == "echo: ws request" await client.close() await server.stop() @pytest.mark.asyncio async def test_ws_concurrent_requests(): server = WsServer( "127.0.0.1", 0, "/", lambda ws: WsClient(ws, 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_ws("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()