Files
hse-2026/homework/02-grpc-messenger/tests/test_server.py
T

286 lines
9.5 KiB
Python
Raw Normal View History

2026-09-17 16:34:16 +03:00
import copy
import json
import os
import pathlib
import queue
import re
import socket
import subprocess
import threading
import time
from datetime import datetime, timezone
from typing import Dict
import pytest
test_message = {'author': 'alice', 'text': 'hello'}
PROTO_DIR = pathlib.Path(__file__).resolve().parent.parent / 'solution' / 'proto'
SOCKET_CONNECT_TIMEOUT_S = 1
SERVICE_READY_TIMEOUT_S = 20
SERVICE_RETRY_INTERVAL_S = 0.5
GRPC_CALL_TIMEOUT_S = 5
GRPC_PROCESS_TIMEOUT_S = 10
GRPC_STREAM_TIMEOUT_S = 60
PROCESS_STOP_TIMEOUT_S = 5
MESSAGE_TIMEOUT_S = 10
def wait_for_socket(host, port):
deadline = time.monotonic() + SERVICE_READY_TIMEOUT_S
last_exception = None
while True:
try:
with socket.create_connection((host, port), timeout=SOCKET_CONNECT_TIMEOUT_S):
pass
return
except OSError as exc:
last_exception = exc
remaining = deadline - time.monotonic()
if remaining <= 0:
pytest.fail(
f'timed out waiting for TCP endpoint {host}:{port}: {last_exception}',
pytrace=False,
)
time.sleep(min(SERVICE_RETRY_INTERVAL_S, remaining))
@pytest.fixture(scope='session')
def server_addr():
addr = os.environ.get('MESSENGER_TEST_SERVER_ADDR', '127.0.0.1:51075')
host = addr.split(':')[0]
port = int(addr.split(':')[1])
wait_for_socket(host, port)
yield addr
def send_message(server_address, message: Dict[str, str]) -> Dict[str, str]:
grpcurl_cmd = ['grpcurl',
'-max-time', str(GRPC_CALL_TIMEOUT_S),
'-import-path', str(PROTO_DIR),
'-proto', 'messenger.proto',
'-d',
json.dumps(message),
'-plaintext',
server_address,
'mes_grpc.MessengerServer/SendMessage']
try:
completed = subprocess.run(
grpcurl_cmd,
capture_output=True,
check=False,
timeout=GRPC_PROCESS_TIMEOUT_S,
)
except subprocess.TimeoutExpired:
pytest.fail(
f'grpcurl did not finish within {GRPC_PROCESS_TIMEOUT_S} seconds',
pytrace=False,
)
assert completed.returncode == 0, completed.stderr
assert len(completed.stderr) == 0, completed.stderr
output_str = completed.stdout.decode('ascii')
output = json.loads(output_str)
message_with_timestamp = copy.deepcopy(message)
message_with_timestamp['sendTime'] = output['sendTime']
return message_with_timestamp
class MessageStream:
def __init__(self, server_address):
grpcurl_cmd = ['grpcurl',
'-max-time', str(GRPC_STREAM_TIMEOUT_S),
'-import-path', str(PROTO_DIR),
'-proto', 'messenger.proto',
'-plaintext',
server_address,
'mes_grpc.MessengerServer/ReadMessages']
self._process = subprocess.Popen(
grpcurl_cmd,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
self._messages = queue.Queue()
self._reader_error = None
self._reader_finished = threading.Event()
self._closed = False
self._reader = threading.Thread(target=self._read_messages, daemon=True)
self._reader.start()
def _read_messages(self):
try:
message_lines = []
for line in self._process.stdout:
message_lines.append(line)
if line.rstrip() == '}':
self._messages.put(json.loads(''.join(message_lines)))
message_lines = []
except Exception as exc:
self._reader_error = exc
finally:
self._reader_finished.set()
def __enter__(self):
return self
def __exit__(self, exc_type, exc_value, traceback):
self.close()
def _failure_detail(self):
if self._reader_error is not None:
return f'reader failed: {self._reader_error}'
returncode = self._process.poll()
if returncode is not None:
stderr = self._process.stderr.read().strip()
return f'grpcurl exited with status {returncode}: {stderr}'
return 'grpcurl is still running but produced no matching message'
def _get_message(self, deadline, timeout_message):
while True:
if self._reader_error is not None:
raise AssertionError(self._failure_detail())
if self._reader_finished.is_set() and self._messages.empty():
raise AssertionError(self._failure_detail())
remaining = deadline - time.monotonic()
if remaining <= 0:
raise AssertionError(f'{timeout_message}: {self._failure_detail()}')
try:
return self._messages.get(timeout=min(remaining, 0.1))
except queue.Empty:
pass
def wait_for_message(self, expected_message, timeout, preceding_messages=None):
deadline = time.monotonic() + timeout
while True:
message = self._get_message(
deadline,
f'timed out waiting for message {expected_message}',
)
if message == expected_message:
return
if preceding_messages is not None:
preceding_messages.append(message)
def read_messages(self, count, timeout):
deadline = time.monotonic() + timeout
messages = []
while len(messages) < count:
messages.append(self._get_message(
deadline,
f'timed out after receiving {len(messages)} of {count} messages',
))
return messages
def close(self):
if self._closed:
return
self._closed = True
if self._process.poll() is None:
self._process.terminate()
try:
self._process.wait(timeout=PROCESS_STOP_TIMEOUT_S)
except subprocess.TimeoutExpired:
self._process.kill()
try:
self._process.wait(timeout=PROCESS_STOP_TIMEOUT_S)
except subprocess.TimeoutExpired:
pytest.fail('grpcurl did not exit after SIGKILL', pytrace=False)
self._reader.join(timeout=PROCESS_STOP_TIMEOUT_S)
assert not self._reader.is_alive()
def wait_for_streams(server_address, streams):
preceding_messages = [[] for _ in streams]
probes = []
for attempt in range(10):
probe = send_message(
server_address,
{'author': 'StreamProbe', 'text': f'probe #{attempt}'},
)
probes.append(probe)
streams_ready = True
for index, stream in enumerate(streams):
try:
stream.wait_for_message(
probe,
timeout=1,
preceding_messages=preceding_messages[index],
)
except AssertionError:
streams_ready = False
if streams_ready:
return [
[message for message in messages if message not in probes]
for messages in preceding_messages
]
raise AssertionError('ReadMessages streams did not become ready')
def timestamp_key(timestamp):
match = re.fullmatch(r'(\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2})(?:\.(\d{1,9}))?Z', timestamp)
assert match is not None, f'invalid protobuf timestamp: {timestamp}'
seconds = int(datetime.strptime(match.group(1), '%Y-%m-%dT%H:%M:%S')
.replace(tzinfo=timezone.utc).timestamp())
nanos = int((match.group(2) or '').ljust(9, '0'))
return seconds, nanos
def test_send_smoke(server_addr):
send_message(server_addr, test_message)
def test_send_returns_ascending_time(server_addr):
outputs = []
for _ in range(10):
outputs.append(send_message(server_addr, test_message))
for output1, output2 in zip(outputs, outputs[1:]):
assert timestamp_key(output1['sendTime']) < timestamp_key(output2['sendTime'])
def test_get_messages_smoke(server_addr):
with MessageStream(server_addr) as stream:
wait_for_streams(server_addr, [stream])
test_message_with_timestamp = send_message(server_addr, test_message)
messages = stream.read_messages(1, timeout=MESSAGE_TIMEOUT_S)
assert len(messages) == 1
assert messages[0] == test_message_with_timestamp
def test_get_only_sends_new(server_addr):
messages1 = []
messages3 = []
n1, n2, n3 = 2, 3, 4
with MessageStream(server_addr) as stream:
messages_before_ready = wait_for_streams(server_addr, [stream])
assert messages_before_ready == [[]]
for _ in range(n1):
message = send_message(server_addr, test_message)
messages1.append(message)
messages = stream.read_messages(n1, timeout=MESSAGE_TIMEOUT_S)
assert len(messages1) == len(messages)
for m1, m2 in zip(messages1, messages):
assert m1 == m2
for _ in range(n2):
send_message(server_addr, test_message)
with MessageStream(server_addr) as stream:
messages_before_ready = wait_for_streams(server_addr, [stream])
assert messages_before_ready == [[]]
for _ in range(n3):
message = send_message(server_addr, test_message)
messages3.append(message)
messages = stream.read_messages(n3, timeout=MESSAGE_TIMEOUT_S)
assert len(messages3) == len(messages)
for m1, m2 in zip(messages3, messages):
assert m1 == m2