Build 060.29: implement Market Data Access and Replay
This commit is contained in:
@@ -0,0 +1,312 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import replace
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryQuery,
|
||||
HistoricalTimeRange,
|
||||
QuoteHistoryQuery,
|
||||
)
|
||||
from src.market_data.access.postgres_candle_revision_history_repository import (
|
||||
PostgresCandleRevisionHistoryRepository,
|
||||
)
|
||||
from src.market_data.access.postgres_quote_history_repository import (
|
||||
PostgresQuoteHistoryRepository,
|
||||
)
|
||||
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.exceptions import ReplayPlanLimitExceededError
|
||||
from src.market_data.replay.models import (
|
||||
ReplayDataType,
|
||||
ReplayPlanRequest,
|
||||
)
|
||||
from src.market_data.replay.postgres_replay_plan_builder import (
|
||||
PostgresReplayPlanBuilder,
|
||||
)
|
||||
from src.market_data.storage.postgres_candle_repository import (
|
||||
PostgresCandleRepository,
|
||||
)
|
||||
from src.market_data.storage.postgres_quote_repository import (
|
||||
PostgresQuoteRepository,
|
||||
)
|
||||
from src.market_data.storage.postgres_trade_repository import (
|
||||
PostgresTradeRepository,
|
||||
)
|
||||
from src.storage.postgres_pool import PostgresConnectionPool
|
||||
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(minutes=10)
|
||||
|
||||
|
||||
def _quote(
|
||||
*,
|
||||
received_at: datetime,
|
||||
source: str = "dzengi_websocket_quote",
|
||||
) -> Quote:
|
||||
return Quote(
|
||||
symbol=SYMBOL,
|
||||
last_price=Decimal("65000.25"),
|
||||
bid_price=Decimal("65000.00"),
|
||||
ask_price=Decimal("65000.50"),
|
||||
exchange_timestamp=received_at - timedelta(milliseconds=1),
|
||||
received_at=received_at,
|
||||
source=source,
|
||||
)
|
||||
|
||||
|
||||
def _candle(
|
||||
*,
|
||||
open_time: datetime,
|
||||
close_price: Decimal = Decimal("105"),
|
||||
interval: str = "1m",
|
||||
) -> Candle:
|
||||
return Candle(
|
||||
symbol=SYMBOL,
|
||||
interval=interval,
|
||||
open_time=open_time,
|
||||
open_price=Decimal("100"),
|
||||
high_price=Decimal("110"),
|
||||
low_price=Decimal("90"),
|
||||
close_price=close_price,
|
||||
volume=Decimal("10"),
|
||||
source="rest_klines:bid",
|
||||
)
|
||||
|
||||
|
||||
def _trade(*, executed_at: datetime) -> Trade:
|
||||
return Trade(
|
||||
symbol=SYMBOL,
|
||||
trade_id=100,
|
||||
price=Decimal("65000.25"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=executed_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source="dzengi_websocket_trade",
|
||||
)
|
||||
|
||||
|
||||
def test_real_quote_and_candle_history_use_public_axes_and_keysets(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
) -> None:
|
||||
quote_writer = PostgresQuoteRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
candle_writer = PostgresCandleRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
quote_reader = PostgresQuoteHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
candle_reader = PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
|
||||
quote_writer.store_quote(
|
||||
venue=VENUE,
|
||||
quote=_quote(received_at=START - timedelta(microseconds=1)),
|
||||
)
|
||||
first_quote = _quote(received_at=START)
|
||||
second_quote = _quote(received_at=START + timedelta(minutes=1))
|
||||
quote_writer.store_quote(venue=VENUE, quote=first_quote)
|
||||
quote_writer.store_quote(
|
||||
venue=VENUE,
|
||||
quote=replace(first_quote, source="dzengi"),
|
||||
)
|
||||
quote_writer.store_quote(venue=VENUE, quote=second_quote)
|
||||
quote_writer.store_quote(
|
||||
venue=VENUE,
|
||||
quote=_quote(received_at=END),
|
||||
)
|
||||
|
||||
first_quote_page = quote_reader.query_quotes(
|
||||
QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=1,
|
||||
)
|
||||
)
|
||||
assert first_quote_page.next_cursor is not None
|
||||
second_quote_page = quote_reader.query_quotes(
|
||||
QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=2,
|
||||
cursor=first_quote_page.next_cursor,
|
||||
)
|
||||
)
|
||||
quote_records = first_quote_page.items + second_quote_page.items
|
||||
|
||||
assert [record.quote for record in quote_records] == [
|
||||
first_quote,
|
||||
second_quote,
|
||||
]
|
||||
assert quote_records[0].observation_sources == (
|
||||
"dzengi_websocket_quote",
|
||||
"dzengi",
|
||||
)
|
||||
assert second_quote_page.next_cursor is None
|
||||
|
||||
candle_writer.store_candle_revision(
|
||||
venue=VENUE,
|
||||
candle=_candle(open_time=START),
|
||||
observed_at=START + timedelta(seconds=10),
|
||||
is_final=False,
|
||||
)
|
||||
candle_writer.store_candle_revision(
|
||||
venue=VENUE,
|
||||
candle=_candle(
|
||||
open_time=START,
|
||||
close_price=Decimal("106"),
|
||||
),
|
||||
observed_at=START + timedelta(seconds=20),
|
||||
is_final=True,
|
||||
)
|
||||
candle_writer.store_candle_revision(
|
||||
venue=VENUE,
|
||||
candle=_candle(open_time=END),
|
||||
observed_at=END + timedelta(seconds=1),
|
||||
is_final=True,
|
||||
)
|
||||
|
||||
first_candle_page = candle_reader.query_candle_revisions(
|
||||
CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="1m",
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=1,
|
||||
)
|
||||
)
|
||||
assert first_candle_page.next_cursor is not None
|
||||
second_candle_page = candle_reader.query_candle_revisions(
|
||||
CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="1m",
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=2,
|
||||
cursor=first_candle_page.next_cursor,
|
||||
)
|
||||
)
|
||||
candle_records = first_candle_page.items + second_candle_page.items
|
||||
|
||||
assert [record.observed_at for record in candle_records] == [
|
||||
START + timedelta(seconds=10),
|
||||
START + timedelta(seconds=20),
|
||||
]
|
||||
assert [record.is_final for record in candle_records] == [False, True]
|
||||
assert second_candle_page.next_cursor is None
|
||||
|
||||
|
||||
def test_real_replay_snapshot_uses_observed_candle_axis_and_global_order(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
) -> None:
|
||||
trade_writer = PostgresTradeRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
quote_writer = PostgresQuoteRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
candle_writer = PostgresCandleRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
candle_reader = PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
builder = PostgresReplayPlanBuilder(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
shared_time = START + timedelta(minutes=1)
|
||||
included_candle = _candle(open_time=START - timedelta(minutes=1))
|
||||
excluded_candle = _candle(
|
||||
open_time=START + timedelta(minutes=3),
|
||||
close_price=Decimal("107"),
|
||||
)
|
||||
|
||||
trade_writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=_trade(executed_at=shared_time),
|
||||
observed_at=shared_time,
|
||||
)
|
||||
quote_writer.store_quote(
|
||||
venue=VENUE,
|
||||
quote=_quote(received_at=shared_time),
|
||||
)
|
||||
candle_writer.store_candle_revision(
|
||||
venue=VENUE,
|
||||
candle=included_candle,
|
||||
observed_at=START + timedelta(minutes=2),
|
||||
is_final=True,
|
||||
)
|
||||
candle_writer.store_candle_revision(
|
||||
venue=VENUE,
|
||||
candle=excluded_candle,
|
||||
observed_at=END + timedelta(minutes=1),
|
||||
is_final=True,
|
||||
)
|
||||
|
||||
request = ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(
|
||||
ReplayDataType.TRADE,
|
||||
ReplayDataType.QUOTE,
|
||||
ReplayDataType.CANDLE_REVISION,
|
||||
),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
candle_intervals=("1m",),
|
||||
max_records=10,
|
||||
)
|
||||
plan = builder.create_plan(request)
|
||||
|
||||
assert [event.data_type for event in plan.events] == [
|
||||
ReplayDataType.TRADE,
|
||||
ReplayDataType.QUOTE,
|
||||
ReplayDataType.CANDLE_REVISION,
|
||||
]
|
||||
assert [event.replay_at for event in plan.events] == [
|
||||
shared_time,
|
||||
shared_time,
|
||||
START + timedelta(minutes=2),
|
||||
]
|
||||
assert plan.events[2].payload == included_candle
|
||||
assert plan.events[2].candle_is_final is True
|
||||
assert [event.order_key for event in plan.events] == sorted(
|
||||
event.order_key for event in plan.events
|
||||
)
|
||||
assert (
|
||||
plan.events[0].replay_sequence
|
||||
< plan.events[1].replay_sequence
|
||||
)
|
||||
assert len({event.replay_sequence for event in plan.events}) == 3
|
||||
|
||||
public_candles = candle_reader.query_candle_revisions(
|
||||
CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="1m",
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=10,
|
||||
)
|
||||
)
|
||||
assert tuple(
|
||||
record.candle for record in public_candles.items
|
||||
) == (excluded_candle,)
|
||||
|
||||
with pytest.raises(ReplayPlanLimitExceededError):
|
||||
builder.create_plan(replace(request, max_records=2))
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,340 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import replace
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import (
|
||||
HistoricalTimeRange,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
PostgresTradeHistoryRepository,
|
||||
TradeHistoryCursor,
|
||||
TradeHistoryQuery,
|
||||
)
|
||||
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,
|
||||
)
|
||||
from src.market_data.storage import (
|
||||
MarketDataStorageConflictError,
|
||||
MarketDataWriteStatus,
|
||||
PostgresTradeRepository,
|
||||
)
|
||||
from src.storage.exceptions import PostgresConnectionPoolError
|
||||
from src.storage.postgres_pool import PostgresConnectionPool
|
||||
from tests.support.postgres_market_data import PostgresTestSettings
|
||||
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
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"
|
||||
|
||||
|
||||
def _trade(
|
||||
*,
|
||||
trade_id: int,
|
||||
executed_at: datetime = START,
|
||||
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 _query(
|
||||
*,
|
||||
limit: int,
|
||||
cursor: TradeHistoryCursor | None = None,
|
||||
) -> TradeHistoryQuery:
|
||||
return TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_real_trade_history_uses_half_open_rollover_keyset_order(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
) -> None:
|
||||
writer = PostgresTradeRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
reader = PostgresTradeHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
rollover_ids = (
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
-1,
|
||||
0,
|
||||
)
|
||||
|
||||
writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=_trade(
|
||||
trade_id=10,
|
||||
executed_at=START - timedelta(microseconds=1),
|
||||
),
|
||||
observed_at=START,
|
||||
)
|
||||
writer.store_trades(
|
||||
venue=VENUE,
|
||||
trades=tuple(
|
||||
_trade(trade_id=trade_id)
|
||||
for trade_id in rollover_ids
|
||||
),
|
||||
observed_at=START + timedelta(seconds=1),
|
||||
)
|
||||
writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=_trade(trade_id=11, executed_at=END),
|
||||
observed_at=END + timedelta(seconds=1),
|
||||
)
|
||||
|
||||
first_page = reader.query_trades(_query(limit=2))
|
||||
|
||||
assert [record.trade.trade_id for record in first_page.items] == [
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
]
|
||||
assert first_page.next_cursor is not None
|
||||
assert (
|
||||
first_page.next_cursor.executed_at,
|
||||
first_page.next_cursor.replay_sequence,
|
||||
) == first_page.items[-1].order_key
|
||||
|
||||
second_page = reader.query_trades(
|
||||
_query(
|
||||
limit=3,
|
||||
cursor=first_page.next_cursor,
|
||||
)
|
||||
)
|
||||
records = first_page.items + second_page.items
|
||||
|
||||
assert [record.trade.trade_id for record in records] == list(
|
||||
rollover_ids
|
||||
)
|
||||
assert [record.event_time for record in records] == [START] * 4
|
||||
assert [record.replay_sequence for record in records] == sorted(
|
||||
record.replay_sequence for record in records
|
||||
)
|
||||
assert len({record.replay_sequence for record in records}) == 4
|
||||
assert second_page.next_cursor is None
|
||||
|
||||
|
||||
def test_real_trade_history_preserves_ordinal_and_returns_provenance(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
) -> None:
|
||||
writer = PostgresTradeRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
reader = PostgresTradeHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
original = _trade(trade_id=100)
|
||||
first_observed_at = START + timedelta(seconds=1)
|
||||
last_observed_at = START + timedelta(seconds=5)
|
||||
|
||||
inserted = writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=original,
|
||||
observed_at=first_observed_at,
|
||||
)
|
||||
original_record = reader.query_trades(_query(limit=10)).items[0]
|
||||
duplicate = writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=original,
|
||||
observed_at=first_observed_at,
|
||||
)
|
||||
provenance = writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=replace(original, source="dzengi"),
|
||||
observed_at=last_observed_at,
|
||||
)
|
||||
updated_record = reader.query_trades(_query(limit=10)).items[0]
|
||||
|
||||
assert inserted.status is MarketDataWriteStatus.INSERTED
|
||||
assert duplicate.status is MarketDataWriteStatus.DUPLICATE
|
||||
assert provenance.status is MarketDataWriteStatus.PROVENANCE_UPDATED
|
||||
assert updated_record.replay_sequence == original_record.replay_sequence
|
||||
assert updated_record.trade == original
|
||||
assert updated_record.first_observed_at == first_observed_at
|
||||
assert updated_record.last_observed_at == last_observed_at
|
||||
assert updated_record.observation_sources == (
|
||||
SOURCE,
|
||||
"dzengi",
|
||||
)
|
||||
|
||||
|
||||
def test_real_trade_history_omits_rolled_back_batch_and_keeps_gap(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
) -> None:
|
||||
writer = PostgresTradeRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
reader = PostgresTradeHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
existing = _trade(trade_id=2)
|
||||
writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=existing,
|
||||
observed_at=START + timedelta(seconds=1),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataStorageConflictError):
|
||||
writer.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
_trade(trade_id=1),
|
||||
replace(existing, price=Decimal("999")),
|
||||
),
|
||||
observed_at=START + timedelta(seconds=2),
|
||||
)
|
||||
|
||||
writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=_trade(trade_id=3),
|
||||
observed_at=START + timedelta(seconds=3),
|
||||
)
|
||||
records = reader.query_trades(_query(limit=10)).items
|
||||
|
||||
assert [record.trade.trade_id for record in records] == [2, 3]
|
||||
assert [record.replay_sequence for record in records] == [1, 4]
|
||||
|
||||
|
||||
def test_two_real_batches_keep_local_order_and_global_unique_sequences(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
) -> None:
|
||||
writer = PostgresTradeRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
reader = PostgresTradeHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
batches = (
|
||||
(
|
||||
_trade(trade_id=SIGNED_TRADE_ID_MAX),
|
||||
_trade(trade_id=SIGNED_TRADE_ID_MIN),
|
||||
),
|
||||
(
|
||||
_trade(trade_id=-1),
|
||||
_trade(trade_id=0),
|
||||
),
|
||||
)
|
||||
start_barrier = threading.Barrier(2)
|
||||
caller_ids: set[int] = set()
|
||||
caller_ids_lock = threading.Lock()
|
||||
|
||||
def store_batch(trades: tuple[Trade, ...]) -> int:
|
||||
with caller_ids_lock:
|
||||
caller_ids.add(threading.get_ident())
|
||||
start_barrier.wait(timeout=5.0)
|
||||
return writer.store_trades(
|
||||
venue=VENUE,
|
||||
trades=trades,
|
||||
observed_at=START + timedelta(seconds=1),
|
||||
).inserted_count
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
inserted_counts = tuple(executor.map(store_batch, batches))
|
||||
|
||||
records = reader.query_trades(_query(limit=10)).items
|
||||
sequence_by_trade_id = {
|
||||
record.trade.trade_id: record.replay_sequence
|
||||
for record in records
|
||||
}
|
||||
|
||||
assert inserted_counts == (2, 2)
|
||||
assert len(caller_ids) == 2
|
||||
assert len(records) == 4
|
||||
assert len(set(sequence_by_trade_id.values())) == 4
|
||||
assert (
|
||||
sequence_by_trade_id[SIGNED_TRADE_ID_MAX]
|
||||
< sequence_by_trade_id[SIGNED_TRADE_ID_MIN]
|
||||
)
|
||||
assert sequence_by_trade_id[-1] < sequence_by_trade_id[0]
|
||||
|
||||
|
||||
def test_real_trade_history_rejects_corrupted_schema_version(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
) -> None:
|
||||
writer = PostgresTradeRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
reader = PostgresTradeHistoryRepository(
|
||||
connection_provider=migrated_postgres_pool.connection,
|
||||
)
|
||||
trade = _trade(trade_id=200)
|
||||
writer.store_trade(
|
||||
venue=VENUE,
|
||||
trade=trade,
|
||||
observed_at=START + timedelta(seconds=1),
|
||||
)
|
||||
|
||||
with migrated_postgres_pool.connection() as connection:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE market_data.trades
|
||||
SET canonical_schema_version = 2
|
||||
WHERE venue = %s
|
||||
AND symbol = %s
|
||||
AND trade_id = %s
|
||||
AND executed_at = %s
|
||||
""",
|
||||
(
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
trade.trade_id,
|
||||
trade.executed_at,
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError):
|
||||
reader.query_trades(_query(limit=10))
|
||||
|
||||
|
||||
def test_real_trade_history_wraps_closed_pool_provider_failure(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
postgres_test_settings: PostgresTestSettings,
|
||||
) -> None:
|
||||
assert migrated_postgres_pool.is_open
|
||||
closed_pool = PostgresConnectionPool(
|
||||
conninfo=postgres_test_settings.dsn,
|
||||
min_size=1,
|
||||
max_size=1,
|
||||
timeout_seconds=5.0,
|
||||
name="trade-history-closed-provider",
|
||||
)
|
||||
closed_pool.open()
|
||||
closed_pool.close()
|
||||
reader = PostgresTradeHistoryRepository(
|
||||
connection_provider=closed_pool.connection,
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
reader.query_trades(_query(limit=10))
|
||||
|
||||
assert isinstance(error_info.value.__cause__, PostgresConnectionPoolError)
|
||||
@@ -0,0 +1,413 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
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,
|
||||
PostgresTradeHistoryRepository,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryRecord,
|
||||
)
|
||||
from src.market_data.acquisition.models.trade import Trade
|
||||
from src.market_data.acquisition.trade_id_sequence import (
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
)
|
||||
from src.market_data.replay import (
|
||||
MarketDataClockProtocol,
|
||||
PostgresReplayPlanBuilder,
|
||||
ReplayConsumerProtocol,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
ReplaySession,
|
||||
ReplaySessionFactory,
|
||||
ReplaySessionState,
|
||||
)
|
||||
from src.market_data.storage import (
|
||||
PostgresTradeRepository,
|
||||
TradeStorageObservationSink,
|
||||
)
|
||||
from src.storage.postgres_pool import PostgresConnectionPool
|
||||
from tests.integration.market_data.acquisition.runtime.loopback_trade_exchange import (
|
||||
LoopbackTradeEnvironment,
|
||||
LoopbackTradeRestServer,
|
||||
LoopbackTradeWebSocketServer,
|
||||
wait_until,
|
||||
)
|
||||
from tests.support.postgres_market_data import (
|
||||
PostgresTestSettings,
|
||||
connect_postgres_test_database,
|
||||
count_other_test_connections,
|
||||
)
|
||||
from tests.support.trade_stream_runtime import (
|
||||
SYMBOL,
|
||||
assert_no_owned_tasks,
|
||||
build_runtime,
|
||||
run_scenario,
|
||||
start_runtime,
|
||||
state_store_from,
|
||||
stop_runtime,
|
||||
)
|
||||
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
VENUE = "dzengi"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DatabaseSnapshot:
|
||||
trades: tuple[tuple[Any, ...], ...]
|
||||
checkpoints: tuple[tuple[Any, ...], ...]
|
||||
replay_sequence: tuple[Any, ...]
|
||||
|
||||
|
||||
class RecordingReplayConsumer:
|
||||
def __init__(self, clock: MarketDataClockProtocol) -> None:
|
||||
self.clock = clock
|
||||
self.events: list[ReplayEvent] = []
|
||||
self.observed_times: list[datetime] = []
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.events.append(event)
|
||||
self.observed_times.append(self.clock.now)
|
||||
|
||||
|
||||
class RecordingReplayConsumerFactory:
|
||||
def __init__(self) -> None:
|
||||
self.plans: list[ReplayPlan] = []
|
||||
self.clocks: list[MarketDataClockProtocol] = []
|
||||
self.consumers: list[RecordingReplayConsumer] = []
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
consumer = RecordingReplayConsumer(clock)
|
||||
self.plans.append(plan)
|
||||
self.clocks.append(clock)
|
||||
self.consumers.append(consumer)
|
||||
return consumer
|
||||
|
||||
|
||||
def _trade_sink(
|
||||
pool: PostgresConnectionPool,
|
||||
) -> TradeStorageObservationSink:
|
||||
return TradeStorageObservationSink(
|
||||
trade_storage=PostgresTradeRepository(
|
||||
connection_provider=pool.connection,
|
||||
),
|
||||
venue=VENUE,
|
||||
)
|
||||
|
||||
|
||||
def _database_snapshot(
|
||||
pool: PostgresConnectionPool,
|
||||
) -> DatabaseSnapshot:
|
||||
with pool.connection() as connection:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version
|
||||
FROM market_data.trades
|
||||
ORDER BY replay_sequence
|
||||
"""
|
||||
)
|
||||
trades = tuple(cursor.fetchall())
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
revision,
|
||||
checkpoint_schema_version
|
||||
FROM market_data.trade_stream_checkpoints
|
||||
ORDER BY venue, symbol
|
||||
"""
|
||||
)
|
||||
checkpoints = tuple(cursor.fetchall())
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT last_value, is_called
|
||||
FROM market_data.replay_sequence
|
||||
"""
|
||||
)
|
||||
replay_sequence = cursor.fetchone()
|
||||
|
||||
assert replay_sequence is not None
|
||||
|
||||
return DatabaseSnapshot(
|
||||
trades=trades,
|
||||
checkpoints=checkpoints,
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
|
||||
def _wait_for_no_other_postgres_connections(
|
||||
settings: PostgresTestSettings,
|
||||
*,
|
||||
timeout_seconds: float = 2.0,
|
||||
) -> None:
|
||||
deadline = time.monotonic() + timeout_seconds
|
||||
|
||||
with connect_postgres_test_database(settings) as control:
|
||||
while True:
|
||||
observed_count = count_other_test_connections(control)
|
||||
|
||||
if observed_count == 0:
|
||||
return
|
||||
|
||||
if time.monotonic() >= deadline:
|
||||
raise AssertionError(
|
||||
"PostgreSQL test connections were not released; "
|
||||
f"observed {observed_count}."
|
||||
)
|
||||
|
||||
time.sleep(0.01)
|
||||
|
||||
|
||||
def _prepare_history_and_replay(
|
||||
*,
|
||||
pool: PostgresConnectionPool,
|
||||
request: ReplayPlanRequest,
|
||||
consumer_factory: RecordingReplayConsumerFactory,
|
||||
) -> tuple[
|
||||
tuple[TradeHistoryRecord, ...],
|
||||
ReplaySession,
|
||||
DatabaseSnapshot,
|
||||
]:
|
||||
history = PostgresTradeHistoryRepository(
|
||||
connection_provider=pool.connection,
|
||||
).query_trades(
|
||||
TradeHistoryQuery(
|
||||
venue=request.venue,
|
||||
symbol=request.symbols[0],
|
||||
time_range=request.time_range,
|
||||
limit=request.max_records,
|
||||
)
|
||||
)
|
||||
session = ReplaySessionFactory(
|
||||
plan_builder=PostgresReplayPlanBuilder(
|
||||
connection_provider=pool.connection,
|
||||
),
|
||||
consumer_factory=consumer_factory,
|
||||
).prepare_session(request)
|
||||
return history.items, session, _database_snapshot(pool)
|
||||
|
||||
|
||||
def test_loopback_runtime_storage_history_and_replay_are_one_read_only_path(
|
||||
migrated_postgres_pool: PostgresConnectionPool,
|
||||
postgres_test_settings: PostgresTestSettings,
|
||||
) -> None:
|
||||
async def scenario() -> None:
|
||||
websocket = LoopbackTradeWebSocketServer()
|
||||
rest = LoopbackTradeRestServer()
|
||||
start_timestamp_ms = time.time_ns() // 1_000_000
|
||||
end_timestamp_ms = start_timestamp_ms + 1
|
||||
|
||||
async with LoopbackTradeEnvironment(
|
||||
websocket=websocket,
|
||||
rest=rest,
|
||||
) as environment:
|
||||
runtime = build_runtime(
|
||||
websocket_url=environment.websocket_url,
|
||||
rest_base_url=rest.base_url,
|
||||
trade_observation_sink=_trade_sink(
|
||||
migrated_postgres_pool
|
||||
),
|
||||
)
|
||||
runtime_task: asyncio.Task[None] | None = None
|
||||
|
||||
try:
|
||||
runtime_task = await start_runtime(runtime)
|
||||
await websocket.wait_for_subscriptions(1)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=SIGNED_TRADE_ID_MAX - 1,
|
||||
timestamp_ms=start_timestamp_ms - 1,
|
||||
price="64555.54",
|
||||
quantity="0.001",
|
||||
)
|
||||
state_store = state_store_from(runtime)
|
||||
await wait_until(
|
||||
lambda: (
|
||||
state_store.contains(SYMBOL)
|
||||
and state_store.get(SYMBOL).last_trade_id
|
||||
== SIGNED_TRADE_ID_MAX - 1
|
||||
)
|
||||
)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=SIGNED_TRADE_ID_MAX,
|
||||
timestamp_ms=start_timestamp_ms,
|
||||
price="64555.55",
|
||||
quantity="0.002",
|
||||
)
|
||||
await wait_until(
|
||||
lambda: state_store.get(SYMBOL).last_trade_id
|
||||
== SIGNED_TRADE_ID_MAX
|
||||
)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=SIGNED_TRADE_ID_MIN,
|
||||
timestamp_ms=start_timestamp_ms,
|
||||
price="64555.56",
|
||||
quantity="0.003",
|
||||
)
|
||||
await wait_until(
|
||||
lambda: state_store.get(SYMBOL).last_trade_id
|
||||
== SIGNED_TRADE_ID_MIN
|
||||
)
|
||||
await websocket.send_trade(
|
||||
0,
|
||||
symbol=SYMBOL,
|
||||
trade_id=SIGNED_TRADE_ID_MIN + 1,
|
||||
timestamp_ms=end_timestamp_ms,
|
||||
price="64555.57",
|
||||
quantity="0.004",
|
||||
)
|
||||
await wait_until(
|
||||
lambda: state_store.get(SYMBOL).last_trade_id
|
||||
== SIGNED_TRADE_ID_MIN + 1
|
||||
)
|
||||
finally:
|
||||
if runtime_task is not None:
|
||||
await stop_runtime(runtime, runtime_task)
|
||||
|
||||
assert websocket.active_handler_count == 0
|
||||
assert rest.thread_is_alive is False
|
||||
|
||||
request = ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=datetime.fromtimestamp(
|
||||
start_timestamp_ms / 1_000,
|
||||
tz=timezone.utc,
|
||||
),
|
||||
end_time=datetime.fromtimestamp(
|
||||
end_timestamp_ms / 1_000,
|
||||
tz=timezone.utc,
|
||||
),
|
||||
),
|
||||
max_records=10,
|
||||
)
|
||||
consumer_factory = RecordingReplayConsumerFactory()
|
||||
history, session, before_replay = await asyncio.to_thread(
|
||||
_prepare_history_and_replay,
|
||||
pool=migrated_postgres_pool,
|
||||
request=request,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert [record.trade.trade_id for record in history] == [
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
]
|
||||
assert [record.trade.price for record in history] == [
|
||||
Decimal("64555.55"),
|
||||
Decimal("64555.56"),
|
||||
]
|
||||
assert [record.trade.quantity for record in history] == [
|
||||
Decimal("0.002"),
|
||||
Decimal("0.003"),
|
||||
]
|
||||
replayed_trades: list[Trade] = []
|
||||
for event in session.plan.events:
|
||||
payload = event.payload
|
||||
|
||||
if not isinstance(payload, Trade):
|
||||
raise AssertionError(
|
||||
"Replay event must contain Trade payload."
|
||||
)
|
||||
|
||||
replayed_trades.append(payload)
|
||||
assert replayed_trades == [record.trade for record in history]
|
||||
assert [
|
||||
(event.payload, event.replay_sequence)
|
||||
for event in session.plan.events
|
||||
] == [
|
||||
(record.trade, record.replay_sequence)
|
||||
for record in history
|
||||
]
|
||||
assert consumer_factory.plans == [session.plan]
|
||||
assert consumer_factory.plans[0] is session.plan
|
||||
assert consumer_factory.clocks == [session.clock]
|
||||
assert before_replay.checkpoints == (
|
||||
(
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
SIGNED_TRADE_ID_MIN + 1,
|
||||
request.time_range.end_time,
|
||||
4,
|
||||
1,
|
||||
),
|
||||
)
|
||||
|
||||
migrated_postgres_pool.close()
|
||||
|
||||
try:
|
||||
_wait_for_no_other_postgres_connections(
|
||||
postgres_test_settings
|
||||
)
|
||||
|
||||
await session.run()
|
||||
finally:
|
||||
migrated_postgres_pool.open()
|
||||
|
||||
consumer = consumer_factory.consumers[0]
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert consumer.clock is session.clock
|
||||
assert tuple(consumer.events) == session.plan.events
|
||||
assert all(
|
||||
actual is expected
|
||||
for actual, expected in zip(
|
||||
consumer.events,
|
||||
session.plan.events,
|
||||
strict=True,
|
||||
)
|
||||
)
|
||||
assert consumer.observed_times == [
|
||||
event.replay_at for event in session.plan.events
|
||||
]
|
||||
assert session.clock.now == session.plan.events[-1].replay_at
|
||||
|
||||
after_replay = await asyncio.to_thread(
|
||||
_database_snapshot,
|
||||
migrated_postgres_pool,
|
||||
)
|
||||
assert after_replay == before_replay
|
||||
await assert_no_owned_tasks()
|
||||
|
||||
run_scenario(scenario())
|
||||
Reference in New Issue
Block a user