Build 060.29: implement Market Data Access and Replay
This commit is contained in:
@@ -0,0 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from src.market_data.access import (
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
MarketDataAccessError,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
MarketDataCursorError,
|
||||
MarketDataHistoricalAccessProtocol,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
QuoteHistoryReaderProtocol,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryReaderProtocol,
|
||||
)
|
||||
|
||||
|
||||
class RecordingHistoricalAccess:
|
||||
def query_trades(
|
||||
self,
|
||||
query: TradeHistoryQuery,
|
||||
) -> TradeHistoryPage:
|
||||
return TradeHistoryPage(query=query, items=())
|
||||
|
||||
def query_quotes(
|
||||
self,
|
||||
query: QuoteHistoryQuery,
|
||||
) -> QuoteHistoryPage:
|
||||
return QuoteHistoryPage(query=query, items=())
|
||||
|
||||
def query_candle_revisions(
|
||||
self,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
return CandleRevisionHistoryPage(query=query, items=())
|
||||
|
||||
|
||||
def test_history_reader_protocols_are_runtime_checkable() -> None:
|
||||
access = RecordingHistoricalAccess()
|
||||
|
||||
assert isinstance(access, TradeHistoryReaderProtocol)
|
||||
assert isinstance(access, QuoteHistoryReaderProtocol)
|
||||
assert isinstance(access, CandleRevisionHistoryReaderProtocol)
|
||||
assert isinstance(access, MarketDataHistoricalAccessProtocol)
|
||||
|
||||
|
||||
def test_access_error_hierarchy_is_specialized() -> None:
|
||||
assert issubclass(MarketDataAccessValidationError, MarketDataAccessError)
|
||||
assert issubclass(MarketDataCursorError, MarketDataAccessValidationError)
|
||||
assert issubclass(MarketDataAccessIntegrityError, MarketDataAccessError)
|
||||
assert issubclass(MarketDataAccessOperationError, MarketDataAccessError)
|
||||
@@ -0,0 +1,882 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import FrozenInstanceError, replace
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import (
|
||||
HISTORY_PAGE_LIMIT_MAX,
|
||||
CandleRevisionHistoryCursor,
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
CandleRevisionHistoryRecord,
|
||||
HistoricalTimeRange,
|
||||
QuoteHistoryCursor,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
QuoteHistoryRecord,
|
||||
TradeHistoryCursor,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryRecord,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 10, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_trade"
|
||||
|
||||
|
||||
def make_trade(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
executed_at: datetime = START + timedelta(minutes=1),
|
||||
symbol: str = SYMBOL,
|
||||
source: str = SOURCE,
|
||||
) -> Trade:
|
||||
return Trade(
|
||||
symbol=symbol,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000.25"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=executed_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source=source,
|
||||
)
|
||||
|
||||
|
||||
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(
|
||||
*,
|
||||
open_time: datetime = START,
|
||||
interval: str = "1m",
|
||||
) -> Candle:
|
||||
return Candle(
|
||||
symbol=SYMBOL,
|
||||
interval=interval,
|
||||
open_time=open_time,
|
||||
open_price=Decimal("64900"),
|
||||
high_price=Decimal("65100"),
|
||||
low_price=Decimal("64800"),
|
||||
close_price=Decimal("65000"),
|
||||
volume=Decimal("12.5"),
|
||||
source="dzengi_rest_candle",
|
||||
)
|
||||
|
||||
|
||||
def make_range() -> HistoricalTimeRange:
|
||||
return HistoricalTimeRange(start_time=START, end_time=END)
|
||||
|
||||
|
||||
def make_trade_query(
|
||||
*,
|
||||
cursor: TradeHistoryCursor | None = None,
|
||||
limit: int = 500,
|
||||
) -> TradeHistoryQuery:
|
||||
return TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_record(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
executed_at: datetime = START + timedelta(minutes=1),
|
||||
replay_sequence: int = 10,
|
||||
) -> TradeHistoryRecord:
|
||||
return TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=make_trade(
|
||||
trade_id=trade_id,
|
||||
executed_at=executed_at,
|
||||
),
|
||||
first_observed_at=executed_at + timedelta(seconds=1),
|
||||
last_observed_at=executed_at + timedelta(seconds=2),
|
||||
observation_sources=(SOURCE,),
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
|
||||
def make_quote_record(
|
||||
*,
|
||||
received_at: datetime = START + timedelta(minutes=2),
|
||||
replay_sequence: int = 20,
|
||||
) -> QuoteHistoryRecord:
|
||||
quote = make_quote(received_at=received_at)
|
||||
return QuoteHistoryRecord(
|
||||
venue=VENUE,
|
||||
quote=quote,
|
||||
observation_sources=(quote.source,),
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
|
||||
def make_candle_record(
|
||||
*,
|
||||
open_time: datetime = START,
|
||||
observed_at: datetime = START + timedelta(seconds=30),
|
||||
replay_sequence: int = 30,
|
||||
interval: str = "1m",
|
||||
) -> CandleRevisionHistoryRecord:
|
||||
candle = make_candle(open_time=open_time, interval=interval)
|
||||
return CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle,
|
||||
observed_at=observed_at,
|
||||
is_final=False,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
|
||||
def test_time_range_is_half_open_and_normalized_to_utc() -> None:
|
||||
offset = timezone(timedelta(hours=3))
|
||||
time_range = HistoricalTimeRange(
|
||||
start_time=START.astimezone(offset),
|
||||
end_time=END.astimezone(offset),
|
||||
)
|
||||
|
||||
assert time_range.start_time == START
|
||||
assert time_range.start_time.tzinfo is timezone.utc
|
||||
assert time_range.end_time == END
|
||||
assert time_range.contains(START)
|
||||
assert time_range.contains(END - timedelta(microseconds=1))
|
||||
assert not time_range.contains(END)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("start_time", "end_time", "error_type"),
|
||||
(
|
||||
(datetime(2026, 8, 2, 10, 0), END, ValueError),
|
||||
(START, datetime(2026, 8, 2, 11, 0), ValueError),
|
||||
(START, START, ValueError),
|
||||
(END, START, ValueError),
|
||||
("2026-08-02", END, TypeError),
|
||||
),
|
||||
)
|
||||
def test_time_range_rejects_invalid_boundaries(
|
||||
start_time: Any,
|
||||
end_time: Any,
|
||||
error_type: type[Exception],
|
||||
) -> None:
|
||||
with pytest.raises(error_type):
|
||||
HistoricalTimeRange(
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
def test_history_records_preserve_exact_canonical_payloads() -> None:
|
||||
trade = make_trade()
|
||||
quote = make_quote()
|
||||
candle = make_candle()
|
||||
|
||||
trade_record = TradeHistoryRecord(
|
||||
venue=" dzengi ",
|
||||
trade=trade,
|
||||
first_observed_at=trade.executed_at,
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=(trade.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
quote_record = QuoteHistoryRecord(
|
||||
venue=VENUE,
|
||||
quote=quote,
|
||||
observation_sources=(quote.source,),
|
||||
replay_sequence=2,
|
||||
)
|
||||
candle_record = CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle,
|
||||
observed_at=candle.open_time,
|
||||
is_final=True,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=3,
|
||||
)
|
||||
|
||||
assert trade_record.venue == VENUE
|
||||
assert trade_record.trade is trade
|
||||
assert quote_record.quote is quote
|
||||
assert candle_record.candle is candle
|
||||
assert trade_record.event_time == trade.executed_at
|
||||
assert quote_record.event_time == quote.received_at
|
||||
assert candle_record.event_time == candle.open_time
|
||||
assert candle_record.replay_at == candle.open_time
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid_sequence", (True, 0, -1, 1.5, "1"))
|
||||
def test_records_reject_invalid_replay_sequence(
|
||||
invalid_sequence: Any,
|
||||
) -> None:
|
||||
with pytest.raises((TypeError, ValueError), match="replay_sequence"):
|
||||
make_trade_record(replay_sequence=invalid_sequence)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"observation_sources",
|
||||
(
|
||||
[],
|
||||
(),
|
||||
("",),
|
||||
(SOURCE, SOURCE),
|
||||
("another_source",),
|
||||
),
|
||||
)
|
||||
def test_trade_record_rejects_invalid_provenance(
|
||||
observation_sources: Any,
|
||||
) -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=trade,
|
||||
first_observed_at=trade.executed_at,
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=observation_sources,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_record_rejects_reversed_observation_times() -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises(ValueError, match="last_observed_at"):
|
||||
TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=trade,
|
||||
first_observed_at=trade.executed_at + timedelta(seconds=1),
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=(trade.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_candle_record_uses_open_time_for_history_and_observed_for_replay() -> None:
|
||||
record = make_candle_record()
|
||||
|
||||
assert record.event_time == START
|
||||
assert record.replay_at == START + timedelta(seconds=30)
|
||||
assert record.interval == "1m"
|
||||
|
||||
|
||||
def test_candle_record_rejects_observation_before_open_time() -> None:
|
||||
with pytest.raises(ValueError, match="observed_at"):
|
||||
make_candle_record(
|
||||
observed_at=START - timedelta(microseconds=1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_final", (0, 1, None, "true"))
|
||||
def test_candle_record_requires_exact_boolean(is_final: Any) -> None:
|
||||
candle = make_candle()
|
||||
|
||||
with pytest.raises(TypeError, match="is_final"):
|
||||
CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle,
|
||||
observed_at=candle.open_time,
|
||||
is_final=is_final,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_cursor_normalizes_scope_and_time() -> None:
|
||||
offset = timezone(timedelta(hours=3))
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=" dzengi ",
|
||||
symbol=" btc/usd_leverage ",
|
||||
time_range=make_range(),
|
||||
executed_at=(START + timedelta(minutes=1)).astimezone(offset),
|
||||
replay_sequence=5,
|
||||
)
|
||||
|
||||
assert cursor.venue == VENUE
|
||||
assert cursor.symbol == SYMBOL
|
||||
assert cursor.executed_at.tzinfo is timezone.utc
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cursor_time",
|
||||
(
|
||||
START - timedelta(microseconds=1),
|
||||
END,
|
||||
END + timedelta(microseconds=1),
|
||||
),
|
||||
)
|
||||
def test_cursor_position_must_belong_to_query_range(
|
||||
cursor_time: datetime,
|
||||
) -> None:
|
||||
with pytest.raises(ValueError, match="query range"):
|
||||
TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=cursor_time,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_cursor_rejects_unknown_version() -> None:
|
||||
with pytest.raises(ValueError, match="version"):
|
||||
TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
version=2,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("limit", (1, HISTORY_PAGE_LIMIT_MAX))
|
||||
def test_trade_query_accepts_limit_boundaries(limit: int) -> None:
|
||||
query = TradeHistoryQuery(
|
||||
venue=" dzengi ",
|
||||
symbol="btc/usd_leverage",
|
||||
time_range=make_range(),
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
assert query.venue == VENUE
|
||||
assert query.symbol == SYMBOL
|
||||
assert query.limit == limit
|
||||
|
||||
|
||||
@pytest.mark.parametrize("limit", (True, 0, -1, 1.5, 1001))
|
||||
def test_query_rejects_invalid_limit(limit: Any) -> None:
|
||||
with pytest.raises((TypeError, ValueError), match="limit"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
|
||||
def test_query_accepts_cursor_when_only_page_limit_changes() -> None:
|
||||
time_range = make_range()
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
query = TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
limit=17,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
assert query.cursor is cursor
|
||||
|
||||
|
||||
def test_query_rejects_cursor_from_another_scope() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="scope"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol="ETH/USD_LEVERAGE",
|
||||
time_range=make_range(),
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_query_rejects_cursor_of_another_data_type() -> None:
|
||||
quote_cursor = QuoteHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
received_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryCursor"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
cursor=quote_cursor, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_candle_query_preserves_interval_case() -> None:
|
||||
query = CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=" 1M ",
|
||||
time_range=make_range(),
|
||||
)
|
||||
|
||||
assert query.interval == "1M"
|
||||
|
||||
|
||||
def test_empty_page_is_valid_without_cursor() -> None:
|
||||
page = TradeHistoryPage(query=make_trade_query(), items=())
|
||||
|
||||
assert page.items == ()
|
||||
assert page.has_more is False
|
||||
|
||||
|
||||
def test_empty_page_rejects_cursor() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="empty page"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(),
|
||||
next_cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_page_accepts_signed_rollover_at_equal_timestamp() -> None:
|
||||
rollover_time = START + timedelta(minutes=1)
|
||||
first = make_trade_record(
|
||||
trade_id=2_147_483_647,
|
||||
executed_at=rollover_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
second = make_trade_record(
|
||||
trade_id=-2_147_483_648,
|
||||
executed_at=rollover_time,
|
||||
replay_sequence=11,
|
||||
)
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=rollover_time,
|
||||
replay_sequence=11,
|
||||
)
|
||||
|
||||
page = TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(first, second),
|
||||
next_cursor=cursor,
|
||||
)
|
||||
|
||||
assert page.items == (first, second)
|
||||
assert page.has_more is True
|
||||
|
||||
|
||||
def test_trade_page_accepts_negative_one_to_zero_boundary() -> None:
|
||||
event_time = START + timedelta(minutes=1)
|
||||
page = TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(
|
||||
make_trade_record(
|
||||
trade_id=-1,
|
||||
executed_at=event_time,
|
||||
replay_sequence=20,
|
||||
),
|
||||
make_trade_record(
|
||||
trade_id=0,
|
||||
executed_at=event_time,
|
||||
replay_sequence=21,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
assert [item.trade.trade_id for item in page.items] == [-1, 0]
|
||||
|
||||
|
||||
def test_page_rejects_reverse_or_duplicate_order_key() -> None:
|
||||
first = make_trade_record(replay_sequence=2)
|
||||
second = make_trade_record(replay_sequence=1)
|
||||
|
||||
with pytest.raises(ValueError, match="strictly ordered"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(first, second),
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_cursor_not_pointing_to_last_item() -> None:
|
||||
item = make_trade_record(replay_sequence=10)
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=item.event_time,
|
||||
replay_sequence=9,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="last page item"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item,),
|
||||
next_cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"page",
|
||||
(
|
||||
QuoteHistoryPage(
|
||||
query=QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
),
|
||||
items=(make_quote_record(),),
|
||||
),
|
||||
CandleRevisionHistoryPage(
|
||||
query=CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="1m",
|
||||
time_range=make_range(),
|
||||
),
|
||||
items=(make_candle_record(),),
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_quote_and_candle_pages_are_typed_and_immutable(page: Any) -> None:
|
||||
assert not hasattr(page, "__dict__")
|
||||
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
setattr(page, "items", ())
|
||||
|
||||
|
||||
def test_cursor_and_query_classes_use_slots_and_are_frozen() -> None:
|
||||
query = QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
)
|
||||
|
||||
assert not hasattr(query, "__dict__")
|
||||
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
setattr(query, "venue", "other")
|
||||
|
||||
|
||||
def test_cursor_window_mismatch_is_rejected() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
shifted = HistoricalTimeRange(
|
||||
start_time=START - timedelta(minutes=1),
|
||||
end_time=END,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="scope"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=shifted,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_candle_cursor_interval_mismatch_is_rejected() -> None:
|
||||
cursor = CandleRevisionHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="1m",
|
||||
time_range=make_range(),
|
||||
open_time=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="interval"):
|
||||
CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="5m",
|
||||
time_range=make_range(),
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_mixed_query_scope() -> None:
|
||||
first = make_trade_record(replay_sequence=1)
|
||||
second = replace(
|
||||
make_trade_record(
|
||||
executed_at=START + timedelta(minutes=2),
|
||||
replay_sequence=2,
|
||||
),
|
||||
venue="other",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="one query scope"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(first, second),
|
||||
)
|
||||
|
||||
|
||||
def test_terminal_page_rejects_item_outside_query_range() -> None:
|
||||
item = make_trade_record(
|
||||
executed_at=START - timedelta(microseconds=1),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="query range"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item,),
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_more_items_than_query_limit() -> None:
|
||||
first = make_trade_record(replay_sequence=1)
|
||||
second = make_trade_record(
|
||||
executed_at=START + timedelta(minutes=2),
|
||||
replay_sequence=2,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="query limit"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(limit=1),
|
||||
items=(first, second),
|
||||
)
|
||||
|
||||
|
||||
def test_page_items_must_follow_incoming_cursor() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START + timedelta(minutes=1),
|
||||
replay_sequence=10,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="follow query cursor"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(cursor=cursor),
|
||||
items=(make_trade_record(replay_sequence=10),),
|
||||
)
|
||||
|
||||
|
||||
def test_records_reject_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()
|
||||
|
||||
trade_subclass = 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,
|
||||
)
|
||||
quote_subclass = 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,
|
||||
)
|
||||
candle_subclass = 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,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="Canonical Trade"):
|
||||
TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=trade_subclass,
|
||||
first_observed_at=trade.executed_at,
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=(trade.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="Canonical Quote"):
|
||||
QuoteHistoryRecord(
|
||||
venue=VENUE,
|
||||
quote=quote_subclass,
|
||||
observation_sources=(quote.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="Canonical Candle"):
|
||||
CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle_subclass,
|
||||
observed_at=candle.open_time,
|
||||
is_final=False,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_cursor_and_query_reject_time_range_subclass() -> None:
|
||||
class HistoricalTimeRangeSubclass(HistoricalTimeRange):
|
||||
pass
|
||||
|
||||
time_range = HistoricalTimeRangeSubclass(START, END)
|
||||
|
||||
with pytest.raises(TypeError, match="HistoricalTimeRange"):
|
||||
TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="HistoricalTimeRange"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
)
|
||||
|
||||
|
||||
def test_query_rejects_cursor_subclass() -> None:
|
||||
class TradeHistoryCursorSubclass(TradeHistoryCursor):
|
||||
pass
|
||||
|
||||
cursor = TradeHistoryCursorSubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryCursor"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_query_record_and_cursor_subclasses() -> None:
|
||||
class TradeHistoryQuerySubclass(TradeHistoryQuery):
|
||||
pass
|
||||
|
||||
class TradeHistoryRecordSubclass(TradeHistoryRecord):
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (END, 1)
|
||||
|
||||
class TradeHistoryCursorSubclass(TradeHistoryCursor):
|
||||
pass
|
||||
|
||||
item = make_trade_record(replay_sequence=10)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryQuery"):
|
||||
TradeHistoryPage(
|
||||
query=TradeHistoryQuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
),
|
||||
items=(item,),
|
||||
)
|
||||
|
||||
item_subclass = TradeHistoryRecordSubclass(
|
||||
venue=item.venue,
|
||||
trade=item.trade,
|
||||
first_observed_at=item.first_observed_at,
|
||||
last_observed_at=item.last_observed_at,
|
||||
observation_sources=item.observation_sources,
|
||||
replay_sequence=item.replay_sequence,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryRecord"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item_subclass,),
|
||||
)
|
||||
|
||||
cursor_subclass = TradeHistoryCursorSubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=item.event_time,
|
||||
replay_sequence=item.replay_sequence,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryCursor"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item,),
|
||||
next_cursor=cursor_subclass,
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_items_tuple_subclass() -> None:
|
||||
class ItemsTupleSubclass(tuple):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="items must be a tuple"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=ItemsTupleSubclass((make_trade_record(),)),
|
||||
)
|
||||
@@ -0,0 +1,195 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access.contracts import (
|
||||
MarketDataHistoricalAccessProtocol,
|
||||
)
|
||||
from src.market_data.access.market_data_historical_access import (
|
||||
MarketDataHistoricalAccess,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
HistoricalTimeRange,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
)
|
||||
|
||||
|
||||
NOW = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
TIME_RANGE = HistoricalTimeRange(
|
||||
start_time=NOW,
|
||||
end_time=NOW + timedelta(hours=1),
|
||||
)
|
||||
|
||||
|
||||
class RecordingTradeReader:
|
||||
def __init__(self, page: TradeHistoryPage) -> None:
|
||||
self.page = page
|
||||
self.queries: list[TradeHistoryQuery] = []
|
||||
|
||||
def query_trades(self, query: TradeHistoryQuery) -> TradeHistoryPage:
|
||||
self.queries.append(query)
|
||||
return self.page
|
||||
|
||||
|
||||
class RecordingQuoteReader:
|
||||
def __init__(self, page: QuoteHistoryPage) -> None:
|
||||
self.page = page
|
||||
self.queries: list[QuoteHistoryQuery] = []
|
||||
|
||||
def query_quotes(self, query: QuoteHistoryQuery) -> QuoteHistoryPage:
|
||||
self.queries.append(query)
|
||||
return self.page
|
||||
|
||||
|
||||
class RecordingCandleReader:
|
||||
def __init__(self, page: CandleRevisionHistoryPage) -> None:
|
||||
self.page = page
|
||||
self.queries: list[CandleRevisionHistoryQuery] = []
|
||||
|
||||
def query_candle_revisions(
|
||||
self,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
self.queries.append(query)
|
||||
return self.page
|
||||
|
||||
|
||||
class BrokenTradeReader:
|
||||
def __init__(self, error: RuntimeError) -> None:
|
||||
self.error = error
|
||||
|
||||
def query_trades(self, query: TradeHistoryQuery) -> TradeHistoryPage:
|
||||
del query
|
||||
raise self.error
|
||||
|
||||
|
||||
def make_dependencies() -> tuple[
|
||||
MarketDataHistoricalAccess,
|
||||
RecordingTradeReader,
|
||||
RecordingQuoteReader,
|
||||
RecordingCandleReader,
|
||||
]:
|
||||
trade_query = TradeHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
quote_query = QuoteHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
candle_query = CandleRevisionHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
interval="1m",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
trade_reader = RecordingTradeReader(
|
||||
TradeHistoryPage(query=trade_query, items=()),
|
||||
)
|
||||
quote_reader = RecordingQuoteReader(
|
||||
QuoteHistoryPage(query=quote_query, items=()),
|
||||
)
|
||||
candle_reader = RecordingCandleReader(
|
||||
CandleRevisionHistoryPage(query=candle_query, items=()),
|
||||
)
|
||||
access = MarketDataHistoricalAccess(
|
||||
trade_reader=trade_reader,
|
||||
quote_reader=quote_reader,
|
||||
candle_revision_reader=candle_reader,
|
||||
)
|
||||
return access, trade_reader, quote_reader, candle_reader
|
||||
|
||||
|
||||
def test_implements_combined_protocol_and_uses_slots() -> None:
|
||||
access, *_ = make_dependencies()
|
||||
|
||||
assert isinstance(access, MarketDataHistoricalAccessProtocol)
|
||||
assert not hasattr(access, "__dict__")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"dependency_name",
|
||||
(
|
||||
"trade_reader",
|
||||
"quote_reader",
|
||||
"candle_revision_reader",
|
||||
),
|
||||
)
|
||||
def test_rejects_dependency_without_required_protocol(
|
||||
dependency_name: str,
|
||||
) -> None:
|
||||
_, trade_reader, quote_reader, candle_reader = make_dependencies()
|
||||
dependencies: dict[str, Any] = {
|
||||
"trade_reader": trade_reader,
|
||||
"quote_reader": quote_reader,
|
||||
"candle_revision_reader": candle_reader,
|
||||
}
|
||||
dependencies[dependency_name] = object()
|
||||
|
||||
with pytest.raises(TypeError, match=dependency_name):
|
||||
MarketDataHistoricalAccess(**dependencies)
|
||||
|
||||
|
||||
def test_delegates_each_query_without_rebuilding_page() -> None:
|
||||
access, trade_reader, quote_reader, candle_reader = make_dependencies()
|
||||
trade_query = TradeHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
limit=10,
|
||||
)
|
||||
quote_query = QuoteHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
limit=20,
|
||||
)
|
||||
candle_query = CandleRevisionHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
interval="1m",
|
||||
time_range=TIME_RANGE,
|
||||
limit=30,
|
||||
)
|
||||
|
||||
trade_page = access.query_trades(trade_query)
|
||||
quote_page = access.query_quotes(quote_query)
|
||||
candle_page = access.query_candle_revisions(candle_query)
|
||||
|
||||
assert trade_page is trade_reader.page
|
||||
assert quote_page is quote_reader.page
|
||||
assert candle_page is candle_reader.page
|
||||
assert trade_reader.queries == [trade_query]
|
||||
assert quote_reader.queries == [quote_query]
|
||||
assert candle_reader.queries == [candle_query]
|
||||
|
||||
|
||||
def test_does_not_swallow_reader_error() -> None:
|
||||
_, _, quote_reader, candle_reader = make_dependencies()
|
||||
expected = RuntimeError("reader failed")
|
||||
broken_reader = BrokenTradeReader(expected)
|
||||
access = MarketDataHistoricalAccess(
|
||||
trade_reader=broken_reader,
|
||||
quote_reader=quote_reader,
|
||||
candle_revision_reader=candle_reader,
|
||||
)
|
||||
query = TradeHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
access.query_trades(query)
|
||||
|
||||
assert captured.value is expected
|
||||
@@ -0,0 +1,620 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
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.contracts import (
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
)
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryCursor,
|
||||
CandleRevisionHistoryQuery,
|
||||
HistoricalTimeRange,
|
||||
)
|
||||
from src.market_data.access.postgres_candle_revision_history_repository import (
|
||||
PostgresCandleRevisionHistoryRepository,
|
||||
)
|
||||
from src.market_data.acquisition.models.candle import Candle
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
INTERVAL = "1m"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_candle"
|
||||
_DEFAULT = object()
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(self, rows: object = ()) -> None:
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.enter_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.execute_error: BaseException | None = None
|
||||
self.fetchall_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
|
||||
if self.execute_error is not None:
|
||||
raise self.execute_error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(self, cursor: RecordingCursor) -> None:
|
||||
self._cursor = cursor
|
||||
self.enter_calls = 0
|
||||
self.cursor_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.cursor_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
|
||||
if self.cursor_error is not None:
|
||||
raise self.cursor_error
|
||||
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def make_query(
|
||||
*,
|
||||
interval: str = INTERVAL,
|
||||
limit: int = 3,
|
||||
cursor: CandleRevisionHistoryCursor | None = None,
|
||||
) -> CandleRevisionHistoryQuery:
|
||||
return CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=interval,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_row(
|
||||
*,
|
||||
venue: object = VENUE,
|
||||
symbol: object = SYMBOL,
|
||||
interval: object = INTERVAL,
|
||||
open_time: object = START + timedelta(minutes=1),
|
||||
observed_at: object = _DEFAULT,
|
||||
open_price: object = Decimal("64000"),
|
||||
high_price: object = Decimal("64200"),
|
||||
low_price: object = Decimal("63900"),
|
||||
close_price: object = Decimal("64150"),
|
||||
volume: object = Decimal("1.25"),
|
||||
is_final: object = False,
|
||||
source: object = SOURCE,
|
||||
observation_sources: object = _DEFAULT,
|
||||
replay_sequence: object = 10,
|
||||
canonical_schema_version: object = 1,
|
||||
) -> tuple[object, ...]:
|
||||
resolved_observed_at = (
|
||||
START + timedelta(minutes=1, seconds=10)
|
||||
if observed_at is _DEFAULT
|
||||
else observed_at
|
||||
)
|
||||
resolved_sources = (
|
||||
[SOURCE]
|
||||
if observation_sources is _DEFAULT
|
||||
else observation_sources
|
||||
)
|
||||
return (
|
||||
venue,
|
||||
symbol,
|
||||
interval,
|
||||
open_time,
|
||||
resolved_observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
source,
|
||||
resolved_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresCandleRevisionHistoryRepository,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
]:
|
||||
cursor = RecordingCursor(rows)
|
||||
connection = RecordingConnection(cursor)
|
||||
provider = RecordingProvider(connection)
|
||||
repository = PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=provider,
|
||||
)
|
||||
return repository, cursor, connection, provider
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
repository, _, _, provider = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(repository, "__dict__")
|
||||
assert isinstance(repository, CandleRevisionHistoryReaderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_query_is_validated_before_connection_borrow() -> None:
|
||||
class QuerySubclass(CandleRevisionHistoryQuery):
|
||||
pass
|
||||
|
||||
repository, _, _, provider = dependencies()
|
||||
query = QuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=INTERVAL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessValidationError, match="query"):
|
||||
repository.query_candle_revisions(query)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_first_page_uses_open_time_half_open_order_and_limit_plus_one() -> None:
|
||||
first_open = START + timedelta(minutes=1)
|
||||
second_open = START + timedelta(minutes=2)
|
||||
repository, cursor, connection, provider = dependencies(
|
||||
[
|
||||
make_row(open_time=first_open, replay_sequence=10),
|
||||
make_row(
|
||||
open_time=second_open,
|
||||
observed_at=second_open + timedelta(seconds=10),
|
||||
is_final=True,
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
query = make_query(limit=3)
|
||||
|
||||
page = repository.query_candle_revisions(query)
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "open_time >= %s" in compact_sql
|
||||
assert "open_time < %s" in compact_sql
|
||||
assert "observed_at >= %s" not in compact_sql
|
||||
assert "(open_time, replay_sequence) >" not in compact_sql
|
||||
assert "ORDER BY open_time ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert parameters == (VENUE, SYMBOL, INTERVAL, START, END, 4)
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert type(page.items[0].candle) is Candle
|
||||
assert page.items[0].event_time is first_open
|
||||
assert page.items[1].is_final is True
|
||||
assert page.next_cursor is None
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert cursor.enter_calls == 1
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [None]
|
||||
|
||||
|
||||
def test_interval_is_case_sensitive_and_not_normalized() -> None:
|
||||
repository, cursor, _, _ = dependencies([])
|
||||
query = make_query(interval="1M")
|
||||
|
||||
repository.query_candle_revisions(query)
|
||||
|
||||
assert cursor.calls[0][1][2] == "1M"
|
||||
|
||||
|
||||
def test_keyset_query_uses_exact_open_time_cursor_tuple() -> None:
|
||||
cursor_time = START + timedelta(minutes=5)
|
||||
incoming = CandleRevisionHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=INTERVAL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
open_time=cursor_time,
|
||||
replay_sequence=25,
|
||||
)
|
||||
repository, cursor, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=cursor_time,
|
||||
observed_at=cursor_time + timedelta(seconds=10),
|
||||
replay_sequence=26,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
repository.query_candle_revisions(make_query(limit=2, cursor=incoming))
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
assert (
|
||||
"(open_time, replay_sequence) > (%s, %s)"
|
||||
in normalized_sql(sql)
|
||||
)
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
INTERVAL,
|
||||
START,
|
||||
END,
|
||||
cursor_time,
|
||||
25,
|
||||
3,
|
||||
)
|
||||
|
||||
|
||||
def test_limit_plus_one_creates_cursor_from_last_returned_item() -> None:
|
||||
open_times = tuple(
|
||||
START + timedelta(minutes=index) for index in (1, 2, 3)
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=open_time,
|
||||
observed_at=open_time + timedelta(seconds=10),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index, open_time in enumerate(open_times)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2)
|
||||
|
||||
page = repository.query_candle_revisions(query)
|
||||
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert page.next_cursor is not None
|
||||
assert page.next_cursor.venue == VENUE
|
||||
assert page.next_cursor.symbol == SYMBOL
|
||||
assert page.next_cursor.interval == INTERVAL
|
||||
assert page.next_cursor.time_range is query.time_range
|
||||
assert page.next_cursor.open_time == open_times[1]
|
||||
assert page.next_cursor.replay_sequence == 11
|
||||
|
||||
|
||||
def test_empty_result_is_valid_without_cursor() -> None:
|
||||
repository, _, _, _ = dependencies([])
|
||||
|
||||
page = repository.query_candle_revisions(make_query())
|
||||
|
||||
assert page.items == ()
|
||||
assert page.next_cursor is None
|
||||
assert page.has_more is False
|
||||
|
||||
|
||||
def test_history_range_uses_open_time_not_observed_at() -> None:
|
||||
open_time = START + timedelta(minutes=1)
|
||||
observed_at = END + timedelta(minutes=10)
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(open_time=open_time, observed_at=observed_at)]
|
||||
)
|
||||
|
||||
page = repository.query_candle_revisions(make_query())
|
||||
|
||||
assert page.items[0].event_time == open_time
|
||||
assert page.items[0].replay_at == observed_at
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides", "error_match"),
|
||||
(
|
||||
({"venue": "other"}, "scope"),
|
||||
({"symbol": "ETH/USD_LEVERAGE"}, "scope"),
|
||||
({"interval": "1M"}, "scope"),
|
||||
({"open_time": END, "observed_at": END}, "range"),
|
||||
),
|
||||
)
|
||||
def test_rows_outside_exact_scope_or_range_are_rejected(
|
||||
overrides: dict[str, object],
|
||||
error_match: str,
|
||||
) -> None:
|
||||
repository, _, _, _ = dependencies([make_row(**overrides)])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match=error_match):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
(
|
||||
{"venue": " dzengi"},
|
||||
{"symbol": "btc/usd_leverage"},
|
||||
{"interval": " 1m"},
|
||||
{"open_time": START.replace(tzinfo=None)},
|
||||
{"observed_at": START.replace(tzinfo=None)},
|
||||
{"observed_at": START},
|
||||
{"open_price": Decimal("0")},
|
||||
{"high_price": Decimal("NaN")},
|
||||
{"low_price": Decimal("65000")},
|
||||
{"close_price": Decimal("65000")},
|
||||
{"volume": Decimal("-0.01")},
|
||||
{"is_final": 1},
|
||||
{"source": " source"},
|
||||
{"observation_sources": ()},
|
||||
{"observation_sources": []},
|
||||
{"observation_sources": [SOURCE, SOURCE]},
|
||||
{"observation_sources": ["recovery", SOURCE]},
|
||||
{"replay_sequence": True},
|
||||
{"replay_sequence": 0},
|
||||
{"canonical_schema_version": 2},
|
||||
),
|
||||
)
|
||||
def test_corrupt_stored_values_raise_integrity_error(
|
||||
overrides: dict[str, object],
|
||||
) -> None:
|
||||
repository, cursor, connection, _ = dependencies(
|
||||
[make_row(**overrides)]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="invalid Canonical Candle",
|
||||
):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
assert connection.exit_exception_types == [
|
||||
MarketDataAccessIntegrityError
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("row", ((), [object()] * 15, (object(),) * 14))
|
||||
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
|
||||
repository, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
|
||||
def test_invalid_rows_collection_and_excess_rows_are_rejected() -> None:
|
||||
repository, _, _, _ = dependencies(iter(()))
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=START + timedelta(minutes=index + 1),
|
||||
observed_at=START + timedelta(minutes=index + 1, seconds=1),
|
||||
replay_sequence=index + 1,
|
||||
)
|
||||
for index in range(3)
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
repository.query_candle_revisions(make_query(limit=1))
|
||||
|
||||
|
||||
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
|
||||
earlier = START + timedelta(minutes=1)
|
||||
later = START + timedelta(minutes=2)
|
||||
unordered, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=later,
|
||||
observed_at=later,
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
open_time=earlier,
|
||||
observed_at=earlier,
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
|
||||
unordered.query_candle_revisions(make_query())
|
||||
|
||||
duplicate, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=earlier,
|
||||
observed_at=earlier,
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
open_time=later,
|
||||
observed_at=later,
|
||||
replay_sequence=10,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="duplicate replay_sequence",
|
||||
):
|
||||
duplicate.query_candle_revisions(make_query())
|
||||
|
||||
|
||||
def test_rows_must_strictly_follow_incoming_cursor() -> None:
|
||||
cursor_time = START + timedelta(minutes=1)
|
||||
incoming = CandleRevisionHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=INTERVAL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
open_time=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=cursor_time,
|
||||
observed_at=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
|
||||
repository.query_candle_revisions(make_query(cursor=incoming))
|
||||
|
||||
|
||||
def test_database_error_is_wrapped_after_context_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = RuntimeError("database failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [RuntimeError]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
|
||||
def test_provider_error_is_wrapped_without_entering_connection() -> None:
|
||||
repository, _, connection, provider = dependencies()
|
||||
provider.error = RuntimeError("provider failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
|
||||
|
||||
def test_base_exception_is_not_wrapped_and_reaches_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [KeyboardInterrupt]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = RuntimeError("cursor cleanup failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
|
||||
cursor = RecordingCursor([make_row()])
|
||||
connection = RecordingConnection(cursor)
|
||||
provider_calls = 0
|
||||
|
||||
def bound_provider() -> Any:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
return nullcontext(connection)
|
||||
|
||||
repository = PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=bound_provider,
|
||||
)
|
||||
|
||||
page = repository.query_candle_revisions(make_query())
|
||||
|
||||
assert len(page.items) == 1
|
||||
assert provider_calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
assert connection.exit_exception_types == []
|
||||
assert cursor.exit_exception_types == [None]
|
||||
@@ -0,0 +1,529 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
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.contracts import QuoteHistoryReaderProtocol
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
HistoricalTimeRange,
|
||||
QuoteHistoryCursor,
|
||||
QuoteHistoryQuery,
|
||||
)
|
||||
from src.market_data.access.postgres_quote_history_repository import (
|
||||
PostgresQuoteHistoryRepository,
|
||||
)
|
||||
from src.market_data.acquisition.models.quote import Quote
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_quote"
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(self, rows: object = ()) -> None:
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.enter_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.execute_error: BaseException | None = None
|
||||
self.fetchall_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
|
||||
if self.execute_error is not None:
|
||||
raise self.execute_error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(self, cursor: RecordingCursor) -> None:
|
||||
self._cursor = cursor
|
||||
self.enter_calls = 0
|
||||
self.cursor_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.cursor_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
|
||||
if self.cursor_error is not None:
|
||||
raise self.cursor_error
|
||||
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def make_query(
|
||||
*,
|
||||
limit: int = 3,
|
||||
cursor: QuoteHistoryCursor | None = None,
|
||||
) -> QuoteHistoryQuery:
|
||||
return QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_row(
|
||||
*,
|
||||
venue: object = VENUE,
|
||||
symbol: object = SYMBOL,
|
||||
received_at: object = START + timedelta(minutes=1),
|
||||
exchange_timestamp: object = START + timedelta(seconds=59),
|
||||
last_price: object = Decimal("64159.45"),
|
||||
bid_price: object = Decimal("64159.40"),
|
||||
ask_price: object = Decimal("64159.50"),
|
||||
source: object = SOURCE,
|
||||
observation_sources: object = None,
|
||||
replay_sequence: object = 10,
|
||||
canonical_schema_version: object = 1,
|
||||
) -> tuple[object, ...]:
|
||||
sources = [SOURCE] if observation_sources is None else observation_sources
|
||||
return (
|
||||
venue,
|
||||
symbol,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
source,
|
||||
sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresQuoteHistoryRepository,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
]:
|
||||
cursor = RecordingCursor(rows)
|
||||
connection = RecordingConnection(cursor)
|
||||
provider = RecordingProvider(connection)
|
||||
repository = PostgresQuoteHistoryRepository(
|
||||
connection_provider=provider,
|
||||
)
|
||||
return repository, cursor, connection, provider
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
repository, _, _, provider = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(repository, "__dict__")
|
||||
assert isinstance(repository, QuoteHistoryReaderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresQuoteHistoryRepository(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_query_is_validated_before_connection_borrow() -> None:
|
||||
class QuoteHistoryQuerySubclass(QuoteHistoryQuery):
|
||||
pass
|
||||
|
||||
repository, _, _, provider = dependencies()
|
||||
query = QuoteHistoryQuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessValidationError, match="query"):
|
||||
repository.query_quotes(query)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_first_page_uses_received_time_half_open_order_and_limit() -> None:
|
||||
first_time = START + timedelta(minutes=1)
|
||||
second_time = START + timedelta(minutes=2)
|
||||
repository, cursor, connection, provider = dependencies(
|
||||
[
|
||||
make_row(received_at=first_time, replay_sequence=10),
|
||||
make_row(received_at=second_time, replay_sequence=11),
|
||||
]
|
||||
)
|
||||
query = make_query(limit=3)
|
||||
|
||||
page = repository.query_quotes(query)
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "received_at >= %s" in compact_sql
|
||||
assert "received_at < %s" in compact_sql
|
||||
assert "(received_at, replay_sequence) >" not in compact_sql
|
||||
assert "ORDER BY received_at ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert parameters == (VENUE, SYMBOL, START, END, 4)
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert type(page.items[0].quote) is Quote
|
||||
assert page.items[0].quote.received_at is first_time
|
||||
assert page.items[0].observation_sources == (SOURCE,)
|
||||
assert page.next_cursor is None
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert cursor.enter_calls == 1
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [None]
|
||||
|
||||
|
||||
def test_keyset_query_uses_exact_received_at_and_sequence_cursor() -> None:
|
||||
cursor_position = START + timedelta(minutes=5)
|
||||
incoming_cursor = QuoteHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
received_at=cursor_position,
|
||||
replay_sequence=25,
|
||||
)
|
||||
repository, recording_cursor, _, _ = dependencies(
|
||||
[make_row(received_at=cursor_position, replay_sequence=26)]
|
||||
)
|
||||
|
||||
repository.query_quotes(make_query(limit=2, cursor=incoming_cursor))
|
||||
|
||||
sql, parameters = recording_cursor.calls[0]
|
||||
assert (
|
||||
"(received_at, replay_sequence) > (%s, %s)"
|
||||
in normalized_sql(sql)
|
||||
)
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
START,
|
||||
END,
|
||||
cursor_position,
|
||||
25,
|
||||
3,
|
||||
)
|
||||
|
||||
|
||||
def test_limit_plus_one_creates_cursor_from_last_returned_quote() -> None:
|
||||
times = tuple(START + timedelta(minutes=index) for index in (1, 2, 3))
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(received_at=event_time, replay_sequence=10 + index)
|
||||
for index, event_time in enumerate(times)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2)
|
||||
|
||||
page = repository.query_quotes(query)
|
||||
|
||||
assert len(page.items) == 2
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert page.next_cursor is not None
|
||||
assert page.next_cursor.venue == query.venue
|
||||
assert page.next_cursor.symbol == query.symbol
|
||||
assert page.next_cursor.time_range is query.time_range
|
||||
assert page.next_cursor.received_at == times[1]
|
||||
assert page.next_cursor.replay_sequence == 11
|
||||
|
||||
|
||||
def test_empty_result_and_exact_start_are_valid() -> None:
|
||||
empty_repository, _, _, _ = dependencies([])
|
||||
|
||||
empty_page = empty_repository.query_quotes(make_query())
|
||||
|
||||
assert empty_page.items == ()
|
||||
assert empty_page.next_cursor is None
|
||||
assert empty_page.has_more is False
|
||||
|
||||
start_repository, _, _, _ = dependencies(
|
||||
[make_row(received_at=START)]
|
||||
)
|
||||
|
||||
start_page = start_repository.query_quotes(make_query())
|
||||
|
||||
assert start_page.items[0].event_time == START
|
||||
|
||||
|
||||
def test_none_exchange_timestamp_is_preserved() -> None:
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(exchange_timestamp=None)]
|
||||
)
|
||||
|
||||
page = repository.query_quotes(make_query())
|
||||
|
||||
assert page.items[0].quote.exchange_timestamp is None
|
||||
|
||||
|
||||
def test_backend_cannot_return_more_than_limit_plus_one() -> None:
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
received_at=START + timedelta(minutes=index + 1),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index in range(3)
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
repository.query_quotes(make_query(limit=1))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
(
|
||||
{"venue": "other"},
|
||||
{"venue": " dzengi "},
|
||||
{"symbol": "btc/usd_leverage"},
|
||||
{"received_at": END},
|
||||
{"received_at": START.replace(tzinfo=None)},
|
||||
{"exchange_timestamp": START.replace(tzinfo=None)},
|
||||
{"last_price": Decimal("0")},
|
||||
{"last_price": Decimal("NaN")},
|
||||
{"last_price": 1},
|
||||
{"bid_price": Decimal("64160")},
|
||||
{"ask_price": Decimal("0")},
|
||||
{"source": " source "},
|
||||
{"observation_sources": ()},
|
||||
{"observation_sources": []},
|
||||
{"observation_sources": [SOURCE, SOURCE]},
|
||||
{"observation_sources": ["recovery", SOURCE]},
|
||||
{"replay_sequence": True},
|
||||
{"replay_sequence": 0},
|
||||
{"canonical_schema_version": 2},
|
||||
),
|
||||
)
|
||||
def test_corrupt_or_out_of_scope_values_raise_integrity_error(
|
||||
overrides: dict[str, object],
|
||||
) -> None:
|
||||
repository, cursor, connection, _ = dependencies(
|
||||
[make_row(**overrides)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
assert connection.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"row",
|
||||
((), [object()] * 11, (object(),) * 10),
|
||||
)
|
||||
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
|
||||
repository, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rows", (None, "rows", object()))
|
||||
def test_invalid_rows_collection_raises_integrity_error(rows: object) -> None:
|
||||
repository, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
|
||||
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
|
||||
later = START + timedelta(minutes=2)
|
||||
earlier = START + timedelta(minutes=1)
|
||||
unordered_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(received_at=later, replay_sequence=10),
|
||||
make_row(received_at=earlier, replay_sequence=11),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
|
||||
unordered_repository.query_quotes(make_query())
|
||||
|
||||
duplicate_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(received_at=earlier, replay_sequence=10),
|
||||
make_row(received_at=later, replay_sequence=10),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="duplicate replay_sequence",
|
||||
):
|
||||
duplicate_repository.query_quotes(make_query())
|
||||
|
||||
|
||||
def test_rows_must_strictly_follow_incoming_cursor() -> None:
|
||||
cursor_time = START + timedelta(minutes=1)
|
||||
incoming = QuoteHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
received_at=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(received_at=cursor_time, replay_sequence=10)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
|
||||
repository.query_quotes(make_query(cursor=incoming))
|
||||
|
||||
|
||||
def test_database_error_is_wrapped_after_context_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = RuntimeError("database failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [RuntimeError]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
|
||||
def test_provider_error_is_wrapped_without_entering_connection() -> None:
|
||||
repository, _, connection, provider = dependencies()
|
||||
provider.error = RuntimeError("provider failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
|
||||
|
||||
def test_keyboard_interrupt_is_not_wrapped_and_reaches_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [KeyboardInterrupt]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = RuntimeError("cursor cleanup failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
|
||||
cursor = RecordingCursor([make_row()])
|
||||
connection = RecordingConnection(cursor)
|
||||
provider_calls = 0
|
||||
|
||||
def bound_provider() -> Any:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
return nullcontext(connection)
|
||||
|
||||
repository = PostgresQuoteHistoryRepository(
|
||||
connection_provider=bound_provider,
|
||||
)
|
||||
|
||||
page = repository.query_quotes(make_query())
|
||||
|
||||
assert len(page.items) == 1
|
||||
assert provider_calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
assert connection.exit_exception_types == []
|
||||
assert cursor.exit_exception_types == [None]
|
||||
@@ -0,0 +1,622 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
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 import (
|
||||
HistoricalTimeRange,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
PostgresTradeHistoryRepository,
|
||||
TradeHistoryCursor,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryReaderProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.acquisition.trade_id_sequence import (
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_trade"
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(self, rows: object = ()) -> None:
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.enter_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.execute_error: BaseException | None = None
|
||||
self.fetchall_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
|
||||
if self.execute_error is not None:
|
||||
raise self.execute_error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(self, cursor: RecordingCursor) -> None:
|
||||
self._cursor = cursor
|
||||
self.enter_calls = 0
|
||||
self.cursor_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.cursor_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
|
||||
if self.cursor_error is not None:
|
||||
raise self.cursor_error
|
||||
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def make_query(
|
||||
*,
|
||||
limit: int = 3,
|
||||
cursor: TradeHistoryCursor | None = None,
|
||||
) -> TradeHistoryQuery:
|
||||
return TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_row(
|
||||
*,
|
||||
venue: object = VENUE,
|
||||
symbol: object = SYMBOL,
|
||||
trade_id: object = 100,
|
||||
executed_at: object = START + timedelta(minutes=1),
|
||||
price: object = Decimal("64159.45"),
|
||||
quantity: object = Decimal("0.125"),
|
||||
aggressor_side: object = "buy",
|
||||
source: object = SOURCE,
|
||||
first_observed_at: object = START + timedelta(minutes=1, seconds=1),
|
||||
last_observed_at: object = START + timedelta(minutes=1, seconds=2),
|
||||
observation_sources: object = None,
|
||||
replay_sequence: object = 10,
|
||||
canonical_schema_version: object = 1,
|
||||
) -> tuple[object, ...]:
|
||||
sources = [SOURCE] if observation_sources is None else observation_sources
|
||||
return (
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresTradeHistoryRepository,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
]:
|
||||
cursor = RecordingCursor(rows)
|
||||
connection = RecordingConnection(cursor)
|
||||
provider = RecordingProvider(connection)
|
||||
repository = PostgresTradeHistoryRepository(
|
||||
connection_provider=provider,
|
||||
)
|
||||
return repository, cursor, connection, provider
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
repository, _, _, provider = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(repository, "__dict__")
|
||||
assert isinstance(repository, TradeHistoryReaderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresTradeHistoryRepository(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_query_is_validated_before_connection_borrow() -> None:
|
||||
class TradeHistoryQuerySubclass(TradeHistoryQuery):
|
||||
pass
|
||||
|
||||
repository, _, _, provider = dependencies()
|
||||
query = TradeHistoryQuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessValidationError, match="query"):
|
||||
repository.query_trades(query)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_first_page_uses_half_open_ordered_limit_plus_one_query() -> None:
|
||||
first_time = START + timedelta(minutes=1)
|
||||
second_time = START + timedelta(minutes=2)
|
||||
repository, cursor, connection, provider = dependencies(
|
||||
[
|
||||
make_row(executed_at=first_time, replay_sequence=10),
|
||||
make_row(
|
||||
trade_id=101,
|
||||
executed_at=second_time,
|
||||
first_observed_at=second_time + timedelta(seconds=1),
|
||||
last_observed_at=second_time + timedelta(seconds=2),
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
query = make_query(limit=3)
|
||||
|
||||
page = repository.query_trades(query)
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "executed_at >= %s" in compact_sql
|
||||
assert "executed_at < %s" in compact_sql
|
||||
assert "(executed_at, replay_sequence) >" not in compact_sql
|
||||
assert "ORDER BY executed_at ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert parameters == (VENUE, SYMBOL, START, END, 4)
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert type(page.items[0].trade) is Trade
|
||||
assert page.items[0].trade.executed_at is first_time
|
||||
assert page.items[0].trade.aggressor_side is TradeAggressorSide.BUY
|
||||
assert page.items[0].observation_sources == (SOURCE,)
|
||||
assert page.next_cursor is None
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert cursor.enter_calls == 1
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [None]
|
||||
|
||||
|
||||
def test_keyset_query_uses_exact_cursor_tuple() -> None:
|
||||
cursor_position = START + timedelta(minutes=5)
|
||||
incoming_cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
executed_at=cursor_position,
|
||||
replay_sequence=25,
|
||||
)
|
||||
row_time = cursor_position
|
||||
repository, recording_cursor, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=row_time,
|
||||
first_observed_at=row_time + timedelta(seconds=1),
|
||||
last_observed_at=row_time + timedelta(seconds=2),
|
||||
replay_sequence=26,
|
||||
)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2, cursor=incoming_cursor)
|
||||
|
||||
repository.query_trades(query)
|
||||
|
||||
sql, parameters = recording_cursor.calls[0]
|
||||
assert (
|
||||
"(executed_at, replay_sequence) > (%s, %s)"
|
||||
in normalized_sql(sql)
|
||||
)
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
START,
|
||||
END,
|
||||
cursor_position,
|
||||
25,
|
||||
3,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_id_rollover_does_not_participate_in_history_order() -> None:
|
||||
event_time = START + timedelta(minutes=1)
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
trade_id=SIGNED_TRADE_ID_MAX,
|
||||
executed_at=event_time,
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
trade_id=SIGNED_TRADE_ID_MIN,
|
||||
executed_at=event_time,
|
||||
replay_sequence=11,
|
||||
),
|
||||
make_row(
|
||||
trade_id=-1,
|
||||
executed_at=event_time,
|
||||
replay_sequence=12,
|
||||
),
|
||||
make_row(
|
||||
trade_id=0,
|
||||
executed_at=event_time,
|
||||
replay_sequence=13,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
page = repository.query_trades(make_query(limit=4))
|
||||
|
||||
assert [item.trade.trade_id for item in page.items] == [
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
-1,
|
||||
0,
|
||||
]
|
||||
|
||||
|
||||
def test_limit_plus_one_creates_cursor_from_last_returned_item() -> None:
|
||||
times = tuple(START + timedelta(minutes=index) for index in (1, 2, 3))
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
trade_id=100 + index,
|
||||
executed_at=event_time,
|
||||
first_observed_at=event_time + timedelta(seconds=1),
|
||||
last_observed_at=event_time + timedelta(seconds=2),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index, event_time in enumerate(times)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2)
|
||||
|
||||
page = repository.query_trades(query)
|
||||
|
||||
assert len(page.items) == 2
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert page.next_cursor is not None
|
||||
assert page.next_cursor.venue == query.venue
|
||||
assert page.next_cursor.symbol == query.symbol
|
||||
assert page.next_cursor.time_range is query.time_range
|
||||
assert page.next_cursor.executed_at == times[1]
|
||||
assert page.next_cursor.replay_sequence == 11
|
||||
|
||||
|
||||
def test_empty_result_is_valid_without_cursor() -> None:
|
||||
repository, _, _, _ = dependencies([])
|
||||
|
||||
page = repository.query_trades(make_query())
|
||||
|
||||
assert page.items == ()
|
||||
assert page.next_cursor is None
|
||||
assert page.has_more is False
|
||||
|
||||
|
||||
def test_half_open_range_includes_exact_start() -> None:
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=START,
|
||||
first_observed_at=START,
|
||||
last_observed_at=START,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
page = repository.query_trades(make_query())
|
||||
|
||||
assert page.items[0].event_time == START
|
||||
|
||||
|
||||
def test_backend_cannot_return_more_than_limit_plus_one() -> None:
|
||||
query = make_query(limit=1)
|
||||
rows = [
|
||||
make_row(
|
||||
trade_id=100 + index,
|
||||
executed_at=START + timedelta(minutes=index + 1),
|
||||
first_observed_at=START + timedelta(minutes=index + 1, seconds=1),
|
||||
last_observed_at=START + timedelta(minutes=index + 1, seconds=2),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index in range(3)
|
||||
]
|
||||
repository, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
repository.query_trades(query)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides", "error_match"),
|
||||
(
|
||||
({"venue": "other"}, "invalid Canonical Trade"),
|
||||
({"symbol": "btc/usd_leverage"}, "invalid Canonical Trade"),
|
||||
({"trade_id": True}, "invalid Canonical Trade"),
|
||||
({"trade_id": SIGNED_TRADE_ID_MAX + 1}, "invalid Canonical Trade"),
|
||||
({"executed_at": END}, "invalid Canonical Trade"),
|
||||
({"price": Decimal("0")}, "invalid Canonical Trade"),
|
||||
({"price": Decimal("NaN")}, "invalid Canonical Trade"),
|
||||
({"quantity": Decimal("0")}, "invalid Canonical Trade"),
|
||||
({"aggressor_side": "hold"}, "invalid Canonical Trade"),
|
||||
({"source": " source "}, "invalid Canonical Trade"),
|
||||
(
|
||||
{"first_observed_at": START.replace(tzinfo=None)},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
(
|
||||
{
|
||||
"first_observed_at": START + timedelta(minutes=3),
|
||||
"last_observed_at": START + timedelta(minutes=2),
|
||||
},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
({"observation_sources": ()}, "invalid Canonical Trade"),
|
||||
({"observation_sources": []}, "invalid Canonical Trade"),
|
||||
(
|
||||
{"observation_sources": [SOURCE, SOURCE]},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
(
|
||||
{"observation_sources": ["recovery", SOURCE]},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
({"replay_sequence": True}, "invalid Canonical Trade"),
|
||||
({"replay_sequence": 0}, "invalid Canonical Trade"),
|
||||
({"canonical_schema_version": 2}, "invalid Canonical Trade"),
|
||||
),
|
||||
)
|
||||
def test_corrupt_stored_values_raise_integrity_error(
|
||||
overrides: dict[str, object],
|
||||
error_match: str,
|
||||
) -> None:
|
||||
repository, cursor, connection, _ = dependencies(
|
||||
[make_row(**overrides)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match=error_match):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
assert connection.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("row", ((), [object()] * 13, (object(),) * 12))
|
||||
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
|
||||
repository, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
|
||||
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
|
||||
later = START + timedelta(minutes=2)
|
||||
earlier = START + timedelta(minutes=1)
|
||||
unordered_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=later,
|
||||
first_observed_at=later + timedelta(seconds=1),
|
||||
last_observed_at=later + timedelta(seconds=2),
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
executed_at=earlier,
|
||||
first_observed_at=earlier + timedelta(seconds=1),
|
||||
last_observed_at=earlier + timedelta(seconds=2),
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
|
||||
unordered_repository.query_trades(make_query())
|
||||
|
||||
duplicate_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=earlier,
|
||||
first_observed_at=earlier + timedelta(seconds=1),
|
||||
last_observed_at=earlier + timedelta(seconds=2),
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
trade_id=101,
|
||||
executed_at=later,
|
||||
first_observed_at=later + timedelta(seconds=1),
|
||||
last_observed_at=later + timedelta(seconds=2),
|
||||
replay_sequence=10,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="duplicate replay_sequence",
|
||||
):
|
||||
duplicate_repository.query_trades(make_query())
|
||||
|
||||
|
||||
def test_rows_must_strictly_follow_incoming_cursor() -> None:
|
||||
cursor_time = START + timedelta(minutes=1)
|
||||
incoming = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
executed_at=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(executed_at=cursor_time, replay_sequence=10)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
|
||||
repository.query_trades(make_query(cursor=incoming))
|
||||
|
||||
|
||||
def test_database_error_is_wrapped_after_context_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = RuntimeError("database failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [RuntimeError]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
|
||||
def test_provider_error_is_wrapped_without_entering_connection() -> None:
|
||||
repository, _, connection, provider = dependencies()
|
||||
provider.error = RuntimeError("provider failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
|
||||
|
||||
def test_keyboard_interrupt_is_not_wrapped_and_reaches_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [KeyboardInterrupt]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = RuntimeError("cursor cleanup failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
|
||||
cursor = RecordingCursor([make_row()])
|
||||
connection = RecordingConnection(cursor)
|
||||
provider_calls = 0
|
||||
|
||||
def bound_provider() -> Any:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
return nullcontext(connection)
|
||||
|
||||
repository = PostgresTradeHistoryRepository(
|
||||
connection_provider=bound_provider,
|
||||
)
|
||||
|
||||
page = repository.query_trades(make_query())
|
||||
|
||||
assert len(page.items) == 1
|
||||
assert provider_calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
assert connection.exit_exception_types == []
|
||||
assert cursor.exit_exception_types == [None]
|
||||
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
|
||||
)
|
||||
@@ -38,6 +38,7 @@ TradeRow = dict[str, Any]
|
||||
@dataclass
|
||||
class TransactionalTradeDatabase:
|
||||
rows: dict[TradeKey, TradeRow] = field(default_factory=dict)
|
||||
next_replay_sequence: int = 1
|
||||
|
||||
|
||||
class TransactionalCursor:
|
||||
@@ -47,6 +48,7 @@ class TransactionalCursor:
|
||||
) -> None:
|
||||
self._connection = connection
|
||||
self._fetchone_result: object = None
|
||||
self._fetchall_result: list[tuple[int]] = []
|
||||
|
||||
def __enter__(self) -> TransactionalCursor:
|
||||
return self
|
||||
@@ -78,6 +80,12 @@ class TransactionalCursor:
|
||||
self._insert(parameters)
|
||||
return
|
||||
|
||||
if normalized.startswith(
|
||||
"SELECT nextval('market_data.replay_sequence'::regclass)"
|
||||
):
|
||||
self._allocate_replay_sequences(parameters)
|
||||
return
|
||||
|
||||
if normalized.startswith("SELECT price, quantity"):
|
||||
self._select(parameters)
|
||||
return
|
||||
@@ -91,21 +99,45 @@ class TransactionalCursor:
|
||||
def fetchone(self) -> object:
|
||||
return self._fetchone_result
|
||||
|
||||
def fetchall(self) -> list[tuple[int]]:
|
||||
return list(self._fetchall_result)
|
||||
|
||||
def _insert(self, parameters: tuple[Any, ...]) -> None:
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
canonical_schema_version,
|
||||
) = parameters
|
||||
if len(parameters) == 12:
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
canonical_schema_version,
|
||||
) = parameters
|
||||
replay_sequence = self._next_replay_sequence()
|
||||
elif len(parameters) == 13:
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
) = parameters
|
||||
else:
|
||||
raise AssertionError("Unexpected Trade INSERT parameters")
|
||||
|
||||
key = (venue, symbol, trade_id, executed_at)
|
||||
working_rows = self._connection.working_rows
|
||||
|
||||
@@ -121,10 +153,26 @@ class TransactionalCursor:
|
||||
"first_observed_at": first_observed_at,
|
||||
"last_observed_at": last_observed_at,
|
||||
"observation_sources": list(observation_sources),
|
||||
"replay_sequence": replay_sequence,
|
||||
"canonical_schema_version": canonical_schema_version,
|
||||
}
|
||||
self._fetchone_result = (1,)
|
||||
|
||||
def _allocate_replay_sequences(
|
||||
self,
|
||||
parameters: tuple[Any, ...],
|
||||
) -> None:
|
||||
(count,) = parameters
|
||||
self._fetchall_result = [
|
||||
(self._next_replay_sequence(),)
|
||||
for _ in range(count)
|
||||
]
|
||||
|
||||
def _next_replay_sequence(self) -> int:
|
||||
value = self._connection.database.next_replay_sequence
|
||||
self._connection.database.next_replay_sequence += 1
|
||||
return value
|
||||
|
||||
def _select(self, parameters: tuple[Any, ...]) -> None:
|
||||
key = parameters
|
||||
row = self._connection.working_rows.get(key)
|
||||
@@ -181,6 +229,10 @@ class TransactionalConnection:
|
||||
def cursor(self) -> TransactionalCursor:
|
||||
return TransactionalCursor(self)
|
||||
|
||||
@property
|
||||
def database(self) -> TransactionalTradeDatabase:
|
||||
return self._database
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingConnectionProvider:
|
||||
@@ -261,6 +313,7 @@ def test_store_trade_inserts_canonical_payload_and_provenance() -> None:
|
||||
"first_observed_at": OBSERVED_AT,
|
||||
"last_observed_at": OBSERVED_AT,
|
||||
"observation_sources": ["dzengi_websocket_trade"],
|
||||
"replay_sequence": 1,
|
||||
"canonical_schema_version": 1,
|
||||
}
|
||||
|
||||
@@ -459,6 +512,122 @@ def test_batch_uses_stable_identity_order() -> None:
|
||||
assert inserted_symbols == ("A", "B", "C")
|
||||
|
||||
|
||||
def test_batch_allocates_replay_sequence_in_input_order() -> None:
|
||||
repository, database, connection, _ = _repository()
|
||||
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
_trade(symbol="C", trade_id=3),
|
||||
_trade(symbol="A", trade_id=1),
|
||||
_trade(symbol="B", trade_id=2),
|
||||
),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
inserted = tuple(
|
||||
(parameters[1], parameters[11])
|
||||
for statement, parameters in connection.calls
|
||||
if statement.startswith("INSERT INTO market_data.trades")
|
||||
)
|
||||
|
||||
assert inserted == (
|
||||
("A", 2),
|
||||
("B", 3),
|
||||
("C", 1),
|
||||
)
|
||||
assert {
|
||||
key[1]: row["replay_sequence"]
|
||||
for key, row in database.rows.items()
|
||||
} == {
|
||||
"A": 2,
|
||||
"B": 3,
|
||||
"C": 1,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("first_trade_id", "second_trade_id"),
|
||||
(
|
||||
(SIGNED_TRADE_ID_MAX, SIGNED_TRADE_ID_MIN),
|
||||
(-1, 0),
|
||||
),
|
||||
)
|
||||
def test_batch_replay_sequence_preserves_signed_rollover_input_order(
|
||||
first_trade_id: int,
|
||||
second_trade_id: int,
|
||||
) -> None:
|
||||
repository, database, _, _ = _repository()
|
||||
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
_trade(trade_id=first_trade_id),
|
||||
_trade(trade_id=second_trade_id),
|
||||
),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
assert database.rows[
|
||||
(VENUE, SYMBOL, first_trade_id, EXECUTED_AT)
|
||||
]["replay_sequence"] == 1
|
||||
assert database.rows[
|
||||
(VENUE, SYMBOL, second_trade_id, EXECUTED_AT)
|
||||
]["replay_sequence"] == 2
|
||||
|
||||
|
||||
def test_batch_duplicate_and_provenance_keep_original_replay_sequence() -> None:
|
||||
repository, database, _, _ = _repository()
|
||||
original = _trade()
|
||||
repository.store_trade(
|
||||
venue=VENUE,
|
||||
trade=original,
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
replace(original, source="dzengi"),
|
||||
original,
|
||||
),
|
||||
observed_at=OBSERVED_AT + timedelta(seconds=1),
|
||||
)
|
||||
|
||||
assert _only_row(database)["replay_sequence"] == 1
|
||||
assert database.next_replay_sequence == 4
|
||||
|
||||
|
||||
def test_failed_batch_keeps_consumed_replay_sequence_gap() -> None:
|
||||
repository, database, _, _ = _repository()
|
||||
existing = _trade(symbol="B", trade_id=2)
|
||||
repository.store_trade(
|
||||
venue=VENUE,
|
||||
trade=existing,
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataStorageConflictError):
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
_trade(symbol="A", trade_id=1),
|
||||
replace(existing, price=Decimal("999")),
|
||||
),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
repository.store_trade(
|
||||
venue=VENUE,
|
||||
trade=_trade(symbol="C", trade_id=3),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
assert database.rows[
|
||||
(VENUE, "C", 3, EXECUTED_AT)
|
||||
]["replay_sequence"] == 4
|
||||
|
||||
|
||||
def test_batch_conflict_rolls_back_preceding_insert() -> None:
|
||||
repository, database, connection, _ = _repository()
|
||||
existing = _trade(symbol="B", trade_id=2)
|
||||
|
||||
@@ -7,6 +7,7 @@ import pytest
|
||||
|
||||
from src.storage.exceptions import StorageMigrationError
|
||||
from src.storage.migrations import (
|
||||
MARKET_DATA_PARTITION_ADVISORY_LOCK_ID,
|
||||
STORAGE_MIGRATION_ADVISORY_LOCK_ID,
|
||||
STORAGE_MIGRATIONS,
|
||||
StorageMigration,
|
||||
@@ -106,6 +107,7 @@ def test_default_migrations_have_stable_order_and_names() -> None:
|
||||
(6, "add_quote_and_candle_observation_sources"),
|
||||
(7, "create_market_data_partition_registry"),
|
||||
(8, "create_trade_stream_checkpoints"),
|
||||
(9, "add_global_market_data_replay_sequence"),
|
||||
)
|
||||
|
||||
|
||||
@@ -158,12 +160,73 @@ def test_default_schema_defines_partitions_identities_and_constraints() -> None:
|
||||
assert "CHECK (checkpoint_schema_version > 0)" in sql
|
||||
|
||||
|
||||
def test_replay_sequence_migration_has_atomic_global_order_contract() -> None:
|
||||
migration = STORAGE_MIGRATIONS[-1]
|
||||
statements = tuple(
|
||||
" ".join(statement.split())
|
||||
for statement in migration.statements
|
||||
)
|
||||
sql = "\n".join(statements)
|
||||
|
||||
assert migration.version == 9
|
||||
assert migration.name == "add_global_market_data_replay_sequence"
|
||||
assert statements[:4] == (
|
||||
(
|
||||
"SELECT pg_advisory_xact_lock("
|
||||
f"{MARKET_DATA_PARTITION_ADVISORY_LOCK_ID}"
|
||||
")"
|
||||
),
|
||||
"LOCK TABLE market_data.trades IN ACCESS EXCLUSIVE MODE",
|
||||
"LOCK TABLE market_data.quotes IN ACCESS EXCLUSIVE MODE",
|
||||
(
|
||||
"LOCK TABLE market_data.candle_revisions "
|
||||
"IN ACCESS EXCLUSIVE MODE"
|
||||
),
|
||||
)
|
||||
assert "CREATE SEQUENCE market_data.replay_sequence AS BIGINT" in sql
|
||||
assert "MINVALUE 1" in sql
|
||||
assert "CACHE 1" in sql
|
||||
assert "NO CYCLE" in sql
|
||||
assert "OWNED BY NONE" in sql
|
||||
assert sql.count("ADD COLUMN replay_sequence BIGINT") == 3
|
||||
assert "CREATE TABLE market_data.replay_sequence" not in sql
|
||||
assert "CREATE TEMPORARY TABLE market_data_replay_sequence_backfill" in sql
|
||||
assert "ON COMMIT DROP" in sql
|
||||
assert "ROW_NUMBER() OVER" in sql
|
||||
assert "UNION ALL" in sql
|
||||
assert "event_time, data_type_rank, venue COLLATE \"C\"" in sql
|
||||
assert "symbol COLLATE \"C\"" in sql
|
||||
assert "candle_interval COLLATE \"C\"" in sql
|
||||
assert "ctid" not in sql.lower()
|
||||
assert sql.count("SET replay_sequence = backfill.replay_sequence") == 3
|
||||
assert "target.trade_id = backfill.trade_id" in sql
|
||||
assert "target.received_at = backfill.received_at" in sql
|
||||
assert "target.interval = backfill.interval" in sql
|
||||
assert "target.open_time = backfill.open_time" in sql
|
||||
assert "target.observed_at = backfill.observed_at" in sql
|
||||
assert "SELECT pg_catalog.setval(" in sql
|
||||
assert "EXISTS ( SELECT 1 FROM market_data_replay_sequence_backfill )" in sql
|
||||
assert sql.count("SET DEFAULT nextval(") == 3
|
||||
assert sql.count("ALTER COLUMN replay_sequence SET NOT NULL") == 3
|
||||
assert sql.count("CHECK (replay_sequence > 0)") == 3
|
||||
assert "CREATE FUNCTION market_data.reject_replay_sequence_change()" in sql
|
||||
assert sql.count("BEFORE UPDATE OF replay_sequence") == 3
|
||||
assert "CREATE INDEX trades_history_keyset_idx" in sql
|
||||
assert "CREATE INDEX quotes_history_keyset_idx" in sql
|
||||
assert "CREATE INDEX candle_revisions_history_keyset_idx" in sql
|
||||
assert "CREATE INDEX candle_revisions_replay_keyset_idx" in sql
|
||||
assert "executed_at, replay_sequence" in sql
|
||||
assert "received_at, replay_sequence" in sql
|
||||
assert "interval, open_time, replay_sequence" in sql
|
||||
assert "interval, observed_at, replay_sequence" in sql
|
||||
|
||||
|
||||
def test_run_locks_and_applies_every_pending_migration_in_order() -> None:
|
||||
runner, cursor, connection, provider = _runner()
|
||||
|
||||
result = runner.run()
|
||||
|
||||
assert result == (1, 2, 3, 4, 5, 6, 7, 8)
|
||||
assert result == (1, 2, 3, 4, 5, 6, 7, 8, 9)
|
||||
assert provider.calls == 1
|
||||
assert connection.entered == 1
|
||||
assert connection.exited == 1
|
||||
@@ -180,7 +243,7 @@ def test_run_locks_and_applies_every_pending_migration_in_order() -> None:
|
||||
)
|
||||
and isinstance(parameters, tuple)
|
||||
)
|
||||
assert inserted_versions == (1, 2, 3, 4, 5, 6, 7, 8)
|
||||
assert inserted_versions == (1, 2, 3, 4, 5, 6, 7, 8, 9)
|
||||
|
||||
|
||||
def test_run_skips_already_applied_migrations() -> None:
|
||||
@@ -209,7 +272,7 @@ def test_run_applies_only_migrations_after_existing_prefix() -> None:
|
||||
|
||||
result = runner.run()
|
||||
|
||||
assert result == (3, 4, 5, 6, 7, 8)
|
||||
assert result == (3, 4, 5, 6, 7, 8, 9)
|
||||
inserted_versions = tuple(
|
||||
parameters[0]
|
||||
for statement, parameters in cursor.calls
|
||||
@@ -218,7 +281,7 @@ def test_run_applies_only_migrations_after_existing_prefix() -> None:
|
||||
)
|
||||
and isinstance(parameters, tuple)
|
||||
)
|
||||
assert inserted_versions == (3, 4, 5, 6, 7, 8)
|
||||
assert inserted_versions == (3, 4, 5, 6, 7, 8, 9)
|
||||
|
||||
|
||||
def test_run_rejects_unknown_applied_version() -> None:
|
||||
|
||||
Reference in New Issue
Block a user