Build 060.29: implement Market Data Access and Replay

This commit is contained in:
2026-08-02 19:20:14 +03:00
parent 8c485e32b1
commit 8c98de9acc
48 changed files with 15614 additions and 58 deletions

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View 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=())

View File

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

View File

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

View File

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

View File

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

View 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")

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

View File

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