Build 060.25: implement Production Runtime Integration
This commit is contained in:
@@ -0,0 +1,106 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.adapters.dzengi.websocket_control_message_handler import (
|
||||
DzengiWebSocketControlMessageHandler,
|
||||
)
|
||||
from src.market_data.acquisition.exceptions import (
|
||||
WebSocketControlMessageError,
|
||||
WebSocketMessageRoutingError,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_inbound_message import (
|
||||
WebSocketControlMessageHandlerProtocol,
|
||||
)
|
||||
|
||||
|
||||
CORRELATION_ID = "trade-subscription-1"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def handler() -> DzengiWebSocketControlMessageHandler:
|
||||
return DzengiWebSocketControlMessageHandler()
|
||||
|
||||
|
||||
def test_implements_public_protocol_and_uses_slots(
|
||||
handler: DzengiWebSocketControlMessageHandler,
|
||||
) -> None:
|
||||
assert isinstance(
|
||||
handler,
|
||||
WebSocketControlMessageHandlerProtocol,
|
||||
)
|
||||
assert not hasattr(handler, "__dict__")
|
||||
|
||||
|
||||
def test_accepts_matching_successful_subscription_ack(
|
||||
handler: DzengiWebSocketControlMessageHandler,
|
||||
) -> None:
|
||||
handler.handle(
|
||||
{
|
||||
"correlationId": CORRELATION_ID,
|
||||
"destination": "trades.subscribe",
|
||||
"status": "OK",
|
||||
},
|
||||
expected_correlation_id=CORRELATION_ID,
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_negative_subscription_ack(
|
||||
handler: DzengiWebSocketControlMessageHandler,
|
||||
) -> None:
|
||||
with pytest.raises(
|
||||
WebSocketControlMessageError,
|
||||
match="отклонил",
|
||||
):
|
||||
handler.handle(
|
||||
{
|
||||
"correlationId": CORRELATION_ID,
|
||||
"status": "ERROR",
|
||||
"payload": {
|
||||
"errorCode": "BAD_REQUEST",
|
||||
},
|
||||
},
|
||||
expected_correlation_id=CORRELATION_ID,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"document",
|
||||
[
|
||||
None,
|
||||
{},
|
||||
{
|
||||
"correlationId": "unknown",
|
||||
"destination": "trades.subscribe",
|
||||
"status": "OK",
|
||||
},
|
||||
{
|
||||
"correlationId": CORRELATION_ID,
|
||||
"destination": "unknown",
|
||||
"status": "OK",
|
||||
},
|
||||
{
|
||||
"correlationId": CORRELATION_ID,
|
||||
"destination": "trades.subscribe",
|
||||
},
|
||||
{
|
||||
"correlationId": CORRELATION_ID,
|
||||
"destination": "trades.subscribe",
|
||||
"status": "",
|
||||
},
|
||||
{
|
||||
"correlationId": CORRELATION_ID,
|
||||
"destination": "trades.subscribe",
|
||||
"status": None,
|
||||
},
|
||||
],
|
||||
)
|
||||
def test_rejects_unknown_or_malformed_control_messages(
|
||||
handler: DzengiWebSocketControlMessageHandler,
|
||||
document: object,
|
||||
) -> None:
|
||||
with pytest.raises(WebSocketMessageRoutingError):
|
||||
handler.handle(
|
||||
document,
|
||||
expected_correlation_id=CORRELATION_ID,
|
||||
)
|
||||
@@ -0,0 +1,137 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.adapters.dzengi.websocket_inbound_message_classifier import (
|
||||
DzengiWebSocketInboundMessageClassifier,
|
||||
)
|
||||
from src.market_data.acquisition.exceptions import (
|
||||
WebSocketMessageRoutingError,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_inbound_message import (
|
||||
WebSocketInboundMessageClassifierProtocol,
|
||||
WebSocketInboundMessageKind,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def classifier() -> DzengiWebSocketInboundMessageClassifier:
|
||||
return DzengiWebSocketInboundMessageClassifier()
|
||||
|
||||
|
||||
def test_implements_public_protocol(
|
||||
classifier: DzengiWebSocketInboundMessageClassifier,
|
||||
) -> None:
|
||||
assert isinstance(
|
||||
classifier,
|
||||
WebSocketInboundMessageClassifierProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_uses_slots(
|
||||
classifier: DzengiWebSocketInboundMessageClassifier,
|
||||
) -> None:
|
||||
assert not hasattr(
|
||||
classifier,
|
||||
"__dict__",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"document",
|
||||
[
|
||||
{
|
||||
"destination": "internal.trade",
|
||||
},
|
||||
{
|
||||
"destination": "ohlc.event",
|
||||
},
|
||||
{
|
||||
"Payload": {},
|
||||
},
|
||||
],
|
||||
)
|
||||
def test_recognizes_market_documents(
|
||||
classifier: DzengiWebSocketInboundMessageClassifier,
|
||||
document: object,
|
||||
) -> None:
|
||||
assert (
|
||||
classifier.classify(document)
|
||||
is WebSocketInboundMessageKind.MARKET
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"correlation_id",
|
||||
[
|
||||
"request-1",
|
||||
1,
|
||||
],
|
||||
)
|
||||
def test_recognizes_control_responses(
|
||||
classifier: DzengiWebSocketInboundMessageClassifier,
|
||||
correlation_id: str | int,
|
||||
) -> None:
|
||||
assert (
|
||||
classifier.classify(
|
||||
{
|
||||
"correlationId": correlation_id,
|
||||
"destination": "trades.subscribe",
|
||||
}
|
||||
)
|
||||
is WebSocketInboundMessageKind.CONTROL
|
||||
)
|
||||
|
||||
|
||||
def test_market_markers_take_priority_over_correlation_id(
|
||||
classifier: DzengiWebSocketInboundMessageClassifier,
|
||||
) -> None:
|
||||
assert (
|
||||
classifier.classify(
|
||||
{
|
||||
"destination": "internal.trade",
|
||||
"correlationId": "request-1",
|
||||
}
|
||||
)
|
||||
is WebSocketInboundMessageKind.MARKET
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"document",
|
||||
[
|
||||
None,
|
||||
[],
|
||||
"message",
|
||||
{},
|
||||
{
|
||||
"destination": "unknown",
|
||||
},
|
||||
{
|
||||
"destination": [],
|
||||
},
|
||||
{
|
||||
"correlationId": "",
|
||||
},
|
||||
{
|
||||
"correlationId": " ",
|
||||
},
|
||||
{
|
||||
"correlationId": None,
|
||||
},
|
||||
{
|
||||
"correlationId": True,
|
||||
},
|
||||
{
|
||||
"correlationId": 1.5,
|
||||
},
|
||||
],
|
||||
)
|
||||
def test_rejects_invalid_or_unknown_documents(
|
||||
classifier: DzengiWebSocketInboundMessageClassifier,
|
||||
document: object,
|
||||
) -> None:
|
||||
with pytest.raises(
|
||||
WebSocketMessageRoutingError,
|
||||
):
|
||||
classifier.classify(document)
|
||||
@@ -0,0 +1,573 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from websockets.protocol import State
|
||||
|
||||
from src.market_data.acquisition.adapters.dzengi.websocket_transport import (
|
||||
DzengiWebSocketTransport,
|
||||
build_dzengi_websocket_url,
|
||||
)
|
||||
from src.market_data.acquisition.exceptions import (
|
||||
WebSocketTransportError,
|
||||
WebSocketTransportNotConnectedError,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_protocol import (
|
||||
WebSocketTransportProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_liveness_probe import (
|
||||
RuntimeLivenessProbeProtocol,
|
||||
)
|
||||
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
incoming: tuple[str | bytes, ...] = (),
|
||||
state: State = State.OPEN,
|
||||
) -> None:
|
||||
self.state = state
|
||||
self.incoming = list(incoming)
|
||||
self.sent_messages: list[str | bytes] = []
|
||||
self.close_calls = 0
|
||||
self.send_error: Exception | None = None
|
||||
self.receive_error: Exception | None = None
|
||||
self.ping_error: Exception | None = None
|
||||
self.pong_error: Exception | None = None
|
||||
self.pong_gate: asyncio.Event | None = None
|
||||
self.ping_calls = 0
|
||||
self.close_error: Exception | None = None
|
||||
|
||||
async def close(
|
||||
self,
|
||||
code: int = 1000,
|
||||
reason: str = "",
|
||||
) -> None:
|
||||
self.close_calls += 1
|
||||
self.state = State.CLOSED
|
||||
|
||||
if self.close_error is not None:
|
||||
raise self.close_error
|
||||
|
||||
async def send(
|
||||
self,
|
||||
message: str | bytes,
|
||||
) -> None:
|
||||
if self.send_error is not None:
|
||||
self.state = State.CLOSED
|
||||
raise self.send_error
|
||||
|
||||
self.sent_messages.append(message)
|
||||
|
||||
async def recv(self) -> str | bytes:
|
||||
if self.receive_error is not None:
|
||||
self.state = State.CLOSED
|
||||
raise self.receive_error
|
||||
|
||||
return self.incoming.pop(0)
|
||||
|
||||
async def ping(self) -> Awaitable[float]:
|
||||
self.ping_calls += 1
|
||||
|
||||
if self.ping_error is not None:
|
||||
raise self.ping_error
|
||||
|
||||
async def wait_for_pong() -> float:
|
||||
if self.pong_gate is not None:
|
||||
await self.pong_gate.wait()
|
||||
|
||||
if self.pong_error is not None:
|
||||
raise self.pong_error
|
||||
|
||||
return 0.01
|
||||
|
||||
return wait_for_pong()
|
||||
|
||||
|
||||
class RecordingConnector:
|
||||
def __init__(
|
||||
self,
|
||||
*connections: FakeConnection,
|
||||
) -> None:
|
||||
self.connections = list(connections)
|
||||
self.calls: list[
|
||||
tuple[str, dict[str, Any]]
|
||||
] = []
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
url: str,
|
||||
**kwargs: Any,
|
||||
) -> FakeConnection:
|
||||
self.calls.append(
|
||||
(
|
||||
url,
|
||||
kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
return self.connections.pop(0)
|
||||
|
||||
|
||||
def create_transport(
|
||||
connection: FakeConnection | None = None,
|
||||
) -> tuple[
|
||||
DzengiWebSocketTransport,
|
||||
FakeConnection,
|
||||
RecordingConnector,
|
||||
]:
|
||||
resolved_connection = connection or FakeConnection()
|
||||
connector = RecordingConnector(
|
||||
resolved_connection,
|
||||
)
|
||||
transport = DzengiWebSocketTransport(
|
||||
url="https://api-adapter.dzengi.com",
|
||||
headers={
|
||||
"Origin": "https://api-adapter.dzengi.com",
|
||||
"X-MBX-APIKEY": "test-key",
|
||||
},
|
||||
open_timeout=11,
|
||||
ping_interval=12,
|
||||
ping_timeout=13,
|
||||
close_timeout=14,
|
||||
connector=connector,
|
||||
)
|
||||
|
||||
return (
|
||||
transport,
|
||||
resolved_connection,
|
||||
connector,
|
||||
)
|
||||
|
||||
|
||||
def test_transport_implements_protocol() -> None:
|
||||
transport, *_ = create_transport()
|
||||
|
||||
assert isinstance(
|
||||
transport,
|
||||
WebSocketTransportProtocol,
|
||||
)
|
||||
assert isinstance(
|
||||
transport,
|
||||
RuntimeLivenessProbeProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_transport_uses_slots() -> None:
|
||||
transport, *_ = create_transport()
|
||||
|
||||
assert not hasattr(transport, "__dict__")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("raw_url", "expected"),
|
||||
[
|
||||
(
|
||||
"https://api-adapter.dzengi.com",
|
||||
"wss://api-adapter.dzengi.com/connect",
|
||||
),
|
||||
(
|
||||
"http://localhost:8080/",
|
||||
"ws://localhost:8080/connect",
|
||||
),
|
||||
(
|
||||
"wss://api-adapter.dzengi.com/connect/",
|
||||
"wss://api-adapter.dzengi.com/connect",
|
||||
),
|
||||
(
|
||||
"ws://localhost:8080/custom?token=abc",
|
||||
"ws://localhost:8080/custom/connect?token=abc",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_build_url_normalizes_supported_urls(
|
||||
raw_url: str,
|
||||
expected: str,
|
||||
) -> None:
|
||||
assert build_dzengi_websocket_url(raw_url) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw_url",
|
||||
[
|
||||
"",
|
||||
" ",
|
||||
"ftp://api-adapter.dzengi.com",
|
||||
"api-adapter.dzengi.com",
|
||||
"wss:///connect",
|
||||
"wss://api-adapter.dzengi.com/#fragment",
|
||||
],
|
||||
)
|
||||
def test_build_url_rejects_invalid_values(
|
||||
raw_url: str,
|
||||
) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
build_dzengi_websocket_url(raw_url)
|
||||
|
||||
|
||||
def test_build_url_rejects_non_string() -> None:
|
||||
with pytest.raises(TypeError):
|
||||
build_dzengi_websocket_url(123) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def test_connect_forwards_connection_options() -> None:
|
||||
transport, _, connector = create_transport()
|
||||
|
||||
asyncio.run(transport.connect())
|
||||
|
||||
assert transport.is_connected is True
|
||||
assert len(connector.calls) == 1
|
||||
|
||||
url, options = connector.calls[0]
|
||||
|
||||
assert url == "wss://api-adapter.dzengi.com/connect"
|
||||
assert options["additional_headers"] == {
|
||||
"Origin": "https://api-adapter.dzengi.com",
|
||||
"X-MBX-APIKEY": "test-key",
|
||||
}
|
||||
assert tuple(options["subprotocols"]) == ("json",)
|
||||
assert options["open_timeout"] == 11.0
|
||||
assert options["ping_interval"] == 12.0
|
||||
assert options["ping_timeout"] == 13.0
|
||||
assert options["close_timeout"] == 14.0
|
||||
|
||||
|
||||
def test_connect_is_idempotent() -> None:
|
||||
transport, _, connector = create_transport()
|
||||
|
||||
async def scenario() -> None:
|
||||
await transport.connect()
|
||||
await transport.connect()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert len(connector.calls) == 1
|
||||
|
||||
|
||||
def test_connect_wraps_connector_error() -> None:
|
||||
async def broken_connector(
|
||||
url: str,
|
||||
**kwargs: Any,
|
||||
) -> FakeConnection:
|
||||
raise OSError("connection failed")
|
||||
|
||||
transport = DzengiWebSocketTransport(
|
||||
url="wss://api-adapter.dzengi.com/connect",
|
||||
connector=broken_connector,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportError,
|
||||
match="connection failed",
|
||||
) as exc_info:
|
||||
asyncio.run(transport.connect())
|
||||
|
||||
assert isinstance(exc_info.value.__cause__, OSError)
|
||||
assert transport.is_connected is False
|
||||
|
||||
|
||||
def test_connect_rejects_non_open_connection() -> None:
|
||||
connection = FakeConnection(
|
||||
state=State.CLOSED,
|
||||
)
|
||||
transport, _, _ = create_transport(connection)
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportError,
|
||||
match="состояние OPEN",
|
||||
):
|
||||
asyncio.run(transport.connect())
|
||||
|
||||
assert connection.close_calls == 1
|
||||
assert transport.is_connected is False
|
||||
|
||||
|
||||
def test_disconnect_closes_connection_and_is_idempotent() -> None:
|
||||
transport, connection, _ = create_transport()
|
||||
|
||||
async def scenario() -> None:
|
||||
await transport.connect()
|
||||
await transport.disconnect()
|
||||
await transport.disconnect()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert connection.close_calls == 1
|
||||
assert transport.is_connected is False
|
||||
|
||||
|
||||
def test_disconnect_clears_connection_after_close_error() -> None:
|
||||
connection = FakeConnection()
|
||||
connection.close_error = RuntimeError("close failed")
|
||||
transport, _, _ = create_transport(connection)
|
||||
|
||||
async def scenario() -> None:
|
||||
await transport.connect()
|
||||
await transport.disconnect()
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportError,
|
||||
match="close failed",
|
||||
):
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.is_connected is False
|
||||
|
||||
|
||||
def test_send_forwards_text_and_binary_messages() -> None:
|
||||
transport, connection, _ = create_transport()
|
||||
|
||||
async def scenario() -> None:
|
||||
await transport.connect()
|
||||
await transport.send("text")
|
||||
await transport.send(b"binary")
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert connection.sent_messages == [
|
||||
"text",
|
||||
b"binary",
|
||||
]
|
||||
|
||||
|
||||
def test_send_requires_open_connection() -> None:
|
||||
transport, _, _ = create_transport()
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportNotConnectedError,
|
||||
):
|
||||
asyncio.run(transport.send("message"))
|
||||
|
||||
|
||||
def test_send_error_is_wrapped_and_discards_closed_connection() -> None:
|
||||
connection = FakeConnection()
|
||||
connection.send_error = RuntimeError("send failed")
|
||||
transport, _, _ = create_transport(connection)
|
||||
|
||||
async def scenario() -> None:
|
||||
await transport.connect()
|
||||
await transport.send("message")
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportError,
|
||||
match="send failed",
|
||||
):
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.is_connected is False
|
||||
|
||||
|
||||
def test_receive_returns_text_and_binary_messages() -> None:
|
||||
connection = FakeConnection(
|
||||
incoming=(
|
||||
"text",
|
||||
b"binary",
|
||||
),
|
||||
)
|
||||
transport, _, _ = create_transport(connection)
|
||||
|
||||
async def scenario() -> tuple[str | bytes, str | bytes]:
|
||||
await transport.connect()
|
||||
|
||||
return (
|
||||
await transport.receive(),
|
||||
await transport.receive(),
|
||||
)
|
||||
|
||||
assert asyncio.run(scenario()) == (
|
||||
"text",
|
||||
b"binary",
|
||||
)
|
||||
|
||||
|
||||
def test_receive_requires_open_connection() -> None:
|
||||
transport, _, _ = create_transport()
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportNotConnectedError,
|
||||
):
|
||||
asyncio.run(transport.receive())
|
||||
|
||||
|
||||
def test_receive_error_is_wrapped_and_discards_closed_connection() -> None:
|
||||
connection = FakeConnection()
|
||||
connection.receive_error = RuntimeError("receive failed")
|
||||
transport, _, _ = create_transport(connection)
|
||||
|
||||
async def scenario() -> None:
|
||||
await transport.connect()
|
||||
await transport.receive()
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportError,
|
||||
match="receive failed",
|
||||
):
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.is_connected is False
|
||||
|
||||
|
||||
def test_probe_returns_true_after_pong() -> None:
|
||||
transport, connection, _ = create_transport()
|
||||
|
||||
async def scenario() -> bool:
|
||||
await transport.connect()
|
||||
return await transport.probe()
|
||||
|
||||
assert asyncio.run(scenario()) is True
|
||||
assert connection.ping_calls == 1
|
||||
|
||||
|
||||
def test_probe_returns_false_without_open_connection() -> None:
|
||||
transport, connection, _ = create_transport()
|
||||
|
||||
assert asyncio.run(transport.probe()) is False
|
||||
assert connection.ping_calls == 0
|
||||
|
||||
|
||||
def test_probe_returns_false_after_pong_timeout() -> None:
|
||||
connection = FakeConnection()
|
||||
connection.pong_gate = asyncio.Event()
|
||||
connector = RecordingConnector(connection)
|
||||
transport = DzengiWebSocketTransport(
|
||||
url="wss://api-adapter.dzengi.com",
|
||||
ping_timeout=0.001,
|
||||
connector=connector,
|
||||
)
|
||||
|
||||
async def scenario() -> bool:
|
||||
await transport.connect()
|
||||
return await transport.probe()
|
||||
|
||||
assert asyncio.run(scenario()) is False
|
||||
assert connection.ping_calls == 1
|
||||
assert transport.is_connected is True
|
||||
|
||||
|
||||
def test_probe_timeout_remains_finite_when_ping_timeout_is_disabled() -> None:
|
||||
connection = FakeConnection()
|
||||
connection.pong_gate = asyncio.Event()
|
||||
connector = RecordingConnector(connection)
|
||||
transport = DzengiWebSocketTransport(
|
||||
url="wss://api-adapter.dzengi.com",
|
||||
ping_timeout=None,
|
||||
probe_timeout=0.001,
|
||||
connector=connector,
|
||||
)
|
||||
|
||||
async def scenario() -> bool:
|
||||
await transport.connect()
|
||||
return await transport.probe()
|
||||
|
||||
assert asyncio.run(scenario()) is False
|
||||
assert connection.ping_calls == 1
|
||||
assert connector.calls[0][1]["ping_timeout"] is None
|
||||
|
||||
|
||||
def test_probe_returns_false_when_connection_closes() -> None:
|
||||
connection = FakeConnection()
|
||||
connection.pong_error = RuntimeError("connection closed")
|
||||
transport, _, _ = create_transport(connection)
|
||||
|
||||
async def scenario() -> bool:
|
||||
await transport.connect()
|
||||
connection.state = State.CLOSED
|
||||
return await transport.probe()
|
||||
|
||||
assert asyncio.run(scenario()) is False
|
||||
assert transport.is_connected is False
|
||||
|
||||
|
||||
def test_probe_wraps_unexpected_ping_error_on_open_connection() -> None:
|
||||
connection = FakeConnection()
|
||||
connection.ping_error = RuntimeError("ping failed")
|
||||
transport, _, _ = create_transport(connection)
|
||||
|
||||
async def scenario() -> None:
|
||||
await transport.connect()
|
||||
await transport.probe()
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportError,
|
||||
match="ping failed",
|
||||
) as error_info:
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert transport.is_connected is True
|
||||
|
||||
|
||||
def test_probe_cancellation_is_not_swallowed() -> None:
|
||||
async def scenario() -> DzengiWebSocketTransport:
|
||||
connection = FakeConnection()
|
||||
connection.pong_gate = asyncio.Event()
|
||||
transport, _, _ = create_transport(connection)
|
||||
await transport.connect()
|
||||
|
||||
probe_task = asyncio.create_task(
|
||||
transport.probe(),
|
||||
)
|
||||
|
||||
while connection.ping_calls == 0:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
probe_task.cancel()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await probe_task
|
||||
|
||||
return transport
|
||||
|
||||
assert asyncio.run(scenario()).is_connected is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("name", "value"),
|
||||
[
|
||||
("open_timeout", 0),
|
||||
("ping_interval", -1),
|
||||
("ping_timeout", True),
|
||||
("probe_timeout", float("inf")),
|
||||
("close_timeout", "10"),
|
||||
],
|
||||
)
|
||||
def test_rejects_invalid_timeout(
|
||||
name: str,
|
||||
value: object,
|
||||
) -> None:
|
||||
options = {
|
||||
name: value,
|
||||
}
|
||||
|
||||
with pytest.raises(
|
||||
(TypeError, ValueError),
|
||||
):
|
||||
DzengiWebSocketTransport(
|
||||
url="wss://api-adapter.dzengi.com",
|
||||
**options, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_copies_headers_from_caller() -> None:
|
||||
headers = {
|
||||
"Origin": "https://api-adapter.dzengi.com",
|
||||
}
|
||||
connection = FakeConnection()
|
||||
connector = RecordingConnector(connection)
|
||||
transport = DzengiWebSocketTransport(
|
||||
url="wss://api-adapter.dzengi.com",
|
||||
headers=headers,
|
||||
connector=connector,
|
||||
)
|
||||
|
||||
headers["Origin"] = "changed"
|
||||
|
||||
asyncio.run(transport.connect())
|
||||
|
||||
assert connector.calls[0][1]["additional_headers"] == {
|
||||
"Origin": "https://api-adapter.dzengi.com",
|
||||
}
|
||||
@@ -0,0 +1,201 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.runtime.acquisition_runtime_event_logging_consumer import (
|
||||
AcquisitionRuntimeEventLoggingConsumer,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.acquisition_runtime_event_publisher import (
|
||||
AcquisitionRuntimeEventConsumerProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_events import (
|
||||
ConnectedEvent,
|
||||
ConnectFailedEvent,
|
||||
DisconnectedEvent,
|
||||
HeartbeatTimeoutEvent,
|
||||
MessageReceivedEvent,
|
||||
MessageSentEvent,
|
||||
ReconnectCompletedEvent,
|
||||
ReconnectFailedEvent,
|
||||
ReconnectStartedEvent,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.transport_messages import (
|
||||
TransportBinaryMessage,
|
||||
TransportTextMessage,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_protocol import (
|
||||
AcquisitionRuntimeEvent,
|
||||
)
|
||||
|
||||
|
||||
TEST_LOGGER_NAME = "tests.acquisition_runtime_events"
|
||||
|
||||
|
||||
def create_consumer() -> AcquisitionRuntimeEventLoggingConsumer:
|
||||
return AcquisitionRuntimeEventLoggingConsumer(
|
||||
event_logger=logging.getLogger(
|
||||
TEST_LOGGER_NAME,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_logging_consumer_satisfies_protocol() -> None:
|
||||
assert isinstance(
|
||||
create_consumer(),
|
||||
AcquisitionRuntimeEventConsumerProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_logging_consumer_uses_slots() -> None:
|
||||
consumer = create_consumer()
|
||||
|
||||
assert not hasattr(consumer, "__dict__")
|
||||
|
||||
|
||||
def test_invalid_logger_is_rejected() -> None:
|
||||
with pytest.raises(
|
||||
TypeError,
|
||||
match="logging.Logger",
|
||||
):
|
||||
AcquisitionRuntimeEventLoggingConsumer(
|
||||
event_logger=object(), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
(
|
||||
"event",
|
||||
"expected_level",
|
||||
"expected_message",
|
||||
),
|
||||
(
|
||||
(
|
||||
ConnectedEvent(),
|
||||
logging.INFO,
|
||||
"connected",
|
||||
),
|
||||
(
|
||||
DisconnectedEvent(),
|
||||
logging.INFO,
|
||||
"disconnected",
|
||||
),
|
||||
(
|
||||
ConnectFailedEvent(
|
||||
reason="connection refused",
|
||||
),
|
||||
logging.ERROR,
|
||||
"connection refused",
|
||||
),
|
||||
(
|
||||
ReconnectStartedEvent(
|
||||
attempt=1,
|
||||
),
|
||||
logging.INFO,
|
||||
"attempt=1",
|
||||
),
|
||||
(
|
||||
ReconnectCompletedEvent(
|
||||
attempt=2,
|
||||
),
|
||||
logging.INFO,
|
||||
"attempt=2",
|
||||
),
|
||||
(
|
||||
ReconnectFailedEvent(
|
||||
attempt=3,
|
||||
reason="timeout",
|
||||
),
|
||||
logging.ERROR,
|
||||
"attempt=3 reason=timeout",
|
||||
),
|
||||
(
|
||||
HeartbeatTimeoutEvent(
|
||||
timeout_seconds=30.0,
|
||||
),
|
||||
logging.WARNING,
|
||||
"timeout_seconds=30.0",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_lifecycle_event_uses_expected_log_level(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
expected_level: int,
|
||||
expected_message: str,
|
||||
) -> None:
|
||||
consumer = create_consumer()
|
||||
|
||||
caplog.set_level(
|
||||
logging.DEBUG,
|
||||
logger=TEST_LOGGER_NAME,
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
consumer.consume(event)
|
||||
)
|
||||
|
||||
assert len(caplog.records) == 1
|
||||
assert caplog.records[0].levelno == expected_level
|
||||
assert expected_message in caplog.records[0].getMessage()
|
||||
|
||||
|
||||
def test_message_events_log_metadata_without_payload(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
consumer = create_consumer()
|
||||
text_payload = "secret-subscription-payload"
|
||||
binary_payload = b"secret-binary-payload"
|
||||
|
||||
caplog.set_level(
|
||||
logging.DEBUG,
|
||||
logger=TEST_LOGGER_NAME,
|
||||
)
|
||||
|
||||
async def consume_messages() -> None:
|
||||
await consumer.consume(
|
||||
MessageReceivedEvent(
|
||||
message=TransportTextMessage(
|
||||
payload=text_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
await consumer.consume(
|
||||
MessageSentEvent(
|
||||
message=TransportBinaryMessage(
|
||||
payload=binary_payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
asyncio.run(consume_messages())
|
||||
|
||||
assert len(caplog.records) == 2
|
||||
assert (
|
||||
"message_type=TransportTextMessage "
|
||||
f"payload_size={len(text_payload)}"
|
||||
in caplog.records[0].getMessage()
|
||||
)
|
||||
assert (
|
||||
"message_type=TransportBinaryMessage "
|
||||
f"payload_size={len(binary_payload)}"
|
||||
in caplog.records[1].getMessage()
|
||||
)
|
||||
assert text_payload not in caplog.text
|
||||
assert binary_payload.decode() not in caplog.text
|
||||
|
||||
|
||||
def test_invalid_event_is_rejected() -> None:
|
||||
consumer = create_consumer()
|
||||
|
||||
with pytest.raises(
|
||||
TypeError,
|
||||
match="AcquisitionRuntimeEvent",
|
||||
):
|
||||
asyncio.run(
|
||||
consumer.consume(
|
||||
object(), # type: ignore[arg-type]
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,679 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.runtime.acquisition_runtime_event_publisher import (
|
||||
AcquisitionRuntimeEventConsumerProtocol,
|
||||
AcquisitionRuntimeEventPublisher,
|
||||
logger as publisher_logger,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.heartbeat import (
|
||||
HeartbeatMonitor,
|
||||
HeartbeatState,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.reconnect import (
|
||||
ReconnectCoordinator,
|
||||
ReconnectState,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_commands import (
|
||||
ConnectCommand,
|
||||
DisconnectCommand,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_events import (
|
||||
ConnectedEvent,
|
||||
DisconnectedEvent,
|
||||
MessageReceivedEvent,
|
||||
ReconnectStartedEvent,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.transport_messages import (
|
||||
TransportTextMessage,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_protocol import (
|
||||
AcquisitionRuntimeCommand,
|
||||
AcquisitionRuntimeEvent,
|
||||
AcquisitionRuntimeEventPublisherProtocol,
|
||||
AcquisitionSubscriptionMessage,
|
||||
)
|
||||
|
||||
|
||||
PUBLISHER_LOGGER_NAME = (
|
||||
"src.market_data.acquisition.runtime."
|
||||
"acquisition_runtime_event_publisher"
|
||||
)
|
||||
|
||||
|
||||
class RecordingConsumer:
|
||||
def __init__(self) -> None:
|
||||
self.events: list[AcquisitionRuntimeEvent] = []
|
||||
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
self.events.append(event)
|
||||
|
||||
|
||||
def test_publisher_satisfies_protocol() -> None:
|
||||
publisher = AcquisitionRuntimeEventPublisher()
|
||||
|
||||
assert isinstance(
|
||||
publisher,
|
||||
AcquisitionRuntimeEventPublisherProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_consumer_satisfies_protocol() -> None:
|
||||
assert isinstance(
|
||||
RecordingConsumer(),
|
||||
AcquisitionRuntimeEventConsumerProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_publisher_uses_slots() -> None:
|
||||
publisher = AcquisitionRuntimeEventPublisher()
|
||||
|
||||
assert not hasattr(publisher, "__dict__")
|
||||
|
||||
|
||||
def test_empty_publisher_accepts_event() -> None:
|
||||
publisher = AcquisitionRuntimeEventPublisher()
|
||||
|
||||
asyncio.run(
|
||||
publisher.publish(
|
||||
ConnectedEvent(),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_constructor_copies_consumer_collection() -> None:
|
||||
first_consumer = RecordingConsumer()
|
||||
consumers = [first_consumer]
|
||||
publisher = AcquisitionRuntimeEventPublisher(consumers)
|
||||
|
||||
second_consumer = RecordingConsumer()
|
||||
consumers.append(second_consumer)
|
||||
|
||||
event = ConnectedEvent()
|
||||
asyncio.run(publisher.publish(event))
|
||||
|
||||
assert first_consumer.events == [event]
|
||||
assert second_consumer.events == []
|
||||
|
||||
|
||||
def test_invalid_consumer_is_rejected() -> None:
|
||||
with pytest.raises(
|
||||
TypeError,
|
||||
match="AcquisitionRuntimeEventConsumerProtocol",
|
||||
):
|
||||
AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
object(), # type: ignore[arg-type]
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_event_is_rejected() -> None:
|
||||
publisher = AcquisitionRuntimeEventPublisher()
|
||||
|
||||
with pytest.raises(
|
||||
TypeError,
|
||||
match="AcquisitionRuntimeEvent",
|
||||
):
|
||||
asyncio.run(
|
||||
publisher.publish(
|
||||
object(), # type: ignore[arg-type]
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_event_is_delivered_to_consumers_in_registration_order() -> None:
|
||||
calls: list[tuple[str, AcquisitionRuntimeEvent]] = []
|
||||
|
||||
class NamedConsumer:
|
||||
def __init__(self, name: str) -> None:
|
||||
self._name = name
|
||||
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
calls.append(
|
||||
(
|
||||
self._name,
|
||||
event,
|
||||
)
|
||||
)
|
||||
|
||||
event = ConnectedEvent()
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
NamedConsumer("first"),
|
||||
NamedConsumer("second"),
|
||||
)
|
||||
)
|
||||
|
||||
asyncio.run(publisher.publish(event))
|
||||
|
||||
assert calls == [
|
||||
(
|
||||
"first",
|
||||
event,
|
||||
),
|
||||
(
|
||||
"second",
|
||||
event,
|
||||
),
|
||||
]
|
||||
assert calls[0][1] is event
|
||||
assert calls[1][1] is event
|
||||
|
||||
|
||||
def test_publish_waits_for_consumer_completion() -> None:
|
||||
completed = False
|
||||
|
||||
class YieldingConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
nonlocal completed
|
||||
|
||||
await asyncio.sleep(0)
|
||||
completed = True
|
||||
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
YieldingConsumer(),
|
||||
)
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
publisher.publish(
|
||||
ConnectedEvent(),
|
||||
)
|
||||
)
|
||||
|
||||
assert completed is True
|
||||
|
||||
|
||||
def test_concurrent_publications_are_serialized() -> None:
|
||||
calls: list[tuple[str, int]] = []
|
||||
|
||||
async def scenario() -> None:
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
|
||||
class BlockingConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
assert isinstance(
|
||||
event,
|
||||
ReconnectStartedEvent,
|
||||
)
|
||||
|
||||
calls.append(
|
||||
(
|
||||
"start",
|
||||
event.attempt,
|
||||
)
|
||||
)
|
||||
|
||||
if event.attempt == 1:
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
|
||||
calls.append(
|
||||
(
|
||||
"end",
|
||||
event.attempt,
|
||||
)
|
||||
)
|
||||
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
BlockingConsumer(),
|
||||
)
|
||||
)
|
||||
|
||||
first_task = asyncio.create_task(
|
||||
publisher.publish(
|
||||
ReconnectStartedEvent(
|
||||
attempt=1,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
await first_started.wait()
|
||||
|
||||
second_task = asyncio.create_task(
|
||||
publisher.publish(
|
||||
ReconnectStartedEvent(
|
||||
attempt=2,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert calls == [
|
||||
(
|
||||
"start",
|
||||
1,
|
||||
),
|
||||
]
|
||||
|
||||
release_first.set()
|
||||
|
||||
await asyncio.gather(
|
||||
first_task,
|
||||
second_task,
|
||||
)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert calls == [
|
||||
(
|
||||
"start",
|
||||
1,
|
||||
),
|
||||
(
|
||||
"end",
|
||||
1,
|
||||
),
|
||||
(
|
||||
"start",
|
||||
2,
|
||||
),
|
||||
(
|
||||
"end",
|
||||
2,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def test_consumer_error_is_logged_and_delivery_continues(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
class BrokenConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
raise RuntimeError(
|
||||
"sensitive consumer failure detail",
|
||||
)
|
||||
|
||||
recording_consumer = RecordingConsumer()
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
BrokenConsumer(),
|
||||
recording_consumer,
|
||||
)
|
||||
)
|
||||
event = ConnectedEvent()
|
||||
|
||||
caplog.set_level(
|
||||
logging.ERROR,
|
||||
logger=PUBLISHER_LOGGER_NAME,
|
||||
)
|
||||
|
||||
asyncio.run(publisher.publish(event))
|
||||
|
||||
assert recording_consumer.events == [event]
|
||||
assert len(caplog.records) == 1
|
||||
assert (
|
||||
"event=ConnectedEvent consumer=BrokenConsumer "
|
||||
"error_type=RuntimeError"
|
||||
in caplog.records[0].getMessage()
|
||||
)
|
||||
assert caplog.records[0].exc_info is None
|
||||
assert (
|
||||
"sensitive consumer failure detail"
|
||||
not in caplog.text
|
||||
)
|
||||
|
||||
|
||||
def test_consumer_error_log_does_not_include_transport_payload(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
class PayloadEchoingBrokenConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
raise RuntimeError(
|
||||
f"consumer rejected {event!r}"
|
||||
)
|
||||
|
||||
payload = "secret-runtime-payload"
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
PayloadEchoingBrokenConsumer(),
|
||||
)
|
||||
)
|
||||
|
||||
caplog.set_level(
|
||||
logging.ERROR,
|
||||
logger=PUBLISHER_LOGGER_NAME,
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
publisher.publish(
|
||||
MessageReceivedEvent(
|
||||
message=TransportTextMessage(
|
||||
payload=payload,
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
assert payload not in caplog.text
|
||||
assert "consumer rejected" not in caplog.text
|
||||
assert "error_type=RuntimeError" in caplog.text
|
||||
|
||||
|
||||
def test_logging_handler_error_does_not_escape_or_stop_delivery() -> None:
|
||||
class BrokenConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
raise RuntimeError(
|
||||
"consumer failed",
|
||||
)
|
||||
|
||||
class RaisingHandler(logging.Handler):
|
||||
def emit(
|
||||
self,
|
||||
record: logging.LogRecord,
|
||||
) -> None:
|
||||
raise RuntimeError(
|
||||
"logging failed",
|
||||
)
|
||||
|
||||
handler = RaisingHandler()
|
||||
recording_consumer = RecordingConsumer()
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
BrokenConsumer(),
|
||||
recording_consumer,
|
||||
)
|
||||
)
|
||||
event = ConnectedEvent()
|
||||
|
||||
publisher_logger.addHandler(handler)
|
||||
|
||||
try:
|
||||
asyncio.run(publisher.publish(event))
|
||||
finally:
|
||||
publisher_logger.removeHandler(handler)
|
||||
|
||||
assert recording_consumer.events == [event]
|
||||
|
||||
|
||||
def test_consumer_cancellation_is_propagated() -> None:
|
||||
class CancelledConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
CancelledConsumer(),
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
asyncio.run(
|
||||
publisher.publish(
|
||||
ConnectedEvent(),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_publisher_can_be_reused_after_consumer_cancellation() -> None:
|
||||
class CancelOnceConsumer:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
self.calls += 1
|
||||
|
||||
if self.calls == 1:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
cancel_once_consumer = CancelOnceConsumer()
|
||||
recording_consumer = RecordingConsumer()
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
cancel_once_consumer,
|
||||
recording_consumer,
|
||||
)
|
||||
)
|
||||
second_event = DisconnectedEvent()
|
||||
|
||||
async def scenario() -> None:
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await publisher.publish(
|
||||
ConnectedEvent(),
|
||||
)
|
||||
|
||||
await publisher.publish(second_event)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert cancel_once_consumer.calls == 2
|
||||
assert recording_consumer.events == [
|
||||
second_event,
|
||||
]
|
||||
|
||||
|
||||
def test_recursive_publication_is_rejected_without_deadlock(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
class RecursiveConsumer:
|
||||
def __init__(self) -> None:
|
||||
self.publisher: (
|
||||
AcquisitionRuntimeEventPublisher | None
|
||||
) = None
|
||||
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
assert self.publisher is not None
|
||||
|
||||
await self.publisher.publish(
|
||||
DisconnectedEvent(),
|
||||
)
|
||||
|
||||
recursive_consumer = RecursiveConsumer()
|
||||
recording_consumer = RecordingConsumer()
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
recursive_consumer,
|
||||
recording_consumer,
|
||||
)
|
||||
)
|
||||
recursive_consumer.publisher = publisher
|
||||
event = ConnectedEvent()
|
||||
|
||||
caplog.set_level(
|
||||
logging.ERROR,
|
||||
logger=PUBLISHER_LOGGER_NAME,
|
||||
)
|
||||
|
||||
asyncio.run(publisher.publish(event))
|
||||
|
||||
assert recording_consumer.events == [event]
|
||||
assert (
|
||||
"event=ConnectedEvent consumer=RecursiveConsumer "
|
||||
"error_type=RuntimeError"
|
||||
in caplog.text
|
||||
)
|
||||
|
||||
|
||||
def test_child_task_recursive_publication_is_rejected() -> None:
|
||||
class ChildTaskRecursiveConsumer:
|
||||
def __init__(self) -> None:
|
||||
self.publisher: (
|
||||
AcquisitionRuntimeEventPublisher | None
|
||||
) = None
|
||||
self.rejected = False
|
||||
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
assert self.publisher is not None
|
||||
|
||||
if not isinstance(event, ConnectedEvent):
|
||||
return
|
||||
|
||||
try:
|
||||
await asyncio.create_task(
|
||||
self.publisher.publish(
|
||||
DisconnectedEvent(),
|
||||
)
|
||||
)
|
||||
except RuntimeError:
|
||||
self.rejected = True
|
||||
|
||||
recursive_consumer = ChildTaskRecursiveConsumer()
|
||||
recording_consumer = RecordingConsumer()
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
recursive_consumer,
|
||||
recording_consumer,
|
||||
)
|
||||
)
|
||||
recursive_consumer.publisher = publisher
|
||||
event = ConnectedEvent()
|
||||
|
||||
asyncio.run(
|
||||
asyncio.wait_for(
|
||||
publisher.publish(event),
|
||||
timeout=1.0,
|
||||
)
|
||||
)
|
||||
|
||||
assert recursive_consumer.rejected is True
|
||||
assert recording_consumer.events == [event]
|
||||
|
||||
|
||||
def test_consumer_error_does_not_interrupt_heartbeat_timeout() -> None:
|
||||
class BrokenConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
raise RuntimeError(
|
||||
"consumer failed",
|
||||
)
|
||||
|
||||
current_time = [0.0]
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
BrokenConsumer(),
|
||||
)
|
||||
)
|
||||
monitor = HeartbeatMonitor(
|
||||
event_publisher=publisher,
|
||||
timeout_seconds=10.0,
|
||||
clock=lambda: current_time[0],
|
||||
)
|
||||
|
||||
monitor.start()
|
||||
current_time[0] = 10.0
|
||||
|
||||
result = asyncio.run(
|
||||
monitor.check_timeout()
|
||||
)
|
||||
|
||||
assert result is True
|
||||
assert monitor.state is HeartbeatState.TIMED_OUT
|
||||
|
||||
|
||||
def test_consumer_error_does_not_interrupt_reconnect() -> None:
|
||||
class BrokenConsumer:
|
||||
async def consume(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
raise RuntimeError(
|
||||
"consumer failed",
|
||||
)
|
||||
|
||||
class RecordingCommandDispatcher:
|
||||
def __init__(self) -> None:
|
||||
self.commands: list[
|
||||
AcquisitionRuntimeCommand
|
||||
] = []
|
||||
|
||||
async def dispatch(
|
||||
self,
|
||||
command: AcquisitionRuntimeCommand,
|
||||
) -> None:
|
||||
self.commands.append(command)
|
||||
|
||||
class RecordingSubscriptionManager:
|
||||
def __init__(self) -> None:
|
||||
self.restore_calls = 0
|
||||
|
||||
async def subscribe(
|
||||
self,
|
||||
subscription_key: str,
|
||||
message: AcquisitionSubscriptionMessage,
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
async def unsubscribe(
|
||||
self,
|
||||
subscription_key: str,
|
||||
message: AcquisitionSubscriptionMessage,
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
async def restore_subscriptions(self) -> None:
|
||||
self.restore_calls += 1
|
||||
|
||||
async def clear_subscriptions(self) -> None:
|
||||
return None
|
||||
|
||||
dispatcher = RecordingCommandDispatcher()
|
||||
subscriptions = RecordingSubscriptionManager()
|
||||
publisher = AcquisitionRuntimeEventPublisher(
|
||||
consumers=(
|
||||
BrokenConsumer(),
|
||||
)
|
||||
)
|
||||
coordinator = ReconnectCoordinator(
|
||||
command_dispatcher=dispatcher,
|
||||
subscription_manager=subscriptions,
|
||||
event_publisher=publisher,
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
coordinator.reconnect()
|
||||
)
|
||||
|
||||
assert coordinator.state is ReconnectState.CONNECTED
|
||||
assert coordinator.attempt == 1
|
||||
assert [
|
||||
type(command)
|
||||
for command in dispatcher.commands
|
||||
] == [
|
||||
DisconnectCommand,
|
||||
ConnectCommand,
|
||||
]
|
||||
assert subscriptions.restore_calls == 1
|
||||
@@ -0,0 +1,115 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.runtime.live_processing_gate import (
|
||||
RuntimeLiveProcessingGate,
|
||||
RuntimeLiveProcessingGateProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_implements_protocol_uses_slots_and_starts_open() -> None:
|
||||
gate = RuntimeLiveProcessingGate()
|
||||
|
||||
assert isinstance(
|
||||
gate,
|
||||
RuntimeLiveProcessingGateProtocol,
|
||||
)
|
||||
assert not hasattr(gate, "__dict__")
|
||||
assert gate.locked is False
|
||||
assert gate.failed is False
|
||||
|
||||
|
||||
def test_serializes_protected_operations() -> None:
|
||||
async def scenario() -> list[str]:
|
||||
gate = RuntimeLiveProcessingGate()
|
||||
order: list[str] = []
|
||||
|
||||
async def contender() -> None:
|
||||
async with gate:
|
||||
order.append("contender")
|
||||
|
||||
async with gate:
|
||||
order.append("owner")
|
||||
contender_task = asyncio.create_task(
|
||||
contender(),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert contender_task.done() is False
|
||||
assert gate.locked is True
|
||||
|
||||
await contender_task
|
||||
return order
|
||||
|
||||
assert asyncio.run(scenario()) == [
|
||||
"owner",
|
||||
"contender",
|
||||
]
|
||||
|
||||
|
||||
def test_releases_gate_after_error() -> None:
|
||||
async def scenario() -> RuntimeLiveProcessingGate:
|
||||
gate = RuntimeLiveProcessingGate()
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="protected failure",
|
||||
):
|
||||
async with gate:
|
||||
raise RuntimeError(
|
||||
"protected failure",
|
||||
)
|
||||
|
||||
return gate
|
||||
|
||||
assert asyncio.run(scenario()).locked is False
|
||||
|
||||
|
||||
def test_failed_gate_rejects_waiters_until_reset() -> None:
|
||||
async def scenario() -> RuntimeLiveProcessingGate:
|
||||
gate = RuntimeLiveProcessingGate()
|
||||
failure = RuntimeError("recovery failed")
|
||||
|
||||
gate.fail(failure)
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="recovery failed",
|
||||
) as error_info:
|
||||
async with gate:
|
||||
raise AssertionError(
|
||||
"failed gate must not admit live processing"
|
||||
)
|
||||
|
||||
assert error_info.value is failure
|
||||
assert gate.locked is False
|
||||
assert gate.failed is True
|
||||
|
||||
gate.reset()
|
||||
|
||||
async with gate:
|
||||
assert gate.locked is True
|
||||
|
||||
return gate
|
||||
|
||||
gate = asyncio.run(scenario())
|
||||
|
||||
assert gate.locked is False
|
||||
assert gate.failed is False
|
||||
|
||||
|
||||
def test_rejects_reset_while_gate_is_locked() -> None:
|
||||
async def scenario() -> None:
|
||||
gate = RuntimeLiveProcessingGate()
|
||||
|
||||
async with gate:
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="cannot be reset while locked",
|
||||
):
|
||||
gate.reset()
|
||||
|
||||
asyncio.run(scenario())
|
||||
@@ -14,6 +14,7 @@ from src.market_data.acquisition.runtime.reconnect import (
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_commands import (
|
||||
ConnectCommand,
|
||||
DisconnectCommand,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_events import (
|
||||
ReconnectCompletedEvent,
|
||||
@@ -115,16 +116,14 @@ def test_initial_state_is_disconnected() -> None:
|
||||
assert coordinator.attempt == 0
|
||||
|
||||
|
||||
def test_reconnect_dispatches_connect_command() -> None:
|
||||
def test_reconnect_dispatches_disconnect_then_connect() -> None:
|
||||
coordinator, dispatcher, *_ = create_coordinator()
|
||||
|
||||
asyncio.run(coordinator.reconnect())
|
||||
|
||||
assert len(dispatcher.commands) == 1
|
||||
assert isinstance(
|
||||
dispatcher.commands[0],
|
||||
ConnectCommand,
|
||||
)
|
||||
assert len(dispatcher.commands) == 2
|
||||
assert isinstance(dispatcher.commands[0], DisconnectCommand)
|
||||
assert isinstance(dispatcher.commands[1], ConnectCommand)
|
||||
|
||||
|
||||
def test_reconnect_restores_subscriptions() -> None:
|
||||
@@ -185,7 +184,9 @@ def test_connect_error_publishes_failed_event() -> None:
|
||||
command: Any,
|
||||
) -> None:
|
||||
self.commands.append(command)
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
if isinstance(command, ConnectCommand):
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
dispatcher = BrokenCommandDispatcher()
|
||||
subscriptions = FakeSubscriptionManager()
|
||||
@@ -218,7 +219,8 @@ def test_connect_error_sets_failed_state() -> None:
|
||||
self,
|
||||
command: Any,
|
||||
) -> None:
|
||||
raise RuntimeError("connection failed")
|
||||
if isinstance(command, ConnectCommand):
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
coordinator = ReconnectCoordinator(
|
||||
command_dispatcher=BrokenCommandDispatcher(),
|
||||
@@ -239,7 +241,8 @@ def test_connect_error_does_not_restore_subscriptions() -> None:
|
||||
self,
|
||||
command: Any,
|
||||
) -> None:
|
||||
raise RuntimeError("connection failed")
|
||||
if isinstance(command, ConnectCommand):
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
subscriptions = FakeSubscriptionManager()
|
||||
|
||||
|
||||
@@ -0,0 +1,711 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.recovery.trade_recovery_result import (
|
||||
TradeRecoveryResult,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.live_processing_gate import (
|
||||
RuntimeLiveProcessingGate,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.reconnect import (
|
||||
ReconnectState,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_reconnect_recovery_coordinator import (
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
RuntimeReconnectRecoveryProtocol,
|
||||
)
|
||||
|
||||
|
||||
BTC = "BTC/USD_LEVERAGE"
|
||||
ETH = "ETH/USD_LEVERAGE"
|
||||
RECOVERY_END_TIME_MS = 1_785_326_405_123
|
||||
|
||||
|
||||
class FakeReconnectCoordinator:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
order: list[str],
|
||||
gate: RuntimeLiveProcessingGate,
|
||||
error: Exception | None = None,
|
||||
entered: asyncio.Event | None = None,
|
||||
release: asyncio.Event | None = None,
|
||||
) -> None:
|
||||
self._order = order
|
||||
self._gate = gate
|
||||
self._error = error
|
||||
self._entered = entered
|
||||
self._release = release
|
||||
self._state = ReconnectState.DISCONNECTED
|
||||
self._attempt = 0
|
||||
|
||||
@property
|
||||
def state(self) -> ReconnectState:
|
||||
return self._state
|
||||
|
||||
@property
|
||||
def attempt(self) -> int:
|
||||
return self._attempt
|
||||
|
||||
async def reconnect(self) -> None:
|
||||
assert self._gate.locked is True
|
||||
|
||||
self._attempt += 1
|
||||
self._state = ReconnectState.CONNECTING
|
||||
self._order.append("reconnect")
|
||||
|
||||
if self._entered is not None:
|
||||
self._entered.set()
|
||||
|
||||
if self._release is not None:
|
||||
await self._release.wait()
|
||||
|
||||
if self._error is not None:
|
||||
self._state = ReconnectState.FAILED
|
||||
raise self._error
|
||||
|
||||
self._state = ReconnectState.CONNECTED
|
||||
|
||||
|
||||
class FakeRecoveryCoordinator:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
order: list[str],
|
||||
gate: RuntimeLiveProcessingGate,
|
||||
error: Exception | None = None,
|
||||
started: threading.Event | None = None,
|
||||
release: threading.Event | None = None,
|
||||
) -> None:
|
||||
self._order = order
|
||||
self._gate = gate
|
||||
self._error = error
|
||||
self._started = started
|
||||
self._release = release
|
||||
self.calls: list[tuple[str, int]] = []
|
||||
self.thread_ids: list[int] = []
|
||||
|
||||
def recover(
|
||||
self,
|
||||
*,
|
||||
symbol: str,
|
||||
recovery_end_time: int,
|
||||
) -> TradeRecoveryResult:
|
||||
assert self._gate.locked is True
|
||||
|
||||
self.calls.append(
|
||||
(
|
||||
symbol,
|
||||
recovery_end_time,
|
||||
)
|
||||
)
|
||||
self.thread_ids.append(
|
||||
threading.get_ident(),
|
||||
)
|
||||
self._order.append(
|
||||
f"recover:{symbol}",
|
||||
)
|
||||
|
||||
if self._started is not None:
|
||||
self._started.set()
|
||||
|
||||
if self._release is not None:
|
||||
if not self._release.wait(timeout=2.0):
|
||||
raise AssertionError(
|
||||
"recovery release was not signalled"
|
||||
)
|
||||
|
||||
if self._error is not None:
|
||||
raise self._error
|
||||
|
||||
return TradeRecoveryResult(
|
||||
symbol=symbol,
|
||||
requested_start_time=recovery_end_time,
|
||||
requested_end_time=recovery_end_time,
|
||||
recovered_trades=(),
|
||||
)
|
||||
|
||||
|
||||
class RecordingClock:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
order: list[str],
|
||||
value: object = RECOVERY_END_TIME_MS,
|
||||
) -> None:
|
||||
self._order = order
|
||||
self._value = value
|
||||
self.calls = 0
|
||||
|
||||
def __call__(self) -> int:
|
||||
self.calls += 1
|
||||
self._order.append("clock")
|
||||
return self._value # type: ignore[return-value]
|
||||
|
||||
|
||||
def create_coordinator(
|
||||
*,
|
||||
symbols: tuple[str, ...] = (BTC,),
|
||||
reconnect_error: Exception | None = None,
|
||||
recovery_error: Exception | None = None,
|
||||
reconnect_entered: asyncio.Event | None = None,
|
||||
reconnect_release: asyncio.Event | None = None,
|
||||
recovery_started: threading.Event | None = None,
|
||||
recovery_release: threading.Event | None = None,
|
||||
clock_value: object = RECOVERY_END_TIME_MS,
|
||||
) -> tuple[
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
RuntimeLiveProcessingGate,
|
||||
FakeReconnectCoordinator,
|
||||
FakeRecoveryCoordinator,
|
||||
RecordingClock,
|
||||
list[str],
|
||||
]:
|
||||
order: list[str] = []
|
||||
gate = RuntimeLiveProcessingGate()
|
||||
reconnect = FakeReconnectCoordinator(
|
||||
order=order,
|
||||
gate=gate,
|
||||
error=reconnect_error,
|
||||
entered=reconnect_entered,
|
||||
release=reconnect_release,
|
||||
)
|
||||
recovery = FakeRecoveryCoordinator(
|
||||
order=order,
|
||||
gate=gate,
|
||||
error=recovery_error,
|
||||
started=recovery_started,
|
||||
release=recovery_release,
|
||||
)
|
||||
clock = RecordingClock(
|
||||
order=order,
|
||||
value=clock_value,
|
||||
)
|
||||
coordinator = RuntimeReconnectRecoveryCoordinator(
|
||||
reconnect_coordinator=reconnect,
|
||||
recovery_coordinator=recovery,
|
||||
live_processing_gate=gate,
|
||||
symbols=symbols,
|
||||
clock=clock,
|
||||
)
|
||||
|
||||
return (
|
||||
coordinator,
|
||||
gate,
|
||||
reconnect,
|
||||
recovery,
|
||||
clock,
|
||||
order,
|
||||
)
|
||||
|
||||
|
||||
async def wait_until(
|
||||
predicate: object,
|
||||
) -> None:
|
||||
for _ in range(100):
|
||||
if callable(predicate) and predicate():
|
||||
return
|
||||
|
||||
await asyncio.sleep(0)
|
||||
|
||||
raise AssertionError("condition was not reached")
|
||||
|
||||
|
||||
def test_implements_protocol_uses_slots_and_delegates_state() -> None:
|
||||
coordinator, _, reconnect, *_ = create_coordinator()
|
||||
|
||||
assert isinstance(
|
||||
coordinator,
|
||||
RuntimeReconnectRecoveryProtocol,
|
||||
)
|
||||
assert not hasattr(coordinator, "__dict__")
|
||||
assert coordinator.state is ReconnectState.DISCONNECTED
|
||||
assert coordinator.attempt == 0
|
||||
assert coordinator.generation == 0
|
||||
assert coordinator.live_processing_gate is coordinator._live_processing_gate
|
||||
assert coordinator.symbols == (BTC,)
|
||||
|
||||
asyncio.run(coordinator.reconnect())
|
||||
|
||||
assert coordinator.state is reconnect.state
|
||||
assert coordinator.state is ReconnectState.CONNECTED
|
||||
assert coordinator.attempt == 1
|
||||
assert coordinator.generation == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("symbols", "error_type"),
|
||||
[
|
||||
([], TypeError),
|
||||
((), ValueError),
|
||||
(("", " "), ValueError),
|
||||
((BTC, 1), TypeError),
|
||||
],
|
||||
)
|
||||
def test_rejects_invalid_symbols(
|
||||
symbols: object,
|
||||
error_type: type[Exception],
|
||||
) -> None:
|
||||
with pytest.raises(error_type):
|
||||
create_coordinator(
|
||||
symbols=symbols, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_reconnect_restore_boundary_and_recovery_order() -> None:
|
||||
(
|
||||
coordinator,
|
||||
gate,
|
||||
_,
|
||||
recovery,
|
||||
clock,
|
||||
order,
|
||||
) = create_coordinator(
|
||||
symbols=(
|
||||
f" {ETH} ",
|
||||
BTC,
|
||||
ETH,
|
||||
),
|
||||
)
|
||||
main_thread_id = threading.get_ident()
|
||||
|
||||
asyncio.run(coordinator.reconnect())
|
||||
|
||||
assert order == [
|
||||
"reconnect",
|
||||
"clock",
|
||||
f"recover:{BTC}",
|
||||
f"recover:{ETH}",
|
||||
]
|
||||
assert recovery.calls == [
|
||||
(
|
||||
BTC,
|
||||
RECOVERY_END_TIME_MS,
|
||||
),
|
||||
(
|
||||
ETH,
|
||||
RECOVERY_END_TIME_MS,
|
||||
),
|
||||
]
|
||||
assert clock.calls == 1
|
||||
assert all(
|
||||
thread_id != main_thread_id
|
||||
for thread_id in recovery.thread_ids
|
||||
)
|
||||
assert gate.locked is False
|
||||
|
||||
|
||||
def test_live_processing_waits_until_recovery_finishes() -> None:
|
||||
async def scenario() -> list[str]:
|
||||
recovery_started = threading.Event()
|
||||
recovery_release = threading.Event()
|
||||
(
|
||||
coordinator,
|
||||
gate,
|
||||
*_,
|
||||
) = create_coordinator(
|
||||
recovery_started=recovery_started,
|
||||
recovery_release=recovery_release,
|
||||
)
|
||||
order: list[str] = []
|
||||
|
||||
coordinator_task = asyncio.create_task(
|
||||
coordinator.reconnect(),
|
||||
)
|
||||
await wait_until(
|
||||
recovery_started.is_set,
|
||||
)
|
||||
|
||||
async def process_live() -> None:
|
||||
async with gate:
|
||||
order.append("live")
|
||||
|
||||
live_task = asyncio.create_task(
|
||||
process_live(),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert live_task.done() is False
|
||||
assert gate.locked is True
|
||||
|
||||
recovery_release.set()
|
||||
await coordinator_task
|
||||
await live_task
|
||||
|
||||
return order
|
||||
|
||||
assert asyncio.run(scenario()) == [
|
||||
"live",
|
||||
]
|
||||
|
||||
|
||||
def test_concurrent_callers_share_one_operation() -> None:
|
||||
async def scenario() -> tuple[
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
FakeReconnectCoordinator,
|
||||
FakeRecoveryCoordinator,
|
||||
]:
|
||||
reconnect_entered = asyncio.Event()
|
||||
reconnect_release = asyncio.Event()
|
||||
(
|
||||
coordinator,
|
||||
_,
|
||||
reconnect,
|
||||
recovery,
|
||||
*_,
|
||||
) = create_coordinator(
|
||||
reconnect_entered=reconnect_entered,
|
||||
reconnect_release=reconnect_release,
|
||||
)
|
||||
|
||||
first = asyncio.create_task(
|
||||
coordinator.reconnect(),
|
||||
)
|
||||
await reconnect_entered.wait()
|
||||
|
||||
second = asyncio.create_task(
|
||||
coordinator.reconnect(),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert reconnect.attempt == 1
|
||||
|
||||
reconnect_release.set()
|
||||
await asyncio.gather(
|
||||
first,
|
||||
second,
|
||||
)
|
||||
|
||||
return (
|
||||
coordinator,
|
||||
reconnect,
|
||||
recovery,
|
||||
)
|
||||
|
||||
coordinator, reconnect, recovery = asyncio.run(
|
||||
scenario()
|
||||
)
|
||||
|
||||
assert reconnect.attempt == 1
|
||||
assert len(recovery.calls) == 1
|
||||
assert coordinator.generation == 1
|
||||
|
||||
|
||||
def test_timeout_and_transport_error_share_one_operation() -> None:
|
||||
async def scenario() -> tuple[
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
FakeReconnectCoordinator,
|
||||
FakeRecoveryCoordinator,
|
||||
]:
|
||||
reconnect_entered = asyncio.Event()
|
||||
reconnect_release = asyncio.Event()
|
||||
(
|
||||
coordinator,
|
||||
_,
|
||||
reconnect,
|
||||
recovery,
|
||||
*_,
|
||||
) = create_coordinator(
|
||||
reconnect_entered=reconnect_entered,
|
||||
reconnect_release=reconnect_release,
|
||||
)
|
||||
observed_generation = coordinator.generation
|
||||
|
||||
timeout_task = asyncio.create_task(
|
||||
coordinator.reconnect(),
|
||||
)
|
||||
await reconnect_entered.wait()
|
||||
|
||||
transport_task = asyncio.create_task(
|
||||
coordinator.reconnect_after_transport_failure(
|
||||
observed_generation=observed_generation,
|
||||
),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert reconnect.attempt == 1
|
||||
reconnect_release.set()
|
||||
|
||||
await asyncio.gather(
|
||||
timeout_task,
|
||||
transport_task,
|
||||
)
|
||||
|
||||
return (
|
||||
coordinator,
|
||||
reconnect,
|
||||
recovery,
|
||||
)
|
||||
|
||||
coordinator, reconnect, recovery = asyncio.run(
|
||||
scenario()
|
||||
)
|
||||
|
||||
assert reconnect.attempt == 1
|
||||
assert len(recovery.calls) == 1
|
||||
assert coordinator.generation == 1
|
||||
|
||||
|
||||
def test_stale_transport_error_reuses_completed_operation() -> None:
|
||||
(
|
||||
coordinator,
|
||||
_,
|
||||
reconnect,
|
||||
recovery,
|
||||
*_,
|
||||
) = create_coordinator()
|
||||
|
||||
observed_generation = coordinator.generation
|
||||
|
||||
asyncio.run(coordinator.reconnect())
|
||||
asyncio.run(
|
||||
coordinator.reconnect_after_transport_failure(
|
||||
observed_generation=observed_generation,
|
||||
)
|
||||
)
|
||||
|
||||
assert reconnect.attempt == 1
|
||||
assert len(recovery.calls) == 1
|
||||
|
||||
|
||||
def test_reconnect_error_skips_clock_and_recovery() -> None:
|
||||
reconnect_error = RuntimeError("reconnect failed")
|
||||
(
|
||||
coordinator,
|
||||
gate,
|
||||
_,
|
||||
recovery,
|
||||
clock,
|
||||
_,
|
||||
) = create_coordinator(
|
||||
reconnect_error=reconnect_error,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="reconnect failed",
|
||||
) as error_info:
|
||||
asyncio.run(coordinator.reconnect())
|
||||
|
||||
assert error_info.value is reconnect_error
|
||||
assert recovery.calls == []
|
||||
assert clock.calls == 0
|
||||
assert gate.locked is False
|
||||
assert gate.failed is True
|
||||
|
||||
|
||||
def test_recovery_error_is_not_wrapped() -> None:
|
||||
recovery_error = RuntimeError("recovery failed")
|
||||
(
|
||||
coordinator,
|
||||
gate,
|
||||
*_,
|
||||
) = create_coordinator(
|
||||
recovery_error=recovery_error,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="recovery failed",
|
||||
) as error_info:
|
||||
asyncio.run(coordinator.reconnect())
|
||||
|
||||
assert error_info.value is recovery_error
|
||||
assert gate.locked is False
|
||||
assert gate.failed is True
|
||||
assert coordinator._recovery_task is None
|
||||
|
||||
|
||||
def test_recovery_error_rejects_buffered_live_processing() -> None:
|
||||
async def scenario() -> tuple[
|
||||
RuntimeLiveProcessingGate,
|
||||
list[str],
|
||||
]:
|
||||
recovery_error = RuntimeError("recovery failed")
|
||||
recovery_started = threading.Event()
|
||||
recovery_release = threading.Event()
|
||||
(
|
||||
coordinator,
|
||||
gate,
|
||||
*_,
|
||||
) = create_coordinator(
|
||||
recovery_error=recovery_error,
|
||||
recovery_started=recovery_started,
|
||||
recovery_release=recovery_release,
|
||||
)
|
||||
processed: list[str] = []
|
||||
coordinator_task = asyncio.create_task(
|
||||
coordinator.reconnect(),
|
||||
)
|
||||
await wait_until(
|
||||
recovery_started.is_set,
|
||||
)
|
||||
|
||||
async def process_live() -> None:
|
||||
async with gate:
|
||||
processed.append("live")
|
||||
|
||||
live_task = asyncio.create_task(
|
||||
process_live(),
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert live_task.done() is False
|
||||
recovery_release.set()
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="recovery failed",
|
||||
):
|
||||
await coordinator_task
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="recovery failed",
|
||||
):
|
||||
await live_task
|
||||
|
||||
return (
|
||||
gate,
|
||||
processed,
|
||||
)
|
||||
|
||||
gate, processed = asyncio.run(scenario())
|
||||
|
||||
assert processed == []
|
||||
assert gate.locked is False
|
||||
assert gate.failed is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("clock_value", "error_type"),
|
||||
[
|
||||
(True, TypeError),
|
||||
(1.5, TypeError),
|
||||
(-1, ValueError),
|
||||
],
|
||||
)
|
||||
def test_rejects_invalid_clock_result(
|
||||
clock_value: object,
|
||||
error_type: type[Exception],
|
||||
) -> None:
|
||||
coordinator, gate, *_ = create_coordinator(
|
||||
clock_value=clock_value,
|
||||
)
|
||||
|
||||
with pytest.raises(error_type):
|
||||
asyncio.run(coordinator.reconnect())
|
||||
|
||||
assert gate.locked is False
|
||||
|
||||
|
||||
def test_cancellation_waits_for_worker_before_opening_gate() -> None:
|
||||
async def scenario() -> tuple[
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
RuntimeLiveProcessingGate,
|
||||
]:
|
||||
recovery_started = threading.Event()
|
||||
recovery_release = threading.Event()
|
||||
(
|
||||
coordinator,
|
||||
gate,
|
||||
*_,
|
||||
) = create_coordinator(
|
||||
recovery_started=recovery_started,
|
||||
recovery_release=recovery_release,
|
||||
)
|
||||
coordinator_task = asyncio.create_task(
|
||||
coordinator.reconnect(),
|
||||
)
|
||||
await wait_until(
|
||||
recovery_started.is_set,
|
||||
)
|
||||
|
||||
coordinator_task.cancel()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert coordinator_task.done() is False
|
||||
assert gate.locked is True
|
||||
assert coordinator._recovery_task is not None
|
||||
|
||||
recovery_release.set()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await coordinator_task
|
||||
|
||||
return (
|
||||
coordinator,
|
||||
gate,
|
||||
)
|
||||
|
||||
coordinator, gate = asyncio.run(scenario())
|
||||
|
||||
assert gate.locked is False
|
||||
assert coordinator._recovery_task is None
|
||||
|
||||
|
||||
def test_repeated_cancellation_waits_for_worker_before_opening_gate() -> None:
|
||||
async def scenario() -> tuple[
|
||||
RuntimeReconnectRecoveryCoordinator,
|
||||
RuntimeLiveProcessingGate,
|
||||
list[str],
|
||||
]:
|
||||
recovery_started = threading.Event()
|
||||
recovery_release = threading.Event()
|
||||
(
|
||||
coordinator,
|
||||
gate,
|
||||
*_,
|
||||
) = create_coordinator(
|
||||
recovery_started=recovery_started,
|
||||
recovery_release=recovery_release,
|
||||
)
|
||||
processed: list[str] = []
|
||||
coordinator_task = asyncio.create_task(
|
||||
coordinator.reconnect(),
|
||||
)
|
||||
await wait_until(
|
||||
recovery_started.is_set,
|
||||
)
|
||||
|
||||
async def process_live() -> None:
|
||||
async with gate:
|
||||
processed.append("live")
|
||||
|
||||
live_task = asyncio.create_task(
|
||||
process_live(),
|
||||
)
|
||||
|
||||
coordinator_task.cancel()
|
||||
await asyncio.sleep(0)
|
||||
coordinator_task.cancel()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert coordinator_task.done() is False
|
||||
assert live_task.done() is False
|
||||
assert gate.locked is True
|
||||
assert coordinator._recovery_task is not None
|
||||
|
||||
recovery_release.set()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await coordinator_task
|
||||
|
||||
await live_task
|
||||
|
||||
return (
|
||||
coordinator,
|
||||
gate,
|
||||
processed,
|
||||
)
|
||||
|
||||
coordinator, gate, processed = asyncio.run(scenario())
|
||||
|
||||
assert processed == ["live"]
|
||||
assert gate.locked is False
|
||||
assert gate.failed is False
|
||||
assert coordinator._recovery_task is None
|
||||
@@ -184,7 +184,9 @@ class FakeStateStore:
|
||||
self._states.clear()
|
||||
|
||||
|
||||
class RecordingWindowPlanner:
|
||||
class RecordingWindowPlanner(
|
||||
TradeRecoveryWindowPlanner,
|
||||
):
|
||||
"""
|
||||
Planner с заранее заданным результатом.
|
||||
"""
|
||||
@@ -290,7 +292,7 @@ def create_coordinator(
|
||||
|
||||
coordinator = RuntimeRecoveryCoordinator(
|
||||
state_store=resolved_state_store,
|
||||
window_planner=resolved_window_planner, # type: ignore[arg-type]
|
||||
window_planner=resolved_window_planner,
|
||||
recovery_controller=resolved_recovery_controller,
|
||||
)
|
||||
|
||||
@@ -924,15 +926,8 @@ def test_controller_error_is_not_swallowed() -> None:
|
||||
|
||||
def test_planner_error_is_not_swallowed() -> None:
|
||||
class BrokenPlanner(
|
||||
TradeRecoveryWindowPlanner,
|
||||
RecordingWindowPlanner,
|
||||
):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
@property
|
||||
def max_window_ms(self) -> int:
|
||||
return 1
|
||||
|
||||
def build_windows(
|
||||
self,
|
||||
*,
|
||||
|
||||
@@ -11,6 +11,9 @@ import pytest
|
||||
from src.market_data.acquisition.runtime.heartbeat import (
|
||||
HeartbeatState,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_liveness_probe import (
|
||||
RuntimeLivenessProbeProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.scheduler import (
|
||||
RuntimeScheduler,
|
||||
RuntimeSchedulerProtocol,
|
||||
@@ -58,9 +61,37 @@ class FakeHeartbeatMonitor:
|
||||
return self._results.pop(0)
|
||||
|
||||
|
||||
class FakeLivenessProbe:
|
||||
def __init__(
|
||||
self,
|
||||
results: list[object] | None = None,
|
||||
*,
|
||||
error: BaseException | None = None,
|
||||
) -> None:
|
||||
self._results = list(
|
||||
results
|
||||
if results is not None
|
||||
else [True]
|
||||
)
|
||||
self._error = error
|
||||
self.calls = 0
|
||||
|
||||
async def probe(self) -> bool:
|
||||
self.calls += 1
|
||||
|
||||
if self._error is not None:
|
||||
raise self._error
|
||||
|
||||
if not self._results:
|
||||
return True
|
||||
|
||||
return self._results.pop(0) # type: ignore[return-value]
|
||||
|
||||
|
||||
class FakeRuntimeSupervisor:
|
||||
def __init__(self) -> None:
|
||||
self.handle_timeout_calls = 0
|
||||
self.notify_activity_calls = 0
|
||||
|
||||
@property
|
||||
def state(self) -> RuntimeSupervisorState:
|
||||
@@ -73,7 +104,7 @@ class FakeRuntimeSupervisor:
|
||||
return None
|
||||
|
||||
def notify_activity(self) -> None:
|
||||
return None
|
||||
self.notify_activity_calls += 1
|
||||
|
||||
async def handle_heartbeat_timeout(self) -> bool:
|
||||
self.handle_timeout_calls += 1
|
||||
@@ -101,6 +132,7 @@ class RecordingSleep:
|
||||
|
||||
def create_scheduler(
|
||||
*,
|
||||
liveness_results: list[object] | None = None,
|
||||
heartbeat_results: list[bool] | None = None,
|
||||
interval_seconds: float = 1.0,
|
||||
sleep: Callable[[float], Awaitable[None]] | None = None,
|
||||
@@ -115,6 +147,9 @@ def create_scheduler(
|
||||
supervisor = FakeRuntimeSupervisor()
|
||||
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(
|
||||
results=liveness_results,
|
||||
),
|
||||
heartbeat_monitor=heartbeat,
|
||||
runtime_supervisor=supervisor,
|
||||
interval_seconds=interval_seconds,
|
||||
@@ -131,6 +166,10 @@ def create_scheduler(
|
||||
def test_scheduler_implements_protocol() -> None:
|
||||
scheduler, *_ = create_scheduler()
|
||||
|
||||
assert isinstance(
|
||||
scheduler._liveness_probe,
|
||||
RuntimeLivenessProbeProtocol,
|
||||
)
|
||||
assert isinstance(
|
||||
scheduler,
|
||||
RuntimeSchedulerProtocol,
|
||||
@@ -149,6 +188,71 @@ def test_initial_state_is_not_running() -> None:
|
||||
assert scheduler.running is False
|
||||
|
||||
|
||||
def test_claim_blocks_start_without_matching_owner() -> None:
|
||||
scheduler, *_ = create_scheduler()
|
||||
owner = object()
|
||||
scheduler.claim(owner)
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="owned by another lifecycle",
|
||||
):
|
||||
asyncio.run(
|
||||
scheduler.start()
|
||||
)
|
||||
|
||||
scheduler.release(owner)
|
||||
|
||||
|
||||
def test_claim_rejects_active_scheduler() -> None:
|
||||
scheduler, *_ = create_scheduler()
|
||||
scheduler._running = True
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="already owned or active",
|
||||
):
|
||||
scheduler.claim(object())
|
||||
|
||||
|
||||
def test_release_rejects_running_owned_scheduler() -> None:
|
||||
scheduler, *_ = create_scheduler()
|
||||
owner = object()
|
||||
scheduler.claim(owner)
|
||||
scheduler._running = True
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="Cannot release",
|
||||
):
|
||||
scheduler.release(owner)
|
||||
|
||||
|
||||
def test_matching_owner_can_run_and_release_scheduler() -> None:
|
||||
scheduler_holder: list[RuntimeScheduler] = []
|
||||
|
||||
async def stop_after_first_iteration(
|
||||
seconds: float,
|
||||
) -> None:
|
||||
scheduler_holder[0].stop()
|
||||
|
||||
scheduler, *_ = create_scheduler(
|
||||
sleep=stop_after_first_iteration,
|
||||
)
|
||||
scheduler_holder.append(scheduler)
|
||||
owner = object()
|
||||
scheduler.claim(owner)
|
||||
|
||||
asyncio.run(
|
||||
scheduler.start(
|
||||
owner=owner,
|
||||
)
|
||||
)
|
||||
scheduler.release(owner)
|
||||
|
||||
assert scheduler.running is False
|
||||
|
||||
|
||||
def test_exposes_interval_seconds() -> None:
|
||||
scheduler, *_ = create_scheduler(
|
||||
interval_seconds=2.5,
|
||||
@@ -187,6 +291,7 @@ def test_rejects_invalid_interval_type(
|
||||
) -> None:
|
||||
with pytest.raises(TypeError):
|
||||
RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=FakeHeartbeatMonitor(),
|
||||
runtime_supervisor=FakeRuntimeSupervisor(),
|
||||
interval_seconds=interval_seconds, # type: ignore[arg-type]
|
||||
@@ -206,6 +311,7 @@ def test_rejects_non_positive_interval(
|
||||
) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=FakeHeartbeatMonitor(),
|
||||
runtime_supervisor=FakeRuntimeSupervisor(),
|
||||
interval_seconds=interval_seconds,
|
||||
@@ -215,6 +321,7 @@ def test_rejects_non_positive_interval(
|
||||
def test_rejects_non_callable_sleep() -> None:
|
||||
with pytest.raises(TypeError):
|
||||
RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=FakeHeartbeatMonitor(),
|
||||
runtime_supervisor=FakeRuntimeSupervisor(),
|
||||
interval_seconds=1.0,
|
||||
@@ -233,9 +340,43 @@ def test_run_once_checks_heartbeat() -> None:
|
||||
|
||||
assert result is False
|
||||
assert heartbeat.check_timeout_calls == 1
|
||||
assert supervisor.notify_activity_calls == 1
|
||||
assert supervisor.handle_timeout_calls == 0
|
||||
|
||||
|
||||
def test_failed_probe_does_not_record_activity() -> None:
|
||||
scheduler, heartbeat, supervisor = create_scheduler(
|
||||
liveness_results=[False],
|
||||
heartbeat_results=[False],
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
scheduler.run_once()
|
||||
)
|
||||
|
||||
assert result is False
|
||||
assert heartbeat.check_timeout_calls == 1
|
||||
assert supervisor.notify_activity_calls == 0
|
||||
assert supervisor.handle_timeout_calls == 0
|
||||
|
||||
|
||||
def test_rejects_non_boolean_probe_result() -> None:
|
||||
scheduler, heartbeat, supervisor = create_scheduler(
|
||||
liveness_results=["alive"],
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
TypeError,
|
||||
match="must return a boolean",
|
||||
):
|
||||
asyncio.run(
|
||||
scheduler.run_once()
|
||||
)
|
||||
|
||||
assert heartbeat.check_timeout_calls == 0
|
||||
assert supervisor.notify_activity_calls == 0
|
||||
|
||||
|
||||
def test_run_once_calls_supervisor_on_timeout() -> None:
|
||||
scheduler, heartbeat, supervisor = create_scheduler(
|
||||
heartbeat_results=[True],
|
||||
@@ -402,6 +543,7 @@ def test_stop_during_run_once_prevents_sleep() -> None:
|
||||
heartbeat = StoppingHeartbeatMonitor()
|
||||
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=heartbeat,
|
||||
runtime_supervisor=supervisor,
|
||||
interval_seconds=1.0,
|
||||
@@ -438,6 +580,7 @@ def test_repeated_start_while_running_does_not_create_second_loop() -> None:
|
||||
scheduler.stop()
|
||||
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=heartbeat,
|
||||
runtime_supervisor=supervisor,
|
||||
interval_seconds=1.0,
|
||||
@@ -491,6 +634,7 @@ def test_heartbeat_error_is_propagated() -> None:
|
||||
raise RuntimeError("heartbeat failed")
|
||||
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=BrokenHeartbeatMonitor(),
|
||||
runtime_supervisor=FakeRuntimeSupervisor(),
|
||||
interval_seconds=1.0,
|
||||
@@ -511,6 +655,7 @@ def test_supervisor_error_is_propagated() -> None:
|
||||
raise RuntimeError("supervisor failed")
|
||||
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=FakeHeartbeatMonitor(
|
||||
results=[True],
|
||||
),
|
||||
@@ -533,6 +678,7 @@ def test_start_resets_running_after_heartbeat_error() -> None:
|
||||
raise RuntimeError("heartbeat failed")
|
||||
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=BrokenHeartbeatMonitor(),
|
||||
runtime_supervisor=FakeRuntimeSupervisor(),
|
||||
interval_seconds=1.0,
|
||||
@@ -552,6 +698,7 @@ def test_start_resets_running_after_supervisor_error() -> None:
|
||||
raise RuntimeError("supervisor failed")
|
||||
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(),
|
||||
heartbeat_monitor=FakeHeartbeatMonitor(
|
||||
results=[True],
|
||||
),
|
||||
@@ -588,3 +735,45 @@ def test_sleep_error_is_propagated_and_resets_running() -> None:
|
||||
|
||||
assert heartbeat.check_timeout_calls == 1
|
||||
assert scheduler.running is False
|
||||
|
||||
|
||||
def test_liveness_error_is_propagated_and_resets_running() -> None:
|
||||
liveness_error = RuntimeError("liveness failed")
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(
|
||||
error=liveness_error,
|
||||
),
|
||||
heartbeat_monitor=FakeHeartbeatMonitor(),
|
||||
runtime_supervisor=FakeRuntimeSupervisor(),
|
||||
interval_seconds=1.0,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="liveness failed",
|
||||
) as error_info:
|
||||
asyncio.run(
|
||||
scheduler.start()
|
||||
)
|
||||
|
||||
assert error_info.value is liveness_error
|
||||
assert scheduler.running is False
|
||||
|
||||
|
||||
def test_liveness_cancellation_is_propagated() -> None:
|
||||
cancellation = asyncio.CancelledError()
|
||||
scheduler = RuntimeScheduler(
|
||||
liveness_probe=FakeLivenessProbe(
|
||||
error=cancellation,
|
||||
),
|
||||
heartbeat_monitor=FakeHeartbeatMonitor(),
|
||||
runtime_supervisor=FakeRuntimeSupervisor(),
|
||||
interval_seconds=1.0,
|
||||
)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
asyncio.run(
|
||||
scheduler.start()
|
||||
)
|
||||
|
||||
assert scheduler.running is False
|
||||
|
||||
@@ -65,6 +65,7 @@ class FakeReconnectCoordinator:
|
||||
def __init__(self) -> None:
|
||||
self.reconnect_calls = 0
|
||||
self._attempt = 0
|
||||
self._generation = 0
|
||||
self._state = ReconnectState.DISCONNECTED
|
||||
|
||||
@property
|
||||
@@ -75,11 +76,26 @@ class FakeReconnectCoordinator:
|
||||
def attempt(self) -> int:
|
||||
return self._attempt
|
||||
|
||||
@property
|
||||
def generation(self) -> int:
|
||||
return self._generation
|
||||
|
||||
async def reconnect(self) -> None:
|
||||
self.reconnect_calls += 1
|
||||
self._attempt += 1
|
||||
self._generation += 1
|
||||
self._state = ReconnectState.CONNECTED
|
||||
|
||||
async def reconnect_after_transport_failure(
|
||||
self,
|
||||
*,
|
||||
observed_generation: int,
|
||||
) -> None:
|
||||
if observed_generation != self._generation:
|
||||
return
|
||||
|
||||
await self.reconnect()
|
||||
|
||||
|
||||
def create_supervisor() -> tuple[
|
||||
RuntimeSupervisor,
|
||||
@@ -271,6 +287,41 @@ def test_timeout_runs_single_reconnect_attempt() -> None:
|
||||
assert reconnect.reconnect_calls == 1
|
||||
|
||||
|
||||
def test_delayed_timeout_reuses_completed_transport_reconnect() -> None:
|
||||
async def scenario() -> tuple[
|
||||
RuntimeSupervisor,
|
||||
FakeHeartbeatMonitor,
|
||||
FakeReconnectCoordinator,
|
||||
bool,
|
||||
]:
|
||||
supervisor, heartbeat, reconnect = create_supervisor()
|
||||
supervisor.start()
|
||||
observed_generation = reconnect.generation
|
||||
|
||||
await reconnect.reconnect_after_transport_failure(
|
||||
observed_generation=observed_generation,
|
||||
)
|
||||
result = await supervisor.handle_heartbeat_timeout()
|
||||
|
||||
return (
|
||||
supervisor,
|
||||
heartbeat,
|
||||
reconnect,
|
||||
result,
|
||||
)
|
||||
|
||||
supervisor, heartbeat, reconnect, result = asyncio.run(
|
||||
scenario()
|
||||
)
|
||||
|
||||
assert result is True
|
||||
assert reconnect.reconnect_calls == 1
|
||||
assert reconnect.generation == 1
|
||||
assert heartbeat.stop_calls == 1
|
||||
assert heartbeat.start_calls == 2
|
||||
assert supervisor.state is RuntimeSupervisorState.RUNNING
|
||||
|
||||
|
||||
def test_successful_reconnect_restarts_heartbeat() -> None:
|
||||
supervisor, heartbeat, _ = create_supervisor()
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,252 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from websockets.protocol import State
|
||||
|
||||
from src.market_data.acquisition.adapters.dzengi.websocket_transport import (
|
||||
DzengiWebSocketTransport,
|
||||
)
|
||||
from src.market_data.acquisition.exceptions import (
|
||||
WebSocketTransportError,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.acquisition_runtime_service import (
|
||||
AcquisitionRuntimeService,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.reconnect import (
|
||||
ReconnectCoordinator,
|
||||
ReconnectState,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_events import (
|
||||
ReconnectCompletedEvent,
|
||||
ReconnectStartedEvent,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.transport_messages import (
|
||||
TransportTextMessage,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_session import (
|
||||
WebSocketSession,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_protocol import (
|
||||
AcquisitionRuntimeEvent,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_subscription_manager import (
|
||||
WebSocketSubscriptionManager,
|
||||
)
|
||||
|
||||
|
||||
SUBSCRIPTION_KEY = "trades:BTC/USD_LEVERAGE"
|
||||
SUBSCRIPTION_PAYLOAD = '{"destination":"trades.subscribe"}'
|
||||
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
send_error: Exception | None = None,
|
||||
) -> None:
|
||||
self.state = State.OPEN
|
||||
self.send_error = send_error
|
||||
self.sent_messages: list[str | bytes] = []
|
||||
self.close_calls = 0
|
||||
|
||||
async def close(
|
||||
self,
|
||||
code: int = 1000,
|
||||
reason: str = "",
|
||||
) -> None:
|
||||
self.close_calls += 1
|
||||
self.state = State.CLOSED
|
||||
|
||||
async def send(
|
||||
self,
|
||||
message: str | bytes,
|
||||
) -> None:
|
||||
if self.send_error is not None:
|
||||
self.state = State.CLOSED
|
||||
raise self.send_error
|
||||
|
||||
self.sent_messages.append(message)
|
||||
|
||||
async def recv(self) -> str | bytes:
|
||||
return ""
|
||||
|
||||
|
||||
class RecordingConnector:
|
||||
def __init__(
|
||||
self,
|
||||
*connections: FakeConnection,
|
||||
) -> None:
|
||||
self._connections = list(connections)
|
||||
self.calls: list[
|
||||
tuple[str, dict[str, Any]]
|
||||
] = []
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
url: str,
|
||||
**kwargs: Any,
|
||||
) -> FakeConnection:
|
||||
self.calls.append(
|
||||
(
|
||||
url,
|
||||
kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
return self._connections.pop(0)
|
||||
|
||||
|
||||
class RecordingEventPublisher:
|
||||
def __init__(self) -> None:
|
||||
self.events: list[AcquisitionRuntimeEvent] = []
|
||||
|
||||
async def publish(
|
||||
self,
|
||||
event: AcquisitionRuntimeEvent,
|
||||
) -> None:
|
||||
self.events.append(event)
|
||||
|
||||
|
||||
def create_runtime(
|
||||
*connections: FakeConnection,
|
||||
) -> tuple[
|
||||
WebSocketSession,
|
||||
WebSocketSubscriptionManager,
|
||||
ReconnectCoordinator,
|
||||
RecordingConnector,
|
||||
RecordingEventPublisher,
|
||||
]:
|
||||
connector = RecordingConnector(
|
||||
*connections,
|
||||
)
|
||||
transport = DzengiWebSocketTransport(
|
||||
url="wss://api-adapter.dzengi.com",
|
||||
connector=connector,
|
||||
)
|
||||
session = WebSocketSession(
|
||||
transport,
|
||||
)
|
||||
subscriptions = WebSocketSubscriptionManager(
|
||||
transport,
|
||||
)
|
||||
publisher = RecordingEventPublisher()
|
||||
runtime_service = AcquisitionRuntimeService(
|
||||
session=session,
|
||||
transport=transport,
|
||||
subscription_manager=subscriptions,
|
||||
event_publisher=publisher,
|
||||
)
|
||||
reconnect = ReconnectCoordinator(
|
||||
command_dispatcher=runtime_service,
|
||||
subscription_manager=subscriptions,
|
||||
event_publisher=publisher,
|
||||
)
|
||||
|
||||
return (
|
||||
session,
|
||||
subscriptions,
|
||||
reconnect,
|
||||
connector,
|
||||
publisher,
|
||||
)
|
||||
|
||||
|
||||
def test_reconnect_replaces_open_connection_before_restore() -> None:
|
||||
first_connection = FakeConnection()
|
||||
second_connection = FakeConnection()
|
||||
(
|
||||
session,
|
||||
subscriptions,
|
||||
reconnect,
|
||||
connector,
|
||||
publisher,
|
||||
) = create_runtime(
|
||||
first_connection,
|
||||
second_connection,
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
await session.start()
|
||||
await subscriptions.subscribe(
|
||||
SUBSCRIPTION_KEY,
|
||||
TransportTextMessage(
|
||||
payload=SUBSCRIPTION_PAYLOAD,
|
||||
),
|
||||
)
|
||||
await reconnect.reconnect()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert len(connector.calls) == 2
|
||||
assert first_connection.close_calls == 1
|
||||
assert first_connection.state is State.CLOSED
|
||||
assert first_connection.sent_messages == [
|
||||
SUBSCRIPTION_PAYLOAD,
|
||||
]
|
||||
assert second_connection.sent_messages == [
|
||||
SUBSCRIPTION_PAYLOAD,
|
||||
]
|
||||
assert session.is_connected is True
|
||||
assert reconnect.state is ReconnectState.CONNECTED
|
||||
assert publisher.events == [
|
||||
ReconnectStartedEvent(attempt=1),
|
||||
ReconnectCompletedEvent(attempt=1),
|
||||
]
|
||||
|
||||
|
||||
def test_reconnect_restores_subscription_after_initial_send_failure() -> None:
|
||||
first_connection = FakeConnection(
|
||||
send_error=RuntimeError(
|
||||
"socket dropped while subscribing",
|
||||
),
|
||||
)
|
||||
second_connection = FakeConnection()
|
||||
(
|
||||
session,
|
||||
subscriptions,
|
||||
reconnect,
|
||||
connector,
|
||||
publisher,
|
||||
) = create_runtime(
|
||||
first_connection,
|
||||
second_connection,
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
await session.start()
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketTransportError,
|
||||
match="socket dropped while subscribing",
|
||||
):
|
||||
await subscriptions.subscribe(
|
||||
SUBSCRIPTION_KEY,
|
||||
TransportTextMessage(
|
||||
payload=SUBSCRIPTION_PAYLOAD,
|
||||
),
|
||||
)
|
||||
|
||||
assert subscriptions.subscription_keys == (
|
||||
SUBSCRIPTION_KEY,
|
||||
)
|
||||
|
||||
await reconnect.reconnect()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert len(connector.calls) == 2
|
||||
assert second_connection.sent_messages == [
|
||||
SUBSCRIPTION_PAYLOAD,
|
||||
]
|
||||
assert subscriptions.subscription_keys == (
|
||||
SUBSCRIPTION_KEY,
|
||||
)
|
||||
assert session.is_connected is True
|
||||
assert reconnect.state is ReconnectState.CONNECTED
|
||||
assert publisher.events == [
|
||||
ReconnectStartedEvent(attempt=1),
|
||||
ReconnectCompletedEvent(attempt=1),
|
||||
]
|
||||
@@ -0,0 +1,196 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.runtime.websocket_protocol import (
|
||||
WebSocketSessionProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_session import (
|
||||
WebSocketSession,
|
||||
)
|
||||
|
||||
|
||||
class FakeTransport:
|
||||
def __init__(self) -> None:
|
||||
self.connected = False
|
||||
self.connect_calls = 0
|
||||
self.disconnect_calls = 0
|
||||
self.connect_error: Exception | None = None
|
||||
self.disconnect_error: Exception | None = None
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
return self.connected
|
||||
|
||||
async def connect(self) -> None:
|
||||
self.connect_calls += 1
|
||||
|
||||
if self.connect_error is not None:
|
||||
raise self.connect_error
|
||||
|
||||
self.connected = True
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
self.disconnect_calls += 1
|
||||
self.connected = False
|
||||
|
||||
if self.disconnect_error is not None:
|
||||
raise self.disconnect_error
|
||||
|
||||
async def send(
|
||||
self,
|
||||
message: str | bytes,
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
async def receive(self) -> str | bytes:
|
||||
return ""
|
||||
|
||||
|
||||
def create_session() -> tuple[
|
||||
WebSocketSession,
|
||||
FakeTransport,
|
||||
]:
|
||||
transport = FakeTransport()
|
||||
session = WebSocketSession(
|
||||
transport,
|
||||
)
|
||||
|
||||
return (
|
||||
session,
|
||||
transport,
|
||||
)
|
||||
|
||||
|
||||
def test_session_implements_protocol() -> None:
|
||||
session, _ = create_session()
|
||||
|
||||
assert isinstance(
|
||||
session,
|
||||
WebSocketSessionProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_session_uses_slots() -> None:
|
||||
session, _ = create_session()
|
||||
|
||||
assert not hasattr(session, "__dict__")
|
||||
|
||||
|
||||
def test_session_is_initially_disconnected() -> None:
|
||||
session, _ = create_session()
|
||||
|
||||
assert session.is_connected is False
|
||||
|
||||
|
||||
def test_start_connects_transport() -> None:
|
||||
session, transport = create_session()
|
||||
|
||||
asyncio.run(session.start())
|
||||
|
||||
assert session.is_connected is True
|
||||
assert transport.connect_calls == 1
|
||||
|
||||
|
||||
def test_start_is_idempotent() -> None:
|
||||
session, transport = create_session()
|
||||
|
||||
async def scenario() -> None:
|
||||
await session.start()
|
||||
await session.start()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert session.is_connected is True
|
||||
assert transport.connect_calls == 1
|
||||
|
||||
|
||||
def test_concurrent_start_creates_one_connection() -> None:
|
||||
session, transport = create_session()
|
||||
|
||||
async def scenario() -> None:
|
||||
await asyncio.gather(
|
||||
session.start(),
|
||||
session.start(),
|
||||
)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert session.is_connected is True
|
||||
assert transport.connect_calls == 1
|
||||
|
||||
|
||||
def test_start_error_leaves_session_disconnected() -> None:
|
||||
session, transport = create_session()
|
||||
transport.connect_error = RuntimeError("start failed")
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="start failed",
|
||||
):
|
||||
asyncio.run(session.start())
|
||||
|
||||
assert session.is_connected is False
|
||||
|
||||
|
||||
def test_start_reconnects_after_remote_disconnect() -> None:
|
||||
session, transport = create_session()
|
||||
|
||||
async def scenario() -> None:
|
||||
await session.start()
|
||||
transport.connected = False
|
||||
|
||||
assert session.is_connected is False
|
||||
|
||||
await session.start()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert session.is_connected is True
|
||||
assert transport.connect_calls == 2
|
||||
|
||||
|
||||
def test_stop_disconnects_transport() -> None:
|
||||
session, transport = create_session()
|
||||
|
||||
async def scenario() -> None:
|
||||
await session.start()
|
||||
await session.stop()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert session.is_connected is False
|
||||
assert transport.disconnect_calls == 1
|
||||
|
||||
|
||||
def test_stop_is_idempotent() -> None:
|
||||
session, transport = create_session()
|
||||
|
||||
async def scenario() -> None:
|
||||
await session.start()
|
||||
await session.stop()
|
||||
await session.stop()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.disconnect_calls == 1
|
||||
assert session.is_connected is False
|
||||
|
||||
|
||||
def test_stop_error_still_clears_session_state() -> None:
|
||||
session, transport = create_session()
|
||||
transport.disconnect_error = RuntimeError("stop failed")
|
||||
|
||||
async def scenario() -> None:
|
||||
await session.start()
|
||||
await session.stop()
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="stop failed",
|
||||
):
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert session.is_connected is False
|
||||
@@ -0,0 +1,419 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.acquisition.exceptions import (
|
||||
WebSocketUnsubscribeNotSupportedError,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.transport_messages import (
|
||||
TransportBinaryMessage,
|
||||
TransportTextMessage,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_protocol import (
|
||||
WebSocketSubscriptionManagerProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.websocket_subscription_manager import (
|
||||
WebSocketSubscriptionManager,
|
||||
)
|
||||
|
||||
|
||||
class RecordingTransport:
|
||||
def __init__(self) -> None:
|
||||
self.messages: list[str | bytes] = []
|
||||
self.send_error: Exception | None = None
|
||||
|
||||
async def connect(self) -> None:
|
||||
return None
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
return None
|
||||
|
||||
async def send(
|
||||
self,
|
||||
message: str | bytes,
|
||||
) -> None:
|
||||
if self.send_error is not None:
|
||||
raise self.send_error
|
||||
|
||||
self.messages.append(message)
|
||||
|
||||
async def receive(self) -> str | bytes:
|
||||
return ""
|
||||
|
||||
|
||||
def create_manager(
|
||||
*,
|
||||
supports_unsubscribe: bool = False,
|
||||
) -> tuple[
|
||||
WebSocketSubscriptionManager,
|
||||
RecordingTransport,
|
||||
]:
|
||||
transport = RecordingTransport()
|
||||
manager = WebSocketSubscriptionManager(
|
||||
transport,
|
||||
supports_unsubscribe=supports_unsubscribe,
|
||||
)
|
||||
|
||||
return (
|
||||
manager,
|
||||
transport,
|
||||
)
|
||||
|
||||
|
||||
def test_manager_implements_protocol() -> None:
|
||||
manager, _ = create_manager()
|
||||
|
||||
assert isinstance(
|
||||
manager,
|
||||
WebSocketSubscriptionManagerProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_manager_uses_slots() -> None:
|
||||
manager, _ = create_manager()
|
||||
|
||||
assert not hasattr(manager, "__dict__")
|
||||
|
||||
|
||||
def test_manager_starts_with_empty_registry() -> None:
|
||||
manager, _ = create_manager()
|
||||
|
||||
assert manager.subscription_keys == ()
|
||||
|
||||
|
||||
def test_subscribe_sends_and_registers_text_message() -> None:
|
||||
manager, transport = create_manager()
|
||||
message = TransportTextMessage(
|
||||
payload='{"destination":"trades.subscribe"}',
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
manager.subscribe(
|
||||
"trades:BTC",
|
||||
message,
|
||||
)
|
||||
)
|
||||
|
||||
assert transport.messages == [
|
||||
message.payload,
|
||||
]
|
||||
assert manager.subscription_keys == (
|
||||
"trades:BTC",
|
||||
)
|
||||
|
||||
|
||||
def test_subscribe_sends_binary_message() -> None:
|
||||
manager, transport = create_manager()
|
||||
message = TransportBinaryMessage(
|
||||
payload=b"\x01\x02",
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
manager.subscribe(
|
||||
"binary",
|
||||
message,
|
||||
)
|
||||
)
|
||||
|
||||
assert transport.messages == [
|
||||
b"\x01\x02",
|
||||
]
|
||||
assert manager.subscription_keys == (
|
||||
"binary",
|
||||
)
|
||||
|
||||
|
||||
def test_duplicate_subscription_key_is_idempotent() -> None:
|
||||
manager, transport = create_manager()
|
||||
first_message = TransportTextMessage(
|
||||
payload="first",
|
||||
)
|
||||
second_message = TransportTextMessage(
|
||||
payload="second",
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
await manager.subscribe(
|
||||
"trades:BTC",
|
||||
first_message,
|
||||
)
|
||||
await manager.subscribe(
|
||||
"trades:BTC",
|
||||
second_message,
|
||||
)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.messages == [
|
||||
"first",
|
||||
]
|
||||
assert manager.subscription_keys == (
|
||||
"trades:BTC",
|
||||
)
|
||||
|
||||
|
||||
def test_failed_subscribe_remains_registered_as_pending() -> None:
|
||||
manager, transport = create_manager()
|
||||
transport.send_error = RuntimeError("send failed")
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="send failed",
|
||||
):
|
||||
asyncio.run(
|
||||
manager.subscribe(
|
||||
"trades:BTC",
|
||||
TransportTextMessage(
|
||||
payload="subscribe",
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
assert manager.subscription_keys == (
|
||||
"trades:BTC",
|
||||
)
|
||||
|
||||
|
||||
def test_pending_subscription_can_be_retried() -> None:
|
||||
manager, transport = create_manager()
|
||||
transport.send_error = RuntimeError("send failed")
|
||||
|
||||
async def scenario() -> None:
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="send failed",
|
||||
):
|
||||
await manager.subscribe(
|
||||
"trades:BTC",
|
||||
TransportTextMessage(
|
||||
payload="first-attempt",
|
||||
),
|
||||
)
|
||||
|
||||
transport.send_error = None
|
||||
|
||||
await manager.subscribe(
|
||||
"trades:BTC",
|
||||
TransportTextMessage(
|
||||
payload="second-attempt",
|
||||
),
|
||||
)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.messages == [
|
||||
"second-attempt",
|
||||
]
|
||||
assert manager.subscription_keys == (
|
||||
"trades:BTC",
|
||||
)
|
||||
|
||||
|
||||
def test_restore_sends_registered_messages_in_order() -> None:
|
||||
manager, transport = create_manager()
|
||||
|
||||
async def scenario() -> None:
|
||||
await manager.subscribe(
|
||||
"first",
|
||||
TransportTextMessage(
|
||||
payload="first-message",
|
||||
),
|
||||
)
|
||||
await manager.subscribe(
|
||||
"second",
|
||||
TransportBinaryMessage(
|
||||
payload=b"second-message",
|
||||
),
|
||||
)
|
||||
|
||||
transport.messages.clear()
|
||||
await manager.restore_subscriptions()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.messages == [
|
||||
"first-message",
|
||||
b"second-message",
|
||||
]
|
||||
|
||||
|
||||
def test_restore_error_preserves_registry() -> None:
|
||||
manager, transport = create_manager()
|
||||
|
||||
async def prepare() -> None:
|
||||
await manager.subscribe(
|
||||
"trades:BTC",
|
||||
TransportTextMessage(
|
||||
payload="subscribe",
|
||||
),
|
||||
)
|
||||
|
||||
asyncio.run(prepare())
|
||||
transport.send_error = RuntimeError("restore failed")
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError,
|
||||
match="restore failed",
|
||||
):
|
||||
asyncio.run(manager.restore_subscriptions())
|
||||
|
||||
assert manager.subscription_keys == (
|
||||
"trades:BTC",
|
||||
)
|
||||
|
||||
|
||||
def test_unsubscribe_is_explicitly_unsupported_by_default() -> None:
|
||||
manager, transport = create_manager()
|
||||
|
||||
async def scenario() -> None:
|
||||
await manager.subscribe(
|
||||
"trades:BTC",
|
||||
TransportTextMessage(
|
||||
payload="subscribe",
|
||||
),
|
||||
)
|
||||
await manager.unsubscribe(
|
||||
"trades:BTC",
|
||||
TransportTextMessage(
|
||||
payload="unsubscribe",
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
WebSocketUnsubscribeNotSupportedError,
|
||||
):
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.messages == [
|
||||
"subscribe",
|
||||
]
|
||||
assert manager.subscription_keys == (
|
||||
"trades:BTC",
|
||||
)
|
||||
|
||||
|
||||
def test_supported_unsubscribe_sends_and_removes_subscription() -> None:
|
||||
manager, transport = create_manager(
|
||||
supports_unsubscribe=True,
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
await manager.subscribe(
|
||||
"generic",
|
||||
TransportTextMessage(
|
||||
payload="subscribe",
|
||||
),
|
||||
)
|
||||
await manager.unsubscribe(
|
||||
"generic",
|
||||
TransportTextMessage(
|
||||
payload="unsubscribe",
|
||||
),
|
||||
)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert transport.messages == [
|
||||
"subscribe",
|
||||
"unsubscribe",
|
||||
]
|
||||
assert manager.subscription_keys == ()
|
||||
|
||||
|
||||
def test_supported_unsubscribe_is_idempotent_for_missing_key() -> None:
|
||||
manager, transport = create_manager(
|
||||
supports_unsubscribe=True,
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
manager.unsubscribe(
|
||||
"missing",
|
||||
TransportTextMessage(
|
||||
payload="unsubscribe",
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
assert transport.messages == []
|
||||
assert manager.subscription_keys == ()
|
||||
|
||||
|
||||
def test_clear_removes_registry_without_sending_messages() -> None:
|
||||
manager, transport = create_manager()
|
||||
|
||||
async def scenario() -> None:
|
||||
await manager.subscribe(
|
||||
"trades:BTC",
|
||||
TransportTextMessage(
|
||||
payload="subscribe",
|
||||
),
|
||||
)
|
||||
|
||||
transport.messages.clear()
|
||||
await manager.clear_subscriptions()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
assert manager.subscription_keys == ()
|
||||
assert transport.messages == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"subscription_key",
|
||||
[
|
||||
"",
|
||||
" ",
|
||||
],
|
||||
)
|
||||
def test_rejects_empty_subscription_key(
|
||||
subscription_key: str,
|
||||
) -> None:
|
||||
manager, _ = create_manager()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
asyncio.run(
|
||||
manager.subscribe(
|
||||
subscription_key,
|
||||
TransportTextMessage(
|
||||
payload="subscribe",
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_non_string_subscription_key() -> None:
|
||||
manager, _ = create_manager()
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
asyncio.run(
|
||||
manager.subscribe(
|
||||
123, # type: ignore[arg-type]
|
||||
TransportTextMessage(
|
||||
payload="subscribe",
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_unsupported_message_type() -> None:
|
||||
manager, _ = create_manager()
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
asyncio.run(
|
||||
manager.subscribe(
|
||||
"trades:BTC",
|
||||
object(), # type: ignore[arg-type]
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_non_boolean_unsubscribe_capability() -> None:
|
||||
transport = RecordingTransport()
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
WebSocketSubscriptionManager(
|
||||
transport,
|
||||
supports_unsubscribe="yes", # type: ignore[arg-type]
|
||||
)
|
||||
@@ -28,6 +28,12 @@ from src.market_data.acquisition.runtime.reconnect import (
|
||||
ReconnectCoordinatorProtocol,
|
||||
ReconnectState,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_reconnect_recovery_coordinator import (
|
||||
RuntimeReconnectRecoveryProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_liveness_probe import (
|
||||
RuntimeLivenessProbeProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.runtime.runtime_events import (
|
||||
ReconnectCompletedEvent,
|
||||
ReconnectStartedEvent,
|
||||
@@ -123,11 +129,17 @@ class FakeSession:
|
||||
|
||||
|
||||
class FakeTransport:
|
||||
def __init__(self) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
probe_results: tuple[bool, ...] = (True,),
|
||||
) -> None:
|
||||
self.connect_calls = 0
|
||||
self.disconnect_calls = 0
|
||||
self.sent_messages: list[str | bytes] = []
|
||||
self.receive_calls = 0
|
||||
self.probe_calls = 0
|
||||
self._probe_results = list(probe_results)
|
||||
|
||||
async def connect(self) -> None:
|
||||
self.connect_calls += 1
|
||||
@@ -145,6 +157,14 @@ class FakeTransport:
|
||||
self.receive_calls += 1
|
||||
return ""
|
||||
|
||||
async def probe(self) -> bool:
|
||||
self.probe_calls += 1
|
||||
|
||||
if not self._probe_results:
|
||||
return True
|
||||
|
||||
return self._probe_results.pop(0)
|
||||
|
||||
|
||||
class FakeSubscriptionManager:
|
||||
def __init__(self) -> None:
|
||||
@@ -264,6 +284,25 @@ class FakeClock:
|
||||
def __call__(self) -> float:
|
||||
return self.value
|
||||
|
||||
def advance(
|
||||
self,
|
||||
seconds: float,
|
||||
) -> None:
|
||||
self.value += seconds
|
||||
|
||||
|
||||
class FakeUnixTimeClock:
|
||||
def __init__(
|
||||
self,
|
||||
value: int = RECOVERY_END_TIME_MS,
|
||||
) -> None:
|
||||
self.value = value
|
||||
self.calls = 0
|
||||
|
||||
def __call__(self) -> int:
|
||||
self.calls += 1
|
||||
return self.value
|
||||
|
||||
|
||||
class RecordingSleep:
|
||||
def __init__(self) -> None:
|
||||
@@ -285,6 +324,7 @@ class CompositionDependencies:
|
||||
message_adapter: FakeMessageAdapter
|
||||
recovery_document_source: StubTradesDocumentSource
|
||||
heartbeat_clock: FakeClock
|
||||
recovery_end_time_clock: FakeUnixTimeClock
|
||||
scheduler_sleep: RecordingSleep
|
||||
|
||||
|
||||
@@ -295,13 +335,16 @@ def create_composition(
|
||||
heartbeat_timeout_seconds: float = 10.0,
|
||||
scheduler_interval_seconds: float = 1.0,
|
||||
max_recovery_window_ms: int = 3_599_999,
|
||||
probe_results: tuple[bool, ...] = (True,),
|
||||
) -> tuple[
|
||||
TradeStreamRuntimeComposition,
|
||||
CompositionDependencies,
|
||||
]:
|
||||
dependencies = CompositionDependencies(
|
||||
session=FakeSession(),
|
||||
transport=FakeTransport(),
|
||||
transport=FakeTransport(
|
||||
probe_results=probe_results,
|
||||
),
|
||||
subscription_manager=FakeSubscriptionManager(),
|
||||
event_publisher=FakeEventPublisher(),
|
||||
message_adapter=FakeMessageAdapter(
|
||||
@@ -311,6 +354,7 @@ def create_composition(
|
||||
recovery_document,
|
||||
),
|
||||
heartbeat_clock=FakeClock(),
|
||||
recovery_end_time_clock=FakeUnixTimeClock(),
|
||||
scheduler_sleep=RecordingSleep(),
|
||||
)
|
||||
|
||||
@@ -323,10 +367,14 @@ def create_composition(
|
||||
recovery_document_source=(
|
||||
dependencies.recovery_document_source
|
||||
),
|
||||
symbols=(SYMBOL,),
|
||||
heartbeat_timeout_seconds=heartbeat_timeout_seconds,
|
||||
scheduler_interval_seconds=scheduler_interval_seconds,
|
||||
max_recovery_window_ms=max_recovery_window_ms,
|
||||
heartbeat_clock=dependencies.heartbeat_clock,
|
||||
recovery_end_time_clock=(
|
||||
dependencies.recovery_end_time_clock
|
||||
),
|
||||
scheduler_sleep=dependencies.scheduler_sleep,
|
||||
)
|
||||
|
||||
@@ -373,6 +421,14 @@ def test_components_implement_public_protocols() -> None:
|
||||
composition.reconnect_coordinator,
|
||||
ReconnectCoordinatorProtocol,
|
||||
)
|
||||
assert isinstance(
|
||||
composition.runtime_reconnect_recovery_coordinator,
|
||||
RuntimeReconnectRecoveryProtocol,
|
||||
)
|
||||
assert isinstance(
|
||||
composition.liveness_probe,
|
||||
RuntimeLivenessProbeProtocol,
|
||||
)
|
||||
assert isinstance(
|
||||
composition.heartbeat_monitor,
|
||||
HeartbeatMonitorProtocol,
|
||||
@@ -460,8 +516,27 @@ def test_runtime_components_share_lifecycle_dependencies() -> None:
|
||||
)
|
||||
assert (
|
||||
composition.runtime_supervisor._reconnect_coordinator
|
||||
is composition.runtime_reconnect_recovery_coordinator
|
||||
)
|
||||
assert (
|
||||
composition.runtime_reconnect_recovery_coordinator
|
||||
._reconnect_coordinator
|
||||
is composition.reconnect_coordinator
|
||||
)
|
||||
assert (
|
||||
composition.runtime_reconnect_recovery_coordinator
|
||||
._recovery_coordinator
|
||||
is composition.runtime_recovery_coordinator
|
||||
)
|
||||
assert (
|
||||
composition.runtime_reconnect_recovery_coordinator
|
||||
.live_processing_gate
|
||||
is composition.live_processing_gate
|
||||
)
|
||||
assert (
|
||||
composition.runtime_reconnect_recovery_coordinator.symbols
|
||||
== (SYMBOL,)
|
||||
)
|
||||
assert (
|
||||
composition.runtime_scheduler._heartbeat_monitor
|
||||
is composition.heartbeat_monitor
|
||||
@@ -470,6 +545,14 @@ def test_runtime_components_share_lifecycle_dependencies() -> None:
|
||||
composition.runtime_scheduler._runtime_supervisor
|
||||
is composition.runtime_supervisor
|
||||
)
|
||||
assert (
|
||||
composition.liveness_probe
|
||||
is dependencies.transport
|
||||
)
|
||||
assert (
|
||||
composition.runtime_scheduler._liveness_probe
|
||||
is composition.liveness_probe
|
||||
)
|
||||
|
||||
|
||||
def test_configuration_is_forwarded() -> None:
|
||||
@@ -507,6 +590,7 @@ def test_creation_has_no_runtime_side_effects() -> None:
|
||||
assert dependencies.transport.disconnect_calls == 0
|
||||
assert dependencies.transport.sent_messages == []
|
||||
assert dependencies.transport.receive_calls == 0
|
||||
assert dependencies.transport.probe_calls == 0
|
||||
|
||||
assert dependencies.subscription_manager.subscriptions == []
|
||||
assert dependencies.subscription_manager.unsubscriptions == []
|
||||
@@ -515,6 +599,7 @@ def test_creation_has_no_runtime_side_effects() -> None:
|
||||
|
||||
assert dependencies.event_publisher.events == []
|
||||
assert dependencies.recovery_document_source.calls == []
|
||||
assert dependencies.recovery_end_time_clock.calls == 0
|
||||
|
||||
assert composition.heartbeat_monitor.state is HeartbeatState.IDLE
|
||||
assert (
|
||||
@@ -528,6 +613,103 @@ def test_creation_has_no_runtime_side_effects() -> None:
|
||||
assert composition.runtime_scheduler.running is False
|
||||
|
||||
|
||||
def test_successful_probe_keeps_quiet_connection_alive() -> None:
|
||||
composition, dependencies = create_composition(
|
||||
heartbeat_timeout_seconds=10.0,
|
||||
probe_results=(True,),
|
||||
)
|
||||
composition.runtime_supervisor.start()
|
||||
dependencies.heartbeat_clock.advance(10.0)
|
||||
|
||||
timed_out = asyncio.run(
|
||||
composition.runtime_scheduler.run_once()
|
||||
)
|
||||
|
||||
assert timed_out is False
|
||||
assert dependencies.transport.probe_calls == 1
|
||||
assert dependencies.session.start_calls == 0
|
||||
assert dependencies.session.stop_calls == 0
|
||||
assert composition.heartbeat_monitor.last_activity_at == 110.0
|
||||
assert (
|
||||
composition.runtime_supervisor.state
|
||||
is RuntimeSupervisorState.RUNNING
|
||||
)
|
||||
|
||||
|
||||
def test_failed_probe_triggers_reconnect_after_timeout() -> None:
|
||||
composition, dependencies = create_composition(
|
||||
heartbeat_timeout_seconds=10.0,
|
||||
probe_results=(False,),
|
||||
)
|
||||
composition.runtime_supervisor.start()
|
||||
dependencies.heartbeat_clock.advance(10.0)
|
||||
|
||||
timed_out = asyncio.run(
|
||||
composition.runtime_scheduler.run_once()
|
||||
)
|
||||
|
||||
assert timed_out is True
|
||||
assert dependencies.transport.probe_calls == 1
|
||||
assert dependencies.session.stop_calls == 1
|
||||
assert dependencies.session.start_calls == 1
|
||||
assert dependencies.subscription_manager.restore_calls == 1
|
||||
assert dependencies.recovery_end_time_clock.calls == 1
|
||||
assert (
|
||||
composition.runtime_supervisor.state
|
||||
is RuntimeSupervisorState.RUNNING
|
||||
)
|
||||
|
||||
|
||||
def test_delayed_timeout_does_not_repeat_completed_transport_reconnect() -> None:
|
||||
async def scenario() -> tuple[
|
||||
TradeStreamRuntimeComposition,
|
||||
CompositionDependencies,
|
||||
bool,
|
||||
]:
|
||||
composition, dependencies = create_composition(
|
||||
heartbeat_timeout_seconds=10.0,
|
||||
probe_results=(False,),
|
||||
)
|
||||
composition.runtime_supervisor.start()
|
||||
observed_generation = (
|
||||
composition.runtime_reconnect_recovery_coordinator.generation
|
||||
)
|
||||
|
||||
await (
|
||||
composition.runtime_reconnect_recovery_coordinator
|
||||
.reconnect_after_transport_failure(
|
||||
observed_generation=observed_generation,
|
||||
)
|
||||
)
|
||||
dependencies.heartbeat_clock.advance(10.0)
|
||||
timed_out = await composition.runtime_scheduler.run_once()
|
||||
|
||||
return (
|
||||
composition,
|
||||
dependencies,
|
||||
timed_out,
|
||||
)
|
||||
|
||||
composition, dependencies, timed_out = asyncio.run(
|
||||
scenario()
|
||||
)
|
||||
|
||||
assert timed_out is True
|
||||
assert dependencies.transport.probe_calls == 1
|
||||
assert dependencies.session.stop_calls == 1
|
||||
assert dependencies.session.start_calls == 1
|
||||
assert dependencies.subscription_manager.restore_calls == 1
|
||||
assert dependencies.recovery_end_time_clock.calls == 1
|
||||
assert (
|
||||
composition.runtime_reconnect_recovery_coordinator.generation
|
||||
== 1
|
||||
)
|
||||
assert (
|
||||
composition.runtime_supervisor.state
|
||||
is RuntimeSupervisorState.RUNNING
|
||||
)
|
||||
|
||||
|
||||
def test_live_checkpoint_is_visible_to_runtime_recovery() -> None:
|
||||
checkpoint_trade = make_trade()
|
||||
|
||||
@@ -642,6 +824,35 @@ def test_reconnect_uses_composed_runtime_dependencies() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_runtime_reconnect_runs_recovery_after_subscription_restore() -> None:
|
||||
composition, dependencies = create_composition(
|
||||
recovery_document=[],
|
||||
)
|
||||
|
||||
composition.trade_stream_acquisition_service.handle_message(
|
||||
{
|
||||
"destination": "internal.trade",
|
||||
}
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
composition.runtime_reconnect_recovery_coordinator.reconnect()
|
||||
)
|
||||
|
||||
assert dependencies.session.stop_calls == 1
|
||||
assert dependencies.session.start_calls == 1
|
||||
assert dependencies.subscription_manager.restore_calls == 1
|
||||
assert dependencies.recovery_end_time_clock.calls == 1
|
||||
assert dependencies.recovery_document_source.calls == [
|
||||
(
|
||||
SYMBOL,
|
||||
CHECKPOINT_TIME_MS,
|
||||
RECOVERY_END_TIME_MS,
|
||||
None,
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def test_separate_compositions_have_independent_state() -> None:
|
||||
first, _ = create_composition()
|
||||
second, _ = create_composition()
|
||||
@@ -690,6 +901,7 @@ def test_invalid_heartbeat_configuration_is_not_hidden(
|
||||
[],
|
||||
),
|
||||
heartbeat_clock=FakeClock(),
|
||||
recovery_end_time_clock=FakeUnixTimeClock(),
|
||||
scheduler_sleep=RecordingSleep(),
|
||||
)
|
||||
|
||||
@@ -705,9 +917,13 @@ def test_invalid_heartbeat_configuration_is_not_hidden(
|
||||
recovery_document_source=(
|
||||
dependencies.recovery_document_source
|
||||
),
|
||||
symbols=(SYMBOL,),
|
||||
heartbeat_timeout_seconds=heartbeat_timeout_seconds, # type: ignore[arg-type]
|
||||
scheduler_interval_seconds=1.0,
|
||||
heartbeat_clock=dependencies.heartbeat_clock,
|
||||
recovery_end_time_clock=(
|
||||
dependencies.recovery_end_time_clock
|
||||
),
|
||||
scheduler_sleep=dependencies.scheduler_sleep,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user