Build 060.25: implement Production Runtime Integration

This commit is contained in:
2026-07-31 00:29:36 +03:00
parent c142145361
commit 60bec1eaf9
50 changed files with 14044 additions and 83 deletions

View File

@@ -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,
)

View File

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

View File

@@ -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",
}

View File

@@ -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]
)
)

View File

@@ -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

View File

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

View File

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

View File

@@ -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

View File

@@ -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,
*,

View File

@@ -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

View File

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

View File

@@ -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),
]

View File

@@ -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

View File

@@ -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]
)

View File

@@ -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,
)