Build 060.26: complete Integration and Regression
This commit is contained in:
@@ -0,0 +1,904 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import threading
|
||||
from collections import deque
|
||||
from collections.abc import Callable
|
||||
from contextlib import AsyncExitStack
|
||||
from dataclasses import dataclass
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from types import TracebackType
|
||||
from typing import Any
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
from websockets.asyncio.server import (
|
||||
Server,
|
||||
ServerConnection,
|
||||
serve,
|
||||
)
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
from websockets.protocol import State
|
||||
from websockets.typing import Subprotocol
|
||||
|
||||
from tests.support.async_wait import wait_until
|
||||
|
||||
RESOURCE_CLEANUP_TIMEOUT_SECONDS = 5.0
|
||||
ENVIRONMENT_CLEANUP_TIMEOUT_SECONDS = 20.0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WebSocketSubscriptionRecord:
|
||||
connection_index: int
|
||||
correlation_id: str
|
||||
symbols: tuple[str, ...]
|
||||
|
||||
|
||||
class LoopbackTradeWebSocketServer:
|
||||
"""Управляемый локальный Dzengi-подобный WebSocket endpoint."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
events: list[str] | None = None,
|
||||
auto_ack: bool = True,
|
||||
) -> None:
|
||||
self._events = events if events is not None else []
|
||||
self._auto_ack = auto_ack
|
||||
self._server: Server | None = None
|
||||
self._host = "127.0.0.1"
|
||||
self._port: int | None = None
|
||||
self._connections: list[ServerConnection] = []
|
||||
self._handler_tasks: set[asyncio.Task[Any]] = set()
|
||||
self._subscriptions: list[WebSocketSubscriptionRecord] = []
|
||||
self._received_documents: list[dict[str, Any]] = []
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
if self._port is None:
|
||||
raise RuntimeError("WebSocket server is not started.")
|
||||
|
||||
return f"ws://{self._host}:{self._port}"
|
||||
|
||||
@property
|
||||
def port(self) -> int:
|
||||
if self._port is None:
|
||||
raise RuntimeError("WebSocket server is not started.")
|
||||
|
||||
return self._port
|
||||
|
||||
@property
|
||||
def connection_count(self) -> int:
|
||||
return len(self._connections)
|
||||
|
||||
@property
|
||||
def active_handler_count(self) -> int:
|
||||
return sum(
|
||||
not task.done()
|
||||
for task in self._handler_tasks
|
||||
)
|
||||
|
||||
@property
|
||||
def active_connection_count(self) -> int:
|
||||
return sum(
|
||||
connection.state is State.OPEN
|
||||
for connection in self._connections
|
||||
)
|
||||
|
||||
@property
|
||||
def subscriptions(
|
||||
self,
|
||||
) -> tuple[WebSocketSubscriptionRecord, ...]:
|
||||
return tuple(self._subscriptions)
|
||||
|
||||
@property
|
||||
def received_documents(
|
||||
self,
|
||||
) -> tuple[dict[str, Any], ...]:
|
||||
return tuple(self._received_documents)
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._server is not None:
|
||||
raise RuntimeError("WebSocket server is already started.")
|
||||
|
||||
server = await serve(
|
||||
self._handle_connection,
|
||||
self._host,
|
||||
0,
|
||||
subprotocols=(Subprotocol("json"),),
|
||||
ping_interval=None,
|
||||
close_timeout=0.2,
|
||||
)
|
||||
self._server = server
|
||||
|
||||
sockets = tuple(server.sockets)
|
||||
|
||||
if not sockets:
|
||||
server.close()
|
||||
await server.wait_closed()
|
||||
self._server = None
|
||||
raise RuntimeError(
|
||||
"WebSocket server did not expose a listening socket."
|
||||
)
|
||||
|
||||
socket = sockets[0]
|
||||
address = socket.getsockname()
|
||||
self._port = int(address[1])
|
||||
|
||||
async def stop(self) -> None:
|
||||
server = self._server
|
||||
|
||||
if server is None:
|
||||
return
|
||||
|
||||
self._server = None
|
||||
self._port = None
|
||||
server.close()
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
server.wait_closed(),
|
||||
timeout=RESOURCE_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
except TimeoutError:
|
||||
for connection in tuple(self._connections):
|
||||
if connection.state is State.OPEN:
|
||||
connection.transport.abort()
|
||||
|
||||
await asyncio.wait_for(
|
||||
server.wait_closed(),
|
||||
timeout=RESOURCE_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
if self.active_handler_count:
|
||||
raise AssertionError(
|
||||
"WebSocket server handlers did not stop."
|
||||
)
|
||||
|
||||
async def wait_for_connections(
|
||||
self,
|
||||
count: int,
|
||||
*,
|
||||
timeout_seconds: float = 3.0,
|
||||
) -> None:
|
||||
await wait_until(
|
||||
lambda: self.connection_count >= count,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
|
||||
async def wait_for_subscriptions(
|
||||
self,
|
||||
count: int,
|
||||
*,
|
||||
timeout_seconds: float = 3.0,
|
||||
) -> None:
|
||||
await wait_until(
|
||||
lambda: len(self._subscriptions) >= count,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
|
||||
async def send_raw(
|
||||
self,
|
||||
connection_index: int,
|
||||
message: str | bytes,
|
||||
) -> None:
|
||||
await self._connections[connection_index].send(message)
|
||||
|
||||
async def send_trade(
|
||||
self,
|
||||
connection_index: int,
|
||||
*,
|
||||
symbol: str,
|
||||
trade_id: int,
|
||||
timestamp_ms: int,
|
||||
price: str = "64555.55",
|
||||
quantity: str = "0.002",
|
||||
buyer: bool = True,
|
||||
) -> None:
|
||||
await self.send_raw(
|
||||
connection_index,
|
||||
json.dumps(
|
||||
{
|
||||
"status": "OK",
|
||||
"destination": "internal.trade",
|
||||
"payload": {
|
||||
"id": trade_id,
|
||||
"price": price,
|
||||
"size": quantity,
|
||||
"ts": timestamp_ms,
|
||||
"symbol": symbol,
|
||||
"buyer": buyer,
|
||||
"orderId": f"order-{trade_id}",
|
||||
},
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
async def abort_connection(
|
||||
self,
|
||||
connection_index: int,
|
||||
) -> None:
|
||||
connection = self._connections[connection_index]
|
||||
connection.transport.abort()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
async def close_connection(
|
||||
self,
|
||||
connection_index: int,
|
||||
*,
|
||||
code: int = 1012,
|
||||
reason: str = "integration test reconnect",
|
||||
) -> None:
|
||||
await self._connections[connection_index].close(
|
||||
code=code,
|
||||
reason=reason,
|
||||
)
|
||||
|
||||
async def _handle_connection(
|
||||
self,
|
||||
connection: ServerConnection,
|
||||
) -> None:
|
||||
task = asyncio.current_task()
|
||||
|
||||
if task is None:
|
||||
raise RuntimeError("WebSocket handler has no asyncio task.")
|
||||
|
||||
self._handler_tasks.add(task)
|
||||
connection_index = len(self._connections)
|
||||
self._connections.append(connection)
|
||||
self._events.append(f"ws.connect:{connection_index}")
|
||||
|
||||
try:
|
||||
async for raw_message in connection:
|
||||
if not isinstance(raw_message, str):
|
||||
continue
|
||||
|
||||
try:
|
||||
document = json.loads(raw_message)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
if not isinstance(document, dict):
|
||||
continue
|
||||
|
||||
self._received_documents.append(document)
|
||||
|
||||
if document.get("destination") != "trades.subscribe":
|
||||
continue
|
||||
|
||||
correlation_id = document.get("correlationId")
|
||||
payload = document.get("payload")
|
||||
|
||||
if (
|
||||
not isinstance(correlation_id, str)
|
||||
or not isinstance(payload, dict)
|
||||
):
|
||||
continue
|
||||
|
||||
raw_symbols = payload.get("symbols")
|
||||
|
||||
if not isinstance(raw_symbols, list) or not all(
|
||||
isinstance(symbol, str)
|
||||
for symbol in raw_symbols
|
||||
):
|
||||
continue
|
||||
|
||||
record = WebSocketSubscriptionRecord(
|
||||
connection_index=connection_index,
|
||||
correlation_id=correlation_id,
|
||||
symbols=tuple(raw_symbols),
|
||||
)
|
||||
self._subscriptions.append(record)
|
||||
self._events.append(
|
||||
f"ws.subscribe:{connection_index}"
|
||||
)
|
||||
|
||||
if self._auto_ack:
|
||||
await connection.send(
|
||||
json.dumps(
|
||||
{
|
||||
"correlationId": correlation_id,
|
||||
"destination": "trades.subscribe",
|
||||
"status": "OK",
|
||||
}
|
||||
)
|
||||
)
|
||||
except ConnectionClosed:
|
||||
pass
|
||||
finally:
|
||||
self._events.append(
|
||||
f"ws.disconnect:{connection_index}"
|
||||
)
|
||||
self._handler_tasks.discard(task)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LoopbackHttpResponse:
|
||||
body: object
|
||||
status: int = 200
|
||||
release: threading.Event | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LoopbackHttpRequest:
|
||||
path: str
|
||||
query: dict[str, tuple[str, ...]]
|
||||
|
||||
|
||||
class _LoopbackThreadingHttpServer(ThreadingHTTPServer):
|
||||
daemon_threads = False
|
||||
block_on_close = True
|
||||
|
||||
|
||||
class LoopbackTradeRestServer:
|
||||
"""Локальный HTTP endpoint для настоящего ExchangeRestClient."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
responses: tuple[LoopbackHttpResponse, ...] = (),
|
||||
events: list[str] | None = None,
|
||||
) -> None:
|
||||
self._events = events if events is not None else []
|
||||
self._responses = deque(responses)
|
||||
self._requests: list[LoopbackHttpRequest] = []
|
||||
self._lock = threading.Lock()
|
||||
self._release_events = {
|
||||
response.release
|
||||
for response in responses
|
||||
if response.release is not None
|
||||
}
|
||||
self._server: _LoopbackThreadingHttpServer | None = None
|
||||
self._thread: threading.Thread | None = None
|
||||
self._host = "127.0.0.1"
|
||||
self._port: int | None = None
|
||||
|
||||
@property
|
||||
def base_url(self) -> str:
|
||||
if self._port is None:
|
||||
raise RuntimeError("HTTP server is not started.")
|
||||
|
||||
return f"http://{self._host}:{self._port}"
|
||||
|
||||
@property
|
||||
def request_count(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._requests)
|
||||
|
||||
@property
|
||||
def requests(self) -> tuple[LoopbackHttpRequest, ...]:
|
||||
with self._lock:
|
||||
return tuple(self._requests)
|
||||
|
||||
@property
|
||||
def thread_is_alive(self) -> bool:
|
||||
thread = self._thread
|
||||
return thread is not None and thread.is_alive()
|
||||
|
||||
def start(self) -> None:
|
||||
if self._server is not None:
|
||||
raise RuntimeError("HTTP server is already started.")
|
||||
|
||||
controller = self
|
||||
|
||||
class RequestHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_GET(self) -> None: # noqa: N802
|
||||
controller._handle_get(self)
|
||||
|
||||
def log_message(
|
||||
self,
|
||||
format: str,
|
||||
*args: object,
|
||||
) -> None:
|
||||
del format, args
|
||||
|
||||
server = _LoopbackThreadingHttpServer(
|
||||
(self._host, 0),
|
||||
RequestHandler,
|
||||
)
|
||||
self._server = server
|
||||
self._port = int(server.server_address[1])
|
||||
self._thread = threading.Thread(
|
||||
target=lambda: server.serve_forever(
|
||||
poll_interval=0.05,
|
||||
),
|
||||
name="loopback-trade-rest-server",
|
||||
)
|
||||
|
||||
try:
|
||||
self._thread.start()
|
||||
except BaseException:
|
||||
server.server_close()
|
||||
self._server = None
|
||||
self._thread = None
|
||||
self._port = None
|
||||
raise
|
||||
|
||||
def stop(self) -> None:
|
||||
server = self._server
|
||||
thread = self._thread
|
||||
|
||||
if server is None:
|
||||
return
|
||||
|
||||
self._server = None
|
||||
self._thread = None
|
||||
self._port = None
|
||||
|
||||
for release_event in self._release_events:
|
||||
release_event.set()
|
||||
|
||||
shutdown_error: BaseException | None = None
|
||||
|
||||
try:
|
||||
server.shutdown()
|
||||
except BaseException as error:
|
||||
shutdown_error = error
|
||||
finally:
|
||||
try:
|
||||
server.server_close()
|
||||
except BaseException as error:
|
||||
if shutdown_error is None:
|
||||
shutdown_error = error
|
||||
else:
|
||||
shutdown_error.add_note(
|
||||
"Loopback HTTP server_close also failed: "
|
||||
f"{type(error).__name__}."
|
||||
)
|
||||
|
||||
if thread is not None:
|
||||
thread.join(timeout=3)
|
||||
|
||||
if thread.is_alive():
|
||||
thread_error = AssertionError(
|
||||
"Loopback HTTP server thread did not stop."
|
||||
)
|
||||
|
||||
if shutdown_error is None:
|
||||
shutdown_error = thread_error
|
||||
else:
|
||||
shutdown_error.add_note(str(thread_error))
|
||||
|
||||
if shutdown_error is not None:
|
||||
raise shutdown_error
|
||||
|
||||
async def wait_for_requests(
|
||||
self,
|
||||
count: int,
|
||||
*,
|
||||
timeout_seconds: float = 3.0,
|
||||
) -> None:
|
||||
await wait_until(
|
||||
lambda: self.request_count >= count,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
|
||||
def _handle_get(
|
||||
self,
|
||||
handler: BaseHTTPRequestHandler,
|
||||
) -> None:
|
||||
parsed = urlsplit(handler.path)
|
||||
query = {
|
||||
key: tuple(values)
|
||||
for key, values in parse_qs(
|
||||
parsed.query,
|
||||
keep_blank_values=True,
|
||||
).items()
|
||||
}
|
||||
request = LoopbackHttpRequest(
|
||||
path=parsed.path,
|
||||
query=query,
|
||||
)
|
||||
|
||||
self._events.append("rest.request")
|
||||
|
||||
with self._lock:
|
||||
self._requests.append(request)
|
||||
response = (
|
||||
self._responses.popleft()
|
||||
if self._responses
|
||||
else LoopbackHttpResponse(body=[])
|
||||
)
|
||||
|
||||
if response.release is not None:
|
||||
response.release.wait()
|
||||
|
||||
if isinstance(response.body, bytes):
|
||||
body = response.body
|
||||
elif isinstance(response.body, str):
|
||||
body = response.body.encode("utf-8")
|
||||
else:
|
||||
body = json.dumps(response.body).encode("utf-8")
|
||||
|
||||
try:
|
||||
handler.send_response(response.status)
|
||||
handler.send_header(
|
||||
"Content-Type",
|
||||
"application/json",
|
||||
)
|
||||
handler.send_header(
|
||||
"Content-Length",
|
||||
str(len(body)),
|
||||
)
|
||||
handler.send_header(
|
||||
"Connection",
|
||||
"close",
|
||||
)
|
||||
handler.end_headers()
|
||||
handler.wfile.write(body)
|
||||
except (BrokenPipeError, ConnectionResetError):
|
||||
pass
|
||||
|
||||
|
||||
class LoopbackTcpFaultProxy:
|
||||
"""TCP relay с управляемой потерей client → server traffic."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
target_host: str,
|
||||
target_port: int,
|
||||
) -> None:
|
||||
self._target_host = target_host
|
||||
self._target_port = target_port
|
||||
self._server: asyncio.Server | None = None
|
||||
self._host = "127.0.0.1"
|
||||
self._port: int | None = None
|
||||
self._connection_count = 0
|
||||
self._blackholed_connections: set[int] = set()
|
||||
self._handler_tasks: set[asyncio.Task[Any]] = set()
|
||||
self._relay_tasks: set[asyncio.Task[Any]] = set()
|
||||
self._writers: set[asyncio.StreamWriter] = set()
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
if self._port is None:
|
||||
raise RuntimeError("TCP fault proxy is not started.")
|
||||
|
||||
return f"ws://{self._host}:{self._port}"
|
||||
|
||||
@property
|
||||
def connection_count(self) -> int:
|
||||
return self._connection_count
|
||||
|
||||
@property
|
||||
def active_handler_count(self) -> int:
|
||||
return sum(
|
||||
not task.done()
|
||||
for task in self._handler_tasks
|
||||
)
|
||||
|
||||
@property
|
||||
def active_relay_count(self) -> int:
|
||||
return sum(
|
||||
not task.done()
|
||||
for task in self._relay_tasks
|
||||
)
|
||||
|
||||
@property
|
||||
def tracked_writer_count(self) -> int:
|
||||
return len(self._writers)
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._server is not None:
|
||||
raise RuntimeError("TCP fault proxy is already started.")
|
||||
|
||||
server = await asyncio.start_server(
|
||||
self._handle_client,
|
||||
self._host,
|
||||
0,
|
||||
)
|
||||
self._server = server
|
||||
sockets = server.sockets
|
||||
|
||||
if not sockets:
|
||||
server.close()
|
||||
await server.wait_closed()
|
||||
self._server = None
|
||||
raise RuntimeError(
|
||||
"TCP fault proxy did not expose a listening socket."
|
||||
)
|
||||
|
||||
socket = sockets[0]
|
||||
address = socket.getsockname()
|
||||
self._port = int(address[1])
|
||||
|
||||
def blackhole_client_to_server(
|
||||
self,
|
||||
connection_index: int,
|
||||
) -> None:
|
||||
self._blackholed_connections.add(connection_index)
|
||||
|
||||
async def stop(self) -> None:
|
||||
server = self._server
|
||||
|
||||
if server is None:
|
||||
return
|
||||
|
||||
self._server = None
|
||||
self._port = None
|
||||
server.close()
|
||||
|
||||
tasks = tuple(
|
||||
task
|
||||
for task in (
|
||||
*self._relay_tasks,
|
||||
*self._handler_tasks,
|
||||
)
|
||||
if not task.done()
|
||||
)
|
||||
writers = tuple(self._writers)
|
||||
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
|
||||
for writer in writers:
|
||||
writer.close()
|
||||
|
||||
async def wait_for_writer(
|
||||
writer: asyncio.StreamWriter,
|
||||
) -> None:
|
||||
try:
|
||||
await writer.wait_closed()
|
||||
except (BrokenPipeError, ConnectionResetError):
|
||||
pass
|
||||
|
||||
await asyncio.wait_for(
|
||||
server.wait_closed(),
|
||||
timeout=RESOURCE_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
if writers:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(
|
||||
*(
|
||||
wait_for_writer(writer)
|
||||
for writer in writers
|
||||
),
|
||||
),
|
||||
timeout=RESOURCE_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
except TimeoutError:
|
||||
for writer in writers:
|
||||
writer.transport.abort()
|
||||
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(
|
||||
*(
|
||||
wait_for_writer(writer)
|
||||
for writer in writers
|
||||
),
|
||||
),
|
||||
timeout=RESOURCE_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
for writer in writers:
|
||||
self._writers.discard(writer)
|
||||
|
||||
if tasks:
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(
|
||||
*tasks,
|
||||
return_exceptions=True,
|
||||
),
|
||||
timeout=RESOURCE_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
if (
|
||||
self.active_handler_count
|
||||
or self.active_relay_count
|
||||
or self.tracked_writer_count
|
||||
):
|
||||
raise AssertionError(
|
||||
"TCP fault proxy resources did not stop."
|
||||
)
|
||||
|
||||
async def _handle_client(
|
||||
self,
|
||||
client_reader: asyncio.StreamReader,
|
||||
client_writer: asyncio.StreamWriter,
|
||||
) -> None:
|
||||
task = asyncio.current_task()
|
||||
|
||||
if task is None:
|
||||
raise RuntimeError("TCP proxy handler has no asyncio task.")
|
||||
|
||||
self._handler_tasks.add(task)
|
||||
connection_index = self._connection_count
|
||||
self._connection_count += 1
|
||||
|
||||
server_writer: asyncio.StreamWriter | None = None
|
||||
relay_tasks: set[asyncio.Task[None]] = set()
|
||||
|
||||
try:
|
||||
server_reader, server_writer = await asyncio.open_connection(
|
||||
self._target_host,
|
||||
self._target_port,
|
||||
)
|
||||
self._writers.update(
|
||||
{
|
||||
client_writer,
|
||||
server_writer,
|
||||
}
|
||||
)
|
||||
|
||||
upstream = asyncio.create_task(
|
||||
self._relay(
|
||||
client_reader,
|
||||
server_writer,
|
||||
should_drop=lambda: (
|
||||
connection_index
|
||||
in self._blackholed_connections
|
||||
),
|
||||
),
|
||||
name="loopback-proxy-upstream",
|
||||
)
|
||||
downstream = asyncio.create_task(
|
||||
self._relay(
|
||||
server_reader,
|
||||
client_writer,
|
||||
should_drop=lambda: False,
|
||||
),
|
||||
name="loopback-proxy-downstream",
|
||||
)
|
||||
relay_tasks = {
|
||||
upstream,
|
||||
downstream,
|
||||
}
|
||||
self._relay_tasks.update(relay_tasks)
|
||||
|
||||
done, pending = await asyncio.wait(
|
||||
relay_tasks,
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
|
||||
for relay_task in pending:
|
||||
relay_task.cancel()
|
||||
|
||||
await asyncio.gather(
|
||||
*done,
|
||||
*pending,
|
||||
return_exceptions=True,
|
||||
)
|
||||
finally:
|
||||
self._relay_tasks.difference_update(relay_tasks)
|
||||
|
||||
for writer in (
|
||||
client_writer,
|
||||
server_writer,
|
||||
):
|
||||
if writer is None:
|
||||
continue
|
||||
|
||||
self._writers.discard(writer)
|
||||
writer.close()
|
||||
|
||||
try:
|
||||
await writer.wait_closed()
|
||||
except (BrokenPipeError, ConnectionResetError):
|
||||
pass
|
||||
|
||||
self._handler_tasks.discard(task)
|
||||
|
||||
@staticmethod
|
||||
async def _relay(
|
||||
reader: asyncio.StreamReader,
|
||||
writer: asyncio.StreamWriter,
|
||||
*,
|
||||
should_drop: Callable[[], bool],
|
||||
) -> None:
|
||||
while True:
|
||||
data = await reader.read(65_536)
|
||||
|
||||
if not data:
|
||||
return
|
||||
|
||||
if should_drop():
|
||||
continue
|
||||
|
||||
writer.write(data)
|
||||
await writer.drain()
|
||||
|
||||
|
||||
class LoopbackTradeEnvironment:
|
||||
"""Exception-safe владелец локальных сетевых ресурсов теста."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
websocket: LoopbackTradeWebSocketServer,
|
||||
rest: LoopbackTradeRestServer,
|
||||
use_fault_proxy: bool = False,
|
||||
) -> None:
|
||||
self.websocket = websocket
|
||||
self.rest = rest
|
||||
self.use_fault_proxy = use_fault_proxy
|
||||
self.proxy: LoopbackTcpFaultProxy | None = None
|
||||
self._exit_stack: AsyncExitStack | None = None
|
||||
|
||||
@property
|
||||
def websocket_url(self) -> str:
|
||||
proxy = self.proxy
|
||||
|
||||
if proxy is not None:
|
||||
return proxy.url
|
||||
|
||||
return self.websocket.url
|
||||
|
||||
async def __aenter__(self) -> LoopbackTradeEnvironment:
|
||||
if self._exit_stack is not None:
|
||||
raise RuntimeError(
|
||||
"Loopback environment is already active."
|
||||
)
|
||||
|
||||
stack = AsyncExitStack()
|
||||
await stack.__aenter__()
|
||||
self._exit_stack = stack
|
||||
|
||||
try:
|
||||
self.rest.start()
|
||||
stack.push_async_callback(self._stop_rest)
|
||||
|
||||
await self.websocket.start()
|
||||
stack.push_async_callback(self._stop_websocket)
|
||||
|
||||
if self.use_fault_proxy:
|
||||
proxy = LoopbackTcpFaultProxy(
|
||||
target_host="127.0.0.1",
|
||||
target_port=self.websocket.port,
|
||||
)
|
||||
self.proxy = proxy
|
||||
await proxy.start()
|
||||
stack.push_async_callback(self._stop_proxy)
|
||||
except BaseException:
|
||||
self._exit_stack = None
|
||||
await stack.aclose()
|
||||
raise
|
||||
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> bool | None:
|
||||
stack = self._exit_stack
|
||||
self._exit_stack = None
|
||||
|
||||
if stack is None:
|
||||
return None
|
||||
|
||||
return await stack.__aexit__(
|
||||
exc_type,
|
||||
exc_value,
|
||||
traceback,
|
||||
)
|
||||
|
||||
async def _stop_proxy(self) -> None:
|
||||
proxy = self.proxy
|
||||
|
||||
if proxy is None:
|
||||
return
|
||||
|
||||
await asyncio.wait_for(
|
||||
proxy.stop(),
|
||||
timeout=ENVIRONMENT_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
async def _stop_websocket(self) -> None:
|
||||
await asyncio.wait_for(
|
||||
self.websocket.stop(),
|
||||
timeout=ENVIRONMENT_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
async def _stop_rest(self) -> None:
|
||||
await asyncio.wait_for(
|
||||
asyncio.to_thread(
|
||||
self.rest.stop,
|
||||
),
|
||||
timeout=ENVIRONMENT_CLEANUP_TIMEOUT_SECONDS,
|
||||
)
|
||||
@@ -0,0 +1,782 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.integration.market_data.acquisition.runtime import (
|
||||
loopback_trade_exchange,
|
||||
)
|
||||
from tests.integration.market_data.acquisition.runtime.loopback_trade_exchange import (
|
||||
LoopbackHttpResponse,
|
||||
LoopbackTcpFaultProxy,
|
||||
LoopbackTradeEnvironment,
|
||||
LoopbackTradeRestServer,
|
||||
LoopbackTradeWebSocketServer,
|
||||
wait_until,
|
||||
)
|
||||
from src.market_data.acquisition.exceptions import (
|
||||
TradeTransportError,
|
||||
WebSocketMessageDecodeError,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_reconnect_recovery_coordinator import (
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.trade_stream_production_runtime import (
|
||||
TradeStreamProductionRuntime,
|
||||
TradeStreamProductionRuntimeState,
|
||||
)
|
||||
from tests.support.trade_stream_runtime import (
|
||||
SYMBOL,
|
||||
assert_no_owned_tasks,
|
||||
build_runtime,
|
||||
reconnect_coordinator_from,
|
||||
run_scenario,
|
||||
start_runtime,
|
||||
state_store_from,
|
||||
stop_runtime,
|
||||
)
|
||||
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
|
||||
def make_recovered_trade(
|
||||
*,
|
||||
trade_id: int,
|
||||
timestamp_ms: int,
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"a": trade_id,
|
||||
"p": "64555.56",
|
||||
"q": "0.003",
|
||||
"T": timestamp_ms,
|
||||
"m": False,
|
||||
}
|
||||
|
||||
|
||||
class HangingStartupRuntime:
|
||||
def __init__(self) -> None:
|
||||
self._state = TradeStreamProductionRuntimeState.STARTING
|
||||
self.stop_calls = 0
|
||||
self.released = asyncio.Event()
|
||||
self.finished = asyncio.Event()
|
||||
|
||||
@property
|
||||
def state(self) -> TradeStreamProductionRuntimeState:
|
||||
return self._state
|
||||
|
||||
async def run(self) -> None:
|
||||
try:
|
||||
await self.released.wait()
|
||||
finally:
|
||||
self.finished.set()
|
||||
|
||||
async def stop(self) -> None:
|
||||
self.stop_calls += 1
|
||||
self.released.set()
|
||||
self._state = TradeStreamProductionRuntimeState.STOPPED
|
||||
|
||||
|
||||
def test_start_runtime_timeout_cleans_task_it_created() -> None:
|
||||
async def scenario() -> None:
|
||||
runtime = HangingStartupRuntime()
|
||||
|
||||
with pytest.raises(TimeoutError):
|
||||
await start_runtime(
|
||||
runtime,
|
||||
timeout_seconds=0.01,
|
||||
)
|
||||
|
||||
assert runtime.stop_calls == 1
|
||||
assert runtime.finished.is_set()
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_environment_cleans_rest_after_websocket_setup_failure(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def scenario() -> None:
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer()
|
||||
|
||||
async def broken_start() -> None:
|
||||
raise RuntimeError("websocket setup failed")
|
||||
|
||||
monkeypatch.setattr(
|
||||
websocket,
|
||||
"start",
|
||||
broken_start,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="websocket setup failed",
|
||||
):
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
):
|
||||
raise AssertionError(
|
||||
"environment entered after setup failure"
|
||||
)
|
||||
|
||||
assert rest.thread_is_alive is False
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_environment_continues_after_websocket_cleanup_error(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def scenario() -> None:
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer()
|
||||
original_stop = websocket.stop
|
||||
|
||||
async def broken_stop() -> None:
|
||||
await original_stop()
|
||||
raise RuntimeError("websocket cleanup failed")
|
||||
|
||||
monkeypatch.setattr(
|
||||
websocket,
|
||||
"stop",
|
||||
broken_stop,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="websocket cleanup failed",
|
||||
):
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
):
|
||||
assert rest.thread_is_alive is True
|
||||
|
||||
assert websocket.active_handler_count == 0
|
||||
assert rest.thread_is_alive is False
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_proxy_waits_for_aborted_writer_before_removing_tracking(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class RecordingServer:
|
||||
def __init__(self) -> None:
|
||||
self.close_calls = 0
|
||||
self.wait_closed_calls = 0
|
||||
|
||||
def close(self) -> None:
|
||||
self.close_calls += 1
|
||||
|
||||
async def wait_closed(self) -> None:
|
||||
self.wait_closed_calls += 1
|
||||
|
||||
class HangingWriter:
|
||||
def __init__(self) -> None:
|
||||
self.close_calls = 0
|
||||
self.wait_closed_calls = 0
|
||||
self.aborted = False
|
||||
self.transport = self
|
||||
|
||||
def close(self) -> None:
|
||||
self.close_calls += 1
|
||||
|
||||
async def wait_closed(self) -> None:
|
||||
self.wait_closed_calls += 1
|
||||
|
||||
if not self.aborted:
|
||||
await asyncio.Event().wait()
|
||||
|
||||
def abort(self) -> None:
|
||||
self.aborted = True
|
||||
|
||||
async def scenario() -> None:
|
||||
proxy = LoopbackTcpFaultProxy(
|
||||
target_host="127.0.0.1",
|
||||
target_port=1,
|
||||
)
|
||||
proxy_graph: Any = proxy
|
||||
server = RecordingServer()
|
||||
writer = HangingWriter()
|
||||
proxy_graph._server = server
|
||||
proxy_graph._port = 1
|
||||
proxy_graph._writers.add(writer)
|
||||
|
||||
monkeypatch.setattr(
|
||||
loopback_trade_exchange,
|
||||
"RESOURCE_CLEANUP_TIMEOUT_SECONDS",
|
||||
0.01,
|
||||
)
|
||||
|
||||
await proxy.stop()
|
||||
|
||||
assert server.close_calls == 1
|
||||
assert server.wait_closed_calls == 1
|
||||
assert writer.close_calls == 1
|
||||
assert writer.wait_closed_calls == 2
|
||||
assert writer.aborted is True
|
||||
assert proxy.tracked_writer_count == 0
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_real_loopback_websocket_updates_shared_checkpoint() -> None:
|
||||
async def scenario() -> None:
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer()
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
) as environment:
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
|
||||
timestamp_ms = time.time_ns() // 1_000_000
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=100,
|
||||
timestamp_ms=timestamp_ms,
|
||||
)
|
||||
|
||||
state_store = state_store_from(runtime)
|
||||
await wait_until(
|
||||
lambda: (
|
||||
state_store.contains(SYMBOL)
|
||||
and state_store.get(SYMBOL).last_trade_id == 100
|
||||
)
|
||||
)
|
||||
|
||||
assert websocket.subscriptions[0].symbols == (SYMBOL,)
|
||||
assert rest.request_count == 0
|
||||
assert runtime_task.done() is False
|
||||
finally:
|
||||
if runtime_task is not None:
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
assert websocket.active_handler_count == 0
|
||||
assert rest.thread_is_alive is False
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_reconnect_restores_then_recovers_before_buffered_live() -> None:
|
||||
async def scenario() -> None:
|
||||
events: list[str] = []
|
||||
recovery_release = threading.Event()
|
||||
base_time_ms = time.time_ns() // 1_000_000 - 1_000
|
||||
websocket = LoopbackTradeWebSocketServer(
|
||||
events=events,
|
||||
)
|
||||
rest = LoopbackTradeRestServer(
|
||||
responses=(
|
||||
LoopbackHttpResponse(
|
||||
body=[
|
||||
make_recovered_trade(
|
||||
trade_id=201,
|
||||
timestamp_ms=base_time_ms + 100,
|
||||
),
|
||||
],
|
||||
release=recovery_release,
|
||||
),
|
||||
),
|
||||
events=events,
|
||||
)
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
) as environment:
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=200,
|
||||
timestamp_ms=base_time_ms,
|
||||
)
|
||||
|
||||
state_store = state_store_from(runtime)
|
||||
await wait_until(
|
||||
lambda: (
|
||||
state_store.contains(SYMBOL)
|
||||
and state_store.get(SYMBOL).last_trade_id == 200
|
||||
)
|
||||
)
|
||||
|
||||
await websocket.abort_connection(0)
|
||||
await websocket.wait_for_subscriptions(2)
|
||||
await rest.wait_for_requests(1)
|
||||
|
||||
await websocket.send_trade(
|
||||
1,
|
||||
symbol=SYMBOL,
|
||||
trade_id=202,
|
||||
timestamp_ms=base_time_ms + 200,
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert state_store.get(SYMBOL).last_trade_id == 200
|
||||
assert events.index("ws.subscribe:1") < events.index(
|
||||
"rest.request"
|
||||
)
|
||||
|
||||
recovery_release.set()
|
||||
await wait_until(
|
||||
lambda: state_store.get(SYMBOL).last_trade_id == 202,
|
||||
)
|
||||
|
||||
assert runtime_task.done() is False
|
||||
assert rest.requests[0].path == "/api/v1/aggTrades"
|
||||
assert rest.requests[0].query["symbol"] == (SYMBOL,)
|
||||
assert (
|
||||
reconnect_coordinator_from(runtime).generation
|
||||
== 1
|
||||
)
|
||||
finally:
|
||||
recovery_release.set()
|
||||
|
||||
if runtime_task is not None:
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_repeated_network_disconnects_advance_one_generation_each() -> None:
|
||||
async def scenario() -> None:
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer()
|
||||
base_time_ms = time.time_ns() // 1_000_000 - 1_000
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
) as environment:
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=300,
|
||||
timestamp_ms=base_time_ms,
|
||||
)
|
||||
|
||||
state_store = state_store_from(runtime)
|
||||
await wait_until(
|
||||
lambda: (
|
||||
state_store.contains(SYMBOL)
|
||||
and state_store.get(SYMBOL).last_trade_id == 300
|
||||
)
|
||||
)
|
||||
|
||||
for generation in range(1, 4):
|
||||
previous_connection = generation - 1
|
||||
|
||||
if generation == 1:
|
||||
await websocket.close_connection(
|
||||
previous_connection,
|
||||
)
|
||||
else:
|
||||
await websocket.abort_connection(
|
||||
previous_connection,
|
||||
)
|
||||
|
||||
await websocket.wait_for_subscriptions(
|
||||
generation + 1,
|
||||
)
|
||||
await rest.wait_for_requests(generation)
|
||||
|
||||
trade_id = 300 + generation
|
||||
await websocket.send_trade(
|
||||
generation,
|
||||
symbol=SYMBOL,
|
||||
trade_id=trade_id,
|
||||
timestamp_ms=(
|
||||
base_time_ms + generation * 100
|
||||
),
|
||||
)
|
||||
await wait_until(
|
||||
lambda trade_id=trade_id: (
|
||||
state_store.get(SYMBOL).last_trade_id
|
||||
== trade_id
|
||||
)
|
||||
)
|
||||
|
||||
assert websocket.connection_count == 4
|
||||
assert len(websocket.subscriptions) == 4
|
||||
assert rest.request_count == 3
|
||||
assert (
|
||||
reconnect_coordinator_from(runtime).generation
|
||||
== 3
|
||||
)
|
||||
assert runtime_task.done() is False
|
||||
finally:
|
||||
if runtime_task is not None:
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_http_recovery_error_rejects_buffered_live_and_cleans_up() -> None:
|
||||
async def scenario() -> None:
|
||||
recovery_release = threading.Event()
|
||||
base_time_ms = time.time_ns() // 1_000_000 - 1_000
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer(
|
||||
responses=(
|
||||
LoopbackHttpResponse(
|
||||
body={
|
||||
"error": "recovery unavailable",
|
||||
},
|
||||
status=503,
|
||||
release=recovery_release,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
) as environment:
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=400,
|
||||
timestamp_ms=base_time_ms,
|
||||
)
|
||||
|
||||
state_store = state_store_from(runtime)
|
||||
await wait_until(
|
||||
lambda: (
|
||||
state_store.contains(SYMBOL)
|
||||
and state_store.get(SYMBOL).last_trade_id == 400
|
||||
)
|
||||
)
|
||||
|
||||
await websocket.abort_connection(0)
|
||||
await websocket.wait_for_subscriptions(2)
|
||||
await rest.wait_for_requests(1)
|
||||
await websocket.send_trade(
|
||||
1,
|
||||
symbol=SYMBOL,
|
||||
trade_id=402,
|
||||
timestamp_ms=base_time_ms + 200,
|
||||
)
|
||||
recovery_release.set()
|
||||
|
||||
with pytest.raises(
|
||||
TradeTransportError,
|
||||
match="Не удалось получить агрегированные сделки",
|
||||
):
|
||||
await asyncio.wait_for(
|
||||
runtime_task,
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
assert state_store.get(SYMBOL).last_trade_id == 400
|
||||
assert (
|
||||
reconnect_coordinator_from(runtime)
|
||||
.live_processing_gate
|
||||
.failed
|
||||
)
|
||||
assert (
|
||||
runtime.state
|
||||
is TradeStreamProductionRuntimeState.FAILED
|
||||
)
|
||||
finally:
|
||||
recovery_release.set()
|
||||
|
||||
if (
|
||||
runtime_task is not None
|
||||
and not runtime_task.done()
|
||||
):
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_invalid_websocket_json_is_terminal_and_releases_resources() -> None:
|
||||
async def scenario() -> None:
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer()
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
) as environment:
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
await websocket.send_raw(0, "{not-json")
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketMessageDecodeError,
|
||||
):
|
||||
await asyncio.wait_for(
|
||||
runtime_task,
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
assert (
|
||||
runtime.state
|
||||
is TradeStreamProductionRuntimeState.FAILED
|
||||
)
|
||||
assert rest.request_count == 0
|
||||
finally:
|
||||
if (
|
||||
runtime_task is not None
|
||||
and not runtime_task.done()
|
||||
):
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_stop_waits_for_real_http_recovery_worker() -> None:
|
||||
async def scenario() -> None:
|
||||
recovery_release = threading.Event()
|
||||
base_time_ms = time.time_ns() // 1_000_000 - 1_000
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer(
|
||||
responses=(
|
||||
LoopbackHttpResponse(
|
||||
body=[
|
||||
make_recovered_trade(
|
||||
trade_id=501,
|
||||
timestamp_ms=base_time_ms + 100,
|
||||
),
|
||||
],
|
||||
release=recovery_release,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
) as environment:
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=500,
|
||||
timestamp_ms=base_time_ms,
|
||||
)
|
||||
|
||||
state_store = state_store_from(runtime)
|
||||
await wait_until(
|
||||
lambda: (
|
||||
state_store.contains(SYMBOL)
|
||||
and state_store.get(SYMBOL).last_trade_id == 500
|
||||
)
|
||||
)
|
||||
|
||||
await websocket.abort_connection(0)
|
||||
await websocket.wait_for_subscriptions(2)
|
||||
await rest.wait_for_requests(1)
|
||||
|
||||
stop_task = asyncio.create_task(runtime.stop())
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert stop_task.done() is False
|
||||
|
||||
recovery_release.set()
|
||||
await asyncio.wait_for(
|
||||
stop_task,
|
||||
timeout=3,
|
||||
)
|
||||
await asyncio.wait_for(
|
||||
runtime_task,
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
assert state_store.get(SYMBOL).last_trade_id == 501
|
||||
assert (
|
||||
runtime.state
|
||||
is TradeStreamProductionRuntimeState.STOPPED
|
||||
)
|
||||
finally:
|
||||
recovery_release.set()
|
||||
|
||||
if (
|
||||
runtime_task is not None
|
||||
and not runtime_task.done()
|
||||
):
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
|
||||
|
||||
def test_ping_timeout_and_receive_failure_share_one_reconnect(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
caller_task_names: list[str] = []
|
||||
original_reconnect = (
|
||||
RuntimeReconnectRecoveryCoordinator
|
||||
.reconnect_after_transport_failure
|
||||
)
|
||||
|
||||
async def recording_reconnect(
|
||||
self: RuntimeReconnectRecoveryCoordinator,
|
||||
*,
|
||||
observed_generation: int,
|
||||
) -> None:
|
||||
task = asyncio.current_task()
|
||||
caller_task_names.append(
|
||||
task.get_name()
|
||||
if task is not None
|
||||
else "<no-task>"
|
||||
)
|
||||
await original_reconnect(
|
||||
self,
|
||||
observed_generation=observed_generation,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
"reconnect_after_transport_failure",
|
||||
recording_reconnect,
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
base_time_ms = time.time_ns() // 1_000_000 - 1_000
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer()
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
use_fault_proxy=True,
|
||||
) as environment:
|
||||
proxy = environment.proxy
|
||||
assert proxy is not None
|
||||
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
probe_timeout_seconds=0.05,
|
||||
close_timeout_seconds=0.05,
|
||||
heartbeat_timeout_seconds=0.15,
|
||||
scheduler_interval_seconds=0.05,
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=600,
|
||||
timestamp_ms=base_time_ms,
|
||||
)
|
||||
|
||||
state_store = state_store_from(runtime)
|
||||
await wait_until(
|
||||
lambda: (
|
||||
state_store.contains(SYMBOL)
|
||||
and state_store.get(SYMBOL).last_trade_id == 600
|
||||
)
|
||||
)
|
||||
|
||||
proxy.blackhole_client_to_server(0)
|
||||
|
||||
await websocket.wait_for_subscriptions(2)
|
||||
await rest.wait_for_requests(1)
|
||||
await websocket.send_trade(
|
||||
1,
|
||||
symbol=SYMBOL,
|
||||
trade_id=601,
|
||||
timestamp_ms=base_time_ms + 100,
|
||||
)
|
||||
await wait_until(
|
||||
lambda: state_store.get(SYMBOL).last_trade_id == 601,
|
||||
)
|
||||
|
||||
assert proxy.connection_count == 2
|
||||
assert websocket.connection_count == 2
|
||||
assert len(websocket.subscriptions) == 2
|
||||
assert rest.request_count == 1
|
||||
assert (
|
||||
reconnect_coordinator_from(runtime).generation
|
||||
== 1
|
||||
)
|
||||
assert sorted(caller_task_names) == [
|
||||
"trade-stream-receive",
|
||||
"trade-stream-scheduler",
|
||||
]
|
||||
assert runtime_task.done() is False
|
||||
finally:
|
||||
if runtime_task is not None:
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
assert proxy.active_handler_count == 0
|
||||
assert proxy.active_relay_count == 0
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
Reference in New Issue
Block a user