Build 060.29: implement Market Data Access and Replay
This commit is contained in:
28
app/tests/unit/market_data/replay/conftest.py
Normal file
28
app/tests/unit/market_data/replay/conftest.py
Normal file
@@ -0,0 +1,28 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.replay import (
|
||||
ReplayDataType,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def empty_replay_plan() -> ReplayPlan:
|
||||
start = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
request = ReplayPlanRequest(
|
||||
venue="dzengi",
|
||||
symbols=("BTC/USD_LEVERAGE",),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=start,
|
||||
end_time=start + timedelta(hours=1),
|
||||
),
|
||||
max_records=100,
|
||||
)
|
||||
return ReplayPlan(request=request, events=())
|
||||
@@ -0,0 +1,214 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
import inspect
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.replay.contracts import (
|
||||
MarketDataClockProtocol,
|
||||
ReplayClockProtocol,
|
||||
)
|
||||
from src.market_data.replay.deterministic_replay_clock import (
|
||||
DeterministicReplayClock,
|
||||
)
|
||||
from src.market_data.replay.exceptions import ReplayClockError
|
||||
|
||||
|
||||
START = datetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
0,
|
||||
123456,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
|
||||
|
||||
def test_matches_protocols_uses_slots_and_is_synchronous() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
assert isinstance(clock, MarketDataClockProtocol)
|
||||
assert isinstance(clock, ReplayClockProtocol)
|
||||
assert not hasattr(clock, "__dict__")
|
||||
assert inspect.iscoroutinefunction(clock.advance_to) is False
|
||||
|
||||
|
||||
def test_starts_at_exact_canonical_utc_time() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
assert clock.now == START
|
||||
assert type(clock.now) is datetime
|
||||
assert clock.now.tzinfo is timezone.utc
|
||||
|
||||
|
||||
def test_normalizes_non_utc_initial_time_without_losing_precision() -> None:
|
||||
offset = timezone(timedelta(hours=3, minutes=30))
|
||||
initial_time = datetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
15,
|
||||
30,
|
||||
0,
|
||||
654321,
|
||||
tzinfo=offset,
|
||||
)
|
||||
|
||||
clock = DeterministicReplayClock(initial_time)
|
||||
|
||||
assert clock.now == datetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
0,
|
||||
654321,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
assert type(clock.now) is datetime
|
||||
|
||||
|
||||
def test_accepts_datetime_subclass_but_stores_base_datetime() -> None:
|
||||
class CompatibleDatetime(datetime):
|
||||
pass
|
||||
|
||||
initial_time = CompatibleDatetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
|
||||
clock = DeterministicReplayClock(initial_time)
|
||||
|
||||
assert clock.now == initial_time
|
||||
assert type(clock.now) is datetime
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, "2026-08-02", 1, object()))
|
||||
def test_rejects_non_datetime_initial_value(invalid: object) -> None:
|
||||
with pytest.raises(TypeError, match="initial_time"):
|
||||
DeterministicReplayClock(invalid) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def test_rejects_naive_initial_time() -> None:
|
||||
with pytest.raises(ValueError, match="timezone"):
|
||||
DeterministicReplayClock(START.replace(tzinfo=None))
|
||||
|
||||
|
||||
def test_advances_forward_with_microsecond_precision() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
expected = START + timedelta(microseconds=1)
|
||||
|
||||
result = clock.advance_to(expected)
|
||||
|
||||
assert result is None
|
||||
assert clock.now == expected
|
||||
|
||||
|
||||
def test_allows_repeated_equal_time_for_distinct_replay_events() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
clock.advance_to(START)
|
||||
clock.advance_to(START)
|
||||
|
||||
assert clock.now == START
|
||||
|
||||
|
||||
def test_equal_instant_with_another_offset_is_idempotent() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
equal_instant = START.astimezone(timezone(timedelta(hours=-4)))
|
||||
|
||||
clock.advance_to(equal_instant)
|
||||
|
||||
assert clock.now == START
|
||||
assert clock.now.tzinfo is timezone.utc
|
||||
|
||||
|
||||
def test_equal_utc_instant_with_another_fold_is_complete_no_op() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
previous = clock.now
|
||||
|
||||
clock.advance_to(START.replace(fold=1))
|
||||
|
||||
assert clock.now is previous
|
||||
assert clock.now.fold == 0
|
||||
|
||||
|
||||
def test_accepts_datetime_subclass_during_advance() -> None:
|
||||
class CompatibleDatetime(datetime):
|
||||
pass
|
||||
|
||||
clock = DeterministicReplayClock(START)
|
||||
later = CompatibleDatetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
1,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
|
||||
clock.advance_to(later)
|
||||
|
||||
assert clock.now == later
|
||||
assert type(clock.now) is datetime
|
||||
|
||||
|
||||
def test_rejects_backward_transition_without_changing_state() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
later = START + timedelta(minutes=1)
|
||||
clock.advance_to(later)
|
||||
|
||||
with pytest.raises(ReplayClockError, match="backwards"):
|
||||
clock.advance_to(START)
|
||||
|
||||
assert clock.now == later
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, "later", 1, object()))
|
||||
def test_invalid_advance_type_does_not_change_state(invalid: object) -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
with pytest.raises(TypeError, match="instant"):
|
||||
clock.advance_to(invalid) # type: ignore[arg-type]
|
||||
|
||||
assert clock.now == START
|
||||
|
||||
|
||||
def test_naive_advance_does_not_change_state() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
with pytest.raises(ValueError, match="timezone"):
|
||||
clock.advance_to(START.replace(tzinfo=None))
|
||||
|
||||
assert clock.now == START
|
||||
|
||||
|
||||
def test_two_clocks_have_independent_state() -> None:
|
||||
first = DeterministicReplayClock(START)
|
||||
second = DeterministicReplayClock(START)
|
||||
|
||||
first.advance_to(START + timedelta(hours=1))
|
||||
|
||||
assert first.now == START + timedelta(hours=1)
|
||||
assert second.now == START
|
||||
|
||||
|
||||
def test_now_is_read_only_and_lifecycle_extensions_are_absent() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(clock, "now", START + timedelta(hours=1))
|
||||
|
||||
assert not hasattr(clock, "reset")
|
||||
assert not hasattr(clock, "advance_by")
|
||||
assert not hasattr(clock, "start")
|
||||
assert not hasattr(clock, "stop")
|
||||
assert clock.now == START
|
||||
@@ -0,0 +1,94 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from src.market_data.replay import (
|
||||
MarketDataClockProtocol,
|
||||
MarketDataReplayError,
|
||||
MarketDataReplayValidationError,
|
||||
ReplayClockError,
|
||||
ReplayClockProtocol,
|
||||
ReplayConsumerProtocol,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanBuilderProtocol,
|
||||
ReplayPlanLimitExceededError,
|
||||
ReplayPlanRequest,
|
||||
ReplaySessionProtocol,
|
||||
ReplaySessionState,
|
||||
ReplaySessionStateError,
|
||||
)
|
||||
|
||||
|
||||
class FakeClock:
|
||||
def __init__(self, now: datetime) -> None:
|
||||
self._now = now
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return self._now
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self._now = instant
|
||||
|
||||
|
||||
class RecordingConsumer:
|
||||
def __init__(self) -> None:
|
||||
self.events: list[ReplayEvent] = []
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.events.append(event)
|
||||
|
||||
|
||||
class RecordingPlanBuilder:
|
||||
def __init__(self, plan: ReplayPlan) -> None:
|
||||
self.plan = plan
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
return self.plan
|
||||
|
||||
|
||||
class RecordingSession:
|
||||
def __init__(self, plan: ReplayPlan, clock: FakeClock) -> None:
|
||||
self._plan = plan
|
||||
self._clock = clock
|
||||
|
||||
@property
|
||||
def state(self) -> ReplaySessionState:
|
||||
return ReplaySessionState.CREATED
|
||||
|
||||
@property
|
||||
def plan(self) -> ReplayPlan:
|
||||
return self._plan
|
||||
|
||||
@property
|
||||
def clock(self) -> MarketDataClockProtocol:
|
||||
return self._clock
|
||||
|
||||
async def run(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def test_replay_protocols_are_runtime_checkable(
|
||||
empty_replay_plan: ReplayPlan,
|
||||
) -> None:
|
||||
clock = FakeClock(empty_replay_plan.request.time_range.start_time)
|
||||
|
||||
assert isinstance(clock, MarketDataClockProtocol)
|
||||
assert isinstance(clock, ReplayClockProtocol)
|
||||
assert isinstance(RecordingConsumer(), ReplayConsumerProtocol)
|
||||
assert isinstance(
|
||||
RecordingPlanBuilder(empty_replay_plan),
|
||||
ReplayPlanBuilderProtocol,
|
||||
)
|
||||
assert isinstance(
|
||||
RecordingSession(empty_replay_plan, clock),
|
||||
ReplaySessionProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_replay_error_hierarchy_is_specialized() -> None:
|
||||
assert issubclass(MarketDataReplayValidationError, MarketDataReplayError)
|
||||
assert issubclass(ReplayPlanLimitExceededError, MarketDataReplayError)
|
||||
assert issubclass(ReplayClockError, MarketDataReplayError)
|
||||
assert issubclass(ReplaySessionStateError, MarketDataReplayError)
|
||||
@@ -0,0 +1,659 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import FrozenInstanceError
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.candle import Candle
|
||||
from src.market_data.acquisition.models.quote import Quote
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.replay import (
|
||||
REPLAY_PLAN_MAX_RECORDS_LIMIT,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
ReplaySessionState,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
|
||||
|
||||
def make_trade(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
executed_at: datetime = START + timedelta(minutes=1),
|
||||
symbol: str = SYMBOL,
|
||||
) -> Trade:
|
||||
return Trade(
|
||||
symbol=symbol,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=executed_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source="dzengi_websocket_trade",
|
||||
)
|
||||
|
||||
|
||||
def make_quote(
|
||||
*,
|
||||
received_at: datetime = START + timedelta(minutes=2),
|
||||
) -> Quote:
|
||||
return Quote(
|
||||
symbol=SYMBOL,
|
||||
last_price=Decimal("65000"),
|
||||
bid_price=Decimal("64999"),
|
||||
ask_price=Decimal("65001"),
|
||||
exchange_timestamp=received_at - timedelta(milliseconds=1),
|
||||
received_at=received_at,
|
||||
source="dzengi_rest_quote",
|
||||
)
|
||||
|
||||
|
||||
def make_candle(*, interval: str = "1m") -> Candle:
|
||||
return Candle(
|
||||
symbol=SYMBOL,
|
||||
interval=interval,
|
||||
open_time=START,
|
||||
open_price=Decimal("64900"),
|
||||
high_price=Decimal("65100"),
|
||||
low_price=Decimal("64800"),
|
||||
close_price=Decimal("65000"),
|
||||
volume=Decimal("10"),
|
||||
source="dzengi_rest_candle",
|
||||
)
|
||||
|
||||
|
||||
def make_request(
|
||||
*,
|
||||
data_types: tuple[ReplayDataType, ...] = (ReplayDataType.TRADE,),
|
||||
candle_intervals: tuple[str, ...] = (),
|
||||
max_records: int = 100,
|
||||
) -> ReplayPlanRequest:
|
||||
return ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=data_types,
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=START,
|
||||
end_time=END,
|
||||
),
|
||||
candle_intervals=candle_intervals,
|
||||
max_records=max_records,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_event(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
replay_at: datetime = START + timedelta(minutes=1),
|
||||
replay_sequence: int = 1,
|
||||
venue: str = VENUE,
|
||||
symbol: str = SYMBOL,
|
||||
) -> ReplayEvent:
|
||||
trade = make_trade(
|
||||
trade_id=trade_id,
|
||||
executed_at=replay_at,
|
||||
symbol=symbol,
|
||||
)
|
||||
return ReplayEvent(
|
||||
venue=venue,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=replay_sequence,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def trade_id_from_event(event: ReplayEvent) -> int:
|
||||
payload = event.payload
|
||||
|
||||
if not isinstance(payload, Trade):
|
||||
raise AssertionError("Replay event must contain Trade payload.")
|
||||
|
||||
return payload.trade_id
|
||||
|
||||
|
||||
def test_trade_event_preserves_payload_identity_and_exact_event_time() -> None:
|
||||
trade = make_trade()
|
||||
event = ReplayEvent(
|
||||
venue=" dzengi ",
|
||||
replay_at=trade.executed_at,
|
||||
replay_sequence=10,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
assert event.venue == VENUE
|
||||
assert event.payload is trade
|
||||
assert event.symbol == SYMBOL
|
||||
assert event.data_type is ReplayDataType.TRADE
|
||||
assert event.order_key == (trade.executed_at, 10)
|
||||
assert event.candle_is_final is None
|
||||
|
||||
|
||||
def test_quote_event_uses_received_time_not_exchange_time() -> None:
|
||||
quote = make_quote()
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=quote.received_at,
|
||||
replay_sequence=11,
|
||||
payload=quote,
|
||||
)
|
||||
|
||||
assert event.data_type is ReplayDataType.QUOTE
|
||||
assert event.replay_at == quote.received_at
|
||||
assert event.replay_at != quote.exchange_timestamp
|
||||
|
||||
|
||||
def test_candle_event_uses_observation_time_and_final_metadata() -> None:
|
||||
candle = make_candle()
|
||||
observed_at = candle.open_time + timedelta(seconds=30)
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=observed_at,
|
||||
replay_sequence=12,
|
||||
payload=candle,
|
||||
candle_is_final=False,
|
||||
)
|
||||
|
||||
assert event.data_type is ReplayDataType.CANDLE_REVISION
|
||||
assert event.replay_at == observed_at
|
||||
assert event.candle_is_final is False
|
||||
|
||||
|
||||
def test_event_normalizes_aware_time_to_utc() -> None:
|
||||
offset = timezone(timedelta(hours=3))
|
||||
trade = make_trade()
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at.astimezone(offset),
|
||||
replay_sequence=1,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
assert event.replay_at.tzinfo is timezone.utc
|
||||
assert event.replay_at == trade.executed_at
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid_sequence", (True, 0, -1, 1.5, "1"))
|
||||
def test_event_rejects_invalid_sequence(invalid_sequence: Any) -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="replay_sequence"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at,
|
||||
replay_sequence=invalid_sequence,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_event_rejects_mismatched_replay_time() -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises(ValueError, match="event time"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at + timedelta(microseconds=1),
|
||||
replay_sequence=1,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def test_candle_event_rejects_time_before_open() -> None:
|
||||
candle = make_candle()
|
||||
|
||||
with pytest.raises(ValueError, match="open_time"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=candle.open_time - timedelta(microseconds=1),
|
||||
replay_sequence=1,
|
||||
payload=candle,
|
||||
candle_is_final=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("candle_is_final", (None, 0, 1, "true"))
|
||||
def test_candle_event_requires_exact_final_boolean(
|
||||
candle_is_final: Any,
|
||||
) -> None:
|
||||
candle = make_candle()
|
||||
|
||||
with pytest.raises(TypeError, match="candle_is_final"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=candle.open_time,
|
||||
replay_sequence=1,
|
||||
payload=candle,
|
||||
candle_is_final=candle_is_final,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_event_rejects_candle_metadata() -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises(ValueError, match="must be None"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at,
|
||||
replay_sequence=1,
|
||||
payload=trade,
|
||||
candle_is_final=False,
|
||||
)
|
||||
|
||||
|
||||
def test_request_normalizes_symbols_and_preserves_interval_case() -> None:
|
||||
request = ReplayPlanRequest(
|
||||
venue=" dzengi ",
|
||||
symbols=(" btc/usd_leverage ",),
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
candle_intervals=(" 1M ",),
|
||||
)
|
||||
|
||||
assert request.venue == VENUE
|
||||
assert request.symbols == (SYMBOL,)
|
||||
assert request.candle_intervals == ("1M",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"symbols",
|
||||
(
|
||||
[],
|
||||
(),
|
||||
("",),
|
||||
("BTC/USD_LEVERAGE", "btc/usd_leverage"),
|
||||
),
|
||||
)
|
||||
def test_request_rejects_invalid_symbols(symbols: Any) -> None:
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=symbols,
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
|
||||
def test_request_requires_intervals_only_for_candles() -> None:
|
||||
with pytest.raises(ValueError, match="required"):
|
||||
make_request(data_types=(ReplayDataType.CANDLE_REVISION,))
|
||||
|
||||
with pytest.raises(ValueError, match="require Candle"):
|
||||
make_request(candle_intervals=("1m",))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"max_records",
|
||||
(True, 0, -1, 1.5, REPLAY_PLAN_MAX_RECORDS_LIMIT + 1),
|
||||
)
|
||||
def test_request_rejects_invalid_max_records(max_records: Any) -> None:
|
||||
with pytest.raises((TypeError, ValueError), match="max_records"):
|
||||
make_request(max_records=max_records)
|
||||
|
||||
|
||||
def test_empty_plan_is_valid_and_preserves_request_scope() -> None:
|
||||
request = make_request()
|
||||
plan = ReplayPlan(request=request, events=())
|
||||
|
||||
assert plan.request is request
|
||||
assert plan.events == ()
|
||||
assert plan.is_empty is True
|
||||
assert len(plan) == 0
|
||||
|
||||
|
||||
def test_plan_preserves_event_and_payload_identity() -> None:
|
||||
request = make_request()
|
||||
event = make_trade_event()
|
||||
plan = ReplayPlan(request=request, events=(event,))
|
||||
|
||||
assert plan.events[0] is event
|
||||
assert plan.events[0].payload is event.payload
|
||||
|
||||
|
||||
def test_plan_accepts_rollover_order_by_time_and_sequence() -> None:
|
||||
same_time = START + timedelta(minutes=1)
|
||||
first = make_trade_event(
|
||||
trade_id=2_147_483_647,
|
||||
replay_at=same_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
second = make_trade_event(
|
||||
trade_id=-2_147_483_648,
|
||||
replay_at=same_time,
|
||||
replay_sequence=11,
|
||||
)
|
||||
|
||||
plan = ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(first, second),
|
||||
)
|
||||
|
||||
assert [trade_id_from_event(event) for event in plan.events] == [
|
||||
2_147_483_647,
|
||||
-2_147_483_648,
|
||||
]
|
||||
|
||||
|
||||
def test_plan_accepts_negative_one_to_zero_at_equal_time() -> None:
|
||||
same_time = START + timedelta(minutes=1)
|
||||
plan = ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(
|
||||
trade_id=-1,
|
||||
replay_at=same_time,
|
||||
replay_sequence=20,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=0,
|
||||
replay_at=same_time,
|
||||
replay_sequence=21,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
assert [trade_id_from_event(event) for event in plan.events] == [-1, 0]
|
||||
|
||||
|
||||
def test_plan_rejects_reverse_order() -> None:
|
||||
with pytest.raises(ValueError, match="strictly ordered"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(replay_sequence=2),
|
||||
make_trade_event(replay_sequence=1),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_reverse_event_time_with_increasing_sequence() -> None:
|
||||
with pytest.raises(ValueError, match="strictly ordered"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=2,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_duplicate_global_sequence() -> None:
|
||||
with pytest.raises(ValueError, match="globally unique"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=1,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_event_outside_time_range() -> None:
|
||||
event = make_trade_event(replay_at=END)
|
||||
|
||||
with pytest.raises(ValueError, match="outside Replay request"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_data_type_outside_request() -> None:
|
||||
quote = make_quote()
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=quote.received_at,
|
||||
replay_sequence=1,
|
||||
payload=quote,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="data type"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_event_from_another_venue() -> None:
|
||||
event = make_trade_event(venue="other")
|
||||
|
||||
with pytest.raises(ValueError, match="event venue"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_event_from_another_symbol() -> None:
|
||||
event = make_trade_event(symbol="ETH/USD_LEVERAGE")
|
||||
|
||||
with pytest.raises(ValueError, match="event symbol"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_candle_interval_outside_request() -> None:
|
||||
candle = make_candle(interval="5m")
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=candle.open_time,
|
||||
replay_sequence=1,
|
||||
payload=candle,
|
||||
candle_is_final=False,
|
||||
)
|
||||
request = make_request(
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
candle_intervals=("1m",),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Candle interval"):
|
||||
ReplayPlan(request=request, events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_more_events_than_request_limit() -> None:
|
||||
event = make_trade_event()
|
||||
|
||||
with pytest.raises(ValueError, match="max_records"):
|
||||
ReplayPlan(
|
||||
request=make_request(max_records=1),
|
||||
events=(event, make_trade_event(replay_sequence=2)),
|
||||
)
|
||||
|
||||
|
||||
def test_replay_models_are_frozen_and_slotted() -> None:
|
||||
event = make_trade_event()
|
||||
request = make_request()
|
||||
plan = ReplayPlan(request=request, events=(event,))
|
||||
|
||||
for model in (event, request, plan):
|
||||
assert not hasattr(model, "__dict__")
|
||||
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
setattr(event, "venue", "other")
|
||||
|
||||
|
||||
def test_session_states_are_explicit_and_complete() -> None:
|
||||
assert tuple(ReplaySessionState) == (
|
||||
ReplaySessionState.CREATED,
|
||||
ReplaySessionState.RUNNING,
|
||||
ReplaySessionState.COMPLETED,
|
||||
ReplaySessionState.FAILED,
|
||||
ReplaySessionState.CANCELLED,
|
||||
)
|
||||
|
||||
|
||||
def test_event_rejects_canonical_payload_subclasses() -> None:
|
||||
class TradeSubclass(Trade):
|
||||
pass
|
||||
|
||||
class QuoteSubclass(Quote):
|
||||
pass
|
||||
|
||||
class CandleSubclass(Candle):
|
||||
pass
|
||||
|
||||
trade = make_trade()
|
||||
quote = make_quote()
|
||||
candle = make_candle()
|
||||
|
||||
payloads = (
|
||||
(
|
||||
TradeSubclass(
|
||||
symbol=trade.symbol,
|
||||
trade_id=trade.trade_id,
|
||||
price=trade.price,
|
||||
quantity=trade.quantity,
|
||||
executed_at=trade.executed_at,
|
||||
aggressor_side=trade.aggressor_side,
|
||||
source=trade.source,
|
||||
),
|
||||
trade.executed_at,
|
||||
None,
|
||||
),
|
||||
(
|
||||
QuoteSubclass(
|
||||
symbol=quote.symbol,
|
||||
last_price=quote.last_price,
|
||||
bid_price=quote.bid_price,
|
||||
ask_price=quote.ask_price,
|
||||
exchange_timestamp=quote.exchange_timestamp,
|
||||
received_at=quote.received_at,
|
||||
source=quote.source,
|
||||
),
|
||||
quote.received_at,
|
||||
None,
|
||||
),
|
||||
(
|
||||
CandleSubclass(
|
||||
symbol=candle.symbol,
|
||||
interval=candle.interval,
|
||||
open_time=candle.open_time,
|
||||
open_price=candle.open_price,
|
||||
high_price=candle.high_price,
|
||||
low_price=candle.low_price,
|
||||
close_price=candle.close_price,
|
||||
volume=candle.volume,
|
||||
source=candle.source,
|
||||
),
|
||||
candle.open_time,
|
||||
False,
|
||||
),
|
||||
)
|
||||
|
||||
for payload, replay_at, candle_is_final in payloads:
|
||||
with pytest.raises(TypeError, match="Canonical Trade, Quote or Candle"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=1,
|
||||
payload=payload,
|
||||
candle_is_final=candle_is_final,
|
||||
)
|
||||
|
||||
|
||||
def test_request_rejects_time_range_subclass() -> None:
|
||||
class HistoricalTimeRangeSubclass(HistoricalTimeRange):
|
||||
def contains(self, instant: datetime) -> bool:
|
||||
return True
|
||||
|
||||
with pytest.raises(TypeError, match="HistoricalTimeRange"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRangeSubclass(START, END),
|
||||
)
|
||||
|
||||
|
||||
def test_request_rejects_tuple_subclasses() -> None:
|
||||
class TupleSubclass(tuple):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="symbols must be a tuple"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=TupleSubclass((SYMBOL,)),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="data_types must be a tuple"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=TupleSubclass((ReplayDataType.TRADE,)),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="candle_intervals must be a tuple"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
candle_intervals=TupleSubclass(("1m",)),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_request_and_event_subclasses() -> None:
|
||||
class ReplayPlanRequestSubclass(ReplayPlanRequest):
|
||||
pass
|
||||
|
||||
class ReplayEventSubclass(ReplayEvent):
|
||||
@property
|
||||
def symbol(self) -> str:
|
||||
return SYMBOL
|
||||
|
||||
@property
|
||||
def data_type(self) -> ReplayDataType:
|
||||
return ReplayDataType.TRADE
|
||||
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (START, 1)
|
||||
|
||||
request_subclass = ReplayPlanRequestSubclass(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="ReplayPlanRequest"):
|
||||
ReplayPlan(request=request_subclass, events=())
|
||||
|
||||
event = make_trade_event()
|
||||
event_subclass = ReplayEventSubclass(
|
||||
venue=event.venue,
|
||||
replay_at=event.replay_at,
|
||||
replay_sequence=event.replay_sequence,
|
||||
payload=event.payload,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="ReplayEvent"):
|
||||
ReplayPlan(request=make_request(), events=(event_subclass,))
|
||||
|
||||
|
||||
def test_plan_rejects_events_tuple_subclass() -> None:
|
||||
class EventsTupleSubclass(tuple):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="events must be a tuple"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=EventsTupleSubclass((make_trade_event(),)),
|
||||
)
|
||||
@@ -0,0 +1,699 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
)
|
||||
from src.market_data.access.models import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.candle import Candle
|
||||
from src.market_data.acquisition.models.quote import Quote
|
||||
from src.market_data.acquisition.models.trade import Trade
|
||||
from src.market_data.replay.contracts import ReplayPlanBuilderProtocol
|
||||
from src.market_data.replay.exceptions import (
|
||||
MarketDataReplayValidationError,
|
||||
ReplayPlanLimitExceededError,
|
||||
)
|
||||
from src.market_data.replay.models import (
|
||||
ReplayDataType,
|
||||
ReplayPlanRequest,
|
||||
)
|
||||
from src.market_data.replay.postgres_replay_plan_builder import (
|
||||
PostgresReplayPlanBuilder,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
TRADE_SOURCE = "dzengi_websocket_trade"
|
||||
QUOTE_SOURCE = "dzengi_rest_quote"
|
||||
CANDLE_SOURCE = "dzengi_rest_candle"
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
events: list[tuple[object, ...]],
|
||||
rows: object = (),
|
||||
) -> None:
|
||||
self._events = events
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, object]] = []
|
||||
self.execute_errors: list[BaseException | None] = []
|
||||
self.fetchall_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self._events.append(("cursor_enter",))
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self._events.append(("cursor_exit", exception_type))
|
||||
return None
|
||||
|
||||
def execute(
|
||||
self,
|
||||
sql: str,
|
||||
parameters: object = None,
|
||||
) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
self._events.append(("execute", normalized_sql(sql), parameters))
|
||||
|
||||
error = (
|
||||
self.execute_errors.pop(0)
|
||||
if self.execute_errors
|
||||
else None
|
||||
)
|
||||
|
||||
if error is not None:
|
||||
raise error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
self._events.append(("fetchall",))
|
||||
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingTransaction:
|
||||
def __init__(self, events: list[tuple[object, ...]]) -> None:
|
||||
self._events = events
|
||||
|
||||
def __enter__(self) -> RecordingTransaction:
|
||||
self._events.append(("transaction_enter",))
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self._events.append(("transaction_exit", exception_type))
|
||||
return None
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
events: list[tuple[object, ...]],
|
||||
cursor: RecordingCursor,
|
||||
) -> None:
|
||||
self._events = events
|
||||
self._cursor = cursor
|
||||
self.transaction_calls = 0
|
||||
self.cursor_calls = 0
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self._events.append(("connection_enter",))
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self._events.append(("connection_exit", exception_type))
|
||||
return None
|
||||
|
||||
def transaction(self) -> RecordingTransaction:
|
||||
self.transaction_calls += 1
|
||||
self._events.append(("transaction",))
|
||||
return RecordingTransaction(self._events)
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
self._events.append(("cursor",))
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
events: list[tuple[object, ...]]
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
self.events.append(("provider",))
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def make_request(
|
||||
*,
|
||||
data_types: tuple[ReplayDataType, ...] = (ReplayDataType.TRADE,),
|
||||
candle_intervals: tuple[str, ...] = (),
|
||||
max_records: int = 10,
|
||||
venue: str = VENUE,
|
||||
) -> ReplayPlanRequest:
|
||||
return ReplayPlanRequest(
|
||||
venue=venue,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=data_types,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
candle_intervals=candle_intervals,
|
||||
max_records=max_records,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_row(
|
||||
*,
|
||||
replay_at: object = START + timedelta(minutes=1),
|
||||
replay_sequence: object = 10,
|
||||
venue: object = VENUE,
|
||||
trade_id: object = 100,
|
||||
price: object = Decimal("65000"),
|
||||
) -> tuple[object, ...]:
|
||||
return (
|
||||
"trade",
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
venue,
|
||||
SYMBOL,
|
||||
trade_id,
|
||||
replay_at,
|
||||
price,
|
||||
Decimal("0.01"),
|
||||
"buy",
|
||||
TRADE_SOURCE,
|
||||
START + timedelta(seconds=1),
|
||||
START + timedelta(seconds=2),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
[TRADE_SOURCE],
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def make_quote_row(
|
||||
*,
|
||||
replay_at: object = START + timedelta(minutes=2),
|
||||
replay_sequence: object = 11,
|
||||
) -> tuple[object, ...]:
|
||||
return (
|
||||
"quote",
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
QUOTE_SOURCE,
|
||||
None,
|
||||
None,
|
||||
replay_at,
|
||||
START + timedelta(minutes=2) - timedelta(milliseconds=1),
|
||||
Decimal("65001"),
|
||||
Decimal("65000"),
|
||||
Decimal("65002"),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
[QUOTE_SOURCE],
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def make_candle_row(
|
||||
*,
|
||||
replay_at: object = START + timedelta(minutes=3),
|
||||
replay_sequence: object = 12,
|
||||
open_time: object = START - timedelta(hours=1),
|
||||
interval: object = "1m",
|
||||
is_final: object = True,
|
||||
) -> tuple[object, ...]:
|
||||
return (
|
||||
"candle_revision",
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
CANDLE_SOURCE,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
interval,
|
||||
open_time,
|
||||
replay_at,
|
||||
Decimal("64900"),
|
||||
Decimal("65100"),
|
||||
Decimal("64800"),
|
||||
Decimal("65000"),
|
||||
Decimal("10"),
|
||||
is_final,
|
||||
[CANDLE_SOURCE],
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresReplayPlanBuilder,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
list[tuple[object, ...]],
|
||||
]:
|
||||
events: list[tuple[object, ...]] = []
|
||||
cursor = RecordingCursor(events=events, rows=rows)
|
||||
connection = RecordingConnection(events=events, cursor=cursor)
|
||||
provider = RecordingProvider(connection=connection, events=events)
|
||||
builder = PostgresReplayPlanBuilder(connection_provider=provider)
|
||||
return builder, cursor, connection, provider, events
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
builder, _, _, provider, _ = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(builder, "__dict__")
|
||||
assert isinstance(builder, ReplayPlanBuilderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresReplayPlanBuilder(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_request_is_validated_before_connection_borrow() -> None:
|
||||
class RequestSubclass(ReplayPlanRequest):
|
||||
pass
|
||||
|
||||
builder, _, _, provider, _ = dependencies()
|
||||
request = RequestSubclass(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="request"):
|
||||
builder.create_plan(request)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_trade_snapshot_uses_one_half_open_parameterized_query() -> None:
|
||||
builder, cursor, connection, provider, events = dependencies()
|
||||
request = make_request(max_records=7, venue="tenant'value")
|
||||
|
||||
plan = builder.create_plan(request)
|
||||
|
||||
assert plan.is_empty is True
|
||||
assert provider.calls == 1
|
||||
assert connection.transaction_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert len(cursor.calls) == 2
|
||||
setup_sql, setup_parameters = cursor.calls[0]
|
||||
sql, parameters = cursor.calls[1]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert normalized_sql(setup_sql) == (
|
||||
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY"
|
||||
)
|
||||
assert setup_parameters is None
|
||||
assert "FROM market_data.trades" in compact_sql
|
||||
assert "FROM market_data.quotes" not in compact_sql
|
||||
assert "FROM market_data.candle_revisions" not in compact_sql
|
||||
assert "executed_at >= %s" in compact_sql
|
||||
assert "executed_at < %s" in compact_sql
|
||||
assert "ORDER BY replay_at ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert "tenant'value" not in sql
|
||||
assert parameters == (
|
||||
"tenant'value",
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
8,
|
||||
)
|
||||
assert events == [
|
||||
("provider",),
|
||||
("connection_enter",),
|
||||
("transaction",),
|
||||
("transaction_enter",),
|
||||
("cursor",),
|
||||
("cursor_enter",),
|
||||
(
|
||||
"execute",
|
||||
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY",
|
||||
None,
|
||||
),
|
||||
("execute", compact_sql, parameters),
|
||||
("fetchall",),
|
||||
("cursor_exit", None),
|
||||
("transaction_exit", None),
|
||||
("connection_exit", None),
|
||||
]
|
||||
|
||||
|
||||
def test_all_requested_types_use_static_union_and_own_time_axes() -> None:
|
||||
builder, cursor, _, _, _ = dependencies()
|
||||
request = make_request(
|
||||
data_types=(
|
||||
ReplayDataType.CANDLE_REVISION,
|
||||
ReplayDataType.TRADE,
|
||||
ReplayDataType.QUOTE,
|
||||
),
|
||||
candle_intervals=("1m", "5m"),
|
||||
max_records=20,
|
||||
)
|
||||
|
||||
builder.create_plan(request)
|
||||
|
||||
sql, parameters = cursor.calls[1]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert compact_sql.count("UNION ALL") == 2
|
||||
assert "executed_at >= %s" in compact_sql
|
||||
assert "received_at >= %s" in compact_sql
|
||||
assert "observed_at >= %s" in compact_sql
|
||||
assert "open_time >= %s" not in compact_sql
|
||||
assert "interval = ANY(%s)" in compact_sql
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
VENUE,
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
VENUE,
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
["1m", "5m"],
|
||||
21,
|
||||
)
|
||||
|
||||
|
||||
def test_materializes_globally_ordered_canonical_events() -> None:
|
||||
event_time = START + timedelta(minutes=2)
|
||||
builder, _, _, _, _ = dependencies(
|
||||
[
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_quote_row(
|
||||
replay_at=event_time,
|
||||
replay_sequence=11,
|
||||
),
|
||||
make_candle_row(
|
||||
replay_at=event_time,
|
||||
replay_sequence=12,
|
||||
),
|
||||
]
|
||||
)
|
||||
request = make_request(
|
||||
data_types=(
|
||||
ReplayDataType.TRADE,
|
||||
ReplayDataType.QUOTE,
|
||||
ReplayDataType.CANDLE_REVISION,
|
||||
),
|
||||
candle_intervals=("1m",),
|
||||
)
|
||||
|
||||
plan = builder.create_plan(request)
|
||||
|
||||
assert [event.replay_sequence for event in plan.events] == [10, 11, 12]
|
||||
assert type(plan.events[0].payload) is Trade
|
||||
assert type(plan.events[1].payload) is Quote
|
||||
candle_payload = plan.events[2].payload
|
||||
assert type(candle_payload) is Candle
|
||||
assert isinstance(candle_payload, Candle)
|
||||
assert plan.events[2].replay_at == event_time
|
||||
assert candle_payload.open_time < START
|
||||
assert plan.events[2].candle_is_final is True
|
||||
|
||||
|
||||
def test_candle_snapshot_filters_by_observed_time_not_open_time() -> None:
|
||||
builder, cursor, _, _, _ = dependencies(
|
||||
[make_candle_row(open_time=START - timedelta(days=1))]
|
||||
)
|
||||
request = make_request(
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
candle_intervals=("1m",),
|
||||
)
|
||||
|
||||
plan = builder.create_plan(request)
|
||||
|
||||
sql, _ = cursor.calls[1]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "observed_at >= %s" in compact_sql
|
||||
assert "observed_at < %s" in compact_sql
|
||||
assert "open_time >= %s" not in compact_sql
|
||||
assert len(plan.events) == 1
|
||||
|
||||
|
||||
def test_limit_plus_one_raises_without_partial_plan() -> None:
|
||||
builder, _, _, _, events = dependencies(
|
||||
[
|
||||
make_trade_row(replay_sequence=10),
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=11,
|
||||
trade_id=101,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(ReplayPlanLimitExceededError, match="max_records"):
|
||||
builder.create_plan(make_request(max_records=1))
|
||||
|
||||
assert events[-3:] == [
|
||||
("cursor_exit", None),
|
||||
("transaction_exit", None),
|
||||
("connection_exit", None),
|
||||
]
|
||||
|
||||
|
||||
def test_more_than_limit_plus_one_is_backend_integrity_error() -> None:
|
||||
builder, _, _, _, _ = dependencies(
|
||||
[
|
||||
make_trade_row(replay_sequence=10),
|
||||
make_trade_row(replay_sequence=11),
|
||||
make_trade_row(replay_sequence=12),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
builder.create_plan(make_request(max_records=1))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rows", (None, 1, "rows", b"rows"))
|
||||
def test_rejects_invalid_rows_container(rows: object) -> None:
|
||||
builder, _, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"row",
|
||||
(
|
||||
(),
|
||||
("trade",),
|
||||
("unknown",) + make_trade_row()[1:],
|
||||
(1,) + make_trade_row()[1:],
|
||||
),
|
||||
)
|
||||
def test_rejects_invalid_snapshot_row_shape_or_type(
|
||||
row: tuple[object, ...],
|
||||
) -> None:
|
||||
builder, _, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"row, replay_request",
|
||||
(
|
||||
(
|
||||
make_trade_row(price=Decimal("0")),
|
||||
make_request(),
|
||||
),
|
||||
(
|
||||
make_quote_row()[:-1] + (2,),
|
||||
make_request(data_types=(ReplayDataType.QUOTE,)),
|
||||
),
|
||||
(
|
||||
make_candle_row(is_final=1),
|
||||
make_request(
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
candle_intervals=("1m",),
|
||||
),
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_shared_mappers_reject_corrupt_canonical_values(
|
||||
row: tuple[object, ...],
|
||||
replay_request: ReplayPlanRequest,
|
||||
) -> None:
|
||||
builder, _, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError):
|
||||
builder.create_plan(replay_request)
|
||||
|
||||
|
||||
def test_rejects_replay_time_that_differs_from_payload_time() -> None:
|
||||
row = list(make_trade_row())
|
||||
row[1] = START + timedelta(minutes=2)
|
||||
builder, _, _, _, _ = dependencies([tuple(row)])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="event"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rows",
|
||||
(
|
||||
(
|
||||
make_trade_row(replay_sequence=10),
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=10,
|
||||
trade_id=101,
|
||||
),
|
||||
),
|
||||
(
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=11,
|
||||
),
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=10,
|
||||
trade_id=101,
|
||||
),
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_rejects_duplicate_sequence_or_unordered_rows(
|
||||
rows: tuple[tuple[object, ...], ...],
|
||||
) -> None:
|
||||
builder, _, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="snapshot"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
def test_rejects_event_outside_request_scope() -> None:
|
||||
builder, _, _, _, _ = dependencies(
|
||||
[make_trade_row(venue="another")]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="snapshot"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure_point", ("provider", "execute", "fetchall"))
|
||||
def test_backend_error_is_wrapped_with_original_cause(
|
||||
failure_point: str,
|
||||
) -> None:
|
||||
builder, cursor, _, provider, events = dependencies()
|
||||
backend_error = RuntimeError("backend failed")
|
||||
|
||||
if failure_point == "provider":
|
||||
provider.error = backend_error
|
||||
elif failure_point == "execute":
|
||||
cursor.execute_errors = [None, backend_error]
|
||||
else:
|
||||
cursor.fetchall_error = backend_error
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as raised:
|
||||
builder.create_plan(make_request())
|
||||
|
||||
assert raised.value.__cause__ is backend_error
|
||||
|
||||
if failure_point != "provider":
|
||||
assert events[-3:] == [
|
||||
("cursor_exit", RuntimeError),
|
||||
("transaction_exit", RuntimeError),
|
||||
("connection_exit", RuntimeError),
|
||||
]
|
||||
|
||||
|
||||
def test_transaction_setup_error_is_wrapped_and_query_is_not_executed() -> None:
|
||||
builder, cursor, _, _, _ = dependencies()
|
||||
setup_error = RuntimeError("cannot configure transaction")
|
||||
cursor.execute_errors = [setup_error]
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as raised:
|
||||
builder.create_plan(make_request())
|
||||
|
||||
assert raised.value.__cause__ is setup_error
|
||||
assert len(cursor.calls) == 1
|
||||
|
||||
|
||||
def test_base_exception_is_not_swallowed_and_contexts_are_closed() -> None:
|
||||
builder, cursor, _, _, events = dependencies()
|
||||
cursor.fetchall_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
assert events[-3:] == [
|
||||
("cursor_exit", KeyboardInterrupt),
|
||||
("transaction_exit", KeyboardInterrupt),
|
||||
("connection_exit", KeyboardInterrupt),
|
||||
]
|
||||
807
app/tests/unit/market_data/replay/test_replay_session.py
Normal file
807
app/tests/unit/market_data/replay/test_replay_session.py
Normal file
@@ -0,0 +1,807 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
import src.market_data.replay as replay_package
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.replay import (
|
||||
MarketDataReplayValidationError,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
ReplaySession,
|
||||
ReplaySessionProtocol,
|
||||
ReplaySessionState,
|
||||
ReplaySessionStateError,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
FIRST_TIME = START + timedelta(minutes=1)
|
||||
SECOND_TIME = START + timedelta(minutes=2)
|
||||
THIRD_TIME = START + timedelta(minutes=3)
|
||||
END = START + timedelta(hours=1)
|
||||
|
||||
|
||||
def make_trade_event(
|
||||
*,
|
||||
trade_id: int,
|
||||
replay_at: datetime,
|
||||
replay_sequence: int,
|
||||
) -> ReplayEvent:
|
||||
trade = Trade(
|
||||
symbol=SYMBOL,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=replay_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source="dzengi_websocket_trade",
|
||||
)
|
||||
return ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=replay_sequence,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def make_plan(*, empty: bool = False) -> ReplayPlan:
|
||||
request = ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=START,
|
||||
end_time=END,
|
||||
),
|
||||
max_records=100,
|
||||
)
|
||||
events = () if empty else (
|
||||
make_trade_event(
|
||||
trade_id=1,
|
||||
replay_at=FIRST_TIME,
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=2,
|
||||
replay_at=FIRST_TIME,
|
||||
replay_sequence=2,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=3,
|
||||
replay_at=SECOND_TIME,
|
||||
replay_sequence=3,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=4,
|
||||
replay_at=THIRD_TIME,
|
||||
replay_sequence=4,
|
||||
),
|
||||
)
|
||||
return ReplayPlan(request=request, events=events)
|
||||
|
||||
|
||||
class RecordingClock:
|
||||
def __init__(self, now: datetime = START) -> None:
|
||||
self._now = now
|
||||
self.advance_calls: list[datetime] = []
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return self._now
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self.advance_calls.append(instant)
|
||||
self._now = instant
|
||||
|
||||
|
||||
class RecordingConsumer:
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
self.clock = clock
|
||||
self.events: list[ReplayEvent] = []
|
||||
self.observed_times: list[datetime] = []
|
||||
self.tasks: list[asyncio.Task[Any] | None] = []
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.events.append(event)
|
||||
self.observed_times.append(self.clock.now)
|
||||
self.tasks.append(asyncio.current_task())
|
||||
|
||||
|
||||
def make_session(
|
||||
*,
|
||||
plan: ReplayPlan | None = None,
|
||||
clock: RecordingClock | None = None,
|
||||
consumer: RecordingConsumer | None = None,
|
||||
) -> tuple[ReplaySession, ReplayPlan, RecordingClock, RecordingConsumer]:
|
||||
resolved_plan = make_plan() if plan is None else plan
|
||||
resolved_clock = (
|
||||
RecordingClock(resolved_plan.request.time_range.start_time)
|
||||
if clock is None
|
||||
else clock
|
||||
)
|
||||
resolved_consumer = (
|
||||
RecordingConsumer(resolved_clock)
|
||||
if consumer is None
|
||||
else consumer
|
||||
)
|
||||
session = ReplaySession(
|
||||
plan=resolved_plan,
|
||||
clock=resolved_clock,
|
||||
consumer=resolved_consumer,
|
||||
)
|
||||
return (
|
||||
session,
|
||||
resolved_plan,
|
||||
resolved_clock,
|
||||
resolved_consumer,
|
||||
)
|
||||
|
||||
|
||||
def test_matches_protocol_uses_slots_and_preserves_dependencies() -> None:
|
||||
session, plan, clock, _ = make_session()
|
||||
|
||||
assert isinstance(session, ReplaySessionProtocol)
|
||||
assert not hasattr(session, "__dict__")
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert session.plan is plan
|
||||
assert session.clock is clock
|
||||
assert replay_package.ReplaySession is ReplaySession
|
||||
assert not hasattr(replay_package, "ReplayEngine")
|
||||
|
||||
|
||||
def test_construction_does_not_advance_clock_or_call_consumer() -> None:
|
||||
session, _, clock, consumer = make_session()
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert clock.advance_calls == []
|
||||
assert consumer.events == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "plan"))
|
||||
def test_rejects_non_plan(invalid: object) -> None:
|
||||
clock = RecordingClock()
|
||||
|
||||
with pytest.raises(TypeError, match="plan"):
|
||||
ReplaySession(
|
||||
plan=invalid, # type: ignore[arg-type]
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_plan_subclass() -> None:
|
||||
class ReplayPlanSubclass(ReplayPlan):
|
||||
pass
|
||||
|
||||
plan = make_plan()
|
||||
subclass = ReplayPlanSubclass(
|
||||
request=plan.request,
|
||||
events=plan.events,
|
||||
)
|
||||
clock = RecordingClock()
|
||||
|
||||
with pytest.raises(TypeError, match="plan"):
|
||||
ReplaySession(
|
||||
plan=subclass,
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_clock_without_protocol() -> None:
|
||||
class InvalidClock:
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return START
|
||||
|
||||
with pytest.raises(TypeError, match="clock"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=InvalidClock(), # type: ignore[arg-type]
|
||||
consumer=RecordingConsumer(RecordingClock()),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_clock_class_before_any_runtime_action() -> None:
|
||||
class ClockClass:
|
||||
advance_calls: list[datetime] = []
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return START
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self.advance_calls.append(instant)
|
||||
|
||||
consumer_clock = RecordingClock()
|
||||
consumer = RecordingConsumer(consumer_clock)
|
||||
|
||||
with pytest.raises(TypeError, match="clock"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=ClockClass, # type: ignore[arg-type]
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
assert ClockClass.advance_calls == []
|
||||
assert consumer.events == []
|
||||
|
||||
|
||||
def test_rejects_asynchronous_clock_advance() -> None:
|
||||
class AsyncClock:
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return START
|
||||
|
||||
async def advance_to(self, instant: datetime) -> None:
|
||||
return None
|
||||
|
||||
clock = AsyncClock()
|
||||
|
||||
with pytest.raises(TypeError, match="synchronous"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock, # type: ignore[arg-type]
|
||||
consumer=RecordingConsumer(RecordingClock()),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_consumer_without_protocol() -> None:
|
||||
class InvalidConsumer:
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=RecordingClock(),
|
||||
consumer=InvalidConsumer(), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_consumer_class_before_advancing_clock() -> None:
|
||||
clock = RecordingClock()
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
assert clock.now == START
|
||||
assert clock.advance_calls == []
|
||||
|
||||
|
||||
def test_rejects_synchronous_consumer() -> None:
|
||||
class SyncConsumer:
|
||||
def consume(self, event: ReplayEvent) -> None:
|
||||
return None
|
||||
|
||||
with pytest.raises(TypeError, match="asynchronous"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=RecordingClock(),
|
||||
consumer=SyncConsumer(), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"now",
|
||||
(
|
||||
START.replace(tzinfo=None),
|
||||
START.astimezone(timezone(timedelta(hours=3))),
|
||||
START.replace(fold=1),
|
||||
),
|
||||
)
|
||||
def test_rejects_non_canonical_clock_time(now: datetime) -> None:
|
||||
clock = RecordingClock(now)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
assert clock.advance_calls == []
|
||||
|
||||
|
||||
def test_rejects_datetime_subclass_from_clock() -> None:
|
||||
class CompatibleDatetime(datetime):
|
||||
pass
|
||||
|
||||
now = CompatibleDatetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
clock = RecordingClock(now)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_non_datetime_clock_time() -> None:
|
||||
class InvalidNowClock:
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return object() # type: ignore[return-value]
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
return None
|
||||
|
||||
clock = InvalidNowClock()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(RecordingClock()),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"now",
|
||||
(START - timedelta(microseconds=1), START + timedelta(microseconds=1)),
|
||||
)
|
||||
def test_rejects_clock_outside_exact_start_without_reset(
|
||||
now: datetime,
|
||||
) -> None:
|
||||
clock = RecordingClock(now)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="start"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
assert clock.now is now
|
||||
assert clock.advance_calls == []
|
||||
|
||||
|
||||
def test_clock_now_error_is_not_swallowed() -> None:
|
||||
expected = RuntimeError("clock now failed")
|
||||
|
||||
class BrokenClock(RecordingClock):
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
raise expected
|
||||
|
||||
clock = BrokenClock()
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
assert captured.value is expected
|
||||
|
||||
|
||||
def test_run_preserves_order_identity_and_advances_before_consumer() -> None:
|
||||
async def scenario() -> None:
|
||||
session, plan, clock, consumer = make_session()
|
||||
caller_task = asyncio.current_task()
|
||||
|
||||
result = await session.run()
|
||||
|
||||
assert result is None
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
assert all(
|
||||
actual is expected
|
||||
for actual, expected in zip(
|
||||
consumer.events,
|
||||
plan.events,
|
||||
strict=True,
|
||||
)
|
||||
)
|
||||
assert clock.advance_calls == [
|
||||
event.replay_at for event in plan.events
|
||||
]
|
||||
assert consumer.observed_times == clock.advance_calls
|
||||
assert all(task is caller_task for task in consumer.tasks)
|
||||
assert clock.now == plan.events[-1].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_empty_plan_completes_without_dependency_calls() -> None:
|
||||
async def scenario() -> None:
|
||||
plan = make_plan(empty=True)
|
||||
session, _, clock, consumer = make_session(plan=plan)
|
||||
initial_time = clock.now
|
||||
|
||||
result = await session.run()
|
||||
|
||||
assert result is None
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert clock.now is initial_time
|
||||
assert clock.advance_calls == []
|
||||
assert consumer.events == []
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_delivery_is_strictly_sequential() -> None:
|
||||
class YieldingConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.active = 0
|
||||
self.maximum_active = 0
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.active += 1
|
||||
self.maximum_active = max(self.maximum_active, self.active)
|
||||
try:
|
||||
await asyncio.sleep(0)
|
||||
await super().consume(event)
|
||||
finally:
|
||||
self.active -= 1
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = YieldingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
await session.run()
|
||||
|
||||
assert consumer.maximum_active == 1
|
||||
assert tuple(consumer.events) == plan.events
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_clock_error_fails_without_delivering_current_or_suffix() -> None:
|
||||
expected = RuntimeError("clock failed")
|
||||
|
||||
class BrokenClock(RecordingClock):
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self.advance_calls.append(instant)
|
||||
if len(self.advance_calls) == 3:
|
||||
raise expected
|
||||
self._now = instant
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = BrokenClock()
|
||||
consumer = RecordingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert tuple(consumer.events) == plan.events[:2]
|
||||
assert clock.advance_calls == [
|
||||
event.replay_at for event in plan.events[:3]
|
||||
]
|
||||
assert clock.now == plan.events[1].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_consumer_error_fails_without_retry_rollback_or_suffix() -> None:
|
||||
expected = RuntimeError("consumer failed")
|
||||
|
||||
class BrokenConsumer(RecordingConsumer):
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
if len(self.events) == 3:
|
||||
raise expected
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = BrokenConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert tuple(consumer.events) == plan.events[:3]
|
||||
assert clock.now == plan.events[2].replay_at
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert tuple(consumer.events) == plan.events[:3]
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_non_exception_base_error_preserves_identity_and_failed_state() -> None:
|
||||
class ReplaySignal(BaseException):
|
||||
pass
|
||||
|
||||
expected = ReplaySignal("stop")
|
||||
|
||||
class BrokenConsumer(RecordingConsumer):
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
raise expected
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=BrokenConsumer(clock),
|
||||
)
|
||||
|
||||
with pytest.raises(ReplaySignal) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert clock.now == plan.events[0].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_consumer_cancellation_preserves_identity_and_cancelled_state() -> None:
|
||||
expected = asyncio.CancelledError("consumer cancelled")
|
||||
|
||||
class CancellingConsumer(RecordingConsumer):
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
raise expected
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = CancellingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.CANCELLED
|
||||
assert consumer.events == [plan.events[0]]
|
||||
assert clock.now == plan.events[0].replay_at
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.CANCELLED
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_external_cancellation_during_consumer_stays_cancelled() -> None:
|
||||
class BlockingConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.entered = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
self.entered.set()
|
||||
await self.release.wait()
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = BlockingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
task = asyncio.create_task(session.run())
|
||||
await consumer.entered.wait()
|
||||
|
||||
assert session.state is ReplaySessionState.RUNNING
|
||||
task.cancel()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert session.state is ReplaySessionState.CANCELLED
|
||||
assert consumer.events == [plan.events[0]]
|
||||
assert clock.now == plan.events[0].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_concurrent_caller_is_rejected_without_damaging_first() -> None:
|
||||
class BlockingConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.entered = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
if len(self.events) == 1:
|
||||
self.entered.set()
|
||||
await self.release.wait()
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = BlockingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
first = asyncio.create_task(session.run())
|
||||
await consumer.entered.wait()
|
||||
|
||||
assert first.done() is False
|
||||
assert session.state is ReplaySessionState.RUNNING
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert first.done() is False
|
||||
assert session.state is ReplaySessionState.RUNNING
|
||||
consumer.release.set()
|
||||
await first
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_caught_reentrant_run_does_not_damage_outer_run() -> None:
|
||||
class ReentrantConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.session: ReplaySession | None = None
|
||||
self.reentrant_errors: list[ReplaySessionStateError] = []
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
assert self.session is not None
|
||||
try:
|
||||
await self.session.run()
|
||||
except ReplaySessionStateError as error:
|
||||
self.reentrant_errors.append(error)
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = ReentrantConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
consumer.session = session
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
assert len(consumer.reentrant_errors) == len(plan.events)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_uncaught_reentrant_run_fails_outer_run() -> None:
|
||||
class ReentrantConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.session: ReplaySession | None = None
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
assert self.session is not None
|
||||
await self.session.run()
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = ReentrantConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
consumer.session = session
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert consumer.events == [plan.events[0]]
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_repeated_run_after_completion_is_rejected() -> None:
|
||||
async def scenario() -> None:
|
||||
session, plan, _, consumer = make_session()
|
||||
await session.run()
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_two_sessions_share_immutable_plan_but_not_state_or_clock() -> None:
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
first, _, first_clock, first_consumer = make_session(plan=plan)
|
||||
second, _, second_clock, second_consumer = make_session(plan=plan)
|
||||
|
||||
await first.run()
|
||||
|
||||
assert first.state is ReplaySessionState.COMPLETED
|
||||
assert second.state is ReplaySessionState.CREATED
|
||||
assert first.plan is second.plan is plan
|
||||
assert first_clock is not second_clock
|
||||
assert second_clock.now == START
|
||||
assert second_consumer.events == []
|
||||
|
||||
await second.run()
|
||||
|
||||
assert tuple(first_consumer.events) == plan.events
|
||||
assert tuple(second_consumer.events) == plan.events
|
||||
assert second.state is ReplaySessionState.COMPLETED
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_lifecycle_and_hidden_task_extensions_are_absent() -> None:
|
||||
session, *_ = make_session()
|
||||
|
||||
assert not hasattr(session, "start")
|
||||
assert not hasattr(session, "stop")
|
||||
assert not hasattr(session, "close")
|
||||
assert not hasattr(session, "reset")
|
||||
assert not hasattr(session, "pause")
|
||||
assert not hasattr(session, "resume")
|
||||
883
app/tests/unit/market_data/replay/test_replay_session_factory.py
Normal file
883
app/tests/unit/market_data/replay/test_replay_session_factory.py
Normal file
@@ -0,0 +1,883 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
|
||||
import src.market_data.replay as replay_package
|
||||
import src.market_data.replay.replay_session_factory as factory_module
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.replay import (
|
||||
DeterministicReplayClock,
|
||||
MarketDataClockProtocol,
|
||||
MarketDataReplayValidationError,
|
||||
ReplayConsumerFactoryProtocol,
|
||||
ReplayConsumerProtocol,
|
||||
ReplayClockProtocol,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanBuilderProtocol,
|
||||
ReplayPlanRequest,
|
||||
ReplaySession,
|
||||
ReplaySessionFactory,
|
||||
ReplaySessionState,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
FIRST_TIME = START + timedelta(minutes=1)
|
||||
SECOND_TIME = START + timedelta(minutes=2)
|
||||
END = START + timedelta(hours=1)
|
||||
|
||||
|
||||
def make_request() -> ReplayPlanRequest:
|
||||
return ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=START,
|
||||
end_time=END,
|
||||
),
|
||||
max_records=100,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_event(
|
||||
*,
|
||||
trade_id: int,
|
||||
replay_at: datetime,
|
||||
replay_sequence: int,
|
||||
) -> ReplayEvent:
|
||||
return ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=replay_sequence,
|
||||
payload=Trade(
|
||||
symbol=SYMBOL,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=replay_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source="dzengi_websocket_trade",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def make_plan(
|
||||
*,
|
||||
request: ReplayPlanRequest | None = None,
|
||||
empty: bool = False,
|
||||
) -> ReplayPlan:
|
||||
resolved_request = make_request() if request is None else request
|
||||
events = () if empty else (
|
||||
make_trade_event(
|
||||
trade_id=1,
|
||||
replay_at=FIRST_TIME,
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=2,
|
||||
replay_at=SECOND_TIME,
|
||||
replay_sequence=2,
|
||||
),
|
||||
)
|
||||
return ReplayPlan(
|
||||
request=resolved_request,
|
||||
events=events,
|
||||
)
|
||||
|
||||
|
||||
class RecordingConsumer:
|
||||
def __init__(self, clock: MarketDataClockProtocol) -> None:
|
||||
self.clock = clock
|
||||
self.events: list[ReplayEvent] = []
|
||||
self.observed_times: list[datetime] = []
|
||||
self.start_calls = 0
|
||||
self.stop_calls = 0
|
||||
self.close_calls = 0
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.events.append(event)
|
||||
self.observed_times.append(self.clock.now)
|
||||
|
||||
def start(self) -> None:
|
||||
self.start_calls += 1
|
||||
|
||||
def stop(self) -> None:
|
||||
self.stop_calls += 1
|
||||
|
||||
def close(self) -> None:
|
||||
self.close_calls += 1
|
||||
|
||||
|
||||
class RecordingPlanBuilder:
|
||||
def __init__(
|
||||
self,
|
||||
plan: ReplayPlan,
|
||||
*,
|
||||
actions: list[str] | None = None,
|
||||
) -> None:
|
||||
self.plan = plan
|
||||
self.actions = actions
|
||||
self.requests: list[ReplayPlanRequest] = []
|
||||
self.thread_ids: list[int] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
if self.actions is not None:
|
||||
self.actions.append("builder")
|
||||
self.requests.append(request)
|
||||
self.thread_ids.append(threading.get_ident())
|
||||
return self.plan
|
||||
|
||||
|
||||
class RecordingConsumerFactory:
|
||||
def __init__(self, *, actions: list[str] | None = None) -> None:
|
||||
self.actions = actions
|
||||
self.plans: list[ReplayPlan] = []
|
||||
self.clocks: list[MarketDataClockProtocol] = []
|
||||
self.consumers: list[RecordingConsumer] = []
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
if self.actions is not None:
|
||||
self.actions.append("consumer")
|
||||
consumer = RecordingConsumer(clock)
|
||||
self.plans.append(plan)
|
||||
self.clocks.append(clock)
|
||||
self.consumers.append(consumer)
|
||||
return consumer
|
||||
|
||||
|
||||
def create_factory(
|
||||
*,
|
||||
plan: ReplayPlan | None = None,
|
||||
) -> tuple[
|
||||
ReplaySessionFactory,
|
||||
ReplayPlan,
|
||||
RecordingPlanBuilder,
|
||||
RecordingConsumerFactory,
|
||||
]:
|
||||
resolved_plan = make_plan() if plan is None else plan
|
||||
plan_builder = RecordingPlanBuilder(resolved_plan)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
return (
|
||||
ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
),
|
||||
resolved_plan,
|
||||
plan_builder,
|
||||
consumer_factory,
|
||||
)
|
||||
|
||||
|
||||
def test_matches_protocols_uses_slots_and_is_exported() -> None:
|
||||
factory, _, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
assert isinstance(plan_builder, ReplayPlanBuilderProtocol)
|
||||
assert isinstance(consumer_factory, ReplayConsumerFactoryProtocol)
|
||||
assert not hasattr(factory, "__dict__")
|
||||
assert replay_package.ReplaySessionFactory is ReplaySessionFactory
|
||||
assert (
|
||||
replay_package.ReplayConsumerFactoryProtocol
|
||||
is ReplayConsumerFactoryProtocol
|
||||
)
|
||||
|
||||
|
||||
def test_constructor_only_preserves_dependencies() -> None:
|
||||
plan = make_plan()
|
||||
plan_builder = RecordingPlanBuilder(plan)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
assert isinstance(factory, ReplaySessionFactory)
|
||||
assert plan_builder.requests == []
|
||||
assert consumer_factory.plans == []
|
||||
assert consumer_factory.clocks == []
|
||||
assert consumer_factory.consumers == []
|
||||
|
||||
|
||||
def test_constructor_does_not_create_clock_session_or_task(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class ForbiddenClock:
|
||||
def __init__(self, initial_time: datetime) -> None:
|
||||
raise AssertionError("Clock must not be created")
|
||||
|
||||
def forbidden_session(**kwargs: object) -> ReplaySession:
|
||||
raise AssertionError("Session must not be created")
|
||||
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"DeterministicReplayClock",
|
||||
ForbiddenClock,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"ReplaySession",
|
||||
forbidden_session,
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
tasks_before = set(asyncio.all_tasks())
|
||||
|
||||
factory, _, _, _ = create_factory()
|
||||
|
||||
assert isinstance(factory, ReplaySessionFactory)
|
||||
assert set(asyncio.all_tasks()) == tasks_before
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "builder"))
|
||||
def test_rejects_invalid_plan_builder(invalid: object) -> None:
|
||||
with pytest.raises(TypeError, match="plan_builder"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=cast(ReplayPlanBuilderProtocol, invalid),
|
||||
consumer_factory=RecordingConsumerFactory(),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_plan_builder_class() -> None:
|
||||
class PlanBuilderClass:
|
||||
def create_plan(
|
||||
self,
|
||||
request: ReplayPlanRequest,
|
||||
) -> ReplayPlan:
|
||||
return make_plan(request=request)
|
||||
|
||||
with pytest.raises(TypeError, match="plan_builder"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=cast(
|
||||
ReplayPlanBuilderProtocol,
|
||||
PlanBuilderClass,
|
||||
),
|
||||
consumer_factory=RecordingConsumerFactory(),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_asynchronous_plan_builder() -> None:
|
||||
class AsyncPlanBuilder:
|
||||
async def create_plan(
|
||||
self,
|
||||
request: ReplayPlanRequest,
|
||||
) -> ReplayPlan:
|
||||
return make_plan(request=request)
|
||||
|
||||
with pytest.raises(TypeError, match="synchronous"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=cast(
|
||||
ReplayPlanBuilderProtocol,
|
||||
AsyncPlanBuilder(),
|
||||
),
|
||||
consumer_factory=RecordingConsumerFactory(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "factory"))
|
||||
def test_rejects_invalid_consumer_factory(invalid: object) -> None:
|
||||
with pytest.raises(TypeError, match="consumer_factory"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan()),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
invalid,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_consumer_factory_class() -> None:
|
||||
class ConsumerFactoryClass:
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
return RecordingConsumer(clock)
|
||||
|
||||
with pytest.raises(TypeError, match="consumer_factory"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan()),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
ConsumerFactoryClass,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_asynchronous_consumer_factory() -> None:
|
||||
class AsyncConsumerFactory:
|
||||
async def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
return RecordingConsumer(clock)
|
||||
|
||||
with pytest.raises(TypeError, match="synchronous"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan()),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
AsyncConsumerFactory(),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_prepare_is_synchronous_and_run_remains_asynchronous() -> None:
|
||||
assert not inspect.iscoroutinefunction(
|
||||
ReplaySessionFactory.prepare_session
|
||||
)
|
||||
assert inspect.iscoroutinefunction(ReplaySession.run)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "request"))
|
||||
def test_rejects_invalid_request_before_dependencies(
|
||||
invalid: object,
|
||||
) -> None:
|
||||
factory, _, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="request"):
|
||||
factory.prepare_session(
|
||||
cast(ReplayPlanRequest, invalid)
|
||||
)
|
||||
|
||||
assert plan_builder.requests == []
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_rejects_request_subclass_before_dependencies() -> None:
|
||||
class ReplayPlanRequestSubclass(ReplayPlanRequest):
|
||||
pass
|
||||
|
||||
request = make_request()
|
||||
subclass = ReplayPlanRequestSubclass(
|
||||
venue=request.venue,
|
||||
symbols=request.symbols,
|
||||
data_types=request.data_types,
|
||||
time_range=request.time_range,
|
||||
candle_intervals=request.candle_intervals,
|
||||
max_records=request.max_records,
|
||||
)
|
||||
factory, _, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="request"):
|
||||
factory.prepare_session(subclass)
|
||||
|
||||
assert plan_builder.requests == []
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_builds_in_order_and_preserves_all_identities(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
actions: list[str] = []
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
plan_builder = RecordingPlanBuilder(plan, actions=actions)
|
||||
consumer_factory = RecordingConsumerFactory(actions=actions)
|
||||
real_session = ReplaySession
|
||||
|
||||
class OrderedClock(DeterministicReplayClock):
|
||||
def __init__(self, initial_time: datetime) -> None:
|
||||
actions.append("clock")
|
||||
super().__init__(initial_time)
|
||||
|
||||
def ordered_session(
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: ReplayClockProtocol,
|
||||
consumer: ReplayConsumerProtocol,
|
||||
) -> ReplaySession:
|
||||
actions.append("session")
|
||||
return real_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"DeterministicReplayClock",
|
||||
OrderedClock,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"ReplaySession",
|
||||
ordered_session,
|
||||
)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
session = factory.prepare_session(request)
|
||||
clock = consumer_factory.clocks[0]
|
||||
|
||||
assert actions == ["builder", "clock", "consumer", "session"]
|
||||
assert plan_builder.requests == [request]
|
||||
assert plan_builder.requests[0] is request
|
||||
assert consumer_factory.plans == [plan]
|
||||
assert consumer_factory.plans[0] is plan
|
||||
assert session.plan is plan
|
||||
assert session.clock is clock
|
||||
assert isinstance(clock, OrderedClock)
|
||||
assert clock.now == START
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
|
||||
|
||||
def test_prepare_runs_builder_in_caller_thread() -> None:
|
||||
factory, plan, plan_builder, consumer_factory = create_factory()
|
||||
caller_thread_id = threading.get_ident()
|
||||
|
||||
session = factory.prepare_session(plan.request)
|
||||
|
||||
assert isinstance(session, ReplaySession)
|
||||
assert plan_builder.thread_ids == [caller_thread_id]
|
||||
assert consumer_factory.consumers[0].events == []
|
||||
|
||||
|
||||
def test_prepare_does_not_start_consumer_or_session_lifecycle() -> None:
|
||||
factory, plan, _, consumer_factory = create_factory()
|
||||
|
||||
async def scenario() -> None:
|
||||
tasks_before = set(asyncio.all_tasks())
|
||||
|
||||
session = factory.prepare_session(plan.request)
|
||||
consumer = consumer_factory.consumers[0]
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert consumer.events == []
|
||||
assert consumer.start_calls == 0
|
||||
assert consumer.stop_calls == 0
|
||||
assert consumer.close_calls == 0
|
||||
assert set(asyncio.all_tasks()) == tasks_before
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_consumer_observes_start_before_run_and_event_time_during_run(
|
||||
) -> None:
|
||||
factory, plan, _, consumer_factory = create_factory()
|
||||
session = factory.prepare_session(plan.request)
|
||||
consumer = consumer_factory.consumers[0]
|
||||
|
||||
assert consumer.clock is session.clock
|
||||
assert consumer.clock.now == START
|
||||
assert consumer.events == []
|
||||
|
||||
asyncio.run(session.run())
|
||||
|
||||
assert tuple(consumer.events) == plan.events
|
||||
assert consumer.observed_times == [FIRST_TIME, SECOND_TIME]
|
||||
assert session.clock.now == SECOND_TIME
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
|
||||
|
||||
def test_empty_plan_builds_created_no_op_session() -> None:
|
||||
request = make_request()
|
||||
plan = make_plan(request=request, empty=True)
|
||||
factory, _, _, consumer_factory = create_factory(plan=plan)
|
||||
|
||||
session = factory.prepare_session(request)
|
||||
consumer = consumer_factory.consumers[0]
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert session.clock.now == START
|
||||
|
||||
asyncio.run(session.run())
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert session.clock.now == START
|
||||
assert consumer.events == []
|
||||
|
||||
|
||||
def test_repeated_prepare_creates_fresh_dependency_graph() -> None:
|
||||
factory, plan, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
first = factory.prepare_session(plan.request)
|
||||
second = factory.prepare_session(plan.request)
|
||||
|
||||
assert first is not second
|
||||
assert first.plan is plan
|
||||
assert second.plan is plan
|
||||
assert first.clock is not second.clock
|
||||
assert consumer_factory.consumers[0] is not (
|
||||
consumer_factory.consumers[1]
|
||||
)
|
||||
assert consumer_factory.clocks == [first.clock, second.clock]
|
||||
assert plan_builder.requests == [plan.request, plan.request]
|
||||
|
||||
asyncio.run(first.run())
|
||||
|
||||
assert first.state is ReplaySessionState.COMPLETED
|
||||
assert second.state is ReplaySessionState.CREATED
|
||||
assert second.clock.now == START
|
||||
assert consumer_factory.consumers[1].events == []
|
||||
|
||||
|
||||
class FatalPreparationError(BaseException):
|
||||
pass
|
||||
|
||||
|
||||
class RaisingPlanBuilder:
|
||||
def __init__(self, error: BaseException) -> None:
|
||||
self.error = error
|
||||
self.requests: list[ReplayPlanRequest] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
self.requests.append(request)
|
||||
raise self.error
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expected",
|
||||
(
|
||||
RuntimeError("builder failed"),
|
||||
FatalPreparationError("builder fatal"),
|
||||
asyncio.CancelledError("builder cancelled"),
|
||||
),
|
||||
)
|
||||
def test_builder_error_is_not_wrapped_and_stops_preparation(
|
||||
expected: BaseException,
|
||||
) -> None:
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
plan_builder = RaisingPlanBuilder(expected)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
request = make_request()
|
||||
|
||||
with pytest.raises(type(expected)) as captured:
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert captured.value is expected
|
||||
assert plan_builder.requests == [request]
|
||||
assert consumer_factory.plans == []
|
||||
assert consumer_factory.clocks == []
|
||||
|
||||
|
||||
class InvalidResultPlanBuilder:
|
||||
def __init__(self, result: object) -> None:
|
||||
self.result = result
|
||||
self.requests: list[ReplayPlanRequest] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
self.requests.append(request)
|
||||
return cast(ReplayPlan, self.result)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "plan"))
|
||||
def test_rejects_invalid_builder_result_before_consumer(
|
||||
invalid: object,
|
||||
) -> None:
|
||||
plan_builder = InvalidResultPlanBuilder(invalid)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
request = make_request()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="ReplayPlan"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert plan_builder.requests == [request]
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_rejects_plan_subclass_before_consumer() -> None:
|
||||
class ReplayPlanSubclass(ReplayPlan):
|
||||
pass
|
||||
|
||||
request = make_request()
|
||||
valid = make_plan(request=request)
|
||||
subclass = ReplayPlanSubclass(
|
||||
request=valid.request,
|
||||
events=valid.events,
|
||||
)
|
||||
plan_builder = InvalidResultPlanBuilder(subclass)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="exact"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_rejects_equal_plan_with_different_request_identity() -> None:
|
||||
request = make_request()
|
||||
copied_request = make_request()
|
||||
assert copied_request == request
|
||||
assert copied_request is not request
|
||||
plan = make_plan(request=copied_request)
|
||||
plan_builder = RecordingPlanBuilder(plan)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="identity"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert plan_builder.requests == [request]
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
class RaisingConsumerFactory:
|
||||
def __init__(self, error: BaseException) -> None:
|
||||
self.error = error
|
||||
self.plans: list[ReplayPlan] = []
|
||||
self.clocks: list[MarketDataClockProtocol] = []
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
self.plans.append(plan)
|
||||
self.clocks.append(clock)
|
||||
raise self.error
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expected",
|
||||
(
|
||||
RuntimeError("consumer factory failed"),
|
||||
FatalPreparationError("consumer factory fatal"),
|
||||
asyncio.CancelledError("consumer factory cancelled"),
|
||||
),
|
||||
)
|
||||
def test_consumer_factory_error_is_not_wrapped_or_retried(
|
||||
expected: BaseException,
|
||||
) -> None:
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
consumer_factory = RaisingConsumerFactory(expected)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(plan),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(type(expected)) as captured:
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert captured.value is expected
|
||||
assert consumer_factory.plans == [plan]
|
||||
assert len(consumer_factory.clocks) == 1
|
||||
assert consumer_factory.clocks[0].now == START
|
||||
|
||||
|
||||
class InvalidConsumerFactory:
|
||||
def __init__(self, result: object) -> None:
|
||||
self.result = result
|
||||
self.calls = 0
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
self.calls += 1
|
||||
return cast(ReplayConsumerProtocol, self.result)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "consumer"))
|
||||
def test_rejects_invalid_consumer_result(invalid: object) -> None:
|
||||
request = make_request()
|
||||
consumer_factory = InvalidConsumerFactory(invalid)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert consumer_factory.calls == 1
|
||||
|
||||
|
||||
def test_rejects_consumer_class_result() -> None:
|
||||
request = make_request()
|
||||
consumer_factory = InvalidConsumerFactory(RecordingConsumer)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
|
||||
def test_rejects_synchronous_consumer_result() -> None:
|
||||
class SynchronousConsumer:
|
||||
def consume(self, event: ReplayEvent) -> None:
|
||||
return None
|
||||
|
||||
request = make_request()
|
||||
consumer_factory = InvalidConsumerFactory(SynchronousConsumer())
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="asynchronous"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
|
||||
def test_keyword_incompatible_consumer_factory_error_is_not_wrapped() -> None:
|
||||
class PositionalOnlyConsumerFactory:
|
||||
def create_consumer(
|
||||
self,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
/,
|
||||
) -> ReplayConsumerProtocol:
|
||||
return RecordingConsumer(clock)
|
||||
|
||||
request = make_request()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
PositionalOnlyConsumerFactory(),
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="keyword"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
|
||||
def test_failed_preparation_does_not_poison_next_call() -> None:
|
||||
class RecoveringConsumerFactory(RecordingConsumerFactory):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.attempts = 0
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
self.attempts += 1
|
||||
if self.attempts == 1:
|
||||
raise RuntimeError("first attempt failed")
|
||||
return super().create_consumer(plan=plan, clock=clock)
|
||||
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
consumer_factory = RecoveringConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(plan),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="first attempt failed"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
session = factory.prepare_session(request)
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert consumer_factory.attempts == 2
|
||||
assert len(consumer_factory.consumers) == 1
|
||||
|
||||
|
||||
def test_two_concurrent_callers_receive_independent_graphs() -> None:
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
builder_barrier = threading.Barrier(2)
|
||||
consumer_barrier = threading.Barrier(2)
|
||||
lock = threading.Lock()
|
||||
|
||||
class ConcurrentPlanBuilder:
|
||||
def __init__(self) -> None:
|
||||
self.callers: list[int] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
with lock:
|
||||
self.callers.append(threading.get_ident())
|
||||
builder_barrier.wait(timeout=5)
|
||||
return plan
|
||||
|
||||
class ConcurrentConsumerFactory:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[
|
||||
tuple[int, MarketDataClockProtocol, RecordingConsumer]
|
||||
] = []
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
consumer = RecordingConsumer(clock)
|
||||
with lock:
|
||||
self.calls.append(
|
||||
(threading.get_ident(), clock, consumer)
|
||||
)
|
||||
consumer_barrier.wait(timeout=5)
|
||||
return consumer
|
||||
|
||||
plan_builder = ConcurrentPlanBuilder()
|
||||
consumer_factory = ConcurrentConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
futures = [
|
||||
executor.submit(factory.prepare_session, request)
|
||||
for _ in range(2)
|
||||
]
|
||||
sessions = [future.result(timeout=10) for future in futures]
|
||||
|
||||
assert len(set(plan_builder.callers)) == 2
|
||||
assert len({call[0] for call in consumer_factory.calls}) == 2
|
||||
assert sessions[0] is not sessions[1]
|
||||
assert sessions[0].clock is not sessions[1].clock
|
||||
assert consumer_factory.calls[0][2] is not (
|
||||
consumer_factory.calls[1][2]
|
||||
)
|
||||
assert {id(call[1]) for call in consumer_factory.calls} == {
|
||||
id(session.clock) for session in sessions
|
||||
}
|
||||
assert all(
|
||||
session.state is ReplaySessionState.CREATED
|
||||
for session in sessions
|
||||
)
|
||||
Reference in New Issue
Block a user