proto-socket/python/test/test_tcp.py
toki 9cc1f1d58f sync: update communicator implementation across all languages
- Align protocol documentation (PROTOCOL.md, README.md, VERSIONING.md)
- Go: add nonce test, update communicator
- Kotlin: update Communicator, TcpClient, TcpServer, add TLS test
- Python: update all modules, add certificate test resources
- TypeScript: update communicator, tcp/ws clients and servers, add tests
- Dart: update communicator, heartbeat mixin, and tests
2026-04-26 05:31:56 +09:00

217 lines
6.5 KiB
Python

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()