Build 060.29: implement Market Data Access and Replay
This commit is contained in:
4
app/requirements-dev.txt
Normal file
4
app/requirements-dev.txt
Normal file
@@ -0,0 +1,4 @@
|
||||
-r requirements.txt
|
||||
|
||||
pytest==9.1.1
|
||||
pyright[nodejs]==1.1.411
|
||||
86
app/src/market_data/access/__init__.py
Normal file
86
app/src/market_data/access/__init__.py
Normal file
@@ -0,0 +1,86 @@
|
||||
from src.market_data.access.contracts import (
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
MarketDataHistoricalAccessProtocol,
|
||||
QuoteHistoryReaderProtocol,
|
||||
TradeHistoryReaderProtocol,
|
||||
)
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessError,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
MarketDataCursorError,
|
||||
)
|
||||
from src.market_data.access.market_data_historical_access import (
|
||||
MarketDataHistoricalAccess,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
HISTORY_CURSOR_VERSION,
|
||||
HISTORY_PAGE_LIMIT_DEFAULT,
|
||||
HISTORY_PAGE_LIMIT_MAX,
|
||||
CandleRevisionHistoryCursor,
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
CandleRevisionHistoryRecord,
|
||||
HistoricalTimeRange,
|
||||
HistoryPage,
|
||||
HistoryQuery,
|
||||
HistoryRecord,
|
||||
QuoteHistoryCursor,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
QuoteHistoryRecord,
|
||||
TradeHistoryCursor,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryRecord,
|
||||
)
|
||||
from src.market_data.access.postgres_candle_revision_history_repository import (
|
||||
PostgresCandleRevisionHistoryRepository,
|
||||
)
|
||||
from src.market_data.access.postgres_history_support import (
|
||||
PostgresHistoryConnectionProvider,
|
||||
)
|
||||
from src.market_data.access.postgres_quote_history_repository import (
|
||||
PostgresQuoteHistoryRepository,
|
||||
)
|
||||
from src.market_data.access.postgres_trade_history_repository import (
|
||||
PostgresTradeHistoryRepository,
|
||||
)
|
||||
|
||||
|
||||
__all__ = (
|
||||
"HISTORY_CURSOR_VERSION",
|
||||
"HISTORY_PAGE_LIMIT_DEFAULT",
|
||||
"HISTORY_PAGE_LIMIT_MAX",
|
||||
"CandleRevisionHistoryCursor",
|
||||
"CandleRevisionHistoryPage",
|
||||
"CandleRevisionHistoryQuery",
|
||||
"CandleRevisionHistoryReaderProtocol",
|
||||
"CandleRevisionHistoryRecord",
|
||||
"HistoricalTimeRange",
|
||||
"HistoryPage",
|
||||
"HistoryQuery",
|
||||
"HistoryRecord",
|
||||
"MarketDataAccessError",
|
||||
"MarketDataAccessIntegrityError",
|
||||
"MarketDataAccessOperationError",
|
||||
"MarketDataAccessValidationError",
|
||||
"MarketDataCursorError",
|
||||
"MarketDataHistoricalAccess",
|
||||
"MarketDataHistoricalAccessProtocol",
|
||||
"PostgresCandleRevisionHistoryRepository",
|
||||
"PostgresHistoryConnectionProvider",
|
||||
"PostgresQuoteHistoryRepository",
|
||||
"PostgresTradeHistoryRepository",
|
||||
"QuoteHistoryCursor",
|
||||
"QuoteHistoryPage",
|
||||
"QuoteHistoryQuery",
|
||||
"QuoteHistoryReaderProtocol",
|
||||
"QuoteHistoryRecord",
|
||||
"TradeHistoryCursor",
|
||||
"TradeHistoryPage",
|
||||
"TradeHistoryQuery",
|
||||
"TradeHistoryReaderProtocol",
|
||||
"TradeHistoryRecord",
|
||||
)
|
||||
55
app/src/market_data/access/contracts.py
Normal file
55
app/src/market_data/access/contracts.py
Normal file
@@ -0,0 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class TradeHistoryReaderProtocol(Protocol):
|
||||
"""Граница чтения одной страницы Canonical Trade History."""
|
||||
|
||||
def query_trades(
|
||||
self,
|
||||
query: TradeHistoryQuery,
|
||||
) -> TradeHistoryPage:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class QuoteHistoryReaderProtocol(Protocol):
|
||||
"""Граница чтения одной страницы Canonical Quote History."""
|
||||
|
||||
def query_quotes(
|
||||
self,
|
||||
query: QuoteHistoryQuery,
|
||||
) -> QuoteHistoryPage:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class CandleRevisionHistoryReaderProtocol(Protocol):
|
||||
"""Граница чтения одной страницы Canonical Candle revisions."""
|
||||
|
||||
def query_candle_revisions(
|
||||
self,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class MarketDataHistoricalAccessProtocol(
|
||||
TradeHistoryReaderProtocol,
|
||||
QuoteHistoryReaderProtocol,
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
Protocol,
|
||||
):
|
||||
"""Объединённая DB-neutral граница Historical Access."""
|
||||
21
app/src/market_data/access/exceptions.py
Normal file
21
app/src/market_data/access/exceptions.py
Normal file
@@ -0,0 +1,21 @@
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class MarketDataAccessError(Exception):
|
||||
"""Базовая ошибка доступа к историческим рыночным данным."""
|
||||
|
||||
|
||||
class MarketDataAccessValidationError(MarketDataAccessError):
|
||||
"""Недопустимый запрос или значение на публичной границе чтения."""
|
||||
|
||||
|
||||
class MarketDataCursorError(MarketDataAccessValidationError):
|
||||
"""Курсор повреждён либо принадлежит другому запросу."""
|
||||
|
||||
|
||||
class MarketDataAccessIntegrityError(MarketDataAccessError):
|
||||
"""Сохранённая строка не соответствует Canonical-контракту."""
|
||||
|
||||
|
||||
class MarketDataAccessOperationError(MarketDataAccessError):
|
||||
"""Ошибка backend-операции Historical Access."""
|
||||
76
app/src/market_data/access/market_data_historical_access.py
Normal file
76
app/src/market_data/access/market_data_historical_access.py
Normal file
@@ -0,0 +1,76 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from src.market_data.access.contracts import (
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
QuoteHistoryReaderProtocol,
|
||||
TradeHistoryReaderProtocol,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
)
|
||||
|
||||
|
||||
class MarketDataHistoricalAccess:
|
||||
"""Независимый от базы данных фасад трёх модулей чтения истории."""
|
||||
|
||||
__slots__ = (
|
||||
"_trade_reader",
|
||||
"_quote_reader",
|
||||
"_candle_revision_reader",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
trade_reader: TradeHistoryReaderProtocol,
|
||||
quote_reader: QuoteHistoryReaderProtocol,
|
||||
candle_revision_reader: CandleRevisionHistoryReaderProtocol,
|
||||
) -> None:
|
||||
if not isinstance(trade_reader, TradeHistoryReaderProtocol):
|
||||
raise TypeError(
|
||||
"trade_reader must implement TradeHistoryReaderProtocol"
|
||||
)
|
||||
|
||||
if not isinstance(quote_reader, QuoteHistoryReaderProtocol):
|
||||
raise TypeError(
|
||||
"quote_reader must implement QuoteHistoryReaderProtocol"
|
||||
)
|
||||
|
||||
if not isinstance(
|
||||
candle_revision_reader,
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
):
|
||||
raise TypeError(
|
||||
"candle_revision_reader must implement "
|
||||
"CandleRevisionHistoryReaderProtocol"
|
||||
)
|
||||
|
||||
self._trade_reader = trade_reader
|
||||
self._quote_reader = quote_reader
|
||||
self._candle_revision_reader = candle_revision_reader
|
||||
|
||||
def query_trades(
|
||||
self,
|
||||
query: TradeHistoryQuery,
|
||||
) -> TradeHistoryPage:
|
||||
"""Передать запрос истории сделок соответствующему модулю."""
|
||||
return self._trade_reader.query_trades(query)
|
||||
|
||||
def query_quotes(
|
||||
self,
|
||||
query: QuoteHistoryQuery,
|
||||
) -> QuoteHistoryPage:
|
||||
"""Передать запрос истории котировок соответствующему модулю."""
|
||||
return self._quote_reader.query_quotes(query)
|
||||
|
||||
def query_candle_revisions(
|
||||
self,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
"""Передать запрос истории свечей соответствующему модулю."""
|
||||
return self._candle_revision_reader.query_candle_revisions(query)
|
||||
901
app/src/market_data/access/models.py
Normal file
901
app/src/market_data/access/models.py
Normal file
@@ -0,0 +1,901 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Callable
|
||||
|
||||
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.acquisition.symbols import normalize_symbol
|
||||
|
||||
|
||||
HISTORY_CURSOR_VERSION = 1
|
||||
HISTORY_PAGE_LIMIT_DEFAULT = 500
|
||||
HISTORY_PAGE_LIMIT_MAX = 1_000
|
||||
|
||||
|
||||
def _normalize_non_empty_text(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
uppercase: bool = False,
|
||||
) -> str:
|
||||
if not isinstance(value, str):
|
||||
raise TypeError(f"{field_name} must be a string")
|
||||
|
||||
normalized = value.strip()
|
||||
|
||||
if not normalized:
|
||||
raise ValueError(f"{field_name} must not be empty")
|
||||
|
||||
return normalized.upper() if uppercase else normalized
|
||||
|
||||
|
||||
def _normalize_symbol(value: object, *, field_name: str) -> str:
|
||||
raw_symbol = _normalize_non_empty_text(
|
||||
value,
|
||||
field_name=field_name,
|
||||
)
|
||||
normalized = normalize_symbol(raw_symbol)
|
||||
|
||||
if not normalized:
|
||||
raise ValueError(f"{field_name} must not be empty")
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def _normalize_utc_datetime(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
) -> datetime:
|
||||
if not isinstance(value, datetime):
|
||||
raise TypeError(f"{field_name} must be a datetime")
|
||||
|
||||
if value.tzinfo is None or value.utcoffset() is None:
|
||||
raise ValueError(f"{field_name} must contain timezone")
|
||||
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _normalize_positive_integer(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
maximum: int | None = None,
|
||||
) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"{field_name} must be an integer")
|
||||
|
||||
if value <= 0:
|
||||
raise ValueError(f"{field_name} must be positive")
|
||||
|
||||
if maximum is not None and value > maximum:
|
||||
raise ValueError(
|
||||
f"{field_name} must not exceed {maximum}"
|
||||
)
|
||||
|
||||
return value
|
||||
|
||||
|
||||
def _normalize_observation_sources(
|
||||
value: object,
|
||||
*,
|
||||
primary_source: object,
|
||||
) -> tuple[str, ...]:
|
||||
if not isinstance(value, tuple):
|
||||
raise TypeError("observation_sources must be a tuple")
|
||||
|
||||
if not value:
|
||||
raise ValueError("observation_sources must not be empty")
|
||||
|
||||
normalized = tuple(
|
||||
_normalize_non_empty_text(
|
||||
source,
|
||||
field_name="observation_sources item",
|
||||
)
|
||||
for source in value
|
||||
)
|
||||
|
||||
if len(set(normalized)) != len(normalized):
|
||||
raise ValueError("observation_sources must be unique")
|
||||
|
||||
normalized_primary = _normalize_non_empty_text(
|
||||
primary_source,
|
||||
field_name="payload.source",
|
||||
)
|
||||
|
||||
if normalized[0] != normalized_primary:
|
||||
raise ValueError(
|
||||
"first observation source must match payload.source"
|
||||
)
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def _validate_canonical_symbol_and_source(
|
||||
*,
|
||||
symbol: object,
|
||||
source: object,
|
||||
) -> None:
|
||||
normalized_symbol = _normalize_symbol(
|
||||
symbol,
|
||||
field_name="payload.symbol",
|
||||
)
|
||||
|
||||
if normalized_symbol != symbol:
|
||||
raise ValueError("payload.symbol must be canonical")
|
||||
|
||||
normalized_source = _normalize_non_empty_text(
|
||||
source,
|
||||
field_name="payload.source",
|
||||
)
|
||||
|
||||
if normalized_source != source:
|
||||
raise ValueError("payload.source must be canonical")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class HistoricalTimeRange:
|
||||
"""Полуоткрытый UTC-диапазон исторического запроса ``[start, end)``."""
|
||||
|
||||
start_time: datetime
|
||||
end_time: datetime
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_start = _normalize_utc_datetime(
|
||||
self.start_time,
|
||||
field_name="start_time",
|
||||
)
|
||||
normalized_end = _normalize_utc_datetime(
|
||||
self.end_time,
|
||||
field_name="end_time",
|
||||
)
|
||||
|
||||
if normalized_start >= normalized_end:
|
||||
raise ValueError("start_time must be earlier than end_time")
|
||||
|
||||
object.__setattr__(self, "start_time", normalized_start)
|
||||
object.__setattr__(self, "end_time", normalized_end)
|
||||
|
||||
def contains(self, instant: datetime) -> bool:
|
||||
"""Проверить принадлежность времени полуоткрытому диапазону."""
|
||||
normalized = _normalize_utc_datetime(
|
||||
instant,
|
||||
field_name="instant",
|
||||
)
|
||||
return self.start_time <= normalized < self.end_time
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TradeHistoryRecord:
|
||||
"""Canonical Trade и неизменяемые технические данные его хранения."""
|
||||
|
||||
venue: str
|
||||
trade: Trade
|
||||
first_observed_at: datetime
|
||||
last_observed_at: datetime
|
||||
observation_sources: tuple[str, ...]
|
||||
replay_sequence: int
|
||||
canonical_schema_version: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_venue = _normalize_non_empty_text(
|
||||
self.venue,
|
||||
field_name="venue",
|
||||
)
|
||||
|
||||
if type(self.trade) is not Trade:
|
||||
raise TypeError("trade must be a Canonical Trade")
|
||||
|
||||
_validate_canonical_symbol_and_source(
|
||||
symbol=self.trade.symbol,
|
||||
source=self.trade.source,
|
||||
)
|
||||
_normalize_utc_datetime(
|
||||
self.trade.executed_at,
|
||||
field_name="trade.executed_at",
|
||||
)
|
||||
normalized_first = _normalize_utc_datetime(
|
||||
self.first_observed_at,
|
||||
field_name="first_observed_at",
|
||||
)
|
||||
normalized_last = _normalize_utc_datetime(
|
||||
self.last_observed_at,
|
||||
field_name="last_observed_at",
|
||||
)
|
||||
|
||||
if normalized_last < normalized_first:
|
||||
raise ValueError(
|
||||
"last_observed_at must not precede first_observed_at"
|
||||
)
|
||||
|
||||
normalized_sources = _normalize_observation_sources(
|
||||
self.observation_sources,
|
||||
primary_source=self.trade.source,
|
||||
)
|
||||
normalized_sequence = _normalize_positive_integer(
|
||||
self.replay_sequence,
|
||||
field_name="replay_sequence",
|
||||
)
|
||||
normalized_version = _normalize_positive_integer(
|
||||
self.canonical_schema_version,
|
||||
field_name="canonical_schema_version",
|
||||
)
|
||||
|
||||
object.__setattr__(self, "venue", normalized_venue)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"first_observed_at",
|
||||
normalized_first,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"last_observed_at",
|
||||
normalized_last,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"observation_sources",
|
||||
normalized_sources,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"replay_sequence",
|
||||
normalized_sequence,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"canonical_schema_version",
|
||||
normalized_version,
|
||||
)
|
||||
|
||||
@property
|
||||
def symbol(self) -> str:
|
||||
return self.trade.symbol
|
||||
|
||||
@property
|
||||
def event_time(self) -> datetime:
|
||||
return self.trade.executed_at.astimezone(timezone.utc)
|
||||
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (self.event_time, self.replay_sequence)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class QuoteHistoryRecord:
|
||||
"""Canonical Quote и неизменяемые технические данные его хранения."""
|
||||
|
||||
venue: str
|
||||
quote: Quote
|
||||
observation_sources: tuple[str, ...]
|
||||
replay_sequence: int
|
||||
canonical_schema_version: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_venue = _normalize_non_empty_text(
|
||||
self.venue,
|
||||
field_name="venue",
|
||||
)
|
||||
|
||||
if type(self.quote) is not Quote:
|
||||
raise TypeError("quote must be a Canonical Quote")
|
||||
|
||||
_validate_canonical_symbol_and_source(
|
||||
symbol=self.quote.symbol,
|
||||
source=self.quote.source,
|
||||
)
|
||||
_normalize_utc_datetime(
|
||||
self.quote.received_at,
|
||||
field_name="quote.received_at",
|
||||
)
|
||||
|
||||
if self.quote.exchange_timestamp is not None:
|
||||
_normalize_utc_datetime(
|
||||
self.quote.exchange_timestamp,
|
||||
field_name="quote.exchange_timestamp",
|
||||
)
|
||||
|
||||
normalized_sources = _normalize_observation_sources(
|
||||
self.observation_sources,
|
||||
primary_source=self.quote.source,
|
||||
)
|
||||
normalized_sequence = _normalize_positive_integer(
|
||||
self.replay_sequence,
|
||||
field_name="replay_sequence",
|
||||
)
|
||||
normalized_version = _normalize_positive_integer(
|
||||
self.canonical_schema_version,
|
||||
field_name="canonical_schema_version",
|
||||
)
|
||||
|
||||
object.__setattr__(self, "venue", normalized_venue)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"observation_sources",
|
||||
normalized_sources,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"replay_sequence",
|
||||
normalized_sequence,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"canonical_schema_version",
|
||||
normalized_version,
|
||||
)
|
||||
|
||||
@property
|
||||
def symbol(self) -> str:
|
||||
return self.quote.symbol
|
||||
|
||||
@property
|
||||
def event_time(self) -> datetime:
|
||||
return self.quote.received_at.astimezone(timezone.utc)
|
||||
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (self.event_time, self.replay_sequence)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CandleRevisionHistoryRecord:
|
||||
"""Canonical Candle и метаданные одной наблюдавшейся ревизии."""
|
||||
|
||||
venue: str
|
||||
candle: Candle
|
||||
observed_at: datetime
|
||||
is_final: bool
|
||||
observation_sources: tuple[str, ...]
|
||||
replay_sequence: int
|
||||
canonical_schema_version: int = 1
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_venue = _normalize_non_empty_text(
|
||||
self.venue,
|
||||
field_name="venue",
|
||||
)
|
||||
|
||||
if type(self.candle) is not Candle:
|
||||
raise TypeError("candle must be a Canonical Candle")
|
||||
|
||||
_validate_canonical_symbol_and_source(
|
||||
symbol=self.candle.symbol,
|
||||
source=self.candle.source,
|
||||
)
|
||||
normalized_interval = _normalize_non_empty_text(
|
||||
self.candle.interval,
|
||||
field_name="candle.interval",
|
||||
)
|
||||
|
||||
if normalized_interval != self.candle.interval:
|
||||
raise ValueError("candle.interval must be canonical")
|
||||
normalized_open_time = _normalize_utc_datetime(
|
||||
self.candle.open_time,
|
||||
field_name="candle.open_time",
|
||||
)
|
||||
normalized_observed_at = _normalize_utc_datetime(
|
||||
self.observed_at,
|
||||
field_name="observed_at",
|
||||
)
|
||||
|
||||
if normalized_observed_at < normalized_open_time:
|
||||
raise ValueError("observed_at must not precede candle.open_time")
|
||||
|
||||
if not isinstance(self.is_final, bool):
|
||||
raise TypeError("is_final must be a boolean")
|
||||
|
||||
normalized_sources = _normalize_observation_sources(
|
||||
self.observation_sources,
|
||||
primary_source=self.candle.source,
|
||||
)
|
||||
normalized_sequence = _normalize_positive_integer(
|
||||
self.replay_sequence,
|
||||
field_name="replay_sequence",
|
||||
)
|
||||
normalized_version = _normalize_positive_integer(
|
||||
self.canonical_schema_version,
|
||||
field_name="canonical_schema_version",
|
||||
)
|
||||
|
||||
object.__setattr__(self, "venue", normalized_venue)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"observed_at",
|
||||
normalized_observed_at,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"observation_sources",
|
||||
normalized_sources,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"replay_sequence",
|
||||
normalized_sequence,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"canonical_schema_version",
|
||||
normalized_version,
|
||||
)
|
||||
|
||||
@property
|
||||
def symbol(self) -> str:
|
||||
return self.candle.symbol
|
||||
|
||||
@property
|
||||
def interval(self) -> str:
|
||||
return self.candle.interval
|
||||
|
||||
@property
|
||||
def event_time(self) -> datetime:
|
||||
"""Вернуть ось Historical Query — время открытия свечи."""
|
||||
return self.candle.open_time.astimezone(timezone.utc)
|
||||
|
||||
@property
|
||||
def replay_at(self) -> datetime:
|
||||
"""Вернуть ось Replay — время наблюдения конкретной ревизии."""
|
||||
return self.observed_at
|
||||
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (self.event_time, self.replay_sequence)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TradeHistoryCursor:
|
||||
"""Структурированный keyset-курсор Trade History."""
|
||||
|
||||
venue: str
|
||||
symbol: str
|
||||
time_range: HistoricalTimeRange
|
||||
executed_at: datetime
|
||||
replay_sequence: int
|
||||
version: int = HISTORY_CURSOR_VERSION
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_venue, normalized_symbol = _normalize_cursor_scope(
|
||||
venue=self.venue,
|
||||
symbol=self.symbol,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
normalized_time = _normalize_cursor_position(
|
||||
cursor_time=self.executed_at,
|
||||
replay_sequence=self.replay_sequence,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
_validate_cursor_version(self.version)
|
||||
|
||||
object.__setattr__(self, "venue", normalized_venue)
|
||||
object.__setattr__(self, "symbol", normalized_symbol)
|
||||
object.__setattr__(self, "executed_at", normalized_time)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class QuoteHistoryCursor:
|
||||
"""Структурированный keyset-курсор Quote History."""
|
||||
|
||||
venue: str
|
||||
symbol: str
|
||||
time_range: HistoricalTimeRange
|
||||
received_at: datetime
|
||||
replay_sequence: int
|
||||
version: int = HISTORY_CURSOR_VERSION
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_venue, normalized_symbol = _normalize_cursor_scope(
|
||||
venue=self.venue,
|
||||
symbol=self.symbol,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
normalized_time = _normalize_cursor_position(
|
||||
cursor_time=self.received_at,
|
||||
replay_sequence=self.replay_sequence,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
_validate_cursor_version(self.version)
|
||||
|
||||
object.__setattr__(self, "venue", normalized_venue)
|
||||
object.__setattr__(self, "symbol", normalized_symbol)
|
||||
object.__setattr__(self, "received_at", normalized_time)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CandleRevisionHistoryCursor:
|
||||
"""Структурированный keyset-курсор Candle Revision History."""
|
||||
|
||||
venue: str
|
||||
symbol: str
|
||||
interval: str
|
||||
time_range: HistoricalTimeRange
|
||||
open_time: datetime
|
||||
replay_sequence: int
|
||||
version: int = HISTORY_CURSOR_VERSION
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_venue, normalized_symbol = _normalize_cursor_scope(
|
||||
venue=self.venue,
|
||||
symbol=self.symbol,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
normalized_interval = _normalize_non_empty_text(
|
||||
self.interval,
|
||||
field_name="interval",
|
||||
)
|
||||
normalized_time = _normalize_cursor_position(
|
||||
cursor_time=self.open_time,
|
||||
replay_sequence=self.replay_sequence,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
_validate_cursor_version(self.version)
|
||||
|
||||
object.__setattr__(self, "venue", normalized_venue)
|
||||
object.__setattr__(self, "symbol", normalized_symbol)
|
||||
object.__setattr__(self, "interval", normalized_interval)
|
||||
object.__setattr__(self, "open_time", normalized_time)
|
||||
|
||||
|
||||
def _normalize_cursor_scope(
|
||||
*,
|
||||
venue: object,
|
||||
symbol: object,
|
||||
time_range: object,
|
||||
) -> tuple[str, str]:
|
||||
if type(time_range) is not HistoricalTimeRange:
|
||||
raise TypeError("time_range must be HistoricalTimeRange")
|
||||
|
||||
return (
|
||||
_normalize_non_empty_text(venue, field_name="venue"),
|
||||
_normalize_symbol(symbol, field_name="symbol"),
|
||||
)
|
||||
|
||||
|
||||
def _normalize_cursor_position(
|
||||
*,
|
||||
cursor_time: object,
|
||||
replay_sequence: object,
|
||||
time_range: HistoricalTimeRange,
|
||||
) -> datetime:
|
||||
normalized_time = _normalize_utc_datetime(
|
||||
cursor_time,
|
||||
field_name="cursor time",
|
||||
)
|
||||
_normalize_positive_integer(
|
||||
replay_sequence,
|
||||
field_name="replay_sequence",
|
||||
)
|
||||
|
||||
if not time_range.contains(normalized_time):
|
||||
raise ValueError("cursor time must belong to query range")
|
||||
|
||||
return normalized_time
|
||||
|
||||
|
||||
def _validate_cursor_version(version: object) -> None:
|
||||
if isinstance(version, bool) or not isinstance(version, int):
|
||||
raise TypeError("version must be an integer")
|
||||
|
||||
if version != HISTORY_CURSOR_VERSION:
|
||||
raise ValueError("unsupported history cursor version")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TradeHistoryQuery:
|
||||
"""Один forward-only запрос истории Trades."""
|
||||
|
||||
venue: str
|
||||
symbol: str
|
||||
time_range: HistoricalTimeRange
|
||||
limit: int = HISTORY_PAGE_LIMIT_DEFAULT
|
||||
cursor: TradeHistoryCursor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
venue, symbol = _normalize_query(
|
||||
venue=self.venue,
|
||||
symbol=self.symbol,
|
||||
time_range=self.time_range,
|
||||
limit=self.limit,
|
||||
)
|
||||
_validate_query_cursor(
|
||||
cursor=self.cursor,
|
||||
cursor_type=TradeHistoryCursor,
|
||||
venue=venue,
|
||||
symbol=symbol,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
object.__setattr__(self, "venue", venue)
|
||||
object.__setattr__(self, "symbol", symbol)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class QuoteHistoryQuery:
|
||||
"""Один forward-only запрос истории Quotes."""
|
||||
|
||||
venue: str
|
||||
symbol: str
|
||||
time_range: HistoricalTimeRange
|
||||
limit: int = HISTORY_PAGE_LIMIT_DEFAULT
|
||||
cursor: QuoteHistoryCursor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
venue, symbol = _normalize_query(
|
||||
venue=self.venue,
|
||||
symbol=self.symbol,
|
||||
time_range=self.time_range,
|
||||
limit=self.limit,
|
||||
)
|
||||
_validate_query_cursor(
|
||||
cursor=self.cursor,
|
||||
cursor_type=QuoteHistoryCursor,
|
||||
venue=venue,
|
||||
symbol=symbol,
|
||||
time_range=self.time_range,
|
||||
)
|
||||
object.__setattr__(self, "venue", venue)
|
||||
object.__setattr__(self, "symbol", symbol)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CandleRevisionHistoryQuery:
|
||||
"""Один forward-only запрос истории ревизий Candles."""
|
||||
|
||||
venue: str
|
||||
symbol: str
|
||||
interval: str
|
||||
time_range: HistoricalTimeRange
|
||||
limit: int = HISTORY_PAGE_LIMIT_DEFAULT
|
||||
cursor: CandleRevisionHistoryCursor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
venue, symbol = _normalize_query(
|
||||
venue=self.venue,
|
||||
symbol=self.symbol,
|
||||
time_range=self.time_range,
|
||||
limit=self.limit,
|
||||
)
|
||||
interval = _normalize_non_empty_text(
|
||||
self.interval,
|
||||
field_name="interval",
|
||||
)
|
||||
_validate_query_cursor(
|
||||
cursor=self.cursor,
|
||||
cursor_type=CandleRevisionHistoryCursor,
|
||||
venue=venue,
|
||||
symbol=symbol,
|
||||
time_range=self.time_range,
|
||||
interval=interval,
|
||||
)
|
||||
object.__setattr__(self, "venue", venue)
|
||||
object.__setattr__(self, "symbol", symbol)
|
||||
object.__setattr__(self, "interval", interval)
|
||||
|
||||
|
||||
def _normalize_query(
|
||||
*,
|
||||
venue: object,
|
||||
symbol: object,
|
||||
time_range: object,
|
||||
limit: object,
|
||||
) -> tuple[str, str]:
|
||||
if type(time_range) is not HistoricalTimeRange:
|
||||
raise TypeError("time_range must be HistoricalTimeRange")
|
||||
|
||||
_normalize_positive_integer(
|
||||
limit,
|
||||
field_name="limit",
|
||||
maximum=HISTORY_PAGE_LIMIT_MAX,
|
||||
)
|
||||
return (
|
||||
_normalize_non_empty_text(venue, field_name="venue"),
|
||||
_normalize_symbol(symbol, field_name="symbol"),
|
||||
)
|
||||
|
||||
|
||||
def _validate_query_cursor(
|
||||
*,
|
||||
cursor: object,
|
||||
cursor_type: type[object],
|
||||
venue: str,
|
||||
symbol: str,
|
||||
time_range: HistoricalTimeRange,
|
||||
interval: str | None = None,
|
||||
) -> None:
|
||||
if cursor is None:
|
||||
return
|
||||
|
||||
if type(cursor) is not cursor_type:
|
||||
raise TypeError(f"cursor must be {cursor_type.__name__} or None")
|
||||
|
||||
if (
|
||||
getattr(cursor, "venue") != venue
|
||||
or getattr(cursor, "symbol") != symbol
|
||||
or getattr(cursor, "time_range") != time_range
|
||||
):
|
||||
raise ValueError("cursor does not belong to query scope")
|
||||
|
||||
if interval is not None and getattr(cursor, "interval") != interval:
|
||||
raise ValueError("cursor does not belong to query interval")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TradeHistoryPage:
|
||||
"""Одна типизированная страница Trade History."""
|
||||
|
||||
query: TradeHistoryQuery
|
||||
items: tuple[TradeHistoryRecord, ...]
|
||||
next_cursor: TradeHistoryCursor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
_validate_page(
|
||||
query=self.query,
|
||||
query_type=TradeHistoryQuery,
|
||||
items=self.items,
|
||||
item_type=TradeHistoryRecord,
|
||||
next_cursor=self.next_cursor,
|
||||
cursor_type=TradeHistoryCursor,
|
||||
cursor_time_name="executed_at",
|
||||
)
|
||||
|
||||
@property
|
||||
def has_more(self) -> bool:
|
||||
return self.next_cursor is not None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class QuoteHistoryPage:
|
||||
"""Одна типизированная страница Quote History."""
|
||||
|
||||
query: QuoteHistoryQuery
|
||||
items: tuple[QuoteHistoryRecord, ...]
|
||||
next_cursor: QuoteHistoryCursor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
_validate_page(
|
||||
query=self.query,
|
||||
query_type=QuoteHistoryQuery,
|
||||
items=self.items,
|
||||
item_type=QuoteHistoryRecord,
|
||||
next_cursor=self.next_cursor,
|
||||
cursor_type=QuoteHistoryCursor,
|
||||
cursor_time_name="received_at",
|
||||
)
|
||||
|
||||
@property
|
||||
def has_more(self) -> bool:
|
||||
return self.next_cursor is not None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CandleRevisionHistoryPage:
|
||||
"""Одна типизированная страница Candle Revision History."""
|
||||
|
||||
query: CandleRevisionHistoryQuery
|
||||
items: tuple[CandleRevisionHistoryRecord, ...]
|
||||
next_cursor: CandleRevisionHistoryCursor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
_validate_page(
|
||||
query=self.query,
|
||||
query_type=CandleRevisionHistoryQuery,
|
||||
items=self.items,
|
||||
item_type=CandleRevisionHistoryRecord,
|
||||
next_cursor=self.next_cursor,
|
||||
cursor_type=CandleRevisionHistoryCursor,
|
||||
cursor_time_name="open_time",
|
||||
)
|
||||
|
||||
@property
|
||||
def has_more(self) -> bool:
|
||||
return self.next_cursor is not None
|
||||
|
||||
|
||||
def _validate_page(
|
||||
*,
|
||||
query: object,
|
||||
query_type: type[object],
|
||||
items: object,
|
||||
item_type: type[object],
|
||||
next_cursor: object,
|
||||
cursor_type: type[object],
|
||||
cursor_time_name: str,
|
||||
) -> None:
|
||||
if type(query) is not query_type:
|
||||
raise TypeError(f"query must be {query_type.__name__}")
|
||||
|
||||
if type(items) is not tuple:
|
||||
raise TypeError("items must be a tuple")
|
||||
|
||||
if len(items) > getattr(query, "limit"):
|
||||
raise ValueError("items must not exceed query limit")
|
||||
|
||||
if any(type(item) is not item_type for item in items):
|
||||
raise TypeError(f"items must contain only {item_type.__name__}")
|
||||
|
||||
if next_cursor is not None and type(next_cursor) is not cursor_type:
|
||||
raise TypeError(
|
||||
f"next_cursor must be {cursor_type.__name__} or None"
|
||||
)
|
||||
|
||||
if not items:
|
||||
if next_cursor is not None:
|
||||
raise ValueError("empty page must not have next_cursor")
|
||||
return
|
||||
|
||||
for previous, current in zip(items, items[1:], strict=False):
|
||||
if getattr(previous, "order_key") >= getattr(current, "order_key"):
|
||||
raise ValueError("page items must be strictly ordered")
|
||||
|
||||
for item in items:
|
||||
if (
|
||||
getattr(item, "venue") != getattr(query, "venue")
|
||||
or getattr(item, "symbol") != getattr(query, "symbol")
|
||||
):
|
||||
raise ValueError("page items must belong to one query scope")
|
||||
|
||||
if getattr(item, "interval", None) != getattr(
|
||||
query,
|
||||
"interval",
|
||||
None,
|
||||
):
|
||||
raise ValueError("page items must belong to one interval")
|
||||
|
||||
if not getattr(query, "time_range").contains(
|
||||
getattr(item, "event_time")
|
||||
):
|
||||
raise ValueError("page item time must belong to query range")
|
||||
|
||||
query_cursor = getattr(query, "cursor")
|
||||
|
||||
if query_cursor is not None:
|
||||
incoming_key = (
|
||||
getattr(query_cursor, cursor_time_name),
|
||||
getattr(query_cursor, "replay_sequence"),
|
||||
)
|
||||
|
||||
if getattr(items[0], "order_key") <= incoming_key:
|
||||
raise ValueError("page items must follow query cursor")
|
||||
|
||||
if next_cursor is None:
|
||||
return
|
||||
|
||||
last = items[-1]
|
||||
|
||||
if (
|
||||
getattr(next_cursor, "venue") != getattr(query, "venue")
|
||||
or getattr(next_cursor, "symbol") != getattr(query, "symbol")
|
||||
or getattr(next_cursor, "time_range")
|
||||
!= getattr(query, "time_range")
|
||||
or getattr(next_cursor, "replay_sequence")
|
||||
!= getattr(last, "replay_sequence")
|
||||
or getattr(next_cursor, cursor_time_name)
|
||||
!= getattr(last, "event_time")
|
||||
):
|
||||
raise ValueError("next_cursor must point to the last page item")
|
||||
|
||||
cursor_interval = getattr(next_cursor, "interval", None)
|
||||
|
||||
if cursor_interval != getattr(query, "interval", None):
|
||||
raise ValueError("next_cursor must belong to page interval")
|
||||
|
||||
|
||||
HistoryRecord = (
|
||||
TradeHistoryRecord
|
||||
| QuoteHistoryRecord
|
||||
| CandleRevisionHistoryRecord
|
||||
)
|
||||
HistoryQuery = (
|
||||
TradeHistoryQuery
|
||||
| QuoteHistoryQuery
|
||||
| CandleRevisionHistoryQuery
|
||||
)
|
||||
HistoryPage = (
|
||||
TradeHistoryPage
|
||||
| QuoteHistoryPage
|
||||
| CandleRevisionHistoryPage
|
||||
)
|
||||
HistoryOrderKey = tuple[datetime, int]
|
||||
HistoryOrderKeyFactory = Callable[[Any], HistoryOrderKey]
|
||||
@@ -0,0 +1,269 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessError,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryCursor,
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
CandleRevisionHistoryRecord,
|
||||
)
|
||||
from src.market_data.access.postgres_history_support import (
|
||||
PostgresHistoryConnectionProvider,
|
||||
candle_history_record_from_row,
|
||||
)
|
||||
|
||||
|
||||
_CANDLE_REVISION_HISTORY_COLUMNS_SQL = """
|
||||
venue,
|
||||
symbol,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
source,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version
|
||||
"""
|
||||
|
||||
_SELECT_CANDLE_REVISION_HISTORY_SQL = f"""
|
||||
SELECT
|
||||
{_CANDLE_REVISION_HISTORY_COLUMNS_SQL}
|
||||
FROM market_data.candle_revisions
|
||||
WHERE venue = %s
|
||||
AND symbol = %s
|
||||
AND interval = %s
|
||||
AND open_time >= %s
|
||||
AND open_time < %s
|
||||
ORDER BY open_time ASC,
|
||||
replay_sequence ASC
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
_SELECT_CANDLE_REVISION_HISTORY_AFTER_CURSOR_SQL = f"""
|
||||
SELECT
|
||||
{_CANDLE_REVISION_HISTORY_COLUMNS_SQL}
|
||||
FROM market_data.candle_revisions
|
||||
WHERE venue = %s
|
||||
AND symbol = %s
|
||||
AND interval = %s
|
||||
AND open_time >= %s
|
||||
AND open_time < %s
|
||||
AND (open_time, replay_sequence) > (%s, %s)
|
||||
ORDER BY open_time ASC,
|
||||
replay_sequence ASC
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
|
||||
class PostgresCandleRevisionHistoryRepository:
|
||||
"""PostgreSQL-адаптер одной страницы истории ревизий свечей."""
|
||||
|
||||
__slots__ = ("_connection_provider",)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
connection_provider: PostgresHistoryConnectionProvider,
|
||||
) -> None:
|
||||
if not callable(connection_provider):
|
||||
raise TypeError("connection_provider must be callable")
|
||||
|
||||
self._connection_provider = connection_provider
|
||||
|
||||
def query_candle_revisions(
|
||||
self,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
"""Прочитать страницу в устойчивом порядке времени открытия."""
|
||||
if type(query) is not CandleRevisionHistoryQuery:
|
||||
raise MarketDataAccessValidationError(
|
||||
"query must be CandleRevisionHistoryQuery"
|
||||
)
|
||||
|
||||
sql, parameters = self._build_query(query)
|
||||
|
||||
try:
|
||||
with self._connection_provider() as connection:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(sql, parameters)
|
||||
rows = cursor.fetchall()
|
||||
return self._page_from_rows(query=query, rows=rows)
|
||||
except MarketDataAccessError:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise MarketDataAccessOperationError(
|
||||
"Failed to query PostgreSQL Candle Revision History."
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _build_query(
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> tuple[str, tuple[object, ...]]:
|
||||
fetch_limit = query.limit + 1
|
||||
cursor = query.cursor
|
||||
|
||||
if cursor is None:
|
||||
return (
|
||||
_SELECT_CANDLE_REVISION_HISTORY_SQL,
|
||||
(
|
||||
query.venue,
|
||||
query.symbol,
|
||||
query.interval,
|
||||
query.time_range.start_time,
|
||||
query.time_range.end_time,
|
||||
fetch_limit,
|
||||
),
|
||||
)
|
||||
|
||||
return (
|
||||
_SELECT_CANDLE_REVISION_HISTORY_AFTER_CURSOR_SQL,
|
||||
(
|
||||
query.venue,
|
||||
query.symbol,
|
||||
query.interval,
|
||||
query.time_range.start_time,
|
||||
query.time_range.end_time,
|
||||
cursor.open_time,
|
||||
cursor.replay_sequence,
|
||||
fetch_limit,
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _page_from_rows(
|
||||
cls,
|
||||
*,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
rows: object,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
if not isinstance(rows, Sequence) or isinstance(
|
||||
rows,
|
||||
(str, bytes, bytearray),
|
||||
):
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Candle Revision History rows."
|
||||
)
|
||||
|
||||
if len(rows) > query.limit + 1:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL exceeded the requested Candle History limit."
|
||||
)
|
||||
|
||||
records = tuple(
|
||||
cls._record_from_row(row=row, query=query)
|
||||
for row in rows
|
||||
)
|
||||
cls._validate_records(records=records, query=query)
|
||||
|
||||
has_more = len(records) > query.limit
|
||||
items = records[: query.limit]
|
||||
next_cursor = (
|
||||
cls._cursor_from_record(query=query, record=items[-1])
|
||||
if has_more
|
||||
else None
|
||||
)
|
||||
|
||||
try:
|
||||
return CandleRevisionHistoryPage(
|
||||
query=query,
|
||||
items=items,
|
||||
next_cursor=next_cursor,
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned inconsistent Candle History page."
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _record_from_row(
|
||||
*,
|
||||
row: object,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryRecord:
|
||||
record = candle_history_record_from_row(row)
|
||||
|
||||
if (
|
||||
record.venue != query.venue
|
||||
or record.symbol != query.symbol
|
||||
or record.interval != query.interval
|
||||
):
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned Candle revision outside query scope."
|
||||
)
|
||||
|
||||
if not query.time_range.contains(record.event_time):
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned Candle revision outside query range."
|
||||
)
|
||||
|
||||
return record
|
||||
|
||||
@staticmethod
|
||||
def _validate_records(
|
||||
*,
|
||||
records: tuple[CandleRevisionHistoryRecord, ...],
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> None:
|
||||
seen_sequences: set[int] = set()
|
||||
previous_key: tuple[datetime, int] | None = None
|
||||
incoming_key = (
|
||||
(
|
||||
query.cursor.open_time,
|
||||
query.cursor.replay_sequence,
|
||||
)
|
||||
if query.cursor is not None
|
||||
else None
|
||||
)
|
||||
|
||||
for record in records:
|
||||
if record.replay_sequence in seen_sequences:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Candle History contains duplicate replay_sequence."
|
||||
)
|
||||
|
||||
if incoming_key is not None and record.order_key <= incoming_key:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Candle History row does not follow query cursor."
|
||||
)
|
||||
|
||||
if previous_key is not None and record.order_key <= previous_key:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Candle History rows are not strictly ordered."
|
||||
)
|
||||
|
||||
seen_sequences.add(record.replay_sequence)
|
||||
previous_key = record.order_key
|
||||
|
||||
@staticmethod
|
||||
def _cursor_from_record(
|
||||
*,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
record: CandleRevisionHistoryRecord,
|
||||
) -> CandleRevisionHistoryCursor:
|
||||
try:
|
||||
return CandleRevisionHistoryCursor(
|
||||
venue=query.venue,
|
||||
symbol=query.symbol,
|
||||
interval=query.interval,
|
||||
time_range=query.time_range,
|
||||
open_time=record.event_time,
|
||||
replay_sequence=record.replay_sequence,
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Cannot create Candle History cursor from stored row."
|
||||
) from error
|
||||
486
app/src/market_data/access/postgres_history_support.py
Normal file
486
app/src/market_data/access/postgres_history_support.py
Normal file
@@ -0,0 +1,486 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from contextlib import AbstractContextManager
|
||||
from datetime import datetime, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from src.market_data.access.exceptions import MarketDataAccessIntegrityError
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryRecord,
|
||||
QuoteHistoryRecord,
|
||||
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,
|
||||
)
|
||||
from src.market_data.acquisition.symbols import normalize_symbol
|
||||
from src.market_data.acquisition.trade_id_sequence import (
|
||||
validate_signed_trade_id,
|
||||
)
|
||||
|
||||
|
||||
CANONICAL_TRADE_SCHEMA_VERSION = 1
|
||||
CANONICAL_QUOTE_SCHEMA_VERSION = 1
|
||||
CANONICAL_CANDLE_SCHEMA_VERSION = 1
|
||||
|
||||
PostgresHistoryConnectionProvider = Callable[
|
||||
[],
|
||||
AbstractContextManager[Any],
|
||||
]
|
||||
|
||||
|
||||
def canonical_text(value: object, *, field_name: str) -> str:
|
||||
"""Проверить каноническую непустую строку без скрытой нормализации."""
|
||||
if type(value) is not str:
|
||||
raise TypeError(f"{field_name} must be a string")
|
||||
|
||||
normalized = value.strip()
|
||||
|
||||
if not normalized or normalized != value:
|
||||
raise ValueError(f"{field_name} must be canonical")
|
||||
|
||||
return value
|
||||
|
||||
|
||||
def canonical_symbol(value: object, *, field_name: str) -> str:
|
||||
"""Проверить уже канонизированный символ сохранённой строки."""
|
||||
raw_symbol = canonical_text(value, field_name=field_name)
|
||||
normalized = normalize_symbol(raw_symbol)
|
||||
|
||||
if normalized != raw_symbol:
|
||||
raise ValueError(f"{field_name} must be canonical")
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def aware_utc_datetime(value: object, *, field_name: str) -> datetime:
|
||||
"""Проверить datetime с часовым поясом и привести его к UTC."""
|
||||
if (
|
||||
type(value) is not datetime
|
||||
or value.tzinfo is None
|
||||
or value.utcoffset() is None
|
||||
):
|
||||
raise TypeError(f"{field_name} must be timezone-aware datetime")
|
||||
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def positive_integer(value: object, *, field_name: str) -> int:
|
||||
"""Проверить точный положительный целочисленный тип."""
|
||||
if type(value) is not int:
|
||||
raise TypeError(f"{field_name} must be an integer")
|
||||
|
||||
if value <= 0:
|
||||
raise ValueError(f"{field_name} must be positive")
|
||||
|
||||
return value
|
||||
|
||||
|
||||
def signed_trade_id(value: object) -> int:
|
||||
"""Проверить знаковый 32-битный идентификатор без изменения."""
|
||||
if type(value) is not int:
|
||||
raise TypeError("stored trade.trade_id must be an integer")
|
||||
|
||||
validate_signed_trade_id(value)
|
||||
return value
|
||||
|
||||
|
||||
def decimal_value(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
allow_zero: bool = False,
|
||||
) -> Decimal:
|
||||
"""Проверить точный конечный Decimal и допустимый знак."""
|
||||
if type(value) is not Decimal or not value.is_finite():
|
||||
raise TypeError(f"{field_name} must be a finite Decimal")
|
||||
|
||||
if allow_zero:
|
||||
if value < 0:
|
||||
raise ValueError(f"{field_name} must not be negative")
|
||||
elif value <= 0:
|
||||
raise ValueError(f"{field_name} must be positive")
|
||||
|
||||
return value
|
||||
|
||||
|
||||
def observation_sources(
|
||||
value: object,
|
||||
*,
|
||||
primary_source: str,
|
||||
field_name: str,
|
||||
) -> tuple[str, ...]:
|
||||
"""Проверить порядок и уникальность источников наблюдения."""
|
||||
if type(value) is not list:
|
||||
raise TypeError(f"{field_name} must be a list")
|
||||
|
||||
if not value:
|
||||
raise ValueError(f"{field_name} must not be empty")
|
||||
|
||||
normalized = tuple(
|
||||
canonical_text(
|
||||
source,
|
||||
field_name=f"{field_name} item",
|
||||
)
|
||||
for source in value
|
||||
)
|
||||
|
||||
if len(set(normalized)) != len(normalized):
|
||||
raise ValueError(f"{field_name} must be unique")
|
||||
|
||||
if normalized[0] != primary_source:
|
||||
raise ValueError(
|
||||
"first observation source must match stored payload.source"
|
||||
)
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def trade_history_record_from_row(row: object) -> TradeHistoryRecord:
|
||||
"""Материализовать строго проверенную строку канонической сделки."""
|
||||
if type(row) is not tuple or len(row) != 13:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Trade History row."
|
||||
)
|
||||
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
stored_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
) = row
|
||||
|
||||
try:
|
||||
normalized_venue = canonical_text(
|
||||
venue,
|
||||
field_name="stored trade.venue",
|
||||
)
|
||||
normalized_symbol = canonical_symbol(
|
||||
symbol,
|
||||
field_name="stored trade.symbol",
|
||||
)
|
||||
normalized_trade_id = signed_trade_id(trade_id)
|
||||
normalized_executed_at = aware_utc_datetime(
|
||||
executed_at,
|
||||
field_name="stored trade.executed_at",
|
||||
)
|
||||
normalized_price = decimal_value(
|
||||
price,
|
||||
field_name="stored trade.price",
|
||||
)
|
||||
normalized_quantity = decimal_value(
|
||||
quantity,
|
||||
field_name="stored trade.quantity",
|
||||
)
|
||||
|
||||
if type(aggressor_side) is not str:
|
||||
raise TypeError("stored trade.aggressor_side must be a string")
|
||||
|
||||
normalized_side = TradeAggressorSide(aggressor_side)
|
||||
normalized_source = canonical_text(
|
||||
source,
|
||||
field_name="stored trade.source",
|
||||
)
|
||||
normalized_first_observed_at = aware_utc_datetime(
|
||||
first_observed_at,
|
||||
field_name="stored trade.first_observed_at",
|
||||
)
|
||||
normalized_last_observed_at = aware_utc_datetime(
|
||||
last_observed_at,
|
||||
field_name="stored trade.last_observed_at",
|
||||
)
|
||||
normalized_sources = observation_sources(
|
||||
stored_sources,
|
||||
primary_source=normalized_source,
|
||||
field_name="stored trade.observation_sources",
|
||||
)
|
||||
normalized_sequence = positive_integer(
|
||||
replay_sequence,
|
||||
field_name="stored trade.replay_sequence",
|
||||
)
|
||||
normalized_version = positive_integer(
|
||||
canonical_schema_version,
|
||||
field_name="stored trade.canonical_schema_version",
|
||||
)
|
||||
|
||||
if normalized_version != CANONICAL_TRADE_SCHEMA_VERSION:
|
||||
raise ValueError("unsupported Canonical Trade schema version")
|
||||
|
||||
trade = Trade(
|
||||
symbol=normalized_symbol,
|
||||
trade_id=normalized_trade_id,
|
||||
price=normalized_price,
|
||||
quantity=normalized_quantity,
|
||||
executed_at=normalized_executed_at,
|
||||
aggressor_side=normalized_side,
|
||||
source=normalized_source,
|
||||
)
|
||||
return TradeHistoryRecord(
|
||||
venue=normalized_venue,
|
||||
trade=trade,
|
||||
first_observed_at=normalized_first_observed_at,
|
||||
last_observed_at=normalized_last_observed_at,
|
||||
observation_sources=normalized_sources,
|
||||
replay_sequence=normalized_sequence,
|
||||
canonical_schema_version=normalized_version,
|
||||
)
|
||||
except (TypeError, ValueError, ArithmeticError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Canonical Trade values."
|
||||
) from error
|
||||
|
||||
|
||||
def quote_history_record_from_row(row: object) -> QuoteHistoryRecord:
|
||||
"""Материализовать строго проверенную строку канонической котировки."""
|
||||
if type(row) is not tuple or len(row) != 11:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Quote History row."
|
||||
)
|
||||
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
source,
|
||||
stored_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
) = row
|
||||
|
||||
try:
|
||||
normalized_venue = canonical_text(
|
||||
venue,
|
||||
field_name="stored quote.venue",
|
||||
)
|
||||
normalized_symbol = canonical_symbol(
|
||||
symbol,
|
||||
field_name="stored quote.symbol",
|
||||
)
|
||||
normalized_received_at = aware_utc_datetime(
|
||||
received_at,
|
||||
field_name="stored quote.received_at",
|
||||
)
|
||||
normalized_exchange_timestamp = (
|
||||
aware_utc_datetime(
|
||||
exchange_timestamp,
|
||||
field_name="stored quote.exchange_timestamp",
|
||||
)
|
||||
if exchange_timestamp is not None
|
||||
else None
|
||||
)
|
||||
normalized_last_price = decimal_value(
|
||||
last_price,
|
||||
field_name="stored quote.last_price",
|
||||
)
|
||||
normalized_bid_price = decimal_value(
|
||||
bid_price,
|
||||
field_name="stored quote.bid_price",
|
||||
)
|
||||
normalized_ask_price = decimal_value(
|
||||
ask_price,
|
||||
field_name="stored quote.ask_price",
|
||||
)
|
||||
|
||||
if normalized_bid_price > normalized_ask_price:
|
||||
raise ValueError("stored quote.bid_price exceeds ask_price")
|
||||
|
||||
normalized_source = canonical_text(
|
||||
source,
|
||||
field_name="stored quote.source",
|
||||
)
|
||||
normalized_sources = observation_sources(
|
||||
stored_sources,
|
||||
primary_source=normalized_source,
|
||||
field_name="stored quote.observation_sources",
|
||||
)
|
||||
normalized_sequence = positive_integer(
|
||||
replay_sequence,
|
||||
field_name="stored quote.replay_sequence",
|
||||
)
|
||||
normalized_version = positive_integer(
|
||||
canonical_schema_version,
|
||||
field_name="stored quote.canonical_schema_version",
|
||||
)
|
||||
|
||||
if normalized_version != CANONICAL_QUOTE_SCHEMA_VERSION:
|
||||
raise ValueError("unsupported Canonical Quote schema version")
|
||||
|
||||
quote = Quote(
|
||||
symbol=normalized_symbol,
|
||||
last_price=normalized_last_price,
|
||||
bid_price=normalized_bid_price,
|
||||
ask_price=normalized_ask_price,
|
||||
exchange_timestamp=normalized_exchange_timestamp,
|
||||
received_at=normalized_received_at,
|
||||
source=normalized_source,
|
||||
)
|
||||
return QuoteHistoryRecord(
|
||||
venue=normalized_venue,
|
||||
quote=quote,
|
||||
observation_sources=normalized_sources,
|
||||
replay_sequence=normalized_sequence,
|
||||
canonical_schema_version=normalized_version,
|
||||
)
|
||||
except (TypeError, ValueError, ArithmeticError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Canonical Quote values."
|
||||
) from error
|
||||
|
||||
|
||||
def candle_history_record_from_row(
|
||||
row: object,
|
||||
) -> CandleRevisionHistoryRecord:
|
||||
"""Материализовать строго проверенную строку ревизии свечи."""
|
||||
if type(row) is not tuple or len(row) != 15:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Candle History row."
|
||||
)
|
||||
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
source,
|
||||
stored_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
) = row
|
||||
|
||||
try:
|
||||
normalized_venue = canonical_text(
|
||||
venue,
|
||||
field_name="stored candle.venue",
|
||||
)
|
||||
normalized_symbol = canonical_symbol(
|
||||
symbol,
|
||||
field_name="stored candle.symbol",
|
||||
)
|
||||
normalized_interval = canonical_text(
|
||||
interval,
|
||||
field_name="stored candle.interval",
|
||||
)
|
||||
normalized_open_time = aware_utc_datetime(
|
||||
open_time,
|
||||
field_name="stored candle.open_time",
|
||||
)
|
||||
normalized_observed_at = aware_utc_datetime(
|
||||
observed_at,
|
||||
field_name="stored candle.observed_at",
|
||||
)
|
||||
|
||||
if normalized_observed_at < normalized_open_time:
|
||||
raise ValueError("stored candle.observed_at precedes open_time")
|
||||
|
||||
normalized_open_price = decimal_value(
|
||||
open_price,
|
||||
field_name="stored candle.open_price",
|
||||
)
|
||||
normalized_high_price = decimal_value(
|
||||
high_price,
|
||||
field_name="stored candle.high_price",
|
||||
)
|
||||
normalized_low_price = decimal_value(
|
||||
low_price,
|
||||
field_name="stored candle.low_price",
|
||||
)
|
||||
normalized_close_price = decimal_value(
|
||||
close_price,
|
||||
field_name="stored candle.close_price",
|
||||
)
|
||||
normalized_volume = decimal_value(
|
||||
volume,
|
||||
field_name="stored candle.volume",
|
||||
allow_zero=True,
|
||||
)
|
||||
|
||||
if normalized_low_price > normalized_high_price:
|
||||
raise ValueError("stored candle.low_price exceeds high_price")
|
||||
|
||||
if not (
|
||||
normalized_low_price
|
||||
<= normalized_open_price
|
||||
<= normalized_high_price
|
||||
):
|
||||
raise ValueError("stored candle.open_price is outside range")
|
||||
|
||||
if not (
|
||||
normalized_low_price
|
||||
<= normalized_close_price
|
||||
<= normalized_high_price
|
||||
):
|
||||
raise ValueError("stored candle.close_price is outside range")
|
||||
|
||||
if type(is_final) is not bool:
|
||||
raise TypeError("stored candle.is_final must be a boolean")
|
||||
|
||||
normalized_source = canonical_text(
|
||||
source,
|
||||
field_name="stored candle.source",
|
||||
)
|
||||
normalized_sources = observation_sources(
|
||||
stored_sources,
|
||||
primary_source=normalized_source,
|
||||
field_name="stored candle.observation_sources",
|
||||
)
|
||||
normalized_sequence = positive_integer(
|
||||
replay_sequence,
|
||||
field_name="stored candle.replay_sequence",
|
||||
)
|
||||
normalized_version = positive_integer(
|
||||
canonical_schema_version,
|
||||
field_name="stored candle.canonical_schema_version",
|
||||
)
|
||||
|
||||
if normalized_version != CANONICAL_CANDLE_SCHEMA_VERSION:
|
||||
raise ValueError("unsupported Canonical Candle schema version")
|
||||
|
||||
candle = Candle(
|
||||
symbol=normalized_symbol,
|
||||
interval=normalized_interval,
|
||||
open_time=normalized_open_time,
|
||||
open_price=normalized_open_price,
|
||||
high_price=normalized_high_price,
|
||||
low_price=normalized_low_price,
|
||||
close_price=normalized_close_price,
|
||||
volume=normalized_volume,
|
||||
source=normalized_source,
|
||||
)
|
||||
return CandleRevisionHistoryRecord(
|
||||
venue=normalized_venue,
|
||||
candle=candle,
|
||||
observed_at=normalized_observed_at,
|
||||
is_final=is_final,
|
||||
observation_sources=normalized_sources,
|
||||
replay_sequence=normalized_sequence,
|
||||
canonical_schema_version=normalized_version,
|
||||
)
|
||||
except (TypeError, ValueError, ArithmeticError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Canonical Candle values."
|
||||
) from error
|
||||
253
app/src/market_data/access/postgres_quote_history_repository.py
Normal file
253
app/src/market_data/access/postgres_quote_history_repository.py
Normal file
@@ -0,0 +1,253 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessError,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
QuoteHistoryCursor,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
QuoteHistoryRecord,
|
||||
)
|
||||
from src.market_data.access.postgres_history_support import (
|
||||
PostgresHistoryConnectionProvider,
|
||||
quote_history_record_from_row,
|
||||
)
|
||||
|
||||
|
||||
_QUOTE_HISTORY_COLUMNS_SQL = """
|
||||
venue,
|
||||
symbol,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
source,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version
|
||||
"""
|
||||
|
||||
_SELECT_QUOTE_HISTORY_SQL = f"""
|
||||
SELECT
|
||||
{_QUOTE_HISTORY_COLUMNS_SQL}
|
||||
FROM market_data.quotes
|
||||
WHERE venue = %s
|
||||
AND symbol = %s
|
||||
AND received_at >= %s
|
||||
AND received_at < %s
|
||||
ORDER BY received_at ASC,
|
||||
replay_sequence ASC
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
_SELECT_QUOTE_HISTORY_AFTER_CURSOR_SQL = f"""
|
||||
SELECT
|
||||
{_QUOTE_HISTORY_COLUMNS_SQL}
|
||||
FROM market_data.quotes
|
||||
WHERE venue = %s
|
||||
AND symbol = %s
|
||||
AND received_at >= %s
|
||||
AND received_at < %s
|
||||
AND (received_at, replay_sequence) > (%s, %s)
|
||||
ORDER BY received_at ASC,
|
||||
replay_sequence ASC
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
|
||||
class PostgresQuoteHistoryRepository:
|
||||
"""PostgreSQL-адаптер одной страницы канонической истории котировок."""
|
||||
|
||||
__slots__ = ("_connection_provider",)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
connection_provider: PostgresHistoryConnectionProvider,
|
||||
) -> None:
|
||||
if not callable(connection_provider):
|
||||
raise TypeError("connection_provider must be callable")
|
||||
|
||||
self._connection_provider = connection_provider
|
||||
|
||||
def query_quotes(self, query: QuoteHistoryQuery) -> QuoteHistoryPage:
|
||||
"""Прочитать ограниченную страницу с устойчивой сортировкой."""
|
||||
if type(query) is not QuoteHistoryQuery:
|
||||
raise MarketDataAccessValidationError(
|
||||
"query must be QuoteHistoryQuery"
|
||||
)
|
||||
|
||||
sql, parameters = self._build_query(query)
|
||||
|
||||
try:
|
||||
with self._connection_provider() as connection:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(sql, parameters)
|
||||
rows = cursor.fetchall()
|
||||
return self._page_from_rows(query=query, rows=rows)
|
||||
except MarketDataAccessError:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise MarketDataAccessOperationError(
|
||||
"Failed to query PostgreSQL Quote History."
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _build_query(
|
||||
query: QuoteHistoryQuery,
|
||||
) -> tuple[str, tuple[object, ...]]:
|
||||
fetch_limit = query.limit + 1
|
||||
cursor = query.cursor
|
||||
|
||||
if cursor is None:
|
||||
return (
|
||||
_SELECT_QUOTE_HISTORY_SQL,
|
||||
(
|
||||
query.venue,
|
||||
query.symbol,
|
||||
query.time_range.start_time,
|
||||
query.time_range.end_time,
|
||||
fetch_limit,
|
||||
),
|
||||
)
|
||||
|
||||
return (
|
||||
_SELECT_QUOTE_HISTORY_AFTER_CURSOR_SQL,
|
||||
(
|
||||
query.venue,
|
||||
query.symbol,
|
||||
query.time_range.start_time,
|
||||
query.time_range.end_time,
|
||||
cursor.received_at,
|
||||
cursor.replay_sequence,
|
||||
fetch_limit,
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _page_from_rows(
|
||||
cls,
|
||||
*,
|
||||
query: QuoteHistoryQuery,
|
||||
rows: object,
|
||||
) -> QuoteHistoryPage:
|
||||
if not isinstance(rows, Sequence) or isinstance(
|
||||
rows,
|
||||
(str, bytes, bytearray),
|
||||
):
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Quote History rows."
|
||||
)
|
||||
|
||||
if len(rows) > query.limit + 1:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL exceeded the requested Quote History limit."
|
||||
)
|
||||
|
||||
records = tuple(
|
||||
cls._record_from_row(row=row, query=query)
|
||||
for row in rows
|
||||
)
|
||||
cls._validate_records(records=records, query=query)
|
||||
|
||||
has_more = len(records) > query.limit
|
||||
items = records[: query.limit]
|
||||
next_cursor = (
|
||||
cls._cursor_from_record(query=query, record=items[-1])
|
||||
if has_more
|
||||
else None
|
||||
)
|
||||
|
||||
try:
|
||||
return QuoteHistoryPage(
|
||||
query=query,
|
||||
items=items,
|
||||
next_cursor=next_cursor,
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned inconsistent Quote History page."
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _record_from_row(
|
||||
*,
|
||||
row: object,
|
||||
query: QuoteHistoryQuery,
|
||||
) -> QuoteHistoryRecord:
|
||||
record = quote_history_record_from_row(row)
|
||||
|
||||
if record.venue != query.venue or record.symbol != query.symbol:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned Quote outside query scope."
|
||||
)
|
||||
|
||||
if not query.time_range.contains(record.event_time):
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned Quote outside query range."
|
||||
)
|
||||
|
||||
return record
|
||||
|
||||
@staticmethod
|
||||
def _validate_records(
|
||||
*,
|
||||
records: tuple[QuoteHistoryRecord, ...],
|
||||
query: QuoteHistoryQuery,
|
||||
) -> None:
|
||||
seen_sequences: set[int] = set()
|
||||
previous_key: tuple[datetime, int] | None = None
|
||||
incoming_key = (
|
||||
(
|
||||
query.cursor.received_at,
|
||||
query.cursor.replay_sequence,
|
||||
)
|
||||
if query.cursor is not None
|
||||
else None
|
||||
)
|
||||
|
||||
for record in records:
|
||||
if record.replay_sequence in seen_sequences:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Quote History contains duplicate replay_sequence."
|
||||
)
|
||||
|
||||
if incoming_key is not None and record.order_key <= incoming_key:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Quote History row does not follow query cursor."
|
||||
)
|
||||
|
||||
if previous_key is not None and record.order_key <= previous_key:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Quote History rows are not strictly ordered."
|
||||
)
|
||||
|
||||
seen_sequences.add(record.replay_sequence)
|
||||
previous_key = record.order_key
|
||||
|
||||
@staticmethod
|
||||
def _cursor_from_record(
|
||||
*,
|
||||
query: QuoteHistoryQuery,
|
||||
record: QuoteHistoryRecord,
|
||||
) -> QuoteHistoryCursor:
|
||||
try:
|
||||
return QuoteHistoryCursor(
|
||||
venue=query.venue,
|
||||
symbol=query.symbol,
|
||||
time_range=query.time_range,
|
||||
received_at=record.event_time,
|
||||
replay_sequence=record.replay_sequence,
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Cannot create Quote History cursor from stored row."
|
||||
) from error
|
||||
257
app/src/market_data/access/postgres_trade_history_repository.py
Normal file
257
app/src/market_data/access/postgres_trade_history_repository.py
Normal file
@@ -0,0 +1,257 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessError,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
TradeHistoryCursor,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryRecord,
|
||||
)
|
||||
from src.market_data.access.postgres_history_support import (
|
||||
PostgresHistoryConnectionProvider,
|
||||
trade_history_record_from_row,
|
||||
)
|
||||
|
||||
_TRADE_HISTORY_COLUMNS_SQL = """
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version
|
||||
"""
|
||||
|
||||
_SELECT_TRADE_HISTORY_SQL = f"""
|
||||
SELECT
|
||||
{_TRADE_HISTORY_COLUMNS_SQL}
|
||||
FROM market_data.trades
|
||||
WHERE venue = %s
|
||||
AND symbol = %s
|
||||
AND executed_at >= %s
|
||||
AND executed_at < %s
|
||||
ORDER BY executed_at ASC,
|
||||
replay_sequence ASC
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
_SELECT_TRADE_HISTORY_AFTER_CURSOR_SQL = f"""
|
||||
SELECT
|
||||
{_TRADE_HISTORY_COLUMNS_SQL}
|
||||
FROM market_data.trades
|
||||
WHERE venue = %s
|
||||
AND symbol = %s
|
||||
AND executed_at >= %s
|
||||
AND executed_at < %s
|
||||
AND (executed_at, replay_sequence) > (%s, %s)
|
||||
ORDER BY executed_at ASC,
|
||||
replay_sequence ASC
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
|
||||
class PostgresTradeHistoryRepository:
|
||||
"""PostgreSQL-адаптер одной страницы канонической истории сделок."""
|
||||
|
||||
__slots__ = ("_connection_provider",)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
connection_provider: PostgresHistoryConnectionProvider,
|
||||
) -> None:
|
||||
if not callable(connection_provider):
|
||||
raise TypeError("connection_provider must be callable")
|
||||
|
||||
self._connection_provider = connection_provider
|
||||
|
||||
def query_trades(self, query: TradeHistoryQuery) -> TradeHistoryPage:
|
||||
"""Прочитать одну bounded страницу в устойчивом keyset-порядке."""
|
||||
if type(query) is not TradeHistoryQuery:
|
||||
raise MarketDataAccessValidationError(
|
||||
"query must be TradeHistoryQuery"
|
||||
)
|
||||
|
||||
sql, parameters = self._build_query(query)
|
||||
|
||||
try:
|
||||
with self._connection_provider() as connection:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(sql, parameters)
|
||||
rows = cursor.fetchall()
|
||||
return self._page_from_rows(query=query, rows=rows)
|
||||
except MarketDataAccessError:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise MarketDataAccessOperationError(
|
||||
"Failed to query PostgreSQL Trade History."
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _build_query(
|
||||
query: TradeHistoryQuery,
|
||||
) -> tuple[str, tuple[object, ...]]:
|
||||
fetch_limit = query.limit + 1
|
||||
cursor = query.cursor
|
||||
|
||||
if cursor is None:
|
||||
return (
|
||||
_SELECT_TRADE_HISTORY_SQL,
|
||||
(
|
||||
query.venue,
|
||||
query.symbol,
|
||||
query.time_range.start_time,
|
||||
query.time_range.end_time,
|
||||
fetch_limit,
|
||||
),
|
||||
)
|
||||
|
||||
return (
|
||||
_SELECT_TRADE_HISTORY_AFTER_CURSOR_SQL,
|
||||
(
|
||||
query.venue,
|
||||
query.symbol,
|
||||
query.time_range.start_time,
|
||||
query.time_range.end_time,
|
||||
cursor.executed_at,
|
||||
cursor.replay_sequence,
|
||||
fetch_limit,
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _page_from_rows(
|
||||
cls,
|
||||
*,
|
||||
query: TradeHistoryQuery,
|
||||
rows: object,
|
||||
) -> TradeHistoryPage:
|
||||
if not isinstance(rows, Sequence) or isinstance(
|
||||
rows,
|
||||
(str, bytes, bytearray),
|
||||
):
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Trade History rows."
|
||||
)
|
||||
|
||||
if len(rows) > query.limit + 1:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL exceeded the requested Trade History limit."
|
||||
)
|
||||
|
||||
records = tuple(
|
||||
cls._record_from_row(row=row, query=query)
|
||||
for row in rows
|
||||
)
|
||||
cls._validate_records(records=records, query=query)
|
||||
|
||||
has_more = len(records) > query.limit
|
||||
items = records[: query.limit]
|
||||
next_cursor = (
|
||||
cls._cursor_from_record(query=query, record=items[-1])
|
||||
if has_more
|
||||
else None
|
||||
)
|
||||
|
||||
try:
|
||||
return TradeHistoryPage(
|
||||
query=query,
|
||||
items=items,
|
||||
next_cursor=next_cursor,
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned inconsistent Trade History page."
|
||||
) from error
|
||||
|
||||
@classmethod
|
||||
def _record_from_row(
|
||||
cls,
|
||||
*,
|
||||
row: object,
|
||||
query: TradeHistoryQuery,
|
||||
) -> TradeHistoryRecord:
|
||||
record = trade_history_record_from_row(row)
|
||||
|
||||
try:
|
||||
if record.venue != query.venue or record.symbol != query.symbol:
|
||||
raise ValueError("stored Trade is outside query scope")
|
||||
|
||||
if not query.time_range.contains(record.event_time):
|
||||
raise ValueError("stored Trade is outside query range")
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Canonical Trade values."
|
||||
) from error
|
||||
|
||||
return record
|
||||
|
||||
@staticmethod
|
||||
def _validate_records(
|
||||
*,
|
||||
records: tuple[TradeHistoryRecord, ...],
|
||||
query: TradeHistoryQuery,
|
||||
) -> None:
|
||||
seen_sequences: set[int] = set()
|
||||
previous_key: tuple[datetime, int] | None = None
|
||||
incoming_key = (
|
||||
(
|
||||
query.cursor.executed_at,
|
||||
query.cursor.replay_sequence,
|
||||
)
|
||||
if query.cursor is not None
|
||||
else None
|
||||
)
|
||||
|
||||
for record in records:
|
||||
if record.replay_sequence in seen_sequences:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Trade History contains duplicate replay_sequence."
|
||||
)
|
||||
|
||||
if incoming_key is not None and record.order_key <= incoming_key:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Trade History row does not follow query cursor."
|
||||
)
|
||||
|
||||
if previous_key is not None and record.order_key <= previous_key:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Trade History rows are not strictly ordered."
|
||||
)
|
||||
|
||||
seen_sequences.add(record.replay_sequence)
|
||||
previous_key = record.order_key
|
||||
|
||||
@staticmethod
|
||||
def _cursor_from_record(
|
||||
*,
|
||||
query: TradeHistoryQuery,
|
||||
record: TradeHistoryRecord,
|
||||
) -> TradeHistoryCursor:
|
||||
try:
|
||||
return TradeHistoryCursor(
|
||||
venue=query.venue,
|
||||
symbol=query.symbol,
|
||||
time_range=query.time_range,
|
||||
executed_at=record.event_time,
|
||||
replay_sequence=record.replay_sequence,
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"Cannot create Trade History cursor from stored row."
|
||||
) from error
|
||||
62
app/src/market_data/replay/__init__.py
Normal file
62
app/src/market_data/replay/__init__.py
Normal file
@@ -0,0 +1,62 @@
|
||||
from src.market_data.replay.contracts import (
|
||||
MarketDataClockProtocol,
|
||||
ReplayClockProtocol,
|
||||
ReplayConsumerFactoryProtocol,
|
||||
ReplayConsumerProtocol,
|
||||
ReplayPlanBuilderProtocol,
|
||||
ReplaySessionProtocol,
|
||||
)
|
||||
from src.market_data.replay.deterministic_replay_clock import (
|
||||
DeterministicReplayClock,
|
||||
)
|
||||
from src.market_data.replay.exceptions import (
|
||||
MarketDataReplayError,
|
||||
MarketDataReplayValidationError,
|
||||
ReplayClockError,
|
||||
ReplayPlanLimitExceededError,
|
||||
ReplaySessionStateError,
|
||||
)
|
||||
from src.market_data.replay.models import (
|
||||
REPLAY_PLAN_MAX_RECORDS_DEFAULT,
|
||||
REPLAY_PLAN_MAX_RECORDS_LIMIT,
|
||||
CanonicalMarketData,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
ReplaySessionState,
|
||||
)
|
||||
from src.market_data.replay.postgres_replay_plan_builder import (
|
||||
PostgresReplayPlanBuilder,
|
||||
)
|
||||
from src.market_data.replay.replay_session import ReplaySession
|
||||
from src.market_data.replay.replay_session_factory import (
|
||||
ReplaySessionFactory,
|
||||
)
|
||||
|
||||
|
||||
__all__ = (
|
||||
"REPLAY_PLAN_MAX_RECORDS_DEFAULT",
|
||||
"REPLAY_PLAN_MAX_RECORDS_LIMIT",
|
||||
"CanonicalMarketData",
|
||||
"DeterministicReplayClock",
|
||||
"MarketDataClockProtocol",
|
||||
"MarketDataReplayError",
|
||||
"MarketDataReplayValidationError",
|
||||
"PostgresReplayPlanBuilder",
|
||||
"ReplayClockError",
|
||||
"ReplayClockProtocol",
|
||||
"ReplayConsumerFactoryProtocol",
|
||||
"ReplayConsumerProtocol",
|
||||
"ReplayDataType",
|
||||
"ReplayEvent",
|
||||
"ReplayPlan",
|
||||
"ReplayPlanBuilderProtocol",
|
||||
"ReplayPlanLimitExceededError",
|
||||
"ReplayPlanRequest",
|
||||
"ReplaySession",
|
||||
"ReplaySessionFactory",
|
||||
"ReplaySessionProtocol",
|
||||
"ReplaySessionState",
|
||||
"ReplaySessionStateError",
|
||||
)
|
||||
77
app/src/market_data/replay/contracts.py
Normal file
77
app/src/market_data/replay/contracts.py
Normal file
@@ -0,0 +1,77 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
from src.market_data.replay.models import (
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
ReplaySessionState,
|
||||
)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class MarketDataClockProtocol(Protocol):
|
||||
"""Доступ consumer к текущему времени без права его изменять."""
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ReplayClockProtocol(MarketDataClockProtocol, Protocol):
|
||||
"""Управляющая граница Virtual Clock одной Replay Session."""
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ReplayConsumerProtocol(Protocol):
|
||||
"""Один последовательный async consumer Canonical Replay events."""
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ReplayConsumerFactoryProtocol(Protocol):
|
||||
"""Граница создания отдельного consumer для одной Replay Session."""
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ReplayPlanBuilderProtocol(Protocol):
|
||||
"""Граница создания bounded snapshot до запуска Replay Session."""
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ReplaySessionProtocol(Protocol):
|
||||
"""Одноразовая исполняемая Replay Session без hidden root task."""
|
||||
|
||||
@property
|
||||
def state(self) -> ReplaySessionState:
|
||||
...
|
||||
|
||||
@property
|
||||
def plan(self) -> ReplayPlan:
|
||||
...
|
||||
|
||||
@property
|
||||
def clock(self) -> MarketDataClockProtocol:
|
||||
...
|
||||
|
||||
async def run(self) -> None:
|
||||
...
|
||||
67
app/src/market_data/replay/deterministic_replay_clock.py
Normal file
67
app/src/market_data/replay/deterministic_replay_clock.py
Normal file
@@ -0,0 +1,67 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from src.market_data.replay.exceptions import ReplayClockError
|
||||
|
||||
|
||||
def _normalize_utc_datetime(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
) -> datetime:
|
||||
"""Проверить время с часовым поясом и канонизировать его в UTC."""
|
||||
if not isinstance(value, datetime):
|
||||
raise TypeError(f"{field_name} must be a datetime")
|
||||
|
||||
if value.tzinfo is None or value.utcoffset() is None:
|
||||
raise ValueError(f"{field_name} must contain timezone")
|
||||
|
||||
normalized = value.astimezone(timezone.utc)
|
||||
|
||||
# Сохраняем совместимость с datetime-подклассами на входе, но не
|
||||
# переносим их изменённое поведение во внутреннее состояние Clock.
|
||||
return datetime(
|
||||
normalized.year,
|
||||
normalized.month,
|
||||
normalized.day,
|
||||
normalized.hour,
|
||||
normalized.minute,
|
||||
normalized.second,
|
||||
normalized.microsecond,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
|
||||
|
||||
class DeterministicReplayClock:
|
||||
"""Синхронные логические часы одного детерминированного Replay."""
|
||||
|
||||
__slots__ = ("_now",)
|
||||
|
||||
def __init__(self, initial_time: datetime) -> None:
|
||||
self._now = _normalize_utc_datetime(
|
||||
initial_time,
|
||||
field_name="initial_time",
|
||||
)
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
"""Вернуть текущий канонический UTC-момент Replay."""
|
||||
return self._now
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
"""Перевести логическое время вперёд или оставить его прежним."""
|
||||
normalized = _normalize_utc_datetime(
|
||||
instant,
|
||||
field_name="instant",
|
||||
)
|
||||
|
||||
if normalized < self._now:
|
||||
raise ReplayClockError(
|
||||
"Replay Clock cannot move backwards."
|
||||
)
|
||||
|
||||
if normalized == self._now:
|
||||
return
|
||||
|
||||
self._now = normalized
|
||||
21
app/src/market_data/replay/exceptions.py
Normal file
21
app/src/market_data/replay/exceptions.py
Normal file
@@ -0,0 +1,21 @@
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class MarketDataReplayError(Exception):
|
||||
"""Базовая ошибка детерминированного Market Data Replay."""
|
||||
|
||||
|
||||
class MarketDataReplayValidationError(MarketDataReplayError):
|
||||
"""Недопустимый Replay request, plan или dependency."""
|
||||
|
||||
|
||||
class ReplayPlanLimitExceededError(MarketDataReplayError):
|
||||
"""Snapshot содержит больше событий, чем разрешает Replay request."""
|
||||
|
||||
|
||||
class ReplayClockError(MarketDataReplayError):
|
||||
"""Virtual Clock получил недопустимый переход времени."""
|
||||
|
||||
|
||||
class ReplaySessionStateError(MarketDataReplayError):
|
||||
"""Операция недопустима в текущем состоянии Replay Session."""
|
||||
371
app/src/market_data/replay/models.py
Normal file
371
app/src/market_data/replay/models.py
Normal file
@@ -0,0 +1,371 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
|
||||
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.acquisition.symbols import normalize_symbol
|
||||
|
||||
|
||||
REPLAY_PLAN_MAX_RECORDS_DEFAULT = 100_000
|
||||
REPLAY_PLAN_MAX_RECORDS_LIMIT = 100_000
|
||||
|
||||
|
||||
class ReplayDataType(Enum):
|
||||
"""Тип Canonical payload в одном Replay Plan."""
|
||||
|
||||
TRADE = "trade"
|
||||
QUOTE = "quote"
|
||||
CANDLE_REVISION = "candle_revision"
|
||||
|
||||
|
||||
class ReplaySessionState(Enum):
|
||||
"""Одноразовый lifecycle одной Replay Session."""
|
||||
|
||||
CREATED = "created"
|
||||
RUNNING = "running"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
CanonicalMarketData = Trade | Quote | Candle
|
||||
|
||||
|
||||
def _normalize_non_empty_text(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
) -> str:
|
||||
if not isinstance(value, str):
|
||||
raise TypeError(f"{field_name} must be a string")
|
||||
|
||||
normalized = value.strip()
|
||||
|
||||
if not normalized:
|
||||
raise ValueError(f"{field_name} must not be empty")
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def _normalize_utc_datetime(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
) -> datetime:
|
||||
if not isinstance(value, datetime):
|
||||
raise TypeError(f"{field_name} must be a datetime")
|
||||
|
||||
if value.tzinfo is None or value.utcoffset() is None:
|
||||
raise ValueError(f"{field_name} must contain timezone")
|
||||
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _normalize_positive_integer(
|
||||
value: object,
|
||||
*,
|
||||
field_name: str,
|
||||
maximum: int | None = None,
|
||||
) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"{field_name} must be an integer")
|
||||
|
||||
if value <= 0:
|
||||
raise ValueError(f"{field_name} must be positive")
|
||||
|
||||
if maximum is not None and value > maximum:
|
||||
raise ValueError(
|
||||
f"{field_name} must not exceed {maximum}"
|
||||
)
|
||||
|
||||
return value
|
||||
|
||||
|
||||
def _canonical_payload_time(payload: CanonicalMarketData) -> datetime:
|
||||
if isinstance(payload, Trade):
|
||||
return _normalize_utc_datetime(
|
||||
payload.executed_at,
|
||||
field_name="payload.executed_at",
|
||||
)
|
||||
|
||||
if isinstance(payload, Quote):
|
||||
return _normalize_utc_datetime(
|
||||
payload.received_at,
|
||||
field_name="payload.received_at",
|
||||
)
|
||||
|
||||
return _normalize_utc_datetime(
|
||||
payload.open_time,
|
||||
field_name="payload.open_time",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReplayEvent:
|
||||
"""Canonical payload с устойчивой позицией во времени Replay."""
|
||||
|
||||
venue: str
|
||||
replay_at: datetime
|
||||
replay_sequence: int
|
||||
payload: CanonicalMarketData
|
||||
candle_is_final: bool | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
normalized_venue = _normalize_non_empty_text(
|
||||
self.venue,
|
||||
field_name="venue",
|
||||
)
|
||||
normalized_replay_at = _normalize_utc_datetime(
|
||||
self.replay_at,
|
||||
field_name="replay_at",
|
||||
)
|
||||
normalized_sequence = _normalize_positive_integer(
|
||||
self.replay_sequence,
|
||||
field_name="replay_sequence",
|
||||
)
|
||||
|
||||
if type(self.payload) not in (Trade, Quote, Candle):
|
||||
raise TypeError(
|
||||
"payload must be a Canonical Trade, Quote or Candle"
|
||||
)
|
||||
|
||||
normalized_symbol = normalize_symbol(
|
||||
_normalize_non_empty_text(
|
||||
self.payload.symbol,
|
||||
field_name="payload.symbol",
|
||||
)
|
||||
)
|
||||
|
||||
if normalized_symbol != self.payload.symbol:
|
||||
raise ValueError("payload.symbol must be canonical")
|
||||
|
||||
normalized_source = _normalize_non_empty_text(
|
||||
self.payload.source,
|
||||
field_name="payload.source",
|
||||
)
|
||||
|
||||
if normalized_source != self.payload.source:
|
||||
raise ValueError("payload.source must be canonical")
|
||||
|
||||
payload_time = _canonical_payload_time(self.payload)
|
||||
|
||||
if isinstance(self.payload, Candle):
|
||||
normalized_interval = _normalize_non_empty_text(
|
||||
self.payload.interval,
|
||||
field_name="payload.interval",
|
||||
)
|
||||
|
||||
if normalized_interval != self.payload.interval:
|
||||
raise ValueError("payload.interval must be canonical")
|
||||
|
||||
if not isinstance(self.candle_is_final, bool):
|
||||
raise TypeError(
|
||||
"candle_is_final must be a boolean for Candle"
|
||||
)
|
||||
|
||||
if normalized_replay_at < payload_time:
|
||||
raise ValueError(
|
||||
"Candle replay_at must not precede open_time"
|
||||
)
|
||||
else:
|
||||
if self.candle_is_final is not None:
|
||||
raise ValueError(
|
||||
"candle_is_final must be None for Trade and Quote"
|
||||
)
|
||||
|
||||
if normalized_replay_at != payload_time:
|
||||
raise ValueError(
|
||||
"replay_at must match Canonical payload event time"
|
||||
)
|
||||
|
||||
object.__setattr__(self, "venue", normalized_venue)
|
||||
object.__setattr__(self, "replay_at", normalized_replay_at)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"replay_sequence",
|
||||
normalized_sequence,
|
||||
)
|
||||
|
||||
@property
|
||||
def symbol(self) -> str:
|
||||
return self.payload.symbol
|
||||
|
||||
@property
|
||||
def data_type(self) -> ReplayDataType:
|
||||
if isinstance(self.payload, Trade):
|
||||
return ReplayDataType.TRADE
|
||||
|
||||
if isinstance(self.payload, Quote):
|
||||
return ReplayDataType.QUOTE
|
||||
|
||||
return ReplayDataType.CANDLE_REVISION
|
||||
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (self.replay_at, self.replay_sequence)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReplayPlanRequest:
|
||||
"""Неизменяемый scope одного bounded Replay snapshot."""
|
||||
|
||||
venue: str
|
||||
symbols: tuple[str, ...]
|
||||
data_types: tuple[ReplayDataType, ...]
|
||||
time_range: HistoricalTimeRange
|
||||
candle_intervals: tuple[str, ...] = ()
|
||||
max_records: int = REPLAY_PLAN_MAX_RECORDS_DEFAULT
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
venue = _normalize_non_empty_text(
|
||||
self.venue,
|
||||
field_name="venue",
|
||||
)
|
||||
|
||||
if type(self.symbols) is not tuple:
|
||||
raise TypeError("symbols must be a tuple")
|
||||
|
||||
if not self.symbols:
|
||||
raise ValueError("symbols must not be empty")
|
||||
|
||||
symbols = tuple(
|
||||
normalize_symbol(
|
||||
_normalize_non_empty_text(
|
||||
symbol,
|
||||
field_name="symbols item",
|
||||
)
|
||||
)
|
||||
for symbol in self.symbols
|
||||
)
|
||||
|
||||
if len(set(symbols)) != len(symbols):
|
||||
raise ValueError("symbols must be unique after normalization")
|
||||
|
||||
if type(self.data_types) is not tuple:
|
||||
raise TypeError("data_types must be a tuple")
|
||||
|
||||
if not self.data_types:
|
||||
raise ValueError("data_types must not be empty")
|
||||
|
||||
if any(
|
||||
not isinstance(data_type, ReplayDataType)
|
||||
for data_type in self.data_types
|
||||
):
|
||||
raise TypeError("data_types must contain only ReplayDataType")
|
||||
|
||||
if len(set(self.data_types)) != len(self.data_types):
|
||||
raise ValueError("data_types must be unique")
|
||||
|
||||
if type(self.time_range) is not HistoricalTimeRange:
|
||||
raise TypeError("time_range must be HistoricalTimeRange")
|
||||
|
||||
if type(self.candle_intervals) is not tuple:
|
||||
raise TypeError("candle_intervals must be a tuple")
|
||||
|
||||
candle_intervals = tuple(
|
||||
_normalize_non_empty_text(
|
||||
interval,
|
||||
field_name="candle_intervals item",
|
||||
)
|
||||
for interval in self.candle_intervals
|
||||
)
|
||||
|
||||
if len(set(candle_intervals)) != len(candle_intervals):
|
||||
raise ValueError("candle_intervals must be unique")
|
||||
|
||||
includes_candles = (
|
||||
ReplayDataType.CANDLE_REVISION in self.data_types
|
||||
)
|
||||
|
||||
if includes_candles and not candle_intervals:
|
||||
raise ValueError(
|
||||
"candle_intervals are required for Candle replay"
|
||||
)
|
||||
|
||||
if not includes_candles and candle_intervals:
|
||||
raise ValueError(
|
||||
"candle_intervals require Candle revision data type"
|
||||
)
|
||||
|
||||
max_records = _normalize_positive_integer(
|
||||
self.max_records,
|
||||
field_name="max_records",
|
||||
maximum=REPLAY_PLAN_MAX_RECORDS_LIMIT,
|
||||
)
|
||||
|
||||
object.__setattr__(self, "venue", venue)
|
||||
object.__setattr__(self, "symbols", symbols)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"candle_intervals",
|
||||
candle_intervals,
|
||||
)
|
||||
object.__setattr__(self, "max_records", max_records)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReplayPlan:
|
||||
"""Полностью материализованный и DB-независимый Replay snapshot."""
|
||||
|
||||
request: ReplayPlanRequest
|
||||
events: tuple[ReplayEvent, ...]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if type(self.request) is not ReplayPlanRequest:
|
||||
raise TypeError("request must be ReplayPlanRequest")
|
||||
|
||||
if type(self.events) is not tuple:
|
||||
raise TypeError("events must be a tuple")
|
||||
|
||||
if len(self.events) > self.request.max_records:
|
||||
raise ValueError("events exceed request.max_records")
|
||||
|
||||
if any(type(event) is not ReplayEvent for event in self.events):
|
||||
raise TypeError("events must contain only ReplayEvent")
|
||||
|
||||
seen_sequences: set[int] = set()
|
||||
|
||||
for index, event in enumerate(self.events):
|
||||
if event.venue != self.request.venue:
|
||||
raise ValueError("event venue is outside Replay request")
|
||||
|
||||
if event.symbol not in self.request.symbols:
|
||||
raise ValueError("event symbol is outside Replay request")
|
||||
|
||||
if event.data_type not in self.request.data_types:
|
||||
raise ValueError("event data type is outside Replay request")
|
||||
|
||||
if not self.request.time_range.contains(event.replay_at):
|
||||
raise ValueError("event time is outside Replay request")
|
||||
|
||||
if isinstance(event.payload, Candle) and (
|
||||
event.payload.interval
|
||||
not in self.request.candle_intervals
|
||||
):
|
||||
raise ValueError(
|
||||
"Candle interval is outside Replay request"
|
||||
)
|
||||
|
||||
if event.replay_sequence in seen_sequences:
|
||||
raise ValueError(
|
||||
"replay_sequence must be globally unique in plan"
|
||||
)
|
||||
|
||||
seen_sequences.add(event.replay_sequence)
|
||||
|
||||
if index > 0 and (
|
||||
self.events[index - 1].order_key >= event.order_key
|
||||
):
|
||||
raise ValueError("events must be strictly ordered")
|
||||
|
||||
@property
|
||||
def is_empty(self) -> bool:
|
||||
return not self.events
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.events)
|
||||
436
app/src/market_data/replay/postgres_replay_plan_builder.py
Normal file
436
app/src/market_data/replay/postgres_replay_plan_builder.py
Normal file
@@ -0,0 +1,436 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessError,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
)
|
||||
from src.market_data.access.postgres_history_support import (
|
||||
PostgresHistoryConnectionProvider,
|
||||
candle_history_record_from_row,
|
||||
quote_history_record_from_row,
|
||||
trade_history_record_from_row,
|
||||
)
|
||||
from src.market_data.replay.exceptions import (
|
||||
MarketDataReplayValidationError,
|
||||
ReplayPlanLimitExceededError,
|
||||
)
|
||||
from src.market_data.replay.models import (
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
)
|
||||
|
||||
|
||||
_SET_SNAPSHOT_TRANSACTION_SQL = (
|
||||
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY"
|
||||
)
|
||||
|
||||
_REPLAY_SNAPSHOT_COLUMNS_SQL = """
|
||||
data_type,
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
trade_price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
observation_sources,
|
||||
canonical_schema_version
|
||||
"""
|
||||
|
||||
_TRADE_SNAPSHOT_BRANCH_SQL = """
|
||||
SELECT
|
||||
'trade'::TEXT AS data_type,
|
||||
executed_at AS replay_at,
|
||||
replay_sequence,
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price AS trade_price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
NULL::TIMESTAMPTZ AS received_at,
|
||||
NULL::TIMESTAMPTZ AS exchange_timestamp,
|
||||
NULL::NUMERIC AS last_price,
|
||||
NULL::NUMERIC AS bid_price,
|
||||
NULL::NUMERIC AS ask_price,
|
||||
NULL::TEXT AS interval,
|
||||
NULL::TIMESTAMPTZ AS open_time,
|
||||
NULL::TIMESTAMPTZ AS observed_at,
|
||||
NULL::NUMERIC AS open_price,
|
||||
NULL::NUMERIC AS high_price,
|
||||
NULL::NUMERIC AS low_price,
|
||||
NULL::NUMERIC AS close_price,
|
||||
NULL::NUMERIC AS volume,
|
||||
NULL::BOOLEAN AS is_final,
|
||||
observation_sources,
|
||||
canonical_schema_version
|
||||
FROM market_data.trades
|
||||
WHERE venue = %s
|
||||
AND symbol = ANY(%s)
|
||||
AND executed_at >= %s
|
||||
AND executed_at < %s
|
||||
"""
|
||||
|
||||
_QUOTE_SNAPSHOT_BRANCH_SQL = """
|
||||
SELECT
|
||||
'quote'::TEXT AS data_type,
|
||||
received_at AS replay_at,
|
||||
replay_sequence,
|
||||
venue,
|
||||
symbol,
|
||||
NULL::INTEGER AS trade_id,
|
||||
NULL::TIMESTAMPTZ AS executed_at,
|
||||
NULL::NUMERIC AS trade_price,
|
||||
NULL::NUMERIC AS quantity,
|
||||
NULL::TEXT AS aggressor_side,
|
||||
source,
|
||||
NULL::TIMESTAMPTZ AS first_observed_at,
|
||||
NULL::TIMESTAMPTZ AS last_observed_at,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
NULL::TEXT AS interval,
|
||||
NULL::TIMESTAMPTZ AS open_time,
|
||||
NULL::TIMESTAMPTZ AS observed_at,
|
||||
NULL::NUMERIC AS open_price,
|
||||
NULL::NUMERIC AS high_price,
|
||||
NULL::NUMERIC AS low_price,
|
||||
NULL::NUMERIC AS close_price,
|
||||
NULL::NUMERIC AS volume,
|
||||
NULL::BOOLEAN AS is_final,
|
||||
observation_sources,
|
||||
canonical_schema_version
|
||||
FROM market_data.quotes
|
||||
WHERE venue = %s
|
||||
AND symbol = ANY(%s)
|
||||
AND received_at >= %s
|
||||
AND received_at < %s
|
||||
"""
|
||||
|
||||
_CANDLE_SNAPSHOT_BRANCH_SQL = """
|
||||
SELECT
|
||||
'candle_revision'::TEXT AS data_type,
|
||||
observed_at AS replay_at,
|
||||
replay_sequence,
|
||||
venue,
|
||||
symbol,
|
||||
NULL::INTEGER AS trade_id,
|
||||
NULL::TIMESTAMPTZ AS executed_at,
|
||||
NULL::NUMERIC AS trade_price,
|
||||
NULL::NUMERIC AS quantity,
|
||||
NULL::TEXT AS aggressor_side,
|
||||
source,
|
||||
NULL::TIMESTAMPTZ AS first_observed_at,
|
||||
NULL::TIMESTAMPTZ AS last_observed_at,
|
||||
NULL::TIMESTAMPTZ AS received_at,
|
||||
NULL::TIMESTAMPTZ AS exchange_timestamp,
|
||||
NULL::NUMERIC AS last_price,
|
||||
NULL::NUMERIC AS bid_price,
|
||||
NULL::NUMERIC AS ask_price,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
observation_sources,
|
||||
canonical_schema_version
|
||||
FROM market_data.candle_revisions
|
||||
WHERE venue = %s
|
||||
AND symbol = ANY(%s)
|
||||
AND observed_at >= %s
|
||||
AND observed_at < %s
|
||||
AND interval = ANY(%s)
|
||||
"""
|
||||
|
||||
_SNAPSHOT_BRANCHES = {
|
||||
ReplayDataType.TRADE: _TRADE_SNAPSHOT_BRANCH_SQL,
|
||||
ReplayDataType.QUOTE: _QUOTE_SNAPSHOT_BRANCH_SQL,
|
||||
ReplayDataType.CANDLE_REVISION: _CANDLE_SNAPSHOT_BRANCH_SQL,
|
||||
}
|
||||
|
||||
_SNAPSHOT_DATA_TYPE_ORDER = (
|
||||
ReplayDataType.TRADE,
|
||||
ReplayDataType.QUOTE,
|
||||
ReplayDataType.CANDLE_REVISION,
|
||||
)
|
||||
|
||||
_SNAPSHOT_ROW_LENGTH = 29
|
||||
|
||||
|
||||
class PostgresReplayPlanBuilder:
|
||||
"""Создать один ограниченный снимок воспроизведения из PostgreSQL."""
|
||||
|
||||
__slots__ = ("_connection_provider",)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
connection_provider: PostgresHistoryConnectionProvider,
|
||||
) -> None:
|
||||
if not callable(connection_provider):
|
||||
raise TypeError("connection_provider must be callable")
|
||||
|
||||
self._connection_provider = connection_provider
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
"""Материализовать план в короткой транзакции только для чтения."""
|
||||
if type(request) is not ReplayPlanRequest:
|
||||
raise MarketDataReplayValidationError(
|
||||
"request must be ReplayPlanRequest"
|
||||
)
|
||||
|
||||
sql, parameters = self._build_snapshot_query(request)
|
||||
|
||||
try:
|
||||
with self._connection_provider() as connection:
|
||||
with connection.transaction():
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(_SET_SNAPSHOT_TRANSACTION_SQL)
|
||||
cursor.execute(sql, parameters)
|
||||
rows = cursor.fetchall()
|
||||
|
||||
return self._plan_from_rows(request=request, rows=rows)
|
||||
except (MarketDataAccessError, ReplayPlanLimitExceededError):
|
||||
raise
|
||||
except Exception as error:
|
||||
raise MarketDataAccessOperationError(
|
||||
"Failed to create PostgreSQL Replay snapshot."
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _build_snapshot_query(
|
||||
request: ReplayPlanRequest,
|
||||
) -> tuple[str, tuple[object, ...]]:
|
||||
requested_types = set(request.data_types)
|
||||
branches: list[str] = []
|
||||
parameters: list[object] = []
|
||||
|
||||
for data_type in _SNAPSHOT_DATA_TYPE_ORDER:
|
||||
if data_type not in requested_types:
|
||||
continue
|
||||
|
||||
branches.append(_SNAPSHOT_BRANCHES[data_type])
|
||||
parameters.extend(
|
||||
(
|
||||
request.venue,
|
||||
list(request.symbols),
|
||||
request.time_range.start_time,
|
||||
request.time_range.end_time,
|
||||
)
|
||||
)
|
||||
|
||||
if data_type is ReplayDataType.CANDLE_REVISION:
|
||||
parameters.append(list(request.candle_intervals))
|
||||
|
||||
union_sql = "\nUNION ALL\n".join(branches)
|
||||
sql = f"""
|
||||
SELECT
|
||||
{_REPLAY_SNAPSHOT_COLUMNS_SQL}
|
||||
FROM (
|
||||
{union_sql}
|
||||
) AS replay_snapshot
|
||||
ORDER BY replay_at ASC,
|
||||
replay_sequence ASC
|
||||
LIMIT %s
|
||||
"""
|
||||
parameters.append(request.max_records + 1)
|
||||
return sql, tuple(parameters)
|
||||
|
||||
@classmethod
|
||||
def _plan_from_rows(
|
||||
cls,
|
||||
*,
|
||||
request: ReplayPlanRequest,
|
||||
rows: object,
|
||||
) -> ReplayPlan:
|
||||
if not isinstance(rows, Sequence) or isinstance(
|
||||
rows,
|
||||
(str, bytes, bytearray),
|
||||
):
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Replay snapshot rows."
|
||||
)
|
||||
|
||||
if len(rows) > request.max_records + 1:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL exceeded the requested Replay snapshot limit."
|
||||
)
|
||||
|
||||
if len(rows) == request.max_records + 1:
|
||||
raise ReplayPlanLimitExceededError(
|
||||
"Replay snapshot exceeds request.max_records."
|
||||
)
|
||||
|
||||
events = tuple(cls._event_from_row(row) for row in rows)
|
||||
|
||||
try:
|
||||
return ReplayPlan(request=request, events=events)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned inconsistent Replay snapshot."
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _event_from_row(row: object) -> ReplayEvent:
|
||||
if type(row) is not tuple or len(row) != _SNAPSHOT_ROW_LENGTH:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Replay snapshot row."
|
||||
)
|
||||
|
||||
(
|
||||
data_type,
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
trade_price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
observation_sources,
|
||||
canonical_schema_version,
|
||||
) = row
|
||||
|
||||
if type(data_type) is not str:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned invalid Replay data type."
|
||||
)
|
||||
|
||||
if data_type == ReplayDataType.TRADE.value:
|
||||
record = trade_history_record_from_row(
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
trade_price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
)
|
||||
payload = record.trade
|
||||
expected_replay_at = record.event_time
|
||||
candle_is_final = None
|
||||
elif data_type == ReplayDataType.QUOTE.value:
|
||||
record = quote_history_record_from_row(
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
source,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
)
|
||||
payload = record.quote
|
||||
expected_replay_at = record.event_time
|
||||
candle_is_final = None
|
||||
elif data_type == ReplayDataType.CANDLE_REVISION.value:
|
||||
record = candle_history_record_from_row(
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
source,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
)
|
||||
payload = record.candle
|
||||
expected_replay_at = record.observed_at
|
||||
candle_is_final = record.is_final
|
||||
else:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned unknown Replay data type."
|
||||
)
|
||||
|
||||
try:
|
||||
if replay_at != expected_replay_at:
|
||||
raise ValueError("Replay event time differs from payload")
|
||||
|
||||
if replay_sequence != record.replay_sequence:
|
||||
raise ValueError("Replay sequence differs from payload")
|
||||
|
||||
return ReplayEvent(
|
||||
venue=record.venue,
|
||||
replay_at=expected_replay_at,
|
||||
replay_sequence=record.replay_sequence,
|
||||
payload=payload,
|
||||
candle_is_final=candle_is_final,
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise MarketDataAccessIntegrityError(
|
||||
"PostgreSQL returned inconsistent Replay event."
|
||||
) from error
|
||||
124
app/src/market_data/replay/replay_session.py
Normal file
124
app/src/market_data/replay/replay_session.py
Normal file
@@ -0,0 +1,124 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from src.market_data.replay.contracts import (
|
||||
MarketDataClockProtocol,
|
||||
ReplayClockProtocol,
|
||||
ReplayConsumerProtocol,
|
||||
)
|
||||
from src.market_data.replay.exceptions import (
|
||||
MarketDataReplayValidationError,
|
||||
ReplaySessionStateError,
|
||||
)
|
||||
from src.market_data.replay.models import (
|
||||
ReplayPlan,
|
||||
ReplaySessionState,
|
||||
)
|
||||
|
||||
|
||||
class ReplaySession:
|
||||
"""
|
||||
Одноразовая последовательная Session детерминированного Replay.
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"_clock",
|
||||
"_consumer",
|
||||
"_plan",
|
||||
"_state",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: ReplayClockProtocol,
|
||||
consumer: ReplayConsumerProtocol,
|
||||
) -> None:
|
||||
if type(plan) is not ReplayPlan:
|
||||
raise TypeError("plan must be a ReplayPlan")
|
||||
|
||||
if isinstance(clock, type) or not isinstance(
|
||||
clock,
|
||||
ReplayClockProtocol,
|
||||
):
|
||||
raise TypeError("clock must implement ReplayClockProtocol")
|
||||
|
||||
if (
|
||||
not callable(clock.advance_to)
|
||||
or inspect.iscoroutinefunction(clock.advance_to)
|
||||
):
|
||||
raise TypeError("clock.advance_to must be synchronous")
|
||||
|
||||
if isinstance(consumer, type) or not isinstance(
|
||||
consumer,
|
||||
ReplayConsumerProtocol,
|
||||
):
|
||||
raise TypeError(
|
||||
"consumer must implement ReplayConsumerProtocol"
|
||||
)
|
||||
|
||||
if not inspect.iscoroutinefunction(consumer.consume):
|
||||
raise TypeError("consumer.consume must be asynchronous")
|
||||
|
||||
clock_now = clock.now
|
||||
|
||||
if (
|
||||
type(clock_now) is not datetime
|
||||
or clock_now.tzinfo is not timezone.utc
|
||||
or clock_now.fold != 0
|
||||
):
|
||||
raise MarketDataReplayValidationError(
|
||||
"Replay Session Clock must contain canonical UTC time."
|
||||
)
|
||||
|
||||
if clock_now != plan.request.time_range.start_time:
|
||||
raise MarketDataReplayValidationError(
|
||||
"Replay Session Clock must start at the Replay request "
|
||||
"start time."
|
||||
)
|
||||
|
||||
self._clock = clock
|
||||
self._consumer = consumer
|
||||
self._plan = plan
|
||||
self._state = ReplaySessionState.CREATED
|
||||
|
||||
@property
|
||||
def state(self) -> ReplaySessionState:
|
||||
"""Вернуть текущее необратимое состояние Session."""
|
||||
return self._state
|
||||
|
||||
@property
|
||||
def plan(self) -> ReplayPlan:
|
||||
"""Вернуть неизменяемый Replay Plan этой Session."""
|
||||
return self._plan
|
||||
|
||||
@property
|
||||
def clock(self) -> MarketDataClockProtocol:
|
||||
"""Предоставить Clock через границу только для чтения."""
|
||||
return self._clock
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Последовательно воспроизвести Plan ровно один раз."""
|
||||
if self._state is not ReplaySessionState.CREATED:
|
||||
raise ReplaySessionStateError(
|
||||
"Replay Session can only run from CREATED state."
|
||||
)
|
||||
|
||||
self._state = ReplaySessionState.RUNNING
|
||||
|
||||
try:
|
||||
for event in self._plan.events:
|
||||
self._clock.advance_to(event.replay_at)
|
||||
await self._consumer.consume(event)
|
||||
except asyncio.CancelledError:
|
||||
self._state = ReplaySessionState.CANCELLED
|
||||
raise
|
||||
except BaseException:
|
||||
self._state = ReplaySessionState.FAILED
|
||||
raise
|
||||
|
||||
self._state = ReplaySessionState.COMPLETED
|
||||
106
app/src/market_data/replay/replay_session_factory.py
Normal file
106
app/src/market_data/replay/replay_session_factory.py
Normal file
@@ -0,0 +1,106 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
|
||||
from src.market_data.replay.contracts import (
|
||||
ReplayConsumerFactoryProtocol,
|
||||
ReplayPlanBuilderProtocol,
|
||||
)
|
||||
from src.market_data.replay.deterministic_replay_clock import (
|
||||
DeterministicReplayClock,
|
||||
)
|
||||
from src.market_data.replay.exceptions import (
|
||||
MarketDataReplayValidationError,
|
||||
)
|
||||
from src.market_data.replay.models import (
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
)
|
||||
from src.market_data.replay.replay_session import ReplaySession
|
||||
|
||||
|
||||
class ReplaySessionFactory:
|
||||
"""Явно подготовить независимый сеанс Replay без его запуска."""
|
||||
|
||||
__slots__ = (
|
||||
"_consumer_factory",
|
||||
"_plan_builder",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
plan_builder: ReplayPlanBuilderProtocol,
|
||||
consumer_factory: ReplayConsumerFactoryProtocol,
|
||||
) -> None:
|
||||
if isinstance(plan_builder, type) or not isinstance(
|
||||
plan_builder,
|
||||
ReplayPlanBuilderProtocol,
|
||||
):
|
||||
raise TypeError(
|
||||
"plan_builder must implement ReplayPlanBuilderProtocol"
|
||||
)
|
||||
|
||||
if (
|
||||
not callable(plan_builder.create_plan)
|
||||
or inspect.iscoroutinefunction(plan_builder.create_plan)
|
||||
):
|
||||
raise TypeError("plan_builder.create_plan must be synchronous")
|
||||
|
||||
if isinstance(consumer_factory, type) or not isinstance(
|
||||
consumer_factory,
|
||||
ReplayConsumerFactoryProtocol,
|
||||
):
|
||||
raise TypeError(
|
||||
"consumer_factory must implement "
|
||||
"ReplayConsumerFactoryProtocol"
|
||||
)
|
||||
|
||||
if (
|
||||
not callable(consumer_factory.create_consumer)
|
||||
or inspect.iscoroutinefunction(
|
||||
consumer_factory.create_consumer
|
||||
)
|
||||
):
|
||||
raise TypeError(
|
||||
"consumer_factory.create_consumer must be synchronous"
|
||||
)
|
||||
|
||||
self._plan_builder = plan_builder
|
||||
self._consumer_factory = consumer_factory
|
||||
|
||||
def prepare_session(
|
||||
self,
|
||||
request: ReplayPlanRequest,
|
||||
) -> ReplaySession:
|
||||
"""Синхронно собрать новый граф зависимостей одного сеанса."""
|
||||
if type(request) is not ReplayPlanRequest:
|
||||
raise MarketDataReplayValidationError(
|
||||
"request must be ReplayPlanRequest"
|
||||
)
|
||||
|
||||
plan = self._plan_builder.create_plan(request)
|
||||
|
||||
if type(plan) is not ReplayPlan:
|
||||
raise MarketDataReplayValidationError(
|
||||
"plan_builder must return an exact ReplayPlan"
|
||||
)
|
||||
|
||||
if plan.request is not request:
|
||||
raise MarketDataReplayValidationError(
|
||||
"Replay Plan must preserve Replay request identity"
|
||||
)
|
||||
|
||||
clock = DeterministicReplayClock(
|
||||
plan.request.time_range.start_time
|
||||
)
|
||||
consumer = self._consumer_factory.create_consumer(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
)
|
||||
|
||||
return ReplaySession(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
@@ -17,9 +17,9 @@ from src.market_data.storage.postgres_repository_support import (
|
||||
PostgresRepositoryConnectionProvider,
|
||||
normalize_aware_datetime,
|
||||
)
|
||||
from src.storage.migrations import MARKET_DATA_PARTITION_ADVISORY_LOCK_ID
|
||||
|
||||
|
||||
MARKET_DATA_PARTITION_ADVISORY_LOCK_ID = 0x445A504152544E
|
||||
_SCHEMA_NAME = "market_data"
|
||||
_TRADE_CHECKPOINT_FK_NAME = "trade_stream_checkpoints_trade_fk"
|
||||
|
||||
|
||||
@@ -61,6 +61,35 @@ ON CONFLICT (venue, symbol, trade_id, executed_at) DO NOTHING
|
||||
RETURNING 1
|
||||
"""
|
||||
|
||||
_INSERT_TRADE_WITH_REPLAY_SEQUENCE_SQL = """
|
||||
INSERT INTO market_data.trades (
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version
|
||||
)
|
||||
VALUES (
|
||||
%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s
|
||||
)
|
||||
ON CONFLICT (venue, symbol, trade_id, executed_at) DO NOTHING
|
||||
RETURNING 1
|
||||
"""
|
||||
|
||||
_ALLOCATE_REPLAY_SEQUENCES_SQL = """
|
||||
SELECT nextval('market_data.replay_sequence'::regclass)
|
||||
FROM generate_series(1, %s)
|
||||
ORDER BY 1
|
||||
"""
|
||||
|
||||
_SELECT_TRADE_FOR_UPDATE_SQL = """
|
||||
SELECT
|
||||
price,
|
||||
@@ -632,14 +661,27 @@ class PostgresTradeRepository:
|
||||
try:
|
||||
with self._connection_provider() as connection:
|
||||
with connection.cursor() as cursor:
|
||||
for prepared in sorted(
|
||||
prepared_trades,
|
||||
key=lambda item: item.identity_order_key,
|
||||
replay_sequences = self._allocate_replay_sequences(
|
||||
cursor=cursor,
|
||||
count=len(prepared_trades),
|
||||
)
|
||||
prepared_with_sequences = tuple(
|
||||
zip(
|
||||
prepared_trades,
|
||||
replay_sequences,
|
||||
strict=True,
|
||||
)
|
||||
)
|
||||
|
||||
for prepared, replay_sequence in sorted(
|
||||
prepared_with_sequences,
|
||||
key=lambda item: item[0].identity_order_key,
|
||||
):
|
||||
status = self._store_prepared_trade(
|
||||
cursor=cursor,
|
||||
venue=normalized_venue,
|
||||
trade=prepared,
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
if status is MarketDataWriteStatus.INSERTED:
|
||||
@@ -1122,25 +1164,34 @@ class PostgresTradeRepository:
|
||||
cursor: Any,
|
||||
venue: str,
|
||||
trade: _PreparedTrade,
|
||||
replay_sequence: int | None = None,
|
||||
) -> MarketDataWriteStatus:
|
||||
cursor.execute(
|
||||
_INSERT_TRADE_SQL,
|
||||
(
|
||||
venue,
|
||||
trade.symbol,
|
||||
trade.trade_id,
|
||||
trade.executed_at,
|
||||
trade.price,
|
||||
trade.quantity,
|
||||
trade.aggressor_side,
|
||||
trade.source,
|
||||
trade.observed_at,
|
||||
trade.observed_at,
|
||||
[trade.source],
|
||||
CANONICAL_TRADE_SCHEMA_VERSION,
|
||||
),
|
||||
parameters = (
|
||||
venue,
|
||||
trade.symbol,
|
||||
trade.trade_id,
|
||||
trade.executed_at,
|
||||
trade.price,
|
||||
trade.quantity,
|
||||
trade.aggressor_side,
|
||||
trade.source,
|
||||
trade.observed_at,
|
||||
trade.observed_at,
|
||||
[trade.source],
|
||||
)
|
||||
|
||||
if replay_sequence is None:
|
||||
statement = _INSERT_TRADE_SQL
|
||||
parameters += (CANONICAL_TRADE_SCHEMA_VERSION,)
|
||||
else:
|
||||
statement = _INSERT_TRADE_WITH_REPLAY_SEQUENCE_SQL
|
||||
parameters += (
|
||||
replay_sequence,
|
||||
CANONICAL_TRADE_SCHEMA_VERSION,
|
||||
)
|
||||
|
||||
cursor.execute(statement, parameters)
|
||||
|
||||
if cursor.fetchone() is not None:
|
||||
return MarketDataWriteStatus.INSERTED
|
||||
|
||||
@@ -1216,6 +1267,53 @@ class PostgresTradeRepository:
|
||||
|
||||
return MarketDataWriteStatus.PROVENANCE_UPDATED
|
||||
|
||||
@staticmethod
|
||||
def _allocate_replay_sequences(
|
||||
*,
|
||||
cursor: Any,
|
||||
count: int,
|
||||
) -> tuple[int, ...]:
|
||||
cursor.execute(
|
||||
_ALLOCATE_REPLAY_SEQUENCES_SQL,
|
||||
(count,),
|
||||
)
|
||||
rows = cursor.fetchall()
|
||||
|
||||
if not isinstance(rows, list | tuple) or len(rows) != count:
|
||||
raise MarketDataStorageOperationError(
|
||||
"PostgreSQL returned invalid Replay sequence allocation."
|
||||
)
|
||||
|
||||
sequences: list[int] = []
|
||||
|
||||
for row in rows:
|
||||
if (
|
||||
not isinstance(row, tuple)
|
||||
or len(row) != 1
|
||||
or isinstance(row[0], bool)
|
||||
or not isinstance(row[0], int)
|
||||
or row[0] <= 0
|
||||
):
|
||||
raise MarketDataStorageOperationError(
|
||||
"PostgreSQL returned invalid Replay sequence value."
|
||||
)
|
||||
|
||||
sequences.append(row[0])
|
||||
|
||||
if any(
|
||||
current <= previous
|
||||
for previous, current in zip(
|
||||
sequences,
|
||||
sequences[1:],
|
||||
strict=False,
|
||||
)
|
||||
):
|
||||
raise MarketDataStorageOperationError(
|
||||
"PostgreSQL returned unordered Replay sequence values."
|
||||
)
|
||||
|
||||
return tuple(sequences)
|
||||
|
||||
@staticmethod
|
||||
def _identity_parameters(
|
||||
*,
|
||||
|
||||
@@ -9,6 +9,7 @@ from src.storage.exceptions import StorageMigrationError
|
||||
|
||||
|
||||
STORAGE_MIGRATION_ADVISORY_LOCK_ID = 0x445A454E545241
|
||||
MARKET_DATA_PARTITION_ADVISORY_LOCK_ID = 0x445A504152544E
|
||||
|
||||
_CREATE_HISTORY_TABLE_SQL = """
|
||||
CREATE TABLE IF NOT EXISTS public.storage_schema_migrations (
|
||||
@@ -330,6 +331,292 @@ STORAGE_MIGRATIONS = (
|
||||
""",
|
||||
),
|
||||
),
|
||||
StorageMigration(
|
||||
version=9,
|
||||
name="add_global_market_data_replay_sequence",
|
||||
statements=(
|
||||
(
|
||||
"SELECT pg_advisory_xact_lock("
|
||||
f"{MARKET_DATA_PARTITION_ADVISORY_LOCK_ID}"
|
||||
")"
|
||||
),
|
||||
"LOCK TABLE market_data.trades IN ACCESS EXCLUSIVE MODE",
|
||||
"LOCK TABLE market_data.quotes IN ACCESS EXCLUSIVE MODE",
|
||||
(
|
||||
"LOCK TABLE market_data.candle_revisions "
|
||||
"IN ACCESS EXCLUSIVE MODE"
|
||||
),
|
||||
"""
|
||||
CREATE SEQUENCE market_data.replay_sequence
|
||||
AS BIGINT
|
||||
INCREMENT BY 1
|
||||
MINVALUE 1
|
||||
START WITH 1
|
||||
CACHE 1
|
||||
NO CYCLE
|
||||
OWNED BY NONE
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.trades
|
||||
ADD COLUMN replay_sequence BIGINT
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.quotes
|
||||
ADD COLUMN replay_sequence BIGINT
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.candle_revisions
|
||||
ADD COLUMN replay_sequence BIGINT
|
||||
""",
|
||||
"""
|
||||
CREATE TEMPORARY TABLE market_data_replay_sequence_backfill
|
||||
ON COMMIT DROP
|
||||
AS
|
||||
SELECT
|
||||
data_type_rank,
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
received_at,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
ROW_NUMBER() OVER (
|
||||
ORDER BY
|
||||
event_time,
|
||||
data_type_rank,
|
||||
venue COLLATE "C",
|
||||
symbol COLLATE "C",
|
||||
trade_executed_at,
|
||||
trade_id,
|
||||
quote_received_at,
|
||||
candle_interval COLLATE "C",
|
||||
candle_open_time,
|
||||
candle_observed_at
|
||||
) AS replay_sequence
|
||||
FROM (
|
||||
SELECT
|
||||
1::SMALLINT AS data_type_rank,
|
||||
executed_at AS event_time,
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
NULL::TIMESTAMPTZ AS received_at,
|
||||
NULL::TEXT AS interval,
|
||||
NULL::TIMESTAMPTZ AS open_time,
|
||||
NULL::TIMESTAMPTZ AS observed_at,
|
||||
executed_at AS trade_executed_at,
|
||||
NULL::TIMESTAMPTZ AS quote_received_at,
|
||||
NULL::TEXT AS candle_interval,
|
||||
NULL::TIMESTAMPTZ AS candle_open_time,
|
||||
NULL::TIMESTAMPTZ AS candle_observed_at
|
||||
FROM market_data.trades
|
||||
|
||||
UNION ALL
|
||||
|
||||
SELECT
|
||||
2::SMALLINT AS data_type_rank,
|
||||
received_at AS event_time,
|
||||
venue,
|
||||
symbol,
|
||||
NULL::INTEGER AS trade_id,
|
||||
NULL::TIMESTAMPTZ AS executed_at,
|
||||
received_at,
|
||||
NULL::TEXT AS interval,
|
||||
NULL::TIMESTAMPTZ AS open_time,
|
||||
NULL::TIMESTAMPTZ AS observed_at,
|
||||
NULL::TIMESTAMPTZ AS trade_executed_at,
|
||||
received_at AS quote_received_at,
|
||||
NULL::TEXT AS candle_interval,
|
||||
NULL::TIMESTAMPTZ AS candle_open_time,
|
||||
NULL::TIMESTAMPTZ AS candle_observed_at
|
||||
FROM market_data.quotes
|
||||
|
||||
UNION ALL
|
||||
|
||||
SELECT
|
||||
3::SMALLINT AS data_type_rank,
|
||||
observed_at AS event_time,
|
||||
venue,
|
||||
symbol,
|
||||
NULL::INTEGER AS trade_id,
|
||||
NULL::TIMESTAMPTZ AS executed_at,
|
||||
NULL::TIMESTAMPTZ AS received_at,
|
||||
interval,
|
||||
open_time,
|
||||
observed_at,
|
||||
NULL::TIMESTAMPTZ AS trade_executed_at,
|
||||
NULL::TIMESTAMPTZ AS quote_received_at,
|
||||
interval AS candle_interval,
|
||||
open_time AS candle_open_time,
|
||||
observed_at AS candle_observed_at
|
||||
FROM market_data.candle_revisions
|
||||
) AS durable_rows
|
||||
""",
|
||||
"""
|
||||
UPDATE market_data.trades AS target
|
||||
SET replay_sequence = backfill.replay_sequence
|
||||
FROM market_data_replay_sequence_backfill AS backfill
|
||||
WHERE backfill.data_type_rank = 1
|
||||
AND target.venue = backfill.venue
|
||||
AND target.symbol = backfill.symbol
|
||||
AND target.trade_id = backfill.trade_id
|
||||
AND target.executed_at = backfill.executed_at
|
||||
""",
|
||||
"""
|
||||
UPDATE market_data.quotes AS target
|
||||
SET replay_sequence = backfill.replay_sequence
|
||||
FROM market_data_replay_sequence_backfill AS backfill
|
||||
WHERE backfill.data_type_rank = 2
|
||||
AND target.venue = backfill.venue
|
||||
AND target.symbol = backfill.symbol
|
||||
AND target.received_at = backfill.received_at
|
||||
""",
|
||||
"""
|
||||
UPDATE market_data.candle_revisions AS target
|
||||
SET replay_sequence = backfill.replay_sequence
|
||||
FROM market_data_replay_sequence_backfill AS backfill
|
||||
WHERE backfill.data_type_rank = 3
|
||||
AND target.venue = backfill.venue
|
||||
AND target.symbol = backfill.symbol
|
||||
AND target.interval = backfill.interval
|
||||
AND target.open_time = backfill.open_time
|
||||
AND target.observed_at = backfill.observed_at
|
||||
""",
|
||||
"""
|
||||
SELECT pg_catalog.setval(
|
||||
'market_data.replay_sequence'::REGCLASS,
|
||||
COALESCE(
|
||||
(
|
||||
SELECT MAX(replay_sequence)
|
||||
FROM market_data_replay_sequence_backfill
|
||||
),
|
||||
1
|
||||
),
|
||||
EXISTS (
|
||||
SELECT 1
|
||||
FROM market_data_replay_sequence_backfill
|
||||
)
|
||||
)
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.trades
|
||||
ALTER COLUMN replay_sequence
|
||||
SET DEFAULT nextval(
|
||||
'market_data.replay_sequence'::REGCLASS
|
||||
)
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.quotes
|
||||
ALTER COLUMN replay_sequence
|
||||
SET DEFAULT nextval(
|
||||
'market_data.replay_sequence'::REGCLASS
|
||||
)
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.candle_revisions
|
||||
ALTER COLUMN replay_sequence
|
||||
SET DEFAULT nextval(
|
||||
'market_data.replay_sequence'::REGCLASS
|
||||
)
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.trades
|
||||
ALTER COLUMN replay_sequence SET NOT NULL,
|
||||
ADD CONSTRAINT trades_replay_sequence_positive
|
||||
CHECK (replay_sequence > 0)
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.quotes
|
||||
ALTER COLUMN replay_sequence SET NOT NULL,
|
||||
ADD CONSTRAINT quotes_replay_sequence_positive
|
||||
CHECK (replay_sequence > 0)
|
||||
""",
|
||||
"""
|
||||
ALTER TABLE market_data.candle_revisions
|
||||
ALTER COLUMN replay_sequence SET NOT NULL,
|
||||
ADD CONSTRAINT candle_revisions_replay_sequence_positive
|
||||
CHECK (replay_sequence > 0)
|
||||
""",
|
||||
"""
|
||||
CREATE FUNCTION market_data.reject_replay_sequence_change()
|
||||
RETURNS TRIGGER
|
||||
LANGUAGE plpgsql
|
||||
AS $function$
|
||||
BEGIN
|
||||
IF NEW.replay_sequence IS DISTINCT FROM OLD.replay_sequence THEN
|
||||
RAISE EXCEPTION USING
|
||||
ERRCODE = '23514',
|
||||
MESSAGE = 'market_data replay_sequence is immutable';
|
||||
END IF;
|
||||
|
||||
RETURN NEW;
|
||||
END;
|
||||
$function$
|
||||
""",
|
||||
"""
|
||||
CREATE TRIGGER trades_replay_sequence_immutable
|
||||
BEFORE UPDATE OF replay_sequence
|
||||
ON market_data.trades
|
||||
FOR EACH ROW
|
||||
EXECUTE FUNCTION market_data.reject_replay_sequence_change()
|
||||
""",
|
||||
"""
|
||||
CREATE TRIGGER quotes_replay_sequence_immutable
|
||||
BEFORE UPDATE OF replay_sequence
|
||||
ON market_data.quotes
|
||||
FOR EACH ROW
|
||||
EXECUTE FUNCTION market_data.reject_replay_sequence_change()
|
||||
""",
|
||||
"""
|
||||
CREATE TRIGGER candle_revisions_replay_sequence_immutable
|
||||
BEFORE UPDATE OF replay_sequence
|
||||
ON market_data.candle_revisions
|
||||
FOR EACH ROW
|
||||
EXECUTE FUNCTION market_data.reject_replay_sequence_change()
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX trades_history_keyset_idx
|
||||
ON market_data.trades (
|
||||
venue,
|
||||
symbol,
|
||||
executed_at,
|
||||
replay_sequence
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX quotes_history_keyset_idx
|
||||
ON market_data.quotes (
|
||||
venue,
|
||||
symbol,
|
||||
received_at,
|
||||
replay_sequence
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX candle_revisions_history_keyset_idx
|
||||
ON market_data.candle_revisions (
|
||||
venue,
|
||||
symbol,
|
||||
interval,
|
||||
open_time,
|
||||
replay_sequence
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX candle_revisions_replay_keyset_idx
|
||||
ON market_data.candle_revisions (
|
||||
venue,
|
||||
symbol,
|
||||
interval,
|
||||
observed_at,
|
||||
replay_sequence
|
||||
)
|
||||
""",
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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())
|
||||
34
app/tests/static/test_python_type_gate.py
Normal file
34
app/tests/static/test_python_type_gate.py
Normal file
@@ -0,0 +1,34 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[3]
|
||||
|
||||
|
||||
def test_python_type_gate_is_clean() -> None:
|
||||
result = subprocess.run(
|
||||
(
|
||||
sys.executable,
|
||||
"-m",
|
||||
"pyright",
|
||||
"--project",
|
||||
str(PROJECT_ROOT / "pyrightconfig.json"),
|
||||
"--warnings",
|
||||
),
|
||||
cwd=PROJECT_ROOT,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False,
|
||||
timeout=30.0,
|
||||
)
|
||||
|
||||
output = "\n".join(
|
||||
part.strip()
|
||||
for part in (result.stdout, result.stderr)
|
||||
if part.strip()
|
||||
)
|
||||
|
||||
assert result.returncode == 0, output
|
||||
@@ -0,0 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from src.market_data.access import (
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
MarketDataAccessError,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
MarketDataCursorError,
|
||||
MarketDataHistoricalAccessProtocol,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
QuoteHistoryReaderProtocol,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryReaderProtocol,
|
||||
)
|
||||
|
||||
|
||||
class RecordingHistoricalAccess:
|
||||
def query_trades(
|
||||
self,
|
||||
query: TradeHistoryQuery,
|
||||
) -> TradeHistoryPage:
|
||||
return TradeHistoryPage(query=query, items=())
|
||||
|
||||
def query_quotes(
|
||||
self,
|
||||
query: QuoteHistoryQuery,
|
||||
) -> QuoteHistoryPage:
|
||||
return QuoteHistoryPage(query=query, items=())
|
||||
|
||||
def query_candle_revisions(
|
||||
self,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
return CandleRevisionHistoryPage(query=query, items=())
|
||||
|
||||
|
||||
def test_history_reader_protocols_are_runtime_checkable() -> None:
|
||||
access = RecordingHistoricalAccess()
|
||||
|
||||
assert isinstance(access, TradeHistoryReaderProtocol)
|
||||
assert isinstance(access, QuoteHistoryReaderProtocol)
|
||||
assert isinstance(access, CandleRevisionHistoryReaderProtocol)
|
||||
assert isinstance(access, MarketDataHistoricalAccessProtocol)
|
||||
|
||||
|
||||
def test_access_error_hierarchy_is_specialized() -> None:
|
||||
assert issubclass(MarketDataAccessValidationError, MarketDataAccessError)
|
||||
assert issubclass(MarketDataCursorError, MarketDataAccessValidationError)
|
||||
assert issubclass(MarketDataAccessIntegrityError, MarketDataAccessError)
|
||||
assert issubclass(MarketDataAccessOperationError, MarketDataAccessError)
|
||||
@@ -0,0 +1,882 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import FrozenInstanceError, replace
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import (
|
||||
HISTORY_PAGE_LIMIT_MAX,
|
||||
CandleRevisionHistoryCursor,
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
CandleRevisionHistoryRecord,
|
||||
HistoricalTimeRange,
|
||||
QuoteHistoryCursor,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
QuoteHistoryRecord,
|
||||
TradeHistoryCursor,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryRecord,
|
||||
)
|
||||
from src.market_data.acquisition.models.candle import Candle
|
||||
from src.market_data.acquisition.models.quote import Quote
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 10, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_trade"
|
||||
|
||||
|
||||
def make_trade(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
executed_at: datetime = START + timedelta(minutes=1),
|
||||
symbol: str = SYMBOL,
|
||||
source: str = SOURCE,
|
||||
) -> Trade:
|
||||
return Trade(
|
||||
symbol=symbol,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000.25"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=executed_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source=source,
|
||||
)
|
||||
|
||||
|
||||
def make_quote(
|
||||
*,
|
||||
received_at: datetime = START + timedelta(minutes=2),
|
||||
) -> Quote:
|
||||
return Quote(
|
||||
symbol=SYMBOL,
|
||||
last_price=Decimal("65000"),
|
||||
bid_price=Decimal("64999"),
|
||||
ask_price=Decimal("65001"),
|
||||
exchange_timestamp=received_at - timedelta(milliseconds=1),
|
||||
received_at=received_at,
|
||||
source="dzengi_rest_quote",
|
||||
)
|
||||
|
||||
|
||||
def make_candle(
|
||||
*,
|
||||
open_time: datetime = START,
|
||||
interval: str = "1m",
|
||||
) -> Candle:
|
||||
return Candle(
|
||||
symbol=SYMBOL,
|
||||
interval=interval,
|
||||
open_time=open_time,
|
||||
open_price=Decimal("64900"),
|
||||
high_price=Decimal("65100"),
|
||||
low_price=Decimal("64800"),
|
||||
close_price=Decimal("65000"),
|
||||
volume=Decimal("12.5"),
|
||||
source="dzengi_rest_candle",
|
||||
)
|
||||
|
||||
|
||||
def make_range() -> HistoricalTimeRange:
|
||||
return HistoricalTimeRange(start_time=START, end_time=END)
|
||||
|
||||
|
||||
def make_trade_query(
|
||||
*,
|
||||
cursor: TradeHistoryCursor | None = None,
|
||||
limit: int = 500,
|
||||
) -> TradeHistoryQuery:
|
||||
return TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_record(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
executed_at: datetime = START + timedelta(minutes=1),
|
||||
replay_sequence: int = 10,
|
||||
) -> TradeHistoryRecord:
|
||||
return TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=make_trade(
|
||||
trade_id=trade_id,
|
||||
executed_at=executed_at,
|
||||
),
|
||||
first_observed_at=executed_at + timedelta(seconds=1),
|
||||
last_observed_at=executed_at + timedelta(seconds=2),
|
||||
observation_sources=(SOURCE,),
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
|
||||
def make_quote_record(
|
||||
*,
|
||||
received_at: datetime = START + timedelta(minutes=2),
|
||||
replay_sequence: int = 20,
|
||||
) -> QuoteHistoryRecord:
|
||||
quote = make_quote(received_at=received_at)
|
||||
return QuoteHistoryRecord(
|
||||
venue=VENUE,
|
||||
quote=quote,
|
||||
observation_sources=(quote.source,),
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
|
||||
def make_candle_record(
|
||||
*,
|
||||
open_time: datetime = START,
|
||||
observed_at: datetime = START + timedelta(seconds=30),
|
||||
replay_sequence: int = 30,
|
||||
interval: str = "1m",
|
||||
) -> CandleRevisionHistoryRecord:
|
||||
candle = make_candle(open_time=open_time, interval=interval)
|
||||
return CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle,
|
||||
observed_at=observed_at,
|
||||
is_final=False,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=replay_sequence,
|
||||
)
|
||||
|
||||
|
||||
def test_time_range_is_half_open_and_normalized_to_utc() -> None:
|
||||
offset = timezone(timedelta(hours=3))
|
||||
time_range = HistoricalTimeRange(
|
||||
start_time=START.astimezone(offset),
|
||||
end_time=END.astimezone(offset),
|
||||
)
|
||||
|
||||
assert time_range.start_time == START
|
||||
assert time_range.start_time.tzinfo is timezone.utc
|
||||
assert time_range.end_time == END
|
||||
assert time_range.contains(START)
|
||||
assert time_range.contains(END - timedelta(microseconds=1))
|
||||
assert not time_range.contains(END)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("start_time", "end_time", "error_type"),
|
||||
(
|
||||
(datetime(2026, 8, 2, 10, 0), END, ValueError),
|
||||
(START, datetime(2026, 8, 2, 11, 0), ValueError),
|
||||
(START, START, ValueError),
|
||||
(END, START, ValueError),
|
||||
("2026-08-02", END, TypeError),
|
||||
),
|
||||
)
|
||||
def test_time_range_rejects_invalid_boundaries(
|
||||
start_time: Any,
|
||||
end_time: Any,
|
||||
error_type: type[Exception],
|
||||
) -> None:
|
||||
with pytest.raises(error_type):
|
||||
HistoricalTimeRange(
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
def test_history_records_preserve_exact_canonical_payloads() -> None:
|
||||
trade = make_trade()
|
||||
quote = make_quote()
|
||||
candle = make_candle()
|
||||
|
||||
trade_record = TradeHistoryRecord(
|
||||
venue=" dzengi ",
|
||||
trade=trade,
|
||||
first_observed_at=trade.executed_at,
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=(trade.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
quote_record = QuoteHistoryRecord(
|
||||
venue=VENUE,
|
||||
quote=quote,
|
||||
observation_sources=(quote.source,),
|
||||
replay_sequence=2,
|
||||
)
|
||||
candle_record = CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle,
|
||||
observed_at=candle.open_time,
|
||||
is_final=True,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=3,
|
||||
)
|
||||
|
||||
assert trade_record.venue == VENUE
|
||||
assert trade_record.trade is trade
|
||||
assert quote_record.quote is quote
|
||||
assert candle_record.candle is candle
|
||||
assert trade_record.event_time == trade.executed_at
|
||||
assert quote_record.event_time == quote.received_at
|
||||
assert candle_record.event_time == candle.open_time
|
||||
assert candle_record.replay_at == candle.open_time
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid_sequence", (True, 0, -1, 1.5, "1"))
|
||||
def test_records_reject_invalid_replay_sequence(
|
||||
invalid_sequence: Any,
|
||||
) -> None:
|
||||
with pytest.raises((TypeError, ValueError), match="replay_sequence"):
|
||||
make_trade_record(replay_sequence=invalid_sequence)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"observation_sources",
|
||||
(
|
||||
[],
|
||||
(),
|
||||
("",),
|
||||
(SOURCE, SOURCE),
|
||||
("another_source",),
|
||||
),
|
||||
)
|
||||
def test_trade_record_rejects_invalid_provenance(
|
||||
observation_sources: Any,
|
||||
) -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=trade,
|
||||
first_observed_at=trade.executed_at,
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=observation_sources,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_record_rejects_reversed_observation_times() -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises(ValueError, match="last_observed_at"):
|
||||
TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=trade,
|
||||
first_observed_at=trade.executed_at + timedelta(seconds=1),
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=(trade.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_candle_record_uses_open_time_for_history_and_observed_for_replay() -> None:
|
||||
record = make_candle_record()
|
||||
|
||||
assert record.event_time == START
|
||||
assert record.replay_at == START + timedelta(seconds=30)
|
||||
assert record.interval == "1m"
|
||||
|
||||
|
||||
def test_candle_record_rejects_observation_before_open_time() -> None:
|
||||
with pytest.raises(ValueError, match="observed_at"):
|
||||
make_candle_record(
|
||||
observed_at=START - timedelta(microseconds=1),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_final", (0, 1, None, "true"))
|
||||
def test_candle_record_requires_exact_boolean(is_final: Any) -> None:
|
||||
candle = make_candle()
|
||||
|
||||
with pytest.raises(TypeError, match="is_final"):
|
||||
CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle,
|
||||
observed_at=candle.open_time,
|
||||
is_final=is_final,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_cursor_normalizes_scope_and_time() -> None:
|
||||
offset = timezone(timedelta(hours=3))
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=" dzengi ",
|
||||
symbol=" btc/usd_leverage ",
|
||||
time_range=make_range(),
|
||||
executed_at=(START + timedelta(minutes=1)).astimezone(offset),
|
||||
replay_sequence=5,
|
||||
)
|
||||
|
||||
assert cursor.venue == VENUE
|
||||
assert cursor.symbol == SYMBOL
|
||||
assert cursor.executed_at.tzinfo is timezone.utc
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cursor_time",
|
||||
(
|
||||
START - timedelta(microseconds=1),
|
||||
END,
|
||||
END + timedelta(microseconds=1),
|
||||
),
|
||||
)
|
||||
def test_cursor_position_must_belong_to_query_range(
|
||||
cursor_time: datetime,
|
||||
) -> None:
|
||||
with pytest.raises(ValueError, match="query range"):
|
||||
TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=cursor_time,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_cursor_rejects_unknown_version() -> None:
|
||||
with pytest.raises(ValueError, match="version"):
|
||||
TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
version=2,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("limit", (1, HISTORY_PAGE_LIMIT_MAX))
|
||||
def test_trade_query_accepts_limit_boundaries(limit: int) -> None:
|
||||
query = TradeHistoryQuery(
|
||||
venue=" dzengi ",
|
||||
symbol="btc/usd_leverage",
|
||||
time_range=make_range(),
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
assert query.venue == VENUE
|
||||
assert query.symbol == SYMBOL
|
||||
assert query.limit == limit
|
||||
|
||||
|
||||
@pytest.mark.parametrize("limit", (True, 0, -1, 1.5, 1001))
|
||||
def test_query_rejects_invalid_limit(limit: Any) -> None:
|
||||
with pytest.raises((TypeError, ValueError), match="limit"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
|
||||
def test_query_accepts_cursor_when_only_page_limit_changes() -> None:
|
||||
time_range = make_range()
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
query = TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
limit=17,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
assert query.cursor is cursor
|
||||
|
||||
|
||||
def test_query_rejects_cursor_from_another_scope() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="scope"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol="ETH/USD_LEVERAGE",
|
||||
time_range=make_range(),
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_query_rejects_cursor_of_another_data_type() -> None:
|
||||
quote_cursor = QuoteHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
received_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryCursor"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
cursor=quote_cursor, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_candle_query_preserves_interval_case() -> None:
|
||||
query = CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=" 1M ",
|
||||
time_range=make_range(),
|
||||
)
|
||||
|
||||
assert query.interval == "1M"
|
||||
|
||||
|
||||
def test_empty_page_is_valid_without_cursor() -> None:
|
||||
page = TradeHistoryPage(query=make_trade_query(), items=())
|
||||
|
||||
assert page.items == ()
|
||||
assert page.has_more is False
|
||||
|
||||
|
||||
def test_empty_page_rejects_cursor() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="empty page"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(),
|
||||
next_cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_page_accepts_signed_rollover_at_equal_timestamp() -> None:
|
||||
rollover_time = START + timedelta(minutes=1)
|
||||
first = make_trade_record(
|
||||
trade_id=2_147_483_647,
|
||||
executed_at=rollover_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
second = make_trade_record(
|
||||
trade_id=-2_147_483_648,
|
||||
executed_at=rollover_time,
|
||||
replay_sequence=11,
|
||||
)
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=rollover_time,
|
||||
replay_sequence=11,
|
||||
)
|
||||
|
||||
page = TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(first, second),
|
||||
next_cursor=cursor,
|
||||
)
|
||||
|
||||
assert page.items == (first, second)
|
||||
assert page.has_more is True
|
||||
|
||||
|
||||
def test_trade_page_accepts_negative_one_to_zero_boundary() -> None:
|
||||
event_time = START + timedelta(minutes=1)
|
||||
page = TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(
|
||||
make_trade_record(
|
||||
trade_id=-1,
|
||||
executed_at=event_time,
|
||||
replay_sequence=20,
|
||||
),
|
||||
make_trade_record(
|
||||
trade_id=0,
|
||||
executed_at=event_time,
|
||||
replay_sequence=21,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
assert [item.trade.trade_id for item in page.items] == [-1, 0]
|
||||
|
||||
|
||||
def test_page_rejects_reverse_or_duplicate_order_key() -> None:
|
||||
first = make_trade_record(replay_sequence=2)
|
||||
second = make_trade_record(replay_sequence=1)
|
||||
|
||||
with pytest.raises(ValueError, match="strictly ordered"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(first, second),
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_cursor_not_pointing_to_last_item() -> None:
|
||||
item = make_trade_record(replay_sequence=10)
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=item.event_time,
|
||||
replay_sequence=9,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="last page item"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item,),
|
||||
next_cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"page",
|
||||
(
|
||||
QuoteHistoryPage(
|
||||
query=QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
),
|
||||
items=(make_quote_record(),),
|
||||
),
|
||||
CandleRevisionHistoryPage(
|
||||
query=CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="1m",
|
||||
time_range=make_range(),
|
||||
),
|
||||
items=(make_candle_record(),),
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_quote_and_candle_pages_are_typed_and_immutable(page: Any) -> None:
|
||||
assert not hasattr(page, "__dict__")
|
||||
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
setattr(page, "items", ())
|
||||
|
||||
|
||||
def test_cursor_and_query_classes_use_slots_and_are_frozen() -> None:
|
||||
query = QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
)
|
||||
|
||||
assert not hasattr(query, "__dict__")
|
||||
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
setattr(query, "venue", "other")
|
||||
|
||||
|
||||
def test_cursor_window_mismatch_is_rejected() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
shifted = HistoricalTimeRange(
|
||||
start_time=START - timedelta(minutes=1),
|
||||
end_time=END,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="scope"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=shifted,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_candle_cursor_interval_mismatch_is_rejected() -> None:
|
||||
cursor = CandleRevisionHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="1m",
|
||||
time_range=make_range(),
|
||||
open_time=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="interval"):
|
||||
CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval="5m",
|
||||
time_range=make_range(),
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_mixed_query_scope() -> None:
|
||||
first = make_trade_record(replay_sequence=1)
|
||||
second = replace(
|
||||
make_trade_record(
|
||||
executed_at=START + timedelta(minutes=2),
|
||||
replay_sequence=2,
|
||||
),
|
||||
venue="other",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="one query scope"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(first, second),
|
||||
)
|
||||
|
||||
|
||||
def test_terminal_page_rejects_item_outside_query_range() -> None:
|
||||
item = make_trade_record(
|
||||
executed_at=START - timedelta(microseconds=1),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="query range"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item,),
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_more_items_than_query_limit() -> None:
|
||||
first = make_trade_record(replay_sequence=1)
|
||||
second = make_trade_record(
|
||||
executed_at=START + timedelta(minutes=2),
|
||||
replay_sequence=2,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="query limit"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(limit=1),
|
||||
items=(first, second),
|
||||
)
|
||||
|
||||
|
||||
def test_page_items_must_follow_incoming_cursor() -> None:
|
||||
cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START + timedelta(minutes=1),
|
||||
replay_sequence=10,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="follow query cursor"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(cursor=cursor),
|
||||
items=(make_trade_record(replay_sequence=10),),
|
||||
)
|
||||
|
||||
|
||||
def test_records_reject_canonical_payload_subclasses() -> None:
|
||||
class TradeSubclass(Trade):
|
||||
pass
|
||||
|
||||
class QuoteSubclass(Quote):
|
||||
pass
|
||||
|
||||
class CandleSubclass(Candle):
|
||||
pass
|
||||
|
||||
trade = make_trade()
|
||||
quote = make_quote()
|
||||
candle = make_candle()
|
||||
|
||||
trade_subclass = TradeSubclass(
|
||||
symbol=trade.symbol,
|
||||
trade_id=trade.trade_id,
|
||||
price=trade.price,
|
||||
quantity=trade.quantity,
|
||||
executed_at=trade.executed_at,
|
||||
aggressor_side=trade.aggressor_side,
|
||||
source=trade.source,
|
||||
)
|
||||
quote_subclass = QuoteSubclass(
|
||||
symbol=quote.symbol,
|
||||
last_price=quote.last_price,
|
||||
bid_price=quote.bid_price,
|
||||
ask_price=quote.ask_price,
|
||||
exchange_timestamp=quote.exchange_timestamp,
|
||||
received_at=quote.received_at,
|
||||
source=quote.source,
|
||||
)
|
||||
candle_subclass = CandleSubclass(
|
||||
symbol=candle.symbol,
|
||||
interval=candle.interval,
|
||||
open_time=candle.open_time,
|
||||
open_price=candle.open_price,
|
||||
high_price=candle.high_price,
|
||||
low_price=candle.low_price,
|
||||
close_price=candle.close_price,
|
||||
volume=candle.volume,
|
||||
source=candle.source,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="Canonical Trade"):
|
||||
TradeHistoryRecord(
|
||||
venue=VENUE,
|
||||
trade=trade_subclass,
|
||||
first_observed_at=trade.executed_at,
|
||||
last_observed_at=trade.executed_at,
|
||||
observation_sources=(trade.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="Canonical Quote"):
|
||||
QuoteHistoryRecord(
|
||||
venue=VENUE,
|
||||
quote=quote_subclass,
|
||||
observation_sources=(quote.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="Canonical Candle"):
|
||||
CandleRevisionHistoryRecord(
|
||||
venue=VENUE,
|
||||
candle=candle_subclass,
|
||||
observed_at=candle.open_time,
|
||||
is_final=False,
|
||||
observation_sources=(candle.source,),
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
|
||||
def test_cursor_and_query_reject_time_range_subclass() -> None:
|
||||
class HistoricalTimeRangeSubclass(HistoricalTimeRange):
|
||||
pass
|
||||
|
||||
time_range = HistoricalTimeRangeSubclass(START, END)
|
||||
|
||||
with pytest.raises(TypeError, match="HistoricalTimeRange"):
|
||||
TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="HistoricalTimeRange"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=time_range,
|
||||
)
|
||||
|
||||
|
||||
def test_query_rejects_cursor_subclass() -> None:
|
||||
class TradeHistoryCursorSubclass(TradeHistoryCursor):
|
||||
pass
|
||||
|
||||
cursor = TradeHistoryCursorSubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=START,
|
||||
replay_sequence=1,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryCursor"):
|
||||
TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_query_record_and_cursor_subclasses() -> None:
|
||||
class TradeHistoryQuerySubclass(TradeHistoryQuery):
|
||||
pass
|
||||
|
||||
class TradeHistoryRecordSubclass(TradeHistoryRecord):
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (END, 1)
|
||||
|
||||
class TradeHistoryCursorSubclass(TradeHistoryCursor):
|
||||
pass
|
||||
|
||||
item = make_trade_record(replay_sequence=10)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryQuery"):
|
||||
TradeHistoryPage(
|
||||
query=TradeHistoryQuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
),
|
||||
items=(item,),
|
||||
)
|
||||
|
||||
item_subclass = TradeHistoryRecordSubclass(
|
||||
venue=item.venue,
|
||||
trade=item.trade,
|
||||
first_observed_at=item.first_observed_at,
|
||||
last_observed_at=item.last_observed_at,
|
||||
observation_sources=item.observation_sources,
|
||||
replay_sequence=item.replay_sequence,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryRecord"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item_subclass,),
|
||||
)
|
||||
|
||||
cursor_subclass = TradeHistoryCursorSubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=make_range(),
|
||||
executed_at=item.event_time,
|
||||
replay_sequence=item.replay_sequence,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="TradeHistoryCursor"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=(item,),
|
||||
next_cursor=cursor_subclass,
|
||||
)
|
||||
|
||||
|
||||
def test_page_rejects_items_tuple_subclass() -> None:
|
||||
class ItemsTupleSubclass(tuple):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="items must be a tuple"):
|
||||
TradeHistoryPage(
|
||||
query=make_trade_query(),
|
||||
items=ItemsTupleSubclass((make_trade_record(),)),
|
||||
)
|
||||
@@ -0,0 +1,195 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access.contracts import (
|
||||
MarketDataHistoricalAccessProtocol,
|
||||
)
|
||||
from src.market_data.access.market_data_historical_access import (
|
||||
MarketDataHistoricalAccess,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryPage,
|
||||
CandleRevisionHistoryQuery,
|
||||
HistoricalTimeRange,
|
||||
QuoteHistoryPage,
|
||||
QuoteHistoryQuery,
|
||||
TradeHistoryPage,
|
||||
TradeHistoryQuery,
|
||||
)
|
||||
|
||||
|
||||
NOW = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
TIME_RANGE = HistoricalTimeRange(
|
||||
start_time=NOW,
|
||||
end_time=NOW + timedelta(hours=1),
|
||||
)
|
||||
|
||||
|
||||
class RecordingTradeReader:
|
||||
def __init__(self, page: TradeHistoryPage) -> None:
|
||||
self.page = page
|
||||
self.queries: list[TradeHistoryQuery] = []
|
||||
|
||||
def query_trades(self, query: TradeHistoryQuery) -> TradeHistoryPage:
|
||||
self.queries.append(query)
|
||||
return self.page
|
||||
|
||||
|
||||
class RecordingQuoteReader:
|
||||
def __init__(self, page: QuoteHistoryPage) -> None:
|
||||
self.page = page
|
||||
self.queries: list[QuoteHistoryQuery] = []
|
||||
|
||||
def query_quotes(self, query: QuoteHistoryQuery) -> QuoteHistoryPage:
|
||||
self.queries.append(query)
|
||||
return self.page
|
||||
|
||||
|
||||
class RecordingCandleReader:
|
||||
def __init__(self, page: CandleRevisionHistoryPage) -> None:
|
||||
self.page = page
|
||||
self.queries: list[CandleRevisionHistoryQuery] = []
|
||||
|
||||
def query_candle_revisions(
|
||||
self,
|
||||
query: CandleRevisionHistoryQuery,
|
||||
) -> CandleRevisionHistoryPage:
|
||||
self.queries.append(query)
|
||||
return self.page
|
||||
|
||||
|
||||
class BrokenTradeReader:
|
||||
def __init__(self, error: RuntimeError) -> None:
|
||||
self.error = error
|
||||
|
||||
def query_trades(self, query: TradeHistoryQuery) -> TradeHistoryPage:
|
||||
del query
|
||||
raise self.error
|
||||
|
||||
|
||||
def make_dependencies() -> tuple[
|
||||
MarketDataHistoricalAccess,
|
||||
RecordingTradeReader,
|
||||
RecordingQuoteReader,
|
||||
RecordingCandleReader,
|
||||
]:
|
||||
trade_query = TradeHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
quote_query = QuoteHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
candle_query = CandleRevisionHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
interval="1m",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
trade_reader = RecordingTradeReader(
|
||||
TradeHistoryPage(query=trade_query, items=()),
|
||||
)
|
||||
quote_reader = RecordingQuoteReader(
|
||||
QuoteHistoryPage(query=quote_query, items=()),
|
||||
)
|
||||
candle_reader = RecordingCandleReader(
|
||||
CandleRevisionHistoryPage(query=candle_query, items=()),
|
||||
)
|
||||
access = MarketDataHistoricalAccess(
|
||||
trade_reader=trade_reader,
|
||||
quote_reader=quote_reader,
|
||||
candle_revision_reader=candle_reader,
|
||||
)
|
||||
return access, trade_reader, quote_reader, candle_reader
|
||||
|
||||
|
||||
def test_implements_combined_protocol_and_uses_slots() -> None:
|
||||
access, *_ = make_dependencies()
|
||||
|
||||
assert isinstance(access, MarketDataHistoricalAccessProtocol)
|
||||
assert not hasattr(access, "__dict__")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"dependency_name",
|
||||
(
|
||||
"trade_reader",
|
||||
"quote_reader",
|
||||
"candle_revision_reader",
|
||||
),
|
||||
)
|
||||
def test_rejects_dependency_without_required_protocol(
|
||||
dependency_name: str,
|
||||
) -> None:
|
||||
_, trade_reader, quote_reader, candle_reader = make_dependencies()
|
||||
dependencies: dict[str, Any] = {
|
||||
"trade_reader": trade_reader,
|
||||
"quote_reader": quote_reader,
|
||||
"candle_revision_reader": candle_reader,
|
||||
}
|
||||
dependencies[dependency_name] = object()
|
||||
|
||||
with pytest.raises(TypeError, match=dependency_name):
|
||||
MarketDataHistoricalAccess(**dependencies)
|
||||
|
||||
|
||||
def test_delegates_each_query_without_rebuilding_page() -> None:
|
||||
access, trade_reader, quote_reader, candle_reader = make_dependencies()
|
||||
trade_query = TradeHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
limit=10,
|
||||
)
|
||||
quote_query = QuoteHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
limit=20,
|
||||
)
|
||||
candle_query = CandleRevisionHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
interval="1m",
|
||||
time_range=TIME_RANGE,
|
||||
limit=30,
|
||||
)
|
||||
|
||||
trade_page = access.query_trades(trade_query)
|
||||
quote_page = access.query_quotes(quote_query)
|
||||
candle_page = access.query_candle_revisions(candle_query)
|
||||
|
||||
assert trade_page is trade_reader.page
|
||||
assert quote_page is quote_reader.page
|
||||
assert candle_page is candle_reader.page
|
||||
assert trade_reader.queries == [trade_query]
|
||||
assert quote_reader.queries == [quote_query]
|
||||
assert candle_reader.queries == [candle_query]
|
||||
|
||||
|
||||
def test_does_not_swallow_reader_error() -> None:
|
||||
_, _, quote_reader, candle_reader = make_dependencies()
|
||||
expected = RuntimeError("reader failed")
|
||||
broken_reader = BrokenTradeReader(expected)
|
||||
access = MarketDataHistoricalAccess(
|
||||
trade_reader=broken_reader,
|
||||
quote_reader=quote_reader,
|
||||
candle_revision_reader=candle_reader,
|
||||
)
|
||||
query = TradeHistoryQuery(
|
||||
venue="dzengi",
|
||||
symbol="BTC/USD_LEVERAGE",
|
||||
time_range=TIME_RANGE,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
access.query_trades(query)
|
||||
|
||||
assert captured.value is expected
|
||||
@@ -0,0 +1,620 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access.contracts import (
|
||||
CandleRevisionHistoryReaderProtocol,
|
||||
)
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
CandleRevisionHistoryCursor,
|
||||
CandleRevisionHistoryQuery,
|
||||
HistoricalTimeRange,
|
||||
)
|
||||
from src.market_data.access.postgres_candle_revision_history_repository import (
|
||||
PostgresCandleRevisionHistoryRepository,
|
||||
)
|
||||
from src.market_data.acquisition.models.candle import Candle
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
INTERVAL = "1m"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_candle"
|
||||
_DEFAULT = object()
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(self, rows: object = ()) -> None:
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.enter_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.execute_error: BaseException | None = None
|
||||
self.fetchall_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
|
||||
if self.execute_error is not None:
|
||||
raise self.execute_error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(self, cursor: RecordingCursor) -> None:
|
||||
self._cursor = cursor
|
||||
self.enter_calls = 0
|
||||
self.cursor_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.cursor_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
|
||||
if self.cursor_error is not None:
|
||||
raise self.cursor_error
|
||||
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def make_query(
|
||||
*,
|
||||
interval: str = INTERVAL,
|
||||
limit: int = 3,
|
||||
cursor: CandleRevisionHistoryCursor | None = None,
|
||||
) -> CandleRevisionHistoryQuery:
|
||||
return CandleRevisionHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=interval,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_row(
|
||||
*,
|
||||
venue: object = VENUE,
|
||||
symbol: object = SYMBOL,
|
||||
interval: object = INTERVAL,
|
||||
open_time: object = START + timedelta(minutes=1),
|
||||
observed_at: object = _DEFAULT,
|
||||
open_price: object = Decimal("64000"),
|
||||
high_price: object = Decimal("64200"),
|
||||
low_price: object = Decimal("63900"),
|
||||
close_price: object = Decimal("64150"),
|
||||
volume: object = Decimal("1.25"),
|
||||
is_final: object = False,
|
||||
source: object = SOURCE,
|
||||
observation_sources: object = _DEFAULT,
|
||||
replay_sequence: object = 10,
|
||||
canonical_schema_version: object = 1,
|
||||
) -> tuple[object, ...]:
|
||||
resolved_observed_at = (
|
||||
START + timedelta(minutes=1, seconds=10)
|
||||
if observed_at is _DEFAULT
|
||||
else observed_at
|
||||
)
|
||||
resolved_sources = (
|
||||
[SOURCE]
|
||||
if observation_sources is _DEFAULT
|
||||
else observation_sources
|
||||
)
|
||||
return (
|
||||
venue,
|
||||
symbol,
|
||||
interval,
|
||||
open_time,
|
||||
resolved_observed_at,
|
||||
open_price,
|
||||
high_price,
|
||||
low_price,
|
||||
close_price,
|
||||
volume,
|
||||
is_final,
|
||||
source,
|
||||
resolved_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresCandleRevisionHistoryRepository,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
]:
|
||||
cursor = RecordingCursor(rows)
|
||||
connection = RecordingConnection(cursor)
|
||||
provider = RecordingProvider(connection)
|
||||
repository = PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=provider,
|
||||
)
|
||||
return repository, cursor, connection, provider
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
repository, _, _, provider = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(repository, "__dict__")
|
||||
assert isinstance(repository, CandleRevisionHistoryReaderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_query_is_validated_before_connection_borrow() -> None:
|
||||
class QuerySubclass(CandleRevisionHistoryQuery):
|
||||
pass
|
||||
|
||||
repository, _, _, provider = dependencies()
|
||||
query = QuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=INTERVAL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessValidationError, match="query"):
|
||||
repository.query_candle_revisions(query)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_first_page_uses_open_time_half_open_order_and_limit_plus_one() -> None:
|
||||
first_open = START + timedelta(minutes=1)
|
||||
second_open = START + timedelta(minutes=2)
|
||||
repository, cursor, connection, provider = dependencies(
|
||||
[
|
||||
make_row(open_time=first_open, replay_sequence=10),
|
||||
make_row(
|
||||
open_time=second_open,
|
||||
observed_at=second_open + timedelta(seconds=10),
|
||||
is_final=True,
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
query = make_query(limit=3)
|
||||
|
||||
page = repository.query_candle_revisions(query)
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "open_time >= %s" in compact_sql
|
||||
assert "open_time < %s" in compact_sql
|
||||
assert "observed_at >= %s" not in compact_sql
|
||||
assert "(open_time, replay_sequence) >" not in compact_sql
|
||||
assert "ORDER BY open_time ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert parameters == (VENUE, SYMBOL, INTERVAL, START, END, 4)
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert type(page.items[0].candle) is Candle
|
||||
assert page.items[0].event_time is first_open
|
||||
assert page.items[1].is_final is True
|
||||
assert page.next_cursor is None
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert cursor.enter_calls == 1
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [None]
|
||||
|
||||
|
||||
def test_interval_is_case_sensitive_and_not_normalized() -> None:
|
||||
repository, cursor, _, _ = dependencies([])
|
||||
query = make_query(interval="1M")
|
||||
|
||||
repository.query_candle_revisions(query)
|
||||
|
||||
assert cursor.calls[0][1][2] == "1M"
|
||||
|
||||
|
||||
def test_keyset_query_uses_exact_open_time_cursor_tuple() -> None:
|
||||
cursor_time = START + timedelta(minutes=5)
|
||||
incoming = CandleRevisionHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=INTERVAL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
open_time=cursor_time,
|
||||
replay_sequence=25,
|
||||
)
|
||||
repository, cursor, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=cursor_time,
|
||||
observed_at=cursor_time + timedelta(seconds=10),
|
||||
replay_sequence=26,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
repository.query_candle_revisions(make_query(limit=2, cursor=incoming))
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
assert (
|
||||
"(open_time, replay_sequence) > (%s, %s)"
|
||||
in normalized_sql(sql)
|
||||
)
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
INTERVAL,
|
||||
START,
|
||||
END,
|
||||
cursor_time,
|
||||
25,
|
||||
3,
|
||||
)
|
||||
|
||||
|
||||
def test_limit_plus_one_creates_cursor_from_last_returned_item() -> None:
|
||||
open_times = tuple(
|
||||
START + timedelta(minutes=index) for index in (1, 2, 3)
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=open_time,
|
||||
observed_at=open_time + timedelta(seconds=10),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index, open_time in enumerate(open_times)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2)
|
||||
|
||||
page = repository.query_candle_revisions(query)
|
||||
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert page.next_cursor is not None
|
||||
assert page.next_cursor.venue == VENUE
|
||||
assert page.next_cursor.symbol == SYMBOL
|
||||
assert page.next_cursor.interval == INTERVAL
|
||||
assert page.next_cursor.time_range is query.time_range
|
||||
assert page.next_cursor.open_time == open_times[1]
|
||||
assert page.next_cursor.replay_sequence == 11
|
||||
|
||||
|
||||
def test_empty_result_is_valid_without_cursor() -> None:
|
||||
repository, _, _, _ = dependencies([])
|
||||
|
||||
page = repository.query_candle_revisions(make_query())
|
||||
|
||||
assert page.items == ()
|
||||
assert page.next_cursor is None
|
||||
assert page.has_more is False
|
||||
|
||||
|
||||
def test_history_range_uses_open_time_not_observed_at() -> None:
|
||||
open_time = START + timedelta(minutes=1)
|
||||
observed_at = END + timedelta(minutes=10)
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(open_time=open_time, observed_at=observed_at)]
|
||||
)
|
||||
|
||||
page = repository.query_candle_revisions(make_query())
|
||||
|
||||
assert page.items[0].event_time == open_time
|
||||
assert page.items[0].replay_at == observed_at
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides", "error_match"),
|
||||
(
|
||||
({"venue": "other"}, "scope"),
|
||||
({"symbol": "ETH/USD_LEVERAGE"}, "scope"),
|
||||
({"interval": "1M"}, "scope"),
|
||||
({"open_time": END, "observed_at": END}, "range"),
|
||||
),
|
||||
)
|
||||
def test_rows_outside_exact_scope_or_range_are_rejected(
|
||||
overrides: dict[str, object],
|
||||
error_match: str,
|
||||
) -> None:
|
||||
repository, _, _, _ = dependencies([make_row(**overrides)])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match=error_match):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
(
|
||||
{"venue": " dzengi"},
|
||||
{"symbol": "btc/usd_leverage"},
|
||||
{"interval": " 1m"},
|
||||
{"open_time": START.replace(tzinfo=None)},
|
||||
{"observed_at": START.replace(tzinfo=None)},
|
||||
{"observed_at": START},
|
||||
{"open_price": Decimal("0")},
|
||||
{"high_price": Decimal("NaN")},
|
||||
{"low_price": Decimal("65000")},
|
||||
{"close_price": Decimal("65000")},
|
||||
{"volume": Decimal("-0.01")},
|
||||
{"is_final": 1},
|
||||
{"source": " source"},
|
||||
{"observation_sources": ()},
|
||||
{"observation_sources": []},
|
||||
{"observation_sources": [SOURCE, SOURCE]},
|
||||
{"observation_sources": ["recovery", SOURCE]},
|
||||
{"replay_sequence": True},
|
||||
{"replay_sequence": 0},
|
||||
{"canonical_schema_version": 2},
|
||||
),
|
||||
)
|
||||
def test_corrupt_stored_values_raise_integrity_error(
|
||||
overrides: dict[str, object],
|
||||
) -> None:
|
||||
repository, cursor, connection, _ = dependencies(
|
||||
[make_row(**overrides)]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="invalid Canonical Candle",
|
||||
):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
assert connection.exit_exception_types == [
|
||||
MarketDataAccessIntegrityError
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("row", ((), [object()] * 15, (object(),) * 14))
|
||||
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
|
||||
repository, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
|
||||
def test_invalid_rows_collection_and_excess_rows_are_rejected() -> None:
|
||||
repository, _, _, _ = dependencies(iter(()))
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=START + timedelta(minutes=index + 1),
|
||||
observed_at=START + timedelta(minutes=index + 1, seconds=1),
|
||||
replay_sequence=index + 1,
|
||||
)
|
||||
for index in range(3)
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
repository.query_candle_revisions(make_query(limit=1))
|
||||
|
||||
|
||||
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
|
||||
earlier = START + timedelta(minutes=1)
|
||||
later = START + timedelta(minutes=2)
|
||||
unordered, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=later,
|
||||
observed_at=later,
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
open_time=earlier,
|
||||
observed_at=earlier,
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
|
||||
unordered.query_candle_revisions(make_query())
|
||||
|
||||
duplicate, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=earlier,
|
||||
observed_at=earlier,
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
open_time=later,
|
||||
observed_at=later,
|
||||
replay_sequence=10,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="duplicate replay_sequence",
|
||||
):
|
||||
duplicate.query_candle_revisions(make_query())
|
||||
|
||||
|
||||
def test_rows_must_strictly_follow_incoming_cursor() -> None:
|
||||
cursor_time = START + timedelta(minutes=1)
|
||||
incoming = CandleRevisionHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
interval=INTERVAL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
open_time=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
open_time=cursor_time,
|
||||
observed_at=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
|
||||
repository.query_candle_revisions(make_query(cursor=incoming))
|
||||
|
||||
|
||||
def test_database_error_is_wrapped_after_context_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = RuntimeError("database failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [RuntimeError]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
|
||||
def test_provider_error_is_wrapped_without_entering_connection() -> None:
|
||||
repository, _, connection, provider = dependencies()
|
||||
provider.error = RuntimeError("provider failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
|
||||
|
||||
def test_base_exception_is_not_wrapped_and_reaches_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [KeyboardInterrupt]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = RuntimeError("cursor cleanup failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_candle_revisions(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
|
||||
cursor = RecordingCursor([make_row()])
|
||||
connection = RecordingConnection(cursor)
|
||||
provider_calls = 0
|
||||
|
||||
def bound_provider() -> Any:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
return nullcontext(connection)
|
||||
|
||||
repository = PostgresCandleRevisionHistoryRepository(
|
||||
connection_provider=bound_provider,
|
||||
)
|
||||
|
||||
page = repository.query_candle_revisions(make_query())
|
||||
|
||||
assert len(page.items) == 1
|
||||
assert provider_calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
assert connection.exit_exception_types == []
|
||||
assert cursor.exit_exception_types == [None]
|
||||
@@ -0,0 +1,529 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access.contracts import QuoteHistoryReaderProtocol
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
)
|
||||
from src.market_data.access.models import (
|
||||
HistoricalTimeRange,
|
||||
QuoteHistoryCursor,
|
||||
QuoteHistoryQuery,
|
||||
)
|
||||
from src.market_data.access.postgres_quote_history_repository import (
|
||||
PostgresQuoteHistoryRepository,
|
||||
)
|
||||
from src.market_data.acquisition.models.quote import Quote
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_quote"
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(self, rows: object = ()) -> None:
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.enter_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.execute_error: BaseException | None = None
|
||||
self.fetchall_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
|
||||
if self.execute_error is not None:
|
||||
raise self.execute_error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(self, cursor: RecordingCursor) -> None:
|
||||
self._cursor = cursor
|
||||
self.enter_calls = 0
|
||||
self.cursor_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.cursor_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
|
||||
if self.cursor_error is not None:
|
||||
raise self.cursor_error
|
||||
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def make_query(
|
||||
*,
|
||||
limit: int = 3,
|
||||
cursor: QuoteHistoryCursor | None = None,
|
||||
) -> QuoteHistoryQuery:
|
||||
return QuoteHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_row(
|
||||
*,
|
||||
venue: object = VENUE,
|
||||
symbol: object = SYMBOL,
|
||||
received_at: object = START + timedelta(minutes=1),
|
||||
exchange_timestamp: object = START + timedelta(seconds=59),
|
||||
last_price: object = Decimal("64159.45"),
|
||||
bid_price: object = Decimal("64159.40"),
|
||||
ask_price: object = Decimal("64159.50"),
|
||||
source: object = SOURCE,
|
||||
observation_sources: object = None,
|
||||
replay_sequence: object = 10,
|
||||
canonical_schema_version: object = 1,
|
||||
) -> tuple[object, ...]:
|
||||
sources = [SOURCE] if observation_sources is None else observation_sources
|
||||
return (
|
||||
venue,
|
||||
symbol,
|
||||
received_at,
|
||||
exchange_timestamp,
|
||||
last_price,
|
||||
bid_price,
|
||||
ask_price,
|
||||
source,
|
||||
sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresQuoteHistoryRepository,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
]:
|
||||
cursor = RecordingCursor(rows)
|
||||
connection = RecordingConnection(cursor)
|
||||
provider = RecordingProvider(connection)
|
||||
repository = PostgresQuoteHistoryRepository(
|
||||
connection_provider=provider,
|
||||
)
|
||||
return repository, cursor, connection, provider
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
repository, _, _, provider = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(repository, "__dict__")
|
||||
assert isinstance(repository, QuoteHistoryReaderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresQuoteHistoryRepository(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_query_is_validated_before_connection_borrow() -> None:
|
||||
class QuoteHistoryQuerySubclass(QuoteHistoryQuery):
|
||||
pass
|
||||
|
||||
repository, _, _, provider = dependencies()
|
||||
query = QuoteHistoryQuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessValidationError, match="query"):
|
||||
repository.query_quotes(query)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_first_page_uses_received_time_half_open_order_and_limit() -> None:
|
||||
first_time = START + timedelta(minutes=1)
|
||||
second_time = START + timedelta(minutes=2)
|
||||
repository, cursor, connection, provider = dependencies(
|
||||
[
|
||||
make_row(received_at=first_time, replay_sequence=10),
|
||||
make_row(received_at=second_time, replay_sequence=11),
|
||||
]
|
||||
)
|
||||
query = make_query(limit=3)
|
||||
|
||||
page = repository.query_quotes(query)
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "received_at >= %s" in compact_sql
|
||||
assert "received_at < %s" in compact_sql
|
||||
assert "(received_at, replay_sequence) >" not in compact_sql
|
||||
assert "ORDER BY received_at ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert parameters == (VENUE, SYMBOL, START, END, 4)
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert type(page.items[0].quote) is Quote
|
||||
assert page.items[0].quote.received_at is first_time
|
||||
assert page.items[0].observation_sources == (SOURCE,)
|
||||
assert page.next_cursor is None
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert cursor.enter_calls == 1
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [None]
|
||||
|
||||
|
||||
def test_keyset_query_uses_exact_received_at_and_sequence_cursor() -> None:
|
||||
cursor_position = START + timedelta(minutes=5)
|
||||
incoming_cursor = QuoteHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
received_at=cursor_position,
|
||||
replay_sequence=25,
|
||||
)
|
||||
repository, recording_cursor, _, _ = dependencies(
|
||||
[make_row(received_at=cursor_position, replay_sequence=26)]
|
||||
)
|
||||
|
||||
repository.query_quotes(make_query(limit=2, cursor=incoming_cursor))
|
||||
|
||||
sql, parameters = recording_cursor.calls[0]
|
||||
assert (
|
||||
"(received_at, replay_sequence) > (%s, %s)"
|
||||
in normalized_sql(sql)
|
||||
)
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
START,
|
||||
END,
|
||||
cursor_position,
|
||||
25,
|
||||
3,
|
||||
)
|
||||
|
||||
|
||||
def test_limit_plus_one_creates_cursor_from_last_returned_quote() -> None:
|
||||
times = tuple(START + timedelta(minutes=index) for index in (1, 2, 3))
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(received_at=event_time, replay_sequence=10 + index)
|
||||
for index, event_time in enumerate(times)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2)
|
||||
|
||||
page = repository.query_quotes(query)
|
||||
|
||||
assert len(page.items) == 2
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert page.next_cursor is not None
|
||||
assert page.next_cursor.venue == query.venue
|
||||
assert page.next_cursor.symbol == query.symbol
|
||||
assert page.next_cursor.time_range is query.time_range
|
||||
assert page.next_cursor.received_at == times[1]
|
||||
assert page.next_cursor.replay_sequence == 11
|
||||
|
||||
|
||||
def test_empty_result_and_exact_start_are_valid() -> None:
|
||||
empty_repository, _, _, _ = dependencies([])
|
||||
|
||||
empty_page = empty_repository.query_quotes(make_query())
|
||||
|
||||
assert empty_page.items == ()
|
||||
assert empty_page.next_cursor is None
|
||||
assert empty_page.has_more is False
|
||||
|
||||
start_repository, _, _, _ = dependencies(
|
||||
[make_row(received_at=START)]
|
||||
)
|
||||
|
||||
start_page = start_repository.query_quotes(make_query())
|
||||
|
||||
assert start_page.items[0].event_time == START
|
||||
|
||||
|
||||
def test_none_exchange_timestamp_is_preserved() -> None:
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(exchange_timestamp=None)]
|
||||
)
|
||||
|
||||
page = repository.query_quotes(make_query())
|
||||
|
||||
assert page.items[0].quote.exchange_timestamp is None
|
||||
|
||||
|
||||
def test_backend_cannot_return_more_than_limit_plus_one() -> None:
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
received_at=START + timedelta(minutes=index + 1),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index in range(3)
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
repository.query_quotes(make_query(limit=1))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
(
|
||||
{"venue": "other"},
|
||||
{"venue": " dzengi "},
|
||||
{"symbol": "btc/usd_leverage"},
|
||||
{"received_at": END},
|
||||
{"received_at": START.replace(tzinfo=None)},
|
||||
{"exchange_timestamp": START.replace(tzinfo=None)},
|
||||
{"last_price": Decimal("0")},
|
||||
{"last_price": Decimal("NaN")},
|
||||
{"last_price": 1},
|
||||
{"bid_price": Decimal("64160")},
|
||||
{"ask_price": Decimal("0")},
|
||||
{"source": " source "},
|
||||
{"observation_sources": ()},
|
||||
{"observation_sources": []},
|
||||
{"observation_sources": [SOURCE, SOURCE]},
|
||||
{"observation_sources": ["recovery", SOURCE]},
|
||||
{"replay_sequence": True},
|
||||
{"replay_sequence": 0},
|
||||
{"canonical_schema_version": 2},
|
||||
),
|
||||
)
|
||||
def test_corrupt_or_out_of_scope_values_raise_integrity_error(
|
||||
overrides: dict[str, object],
|
||||
) -> None:
|
||||
repository, cursor, connection, _ = dependencies(
|
||||
[make_row(**overrides)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
assert connection.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"row",
|
||||
((), [object()] * 11, (object(),) * 10),
|
||||
)
|
||||
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
|
||||
repository, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rows", (None, "rows", object()))
|
||||
def test_invalid_rows_collection_raises_integrity_error(rows: object) -> None:
|
||||
repository, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
|
||||
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
|
||||
later = START + timedelta(minutes=2)
|
||||
earlier = START + timedelta(minutes=1)
|
||||
unordered_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(received_at=later, replay_sequence=10),
|
||||
make_row(received_at=earlier, replay_sequence=11),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
|
||||
unordered_repository.query_quotes(make_query())
|
||||
|
||||
duplicate_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(received_at=earlier, replay_sequence=10),
|
||||
make_row(received_at=later, replay_sequence=10),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="duplicate replay_sequence",
|
||||
):
|
||||
duplicate_repository.query_quotes(make_query())
|
||||
|
||||
|
||||
def test_rows_must_strictly_follow_incoming_cursor() -> None:
|
||||
cursor_time = START + timedelta(minutes=1)
|
||||
incoming = QuoteHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
received_at=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(received_at=cursor_time, replay_sequence=10)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
|
||||
repository.query_quotes(make_query(cursor=incoming))
|
||||
|
||||
|
||||
def test_database_error_is_wrapped_after_context_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = RuntimeError("database failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [RuntimeError]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
|
||||
def test_provider_error_is_wrapped_without_entering_connection() -> None:
|
||||
repository, _, connection, provider = dependencies()
|
||||
provider.error = RuntimeError("provider failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
|
||||
|
||||
def test_keyboard_interrupt_is_not_wrapped_and_reaches_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [KeyboardInterrupt]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = RuntimeError("cursor cleanup failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_quotes(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
|
||||
cursor = RecordingCursor([make_row()])
|
||||
connection = RecordingConnection(cursor)
|
||||
provider_calls = 0
|
||||
|
||||
def bound_provider() -> Any:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
return nullcontext(connection)
|
||||
|
||||
repository = PostgresQuoteHistoryRepository(
|
||||
connection_provider=bound_provider,
|
||||
)
|
||||
|
||||
page = repository.query_quotes(make_query())
|
||||
|
||||
assert len(page.items) == 1
|
||||
assert provider_calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
assert connection.exit_exception_types == []
|
||||
assert cursor.exit_exception_types == [None]
|
||||
@@ -0,0 +1,622 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import (
|
||||
HistoricalTimeRange,
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
MarketDataAccessValidationError,
|
||||
PostgresTradeHistoryRepository,
|
||||
TradeHistoryCursor,
|
||||
TradeHistoryQuery,
|
||||
TradeHistoryReaderProtocol,
|
||||
)
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.acquisition.trade_id_sequence import (
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
SOURCE = "dzengi_websocket_trade"
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(self, rows: object = ()) -> None:
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.enter_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.execute_error: BaseException | None = None
|
||||
self.fetchall_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
|
||||
if self.execute_error is not None:
|
||||
raise self.execute_error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(self, cursor: RecordingCursor) -> None:
|
||||
self._cursor = cursor
|
||||
self.enter_calls = 0
|
||||
self.cursor_calls = 0
|
||||
self.exit_exception_types: list[type[BaseException] | None] = []
|
||||
self.cursor_error: BaseException | None = None
|
||||
self.exit_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self.enter_calls += 1
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self.exit_exception_types.append(exception_type)
|
||||
|
||||
if self.exit_error is not None:
|
||||
raise self.exit_error
|
||||
|
||||
return None
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
|
||||
if self.cursor_error is not None:
|
||||
raise self.cursor_error
|
||||
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def make_query(
|
||||
*,
|
||||
limit: int = 3,
|
||||
cursor: TradeHistoryCursor | None = None,
|
||||
) -> TradeHistoryQuery:
|
||||
return TradeHistoryQuery(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
limit=limit,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
|
||||
def make_row(
|
||||
*,
|
||||
venue: object = VENUE,
|
||||
symbol: object = SYMBOL,
|
||||
trade_id: object = 100,
|
||||
executed_at: object = START + timedelta(minutes=1),
|
||||
price: object = Decimal("64159.45"),
|
||||
quantity: object = Decimal("0.125"),
|
||||
aggressor_side: object = "buy",
|
||||
source: object = SOURCE,
|
||||
first_observed_at: object = START + timedelta(minutes=1, seconds=1),
|
||||
last_observed_at: object = START + timedelta(minutes=1, seconds=2),
|
||||
observation_sources: object = None,
|
||||
replay_sequence: object = 10,
|
||||
canonical_schema_version: object = 1,
|
||||
) -> tuple[object, ...]:
|
||||
sources = [SOURCE] if observation_sources is None else observation_sources
|
||||
return (
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresTradeHistoryRepository,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
]:
|
||||
cursor = RecordingCursor(rows)
|
||||
connection = RecordingConnection(cursor)
|
||||
provider = RecordingProvider(connection)
|
||||
repository = PostgresTradeHistoryRepository(
|
||||
connection_provider=provider,
|
||||
)
|
||||
return repository, cursor, connection, provider
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
repository, _, _, provider = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(repository, "__dict__")
|
||||
assert isinstance(repository, TradeHistoryReaderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresTradeHistoryRepository(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_query_is_validated_before_connection_borrow() -> None:
|
||||
class TradeHistoryQuerySubclass(TradeHistoryQuery):
|
||||
pass
|
||||
|
||||
repository, _, _, provider = dependencies()
|
||||
query = TradeHistoryQuerySubclass(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessValidationError, match="query"):
|
||||
repository.query_trades(query)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_first_page_uses_half_open_ordered_limit_plus_one_query() -> None:
|
||||
first_time = START + timedelta(minutes=1)
|
||||
second_time = START + timedelta(minutes=2)
|
||||
repository, cursor, connection, provider = dependencies(
|
||||
[
|
||||
make_row(executed_at=first_time, replay_sequence=10),
|
||||
make_row(
|
||||
trade_id=101,
|
||||
executed_at=second_time,
|
||||
first_observed_at=second_time + timedelta(seconds=1),
|
||||
last_observed_at=second_time + timedelta(seconds=2),
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
query = make_query(limit=3)
|
||||
|
||||
page = repository.query_trades(query)
|
||||
|
||||
sql, parameters = cursor.calls[0]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "executed_at >= %s" in compact_sql
|
||||
assert "executed_at < %s" in compact_sql
|
||||
assert "(executed_at, replay_sequence) >" not in compact_sql
|
||||
assert "ORDER BY executed_at ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert parameters == (VENUE, SYMBOL, START, END, 4)
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert type(page.items[0].trade) is Trade
|
||||
assert page.items[0].trade.executed_at is first_time
|
||||
assert page.items[0].trade.aggressor_side is TradeAggressorSide.BUY
|
||||
assert page.items[0].observation_sources == (SOURCE,)
|
||||
assert page.next_cursor is None
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert cursor.enter_calls == 1
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [None]
|
||||
|
||||
|
||||
def test_keyset_query_uses_exact_cursor_tuple() -> None:
|
||||
cursor_position = START + timedelta(minutes=5)
|
||||
incoming_cursor = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
executed_at=cursor_position,
|
||||
replay_sequence=25,
|
||||
)
|
||||
row_time = cursor_position
|
||||
repository, recording_cursor, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=row_time,
|
||||
first_observed_at=row_time + timedelta(seconds=1),
|
||||
last_observed_at=row_time + timedelta(seconds=2),
|
||||
replay_sequence=26,
|
||||
)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2, cursor=incoming_cursor)
|
||||
|
||||
repository.query_trades(query)
|
||||
|
||||
sql, parameters = recording_cursor.calls[0]
|
||||
assert (
|
||||
"(executed_at, replay_sequence) > (%s, %s)"
|
||||
in normalized_sql(sql)
|
||||
)
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
START,
|
||||
END,
|
||||
cursor_position,
|
||||
25,
|
||||
3,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_id_rollover_does_not_participate_in_history_order() -> None:
|
||||
event_time = START + timedelta(minutes=1)
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
trade_id=SIGNED_TRADE_ID_MAX,
|
||||
executed_at=event_time,
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
trade_id=SIGNED_TRADE_ID_MIN,
|
||||
executed_at=event_time,
|
||||
replay_sequence=11,
|
||||
),
|
||||
make_row(
|
||||
trade_id=-1,
|
||||
executed_at=event_time,
|
||||
replay_sequence=12,
|
||||
),
|
||||
make_row(
|
||||
trade_id=0,
|
||||
executed_at=event_time,
|
||||
replay_sequence=13,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
page = repository.query_trades(make_query(limit=4))
|
||||
|
||||
assert [item.trade.trade_id for item in page.items] == [
|
||||
SIGNED_TRADE_ID_MAX,
|
||||
SIGNED_TRADE_ID_MIN,
|
||||
-1,
|
||||
0,
|
||||
]
|
||||
|
||||
|
||||
def test_limit_plus_one_creates_cursor_from_last_returned_item() -> None:
|
||||
times = tuple(START + timedelta(minutes=index) for index in (1, 2, 3))
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
trade_id=100 + index,
|
||||
executed_at=event_time,
|
||||
first_observed_at=event_time + timedelta(seconds=1),
|
||||
last_observed_at=event_time + timedelta(seconds=2),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index, event_time in enumerate(times)
|
||||
]
|
||||
)
|
||||
query = make_query(limit=2)
|
||||
|
||||
page = repository.query_trades(query)
|
||||
|
||||
assert len(page.items) == 2
|
||||
assert [item.replay_sequence for item in page.items] == [10, 11]
|
||||
assert page.next_cursor is not None
|
||||
assert page.next_cursor.venue == query.venue
|
||||
assert page.next_cursor.symbol == query.symbol
|
||||
assert page.next_cursor.time_range is query.time_range
|
||||
assert page.next_cursor.executed_at == times[1]
|
||||
assert page.next_cursor.replay_sequence == 11
|
||||
|
||||
|
||||
def test_empty_result_is_valid_without_cursor() -> None:
|
||||
repository, _, _, _ = dependencies([])
|
||||
|
||||
page = repository.query_trades(make_query())
|
||||
|
||||
assert page.items == ()
|
||||
assert page.next_cursor is None
|
||||
assert page.has_more is False
|
||||
|
||||
|
||||
def test_half_open_range_includes_exact_start() -> None:
|
||||
repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=START,
|
||||
first_observed_at=START,
|
||||
last_observed_at=START,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
page = repository.query_trades(make_query())
|
||||
|
||||
assert page.items[0].event_time == START
|
||||
|
||||
|
||||
def test_backend_cannot_return_more_than_limit_plus_one() -> None:
|
||||
query = make_query(limit=1)
|
||||
rows = [
|
||||
make_row(
|
||||
trade_id=100 + index,
|
||||
executed_at=START + timedelta(minutes=index + 1),
|
||||
first_observed_at=START + timedelta(minutes=index + 1, seconds=1),
|
||||
last_observed_at=START + timedelta(minutes=index + 1, seconds=2),
|
||||
replay_sequence=10 + index,
|
||||
)
|
||||
for index in range(3)
|
||||
]
|
||||
repository, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
repository.query_trades(query)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides", "error_match"),
|
||||
(
|
||||
({"venue": "other"}, "invalid Canonical Trade"),
|
||||
({"symbol": "btc/usd_leverage"}, "invalid Canonical Trade"),
|
||||
({"trade_id": True}, "invalid Canonical Trade"),
|
||||
({"trade_id": SIGNED_TRADE_ID_MAX + 1}, "invalid Canonical Trade"),
|
||||
({"executed_at": END}, "invalid Canonical Trade"),
|
||||
({"price": Decimal("0")}, "invalid Canonical Trade"),
|
||||
({"price": Decimal("NaN")}, "invalid Canonical Trade"),
|
||||
({"quantity": Decimal("0")}, "invalid Canonical Trade"),
|
||||
({"aggressor_side": "hold"}, "invalid Canonical Trade"),
|
||||
({"source": " source "}, "invalid Canonical Trade"),
|
||||
(
|
||||
{"first_observed_at": START.replace(tzinfo=None)},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
(
|
||||
{
|
||||
"first_observed_at": START + timedelta(minutes=3),
|
||||
"last_observed_at": START + timedelta(minutes=2),
|
||||
},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
({"observation_sources": ()}, "invalid Canonical Trade"),
|
||||
({"observation_sources": []}, "invalid Canonical Trade"),
|
||||
(
|
||||
{"observation_sources": [SOURCE, SOURCE]},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
(
|
||||
{"observation_sources": ["recovery", SOURCE]},
|
||||
"invalid Canonical Trade",
|
||||
),
|
||||
({"replay_sequence": True}, "invalid Canonical Trade"),
|
||||
({"replay_sequence": 0}, "invalid Canonical Trade"),
|
||||
({"canonical_schema_version": 2}, "invalid Canonical Trade"),
|
||||
),
|
||||
)
|
||||
def test_corrupt_stored_values_raise_integrity_error(
|
||||
overrides: dict[str, object],
|
||||
error_match: str,
|
||||
) -> None:
|
||||
repository, cursor, connection, _ = dependencies(
|
||||
[make_row(**overrides)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match=error_match):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
assert connection.exit_exception_types == [MarketDataAccessIntegrityError]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("row", ((), [object()] * 13, (object(),) * 12))
|
||||
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
|
||||
repository, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
|
||||
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
|
||||
later = START + timedelta(minutes=2)
|
||||
earlier = START + timedelta(minutes=1)
|
||||
unordered_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=later,
|
||||
first_observed_at=later + timedelta(seconds=1),
|
||||
last_observed_at=later + timedelta(seconds=2),
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
executed_at=earlier,
|
||||
first_observed_at=earlier + timedelta(seconds=1),
|
||||
last_observed_at=earlier + timedelta(seconds=2),
|
||||
replay_sequence=11,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
|
||||
unordered_repository.query_trades(make_query())
|
||||
|
||||
duplicate_repository, _, _, _ = dependencies(
|
||||
[
|
||||
make_row(
|
||||
executed_at=earlier,
|
||||
first_observed_at=earlier + timedelta(seconds=1),
|
||||
last_observed_at=earlier + timedelta(seconds=2),
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_row(
|
||||
trade_id=101,
|
||||
executed_at=later,
|
||||
first_observed_at=later + timedelta(seconds=1),
|
||||
last_observed_at=later + timedelta(seconds=2),
|
||||
replay_sequence=10,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MarketDataAccessIntegrityError,
|
||||
match="duplicate replay_sequence",
|
||||
):
|
||||
duplicate_repository.query_trades(make_query())
|
||||
|
||||
|
||||
def test_rows_must_strictly_follow_incoming_cursor() -> None:
|
||||
cursor_time = START + timedelta(minutes=1)
|
||||
incoming = TradeHistoryCursor(
|
||||
venue=VENUE,
|
||||
symbol=SYMBOL,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
executed_at=cursor_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
repository, _, _, _ = dependencies(
|
||||
[make_row(executed_at=cursor_time, replay_sequence=10)]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
|
||||
repository.query_trades(make_query(cursor=incoming))
|
||||
|
||||
|
||||
def test_database_error_is_wrapped_after_context_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = RuntimeError("database failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [RuntimeError]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
|
||||
def test_provider_error_is_wrapped_without_entering_connection() -> None:
|
||||
repository, _, connection, provider = dependencies()
|
||||
provider.error = RuntimeError("provider failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert provider.calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
|
||||
|
||||
def test_keyboard_interrupt_is_not_wrapped_and_reaches_cleanup() -> None:
|
||||
repository, cursor, connection, _ = dependencies()
|
||||
cursor.execute_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [KeyboardInterrupt]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = RuntimeError("cursor cleanup failed")
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as error_info:
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert isinstance(error_info.value.__cause__, RuntimeError)
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [RuntimeError]
|
||||
|
||||
repository, cursor, connection, _ = dependencies([])
|
||||
cursor.exit_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
repository.query_trades(make_query())
|
||||
|
||||
assert cursor.exit_exception_types == [None]
|
||||
assert connection.exit_exception_types == [KeyboardInterrupt]
|
||||
|
||||
|
||||
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
|
||||
cursor = RecordingCursor([make_row()])
|
||||
connection = RecordingConnection(cursor)
|
||||
provider_calls = 0
|
||||
|
||||
def bound_provider() -> Any:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
return nullcontext(connection)
|
||||
|
||||
repository = PostgresTradeHistoryRepository(
|
||||
connection_provider=bound_provider,
|
||||
)
|
||||
|
||||
page = repository.query_trades(make_query())
|
||||
|
||||
assert len(page.items) == 1
|
||||
assert provider_calls == 1
|
||||
assert connection.enter_calls == 0
|
||||
assert connection.exit_exception_types == []
|
||||
assert cursor.exit_exception_types == [None]
|
||||
28
app/tests/unit/market_data/replay/conftest.py
Normal file
28
app/tests/unit/market_data/replay/conftest.py
Normal file
@@ -0,0 +1,28 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.replay import (
|
||||
ReplayDataType,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def empty_replay_plan() -> ReplayPlan:
|
||||
start = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
request = ReplayPlanRequest(
|
||||
venue="dzengi",
|
||||
symbols=("BTC/USD_LEVERAGE",),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=start,
|
||||
end_time=start + timedelta(hours=1),
|
||||
),
|
||||
max_records=100,
|
||||
)
|
||||
return ReplayPlan(request=request, events=())
|
||||
@@ -0,0 +1,214 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
import inspect
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.replay.contracts import (
|
||||
MarketDataClockProtocol,
|
||||
ReplayClockProtocol,
|
||||
)
|
||||
from src.market_data.replay.deterministic_replay_clock import (
|
||||
DeterministicReplayClock,
|
||||
)
|
||||
from src.market_data.replay.exceptions import ReplayClockError
|
||||
|
||||
|
||||
START = datetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
0,
|
||||
123456,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
|
||||
|
||||
def test_matches_protocols_uses_slots_and_is_synchronous() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
assert isinstance(clock, MarketDataClockProtocol)
|
||||
assert isinstance(clock, ReplayClockProtocol)
|
||||
assert not hasattr(clock, "__dict__")
|
||||
assert inspect.iscoroutinefunction(clock.advance_to) is False
|
||||
|
||||
|
||||
def test_starts_at_exact_canonical_utc_time() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
assert clock.now == START
|
||||
assert type(clock.now) is datetime
|
||||
assert clock.now.tzinfo is timezone.utc
|
||||
|
||||
|
||||
def test_normalizes_non_utc_initial_time_without_losing_precision() -> None:
|
||||
offset = timezone(timedelta(hours=3, minutes=30))
|
||||
initial_time = datetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
15,
|
||||
30,
|
||||
0,
|
||||
654321,
|
||||
tzinfo=offset,
|
||||
)
|
||||
|
||||
clock = DeterministicReplayClock(initial_time)
|
||||
|
||||
assert clock.now == datetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
0,
|
||||
654321,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
assert type(clock.now) is datetime
|
||||
|
||||
|
||||
def test_accepts_datetime_subclass_but_stores_base_datetime() -> None:
|
||||
class CompatibleDatetime(datetime):
|
||||
pass
|
||||
|
||||
initial_time = CompatibleDatetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
|
||||
clock = DeterministicReplayClock(initial_time)
|
||||
|
||||
assert clock.now == initial_time
|
||||
assert type(clock.now) is datetime
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, "2026-08-02", 1, object()))
|
||||
def test_rejects_non_datetime_initial_value(invalid: object) -> None:
|
||||
with pytest.raises(TypeError, match="initial_time"):
|
||||
DeterministicReplayClock(invalid) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def test_rejects_naive_initial_time() -> None:
|
||||
with pytest.raises(ValueError, match="timezone"):
|
||||
DeterministicReplayClock(START.replace(tzinfo=None))
|
||||
|
||||
|
||||
def test_advances_forward_with_microsecond_precision() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
expected = START + timedelta(microseconds=1)
|
||||
|
||||
result = clock.advance_to(expected)
|
||||
|
||||
assert result is None
|
||||
assert clock.now == expected
|
||||
|
||||
|
||||
def test_allows_repeated_equal_time_for_distinct_replay_events() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
clock.advance_to(START)
|
||||
clock.advance_to(START)
|
||||
|
||||
assert clock.now == START
|
||||
|
||||
|
||||
def test_equal_instant_with_another_offset_is_idempotent() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
equal_instant = START.astimezone(timezone(timedelta(hours=-4)))
|
||||
|
||||
clock.advance_to(equal_instant)
|
||||
|
||||
assert clock.now == START
|
||||
assert clock.now.tzinfo is timezone.utc
|
||||
|
||||
|
||||
def test_equal_utc_instant_with_another_fold_is_complete_no_op() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
previous = clock.now
|
||||
|
||||
clock.advance_to(START.replace(fold=1))
|
||||
|
||||
assert clock.now is previous
|
||||
assert clock.now.fold == 0
|
||||
|
||||
|
||||
def test_accepts_datetime_subclass_during_advance() -> None:
|
||||
class CompatibleDatetime(datetime):
|
||||
pass
|
||||
|
||||
clock = DeterministicReplayClock(START)
|
||||
later = CompatibleDatetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
1,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
|
||||
clock.advance_to(later)
|
||||
|
||||
assert clock.now == later
|
||||
assert type(clock.now) is datetime
|
||||
|
||||
|
||||
def test_rejects_backward_transition_without_changing_state() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
later = START + timedelta(minutes=1)
|
||||
clock.advance_to(later)
|
||||
|
||||
with pytest.raises(ReplayClockError, match="backwards"):
|
||||
clock.advance_to(START)
|
||||
|
||||
assert clock.now == later
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, "later", 1, object()))
|
||||
def test_invalid_advance_type_does_not_change_state(invalid: object) -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
with pytest.raises(TypeError, match="instant"):
|
||||
clock.advance_to(invalid) # type: ignore[arg-type]
|
||||
|
||||
assert clock.now == START
|
||||
|
||||
|
||||
def test_naive_advance_does_not_change_state() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
with pytest.raises(ValueError, match="timezone"):
|
||||
clock.advance_to(START.replace(tzinfo=None))
|
||||
|
||||
assert clock.now == START
|
||||
|
||||
|
||||
def test_two_clocks_have_independent_state() -> None:
|
||||
first = DeterministicReplayClock(START)
|
||||
second = DeterministicReplayClock(START)
|
||||
|
||||
first.advance_to(START + timedelta(hours=1))
|
||||
|
||||
assert first.now == START + timedelta(hours=1)
|
||||
assert second.now == START
|
||||
|
||||
|
||||
def test_now_is_read_only_and_lifecycle_extensions_are_absent() -> None:
|
||||
clock = DeterministicReplayClock(START)
|
||||
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(clock, "now", START + timedelta(hours=1))
|
||||
|
||||
assert not hasattr(clock, "reset")
|
||||
assert not hasattr(clock, "advance_by")
|
||||
assert not hasattr(clock, "start")
|
||||
assert not hasattr(clock, "stop")
|
||||
assert clock.now == START
|
||||
@@ -0,0 +1,94 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from src.market_data.replay import (
|
||||
MarketDataClockProtocol,
|
||||
MarketDataReplayError,
|
||||
MarketDataReplayValidationError,
|
||||
ReplayClockError,
|
||||
ReplayClockProtocol,
|
||||
ReplayConsumerProtocol,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanBuilderProtocol,
|
||||
ReplayPlanLimitExceededError,
|
||||
ReplayPlanRequest,
|
||||
ReplaySessionProtocol,
|
||||
ReplaySessionState,
|
||||
ReplaySessionStateError,
|
||||
)
|
||||
|
||||
|
||||
class FakeClock:
|
||||
def __init__(self, now: datetime) -> None:
|
||||
self._now = now
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return self._now
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self._now = instant
|
||||
|
||||
|
||||
class RecordingConsumer:
|
||||
def __init__(self) -> None:
|
||||
self.events: list[ReplayEvent] = []
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.events.append(event)
|
||||
|
||||
|
||||
class RecordingPlanBuilder:
|
||||
def __init__(self, plan: ReplayPlan) -> None:
|
||||
self.plan = plan
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
return self.plan
|
||||
|
||||
|
||||
class RecordingSession:
|
||||
def __init__(self, plan: ReplayPlan, clock: FakeClock) -> None:
|
||||
self._plan = plan
|
||||
self._clock = clock
|
||||
|
||||
@property
|
||||
def state(self) -> ReplaySessionState:
|
||||
return ReplaySessionState.CREATED
|
||||
|
||||
@property
|
||||
def plan(self) -> ReplayPlan:
|
||||
return self._plan
|
||||
|
||||
@property
|
||||
def clock(self) -> MarketDataClockProtocol:
|
||||
return self._clock
|
||||
|
||||
async def run(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def test_replay_protocols_are_runtime_checkable(
|
||||
empty_replay_plan: ReplayPlan,
|
||||
) -> None:
|
||||
clock = FakeClock(empty_replay_plan.request.time_range.start_time)
|
||||
|
||||
assert isinstance(clock, MarketDataClockProtocol)
|
||||
assert isinstance(clock, ReplayClockProtocol)
|
||||
assert isinstance(RecordingConsumer(), ReplayConsumerProtocol)
|
||||
assert isinstance(
|
||||
RecordingPlanBuilder(empty_replay_plan),
|
||||
ReplayPlanBuilderProtocol,
|
||||
)
|
||||
assert isinstance(
|
||||
RecordingSession(empty_replay_plan, clock),
|
||||
ReplaySessionProtocol,
|
||||
)
|
||||
|
||||
|
||||
def test_replay_error_hierarchy_is_specialized() -> None:
|
||||
assert issubclass(MarketDataReplayValidationError, MarketDataReplayError)
|
||||
assert issubclass(ReplayPlanLimitExceededError, MarketDataReplayError)
|
||||
assert issubclass(ReplayClockError, MarketDataReplayError)
|
||||
assert issubclass(ReplaySessionStateError, MarketDataReplayError)
|
||||
@@ -0,0 +1,659 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import FrozenInstanceError
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.candle import Candle
|
||||
from src.market_data.acquisition.models.quote import Quote
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.replay import (
|
||||
REPLAY_PLAN_MAX_RECORDS_LIMIT,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
ReplaySessionState,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
|
||||
|
||||
def make_trade(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
executed_at: datetime = START + timedelta(minutes=1),
|
||||
symbol: str = SYMBOL,
|
||||
) -> Trade:
|
||||
return Trade(
|
||||
symbol=symbol,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=executed_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source="dzengi_websocket_trade",
|
||||
)
|
||||
|
||||
|
||||
def make_quote(
|
||||
*,
|
||||
received_at: datetime = START + timedelta(minutes=2),
|
||||
) -> Quote:
|
||||
return Quote(
|
||||
symbol=SYMBOL,
|
||||
last_price=Decimal("65000"),
|
||||
bid_price=Decimal("64999"),
|
||||
ask_price=Decimal("65001"),
|
||||
exchange_timestamp=received_at - timedelta(milliseconds=1),
|
||||
received_at=received_at,
|
||||
source="dzengi_rest_quote",
|
||||
)
|
||||
|
||||
|
||||
def make_candle(*, interval: str = "1m") -> Candle:
|
||||
return Candle(
|
||||
symbol=SYMBOL,
|
||||
interval=interval,
|
||||
open_time=START,
|
||||
open_price=Decimal("64900"),
|
||||
high_price=Decimal("65100"),
|
||||
low_price=Decimal("64800"),
|
||||
close_price=Decimal("65000"),
|
||||
volume=Decimal("10"),
|
||||
source="dzengi_rest_candle",
|
||||
)
|
||||
|
||||
|
||||
def make_request(
|
||||
*,
|
||||
data_types: tuple[ReplayDataType, ...] = (ReplayDataType.TRADE,),
|
||||
candle_intervals: tuple[str, ...] = (),
|
||||
max_records: int = 100,
|
||||
) -> ReplayPlanRequest:
|
||||
return ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=data_types,
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=START,
|
||||
end_time=END,
|
||||
),
|
||||
candle_intervals=candle_intervals,
|
||||
max_records=max_records,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_event(
|
||||
*,
|
||||
trade_id: int = 100,
|
||||
replay_at: datetime = START + timedelta(minutes=1),
|
||||
replay_sequence: int = 1,
|
||||
venue: str = VENUE,
|
||||
symbol: str = SYMBOL,
|
||||
) -> ReplayEvent:
|
||||
trade = make_trade(
|
||||
trade_id=trade_id,
|
||||
executed_at=replay_at,
|
||||
symbol=symbol,
|
||||
)
|
||||
return ReplayEvent(
|
||||
venue=venue,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=replay_sequence,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def trade_id_from_event(event: ReplayEvent) -> int:
|
||||
payload = event.payload
|
||||
|
||||
if not isinstance(payload, Trade):
|
||||
raise AssertionError("Replay event must contain Trade payload.")
|
||||
|
||||
return payload.trade_id
|
||||
|
||||
|
||||
def test_trade_event_preserves_payload_identity_and_exact_event_time() -> None:
|
||||
trade = make_trade()
|
||||
event = ReplayEvent(
|
||||
venue=" dzengi ",
|
||||
replay_at=trade.executed_at,
|
||||
replay_sequence=10,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
assert event.venue == VENUE
|
||||
assert event.payload is trade
|
||||
assert event.symbol == SYMBOL
|
||||
assert event.data_type is ReplayDataType.TRADE
|
||||
assert event.order_key == (trade.executed_at, 10)
|
||||
assert event.candle_is_final is None
|
||||
|
||||
|
||||
def test_quote_event_uses_received_time_not_exchange_time() -> None:
|
||||
quote = make_quote()
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=quote.received_at,
|
||||
replay_sequence=11,
|
||||
payload=quote,
|
||||
)
|
||||
|
||||
assert event.data_type is ReplayDataType.QUOTE
|
||||
assert event.replay_at == quote.received_at
|
||||
assert event.replay_at != quote.exchange_timestamp
|
||||
|
||||
|
||||
def test_candle_event_uses_observation_time_and_final_metadata() -> None:
|
||||
candle = make_candle()
|
||||
observed_at = candle.open_time + timedelta(seconds=30)
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=observed_at,
|
||||
replay_sequence=12,
|
||||
payload=candle,
|
||||
candle_is_final=False,
|
||||
)
|
||||
|
||||
assert event.data_type is ReplayDataType.CANDLE_REVISION
|
||||
assert event.replay_at == observed_at
|
||||
assert event.candle_is_final is False
|
||||
|
||||
|
||||
def test_event_normalizes_aware_time_to_utc() -> None:
|
||||
offset = timezone(timedelta(hours=3))
|
||||
trade = make_trade()
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at.astimezone(offset),
|
||||
replay_sequence=1,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
assert event.replay_at.tzinfo is timezone.utc
|
||||
assert event.replay_at == trade.executed_at
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid_sequence", (True, 0, -1, 1.5, "1"))
|
||||
def test_event_rejects_invalid_sequence(invalid_sequence: Any) -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="replay_sequence"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at,
|
||||
replay_sequence=invalid_sequence,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_event_rejects_mismatched_replay_time() -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises(ValueError, match="event time"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at + timedelta(microseconds=1),
|
||||
replay_sequence=1,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def test_candle_event_rejects_time_before_open() -> None:
|
||||
candle = make_candle()
|
||||
|
||||
with pytest.raises(ValueError, match="open_time"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=candle.open_time - timedelta(microseconds=1),
|
||||
replay_sequence=1,
|
||||
payload=candle,
|
||||
candle_is_final=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("candle_is_final", (None, 0, 1, "true"))
|
||||
def test_candle_event_requires_exact_final_boolean(
|
||||
candle_is_final: Any,
|
||||
) -> None:
|
||||
candle = make_candle()
|
||||
|
||||
with pytest.raises(TypeError, match="candle_is_final"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=candle.open_time,
|
||||
replay_sequence=1,
|
||||
payload=candle,
|
||||
candle_is_final=candle_is_final,
|
||||
)
|
||||
|
||||
|
||||
def test_trade_event_rejects_candle_metadata() -> None:
|
||||
trade = make_trade()
|
||||
|
||||
with pytest.raises(ValueError, match="must be None"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=trade.executed_at,
|
||||
replay_sequence=1,
|
||||
payload=trade,
|
||||
candle_is_final=False,
|
||||
)
|
||||
|
||||
|
||||
def test_request_normalizes_symbols_and_preserves_interval_case() -> None:
|
||||
request = ReplayPlanRequest(
|
||||
venue=" dzengi ",
|
||||
symbols=(" btc/usd_leverage ",),
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
candle_intervals=(" 1M ",),
|
||||
)
|
||||
|
||||
assert request.venue == VENUE
|
||||
assert request.symbols == (SYMBOL,)
|
||||
assert request.candle_intervals == ("1M",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"symbols",
|
||||
(
|
||||
[],
|
||||
(),
|
||||
("",),
|
||||
("BTC/USD_LEVERAGE", "btc/usd_leverage"),
|
||||
),
|
||||
)
|
||||
def test_request_rejects_invalid_symbols(symbols: Any) -> None:
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=symbols,
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
|
||||
def test_request_requires_intervals_only_for_candles() -> None:
|
||||
with pytest.raises(ValueError, match="required"):
|
||||
make_request(data_types=(ReplayDataType.CANDLE_REVISION,))
|
||||
|
||||
with pytest.raises(ValueError, match="require Candle"):
|
||||
make_request(candle_intervals=("1m",))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"max_records",
|
||||
(True, 0, -1, 1.5, REPLAY_PLAN_MAX_RECORDS_LIMIT + 1),
|
||||
)
|
||||
def test_request_rejects_invalid_max_records(max_records: Any) -> None:
|
||||
with pytest.raises((TypeError, ValueError), match="max_records"):
|
||||
make_request(max_records=max_records)
|
||||
|
||||
|
||||
def test_empty_plan_is_valid_and_preserves_request_scope() -> None:
|
||||
request = make_request()
|
||||
plan = ReplayPlan(request=request, events=())
|
||||
|
||||
assert plan.request is request
|
||||
assert plan.events == ()
|
||||
assert plan.is_empty is True
|
||||
assert len(plan) == 0
|
||||
|
||||
|
||||
def test_plan_preserves_event_and_payload_identity() -> None:
|
||||
request = make_request()
|
||||
event = make_trade_event()
|
||||
plan = ReplayPlan(request=request, events=(event,))
|
||||
|
||||
assert plan.events[0] is event
|
||||
assert plan.events[0].payload is event.payload
|
||||
|
||||
|
||||
def test_plan_accepts_rollover_order_by_time_and_sequence() -> None:
|
||||
same_time = START + timedelta(minutes=1)
|
||||
first = make_trade_event(
|
||||
trade_id=2_147_483_647,
|
||||
replay_at=same_time,
|
||||
replay_sequence=10,
|
||||
)
|
||||
second = make_trade_event(
|
||||
trade_id=-2_147_483_648,
|
||||
replay_at=same_time,
|
||||
replay_sequence=11,
|
||||
)
|
||||
|
||||
plan = ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(first, second),
|
||||
)
|
||||
|
||||
assert [trade_id_from_event(event) for event in plan.events] == [
|
||||
2_147_483_647,
|
||||
-2_147_483_648,
|
||||
]
|
||||
|
||||
|
||||
def test_plan_accepts_negative_one_to_zero_at_equal_time() -> None:
|
||||
same_time = START + timedelta(minutes=1)
|
||||
plan = ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(
|
||||
trade_id=-1,
|
||||
replay_at=same_time,
|
||||
replay_sequence=20,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=0,
|
||||
replay_at=same_time,
|
||||
replay_sequence=21,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
assert [trade_id_from_event(event) for event in plan.events] == [-1, 0]
|
||||
|
||||
|
||||
def test_plan_rejects_reverse_order() -> None:
|
||||
with pytest.raises(ValueError, match="strictly ordered"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(replay_sequence=2),
|
||||
make_trade_event(replay_sequence=1),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_reverse_event_time_with_increasing_sequence() -> None:
|
||||
with pytest.raises(ValueError, match="strictly ordered"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=2,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_duplicate_global_sequence() -> None:
|
||||
with pytest.raises(ValueError, match="globally unique"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=(
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=1,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_event_outside_time_range() -> None:
|
||||
event = make_trade_event(replay_at=END)
|
||||
|
||||
with pytest.raises(ValueError, match="outside Replay request"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_data_type_outside_request() -> None:
|
||||
quote = make_quote()
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=quote.received_at,
|
||||
replay_sequence=1,
|
||||
payload=quote,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="data type"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_event_from_another_venue() -> None:
|
||||
event = make_trade_event(venue="other")
|
||||
|
||||
with pytest.raises(ValueError, match="event venue"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_event_from_another_symbol() -> None:
|
||||
event = make_trade_event(symbol="ETH/USD_LEVERAGE")
|
||||
|
||||
with pytest.raises(ValueError, match="event symbol"):
|
||||
ReplayPlan(request=make_request(), events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_candle_interval_outside_request() -> None:
|
||||
candle = make_candle(interval="5m")
|
||||
event = ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=candle.open_time,
|
||||
replay_sequence=1,
|
||||
payload=candle,
|
||||
candle_is_final=False,
|
||||
)
|
||||
request = make_request(
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
candle_intervals=("1m",),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Candle interval"):
|
||||
ReplayPlan(request=request, events=(event,))
|
||||
|
||||
|
||||
def test_plan_rejects_more_events_than_request_limit() -> None:
|
||||
event = make_trade_event()
|
||||
|
||||
with pytest.raises(ValueError, match="max_records"):
|
||||
ReplayPlan(
|
||||
request=make_request(max_records=1),
|
||||
events=(event, make_trade_event(replay_sequence=2)),
|
||||
)
|
||||
|
||||
|
||||
def test_replay_models_are_frozen_and_slotted() -> None:
|
||||
event = make_trade_event()
|
||||
request = make_request()
|
||||
plan = ReplayPlan(request=request, events=(event,))
|
||||
|
||||
for model in (event, request, plan):
|
||||
assert not hasattr(model, "__dict__")
|
||||
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
setattr(event, "venue", "other")
|
||||
|
||||
|
||||
def test_session_states_are_explicit_and_complete() -> None:
|
||||
assert tuple(ReplaySessionState) == (
|
||||
ReplaySessionState.CREATED,
|
||||
ReplaySessionState.RUNNING,
|
||||
ReplaySessionState.COMPLETED,
|
||||
ReplaySessionState.FAILED,
|
||||
ReplaySessionState.CANCELLED,
|
||||
)
|
||||
|
||||
|
||||
def test_event_rejects_canonical_payload_subclasses() -> None:
|
||||
class TradeSubclass(Trade):
|
||||
pass
|
||||
|
||||
class QuoteSubclass(Quote):
|
||||
pass
|
||||
|
||||
class CandleSubclass(Candle):
|
||||
pass
|
||||
|
||||
trade = make_trade()
|
||||
quote = make_quote()
|
||||
candle = make_candle()
|
||||
|
||||
payloads = (
|
||||
(
|
||||
TradeSubclass(
|
||||
symbol=trade.symbol,
|
||||
trade_id=trade.trade_id,
|
||||
price=trade.price,
|
||||
quantity=trade.quantity,
|
||||
executed_at=trade.executed_at,
|
||||
aggressor_side=trade.aggressor_side,
|
||||
source=trade.source,
|
||||
),
|
||||
trade.executed_at,
|
||||
None,
|
||||
),
|
||||
(
|
||||
QuoteSubclass(
|
||||
symbol=quote.symbol,
|
||||
last_price=quote.last_price,
|
||||
bid_price=quote.bid_price,
|
||||
ask_price=quote.ask_price,
|
||||
exchange_timestamp=quote.exchange_timestamp,
|
||||
received_at=quote.received_at,
|
||||
source=quote.source,
|
||||
),
|
||||
quote.received_at,
|
||||
None,
|
||||
),
|
||||
(
|
||||
CandleSubclass(
|
||||
symbol=candle.symbol,
|
||||
interval=candle.interval,
|
||||
open_time=candle.open_time,
|
||||
open_price=candle.open_price,
|
||||
high_price=candle.high_price,
|
||||
low_price=candle.low_price,
|
||||
close_price=candle.close_price,
|
||||
volume=candle.volume,
|
||||
source=candle.source,
|
||||
),
|
||||
candle.open_time,
|
||||
False,
|
||||
),
|
||||
)
|
||||
|
||||
for payload, replay_at, candle_is_final in payloads:
|
||||
with pytest.raises(TypeError, match="Canonical Trade, Quote or Candle"):
|
||||
ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=1,
|
||||
payload=payload,
|
||||
candle_is_final=candle_is_final,
|
||||
)
|
||||
|
||||
|
||||
def test_request_rejects_time_range_subclass() -> None:
|
||||
class HistoricalTimeRangeSubclass(HistoricalTimeRange):
|
||||
def contains(self, instant: datetime) -> bool:
|
||||
return True
|
||||
|
||||
with pytest.raises(TypeError, match="HistoricalTimeRange"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRangeSubclass(START, END),
|
||||
)
|
||||
|
||||
|
||||
def test_request_rejects_tuple_subclasses() -> None:
|
||||
class TupleSubclass(tuple):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="symbols must be a tuple"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=TupleSubclass((SYMBOL,)),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="data_types must be a tuple"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=TupleSubclass((ReplayDataType.TRADE,)),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="candle_intervals must be a tuple"):
|
||||
ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
candle_intervals=TupleSubclass(("1m",)),
|
||||
)
|
||||
|
||||
|
||||
def test_plan_rejects_request_and_event_subclasses() -> None:
|
||||
class ReplayPlanRequestSubclass(ReplayPlanRequest):
|
||||
pass
|
||||
|
||||
class ReplayEventSubclass(ReplayEvent):
|
||||
@property
|
||||
def symbol(self) -> str:
|
||||
return SYMBOL
|
||||
|
||||
@property
|
||||
def data_type(self) -> ReplayDataType:
|
||||
return ReplayDataType.TRADE
|
||||
|
||||
@property
|
||||
def order_key(self) -> tuple[datetime, int]:
|
||||
return (START, 1)
|
||||
|
||||
request_subclass = ReplayPlanRequestSubclass(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="ReplayPlanRequest"):
|
||||
ReplayPlan(request=request_subclass, events=())
|
||||
|
||||
event = make_trade_event()
|
||||
event_subclass = ReplayEventSubclass(
|
||||
venue=event.venue,
|
||||
replay_at=event.replay_at,
|
||||
replay_sequence=event.replay_sequence,
|
||||
payload=event.payload,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="ReplayEvent"):
|
||||
ReplayPlan(request=make_request(), events=(event_subclass,))
|
||||
|
||||
|
||||
def test_plan_rejects_events_tuple_subclass() -> None:
|
||||
class EventsTupleSubclass(tuple):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="events must be a tuple"):
|
||||
ReplayPlan(
|
||||
request=make_request(),
|
||||
events=EventsTupleSubclass((make_trade_event(),)),
|
||||
)
|
||||
@@ -0,0 +1,699 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.market_data.access.exceptions import (
|
||||
MarketDataAccessIntegrityError,
|
||||
MarketDataAccessOperationError,
|
||||
)
|
||||
from src.market_data.access.models import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.candle import Candle
|
||||
from src.market_data.acquisition.models.quote import Quote
|
||||
from src.market_data.acquisition.models.trade import Trade
|
||||
from src.market_data.replay.contracts import ReplayPlanBuilderProtocol
|
||||
from src.market_data.replay.exceptions import (
|
||||
MarketDataReplayValidationError,
|
||||
ReplayPlanLimitExceededError,
|
||||
)
|
||||
from src.market_data.replay.models import (
|
||||
ReplayDataType,
|
||||
ReplayPlanRequest,
|
||||
)
|
||||
from src.market_data.replay.postgres_replay_plan_builder import (
|
||||
PostgresReplayPlanBuilder,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
END = START + timedelta(hours=1)
|
||||
TRADE_SOURCE = "dzengi_websocket_trade"
|
||||
QUOTE_SOURCE = "dzengi_rest_quote"
|
||||
CANDLE_SOURCE = "dzengi_rest_candle"
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
events: list[tuple[object, ...]],
|
||||
rows: object = (),
|
||||
) -> None:
|
||||
self._events = events
|
||||
self.rows = rows
|
||||
self.calls: list[tuple[str, object]] = []
|
||||
self.execute_errors: list[BaseException | None] = []
|
||||
self.fetchall_error: BaseException | None = None
|
||||
|
||||
def __enter__(self) -> RecordingCursor:
|
||||
self._events.append(("cursor_enter",))
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self._events.append(("cursor_exit", exception_type))
|
||||
return None
|
||||
|
||||
def execute(
|
||||
self,
|
||||
sql: str,
|
||||
parameters: object = None,
|
||||
) -> None:
|
||||
self.calls.append((sql, parameters))
|
||||
self._events.append(("execute", normalized_sql(sql), parameters))
|
||||
|
||||
error = (
|
||||
self.execute_errors.pop(0)
|
||||
if self.execute_errors
|
||||
else None
|
||||
)
|
||||
|
||||
if error is not None:
|
||||
raise error
|
||||
|
||||
def fetchall(self) -> object:
|
||||
self._events.append(("fetchall",))
|
||||
|
||||
if self.fetchall_error is not None:
|
||||
raise self.fetchall_error
|
||||
|
||||
return self.rows
|
||||
|
||||
|
||||
class RecordingTransaction:
|
||||
def __init__(self, events: list[tuple[object, ...]]) -> None:
|
||||
self._events = events
|
||||
|
||||
def __enter__(self) -> RecordingTransaction:
|
||||
self._events.append(("transaction_enter",))
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self._events.append(("transaction_exit", exception_type))
|
||||
return None
|
||||
|
||||
|
||||
class RecordingConnection:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
events: list[tuple[object, ...]],
|
||||
cursor: RecordingCursor,
|
||||
) -> None:
|
||||
self._events = events
|
||||
self._cursor = cursor
|
||||
self.transaction_calls = 0
|
||||
self.cursor_calls = 0
|
||||
|
||||
def __enter__(self) -> RecordingConnection:
|
||||
self._events.append(("connection_enter",))
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exception_type: type[BaseException] | None,
|
||||
exception: BaseException | None,
|
||||
traceback: object,
|
||||
) -> None:
|
||||
self._events.append(("connection_exit", exception_type))
|
||||
return None
|
||||
|
||||
def transaction(self) -> RecordingTransaction:
|
||||
self.transaction_calls += 1
|
||||
self._events.append(("transaction",))
|
||||
return RecordingTransaction(self._events)
|
||||
|
||||
def cursor(self) -> RecordingCursor:
|
||||
self.cursor_calls += 1
|
||||
self._events.append(("cursor",))
|
||||
return self._cursor
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingProvider:
|
||||
connection: RecordingConnection
|
||||
events: list[tuple[object, ...]]
|
||||
calls: int = 0
|
||||
error: BaseException | None = None
|
||||
|
||||
def __call__(self) -> RecordingConnection:
|
||||
self.calls += 1
|
||||
self.events.append(("provider",))
|
||||
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
return self.connection
|
||||
|
||||
|
||||
def normalized_sql(sql: str) -> str:
|
||||
return " ".join(sql.split())
|
||||
|
||||
|
||||
def make_request(
|
||||
*,
|
||||
data_types: tuple[ReplayDataType, ...] = (ReplayDataType.TRADE,),
|
||||
candle_intervals: tuple[str, ...] = (),
|
||||
max_records: int = 10,
|
||||
venue: str = VENUE,
|
||||
) -> ReplayPlanRequest:
|
||||
return ReplayPlanRequest(
|
||||
venue=venue,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=data_types,
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
candle_intervals=candle_intervals,
|
||||
max_records=max_records,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_row(
|
||||
*,
|
||||
replay_at: object = START + timedelta(minutes=1),
|
||||
replay_sequence: object = 10,
|
||||
venue: object = VENUE,
|
||||
trade_id: object = 100,
|
||||
price: object = Decimal("65000"),
|
||||
) -> tuple[object, ...]:
|
||||
return (
|
||||
"trade",
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
venue,
|
||||
SYMBOL,
|
||||
trade_id,
|
||||
replay_at,
|
||||
price,
|
||||
Decimal("0.01"),
|
||||
"buy",
|
||||
TRADE_SOURCE,
|
||||
START + timedelta(seconds=1),
|
||||
START + timedelta(seconds=2),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
[TRADE_SOURCE],
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def make_quote_row(
|
||||
*,
|
||||
replay_at: object = START + timedelta(minutes=2),
|
||||
replay_sequence: object = 11,
|
||||
) -> tuple[object, ...]:
|
||||
return (
|
||||
"quote",
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
QUOTE_SOURCE,
|
||||
None,
|
||||
None,
|
||||
replay_at,
|
||||
START + timedelta(minutes=2) - timedelta(milliseconds=1),
|
||||
Decimal("65001"),
|
||||
Decimal("65000"),
|
||||
Decimal("65002"),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
[QUOTE_SOURCE],
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def make_candle_row(
|
||||
*,
|
||||
replay_at: object = START + timedelta(minutes=3),
|
||||
replay_sequence: object = 12,
|
||||
open_time: object = START - timedelta(hours=1),
|
||||
interval: object = "1m",
|
||||
is_final: object = True,
|
||||
) -> tuple[object, ...]:
|
||||
return (
|
||||
"candle_revision",
|
||||
replay_at,
|
||||
replay_sequence,
|
||||
VENUE,
|
||||
SYMBOL,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
CANDLE_SOURCE,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
interval,
|
||||
open_time,
|
||||
replay_at,
|
||||
Decimal("64900"),
|
||||
Decimal("65100"),
|
||||
Decimal("64800"),
|
||||
Decimal("65000"),
|
||||
Decimal("10"),
|
||||
is_final,
|
||||
[CANDLE_SOURCE],
|
||||
1,
|
||||
)
|
||||
|
||||
|
||||
def dependencies(
|
||||
rows: object = (),
|
||||
) -> tuple[
|
||||
PostgresReplayPlanBuilder,
|
||||
RecordingCursor,
|
||||
RecordingConnection,
|
||||
RecordingProvider,
|
||||
list[tuple[object, ...]],
|
||||
]:
|
||||
events: list[tuple[object, ...]] = []
|
||||
cursor = RecordingCursor(events=events, rows=rows)
|
||||
connection = RecordingConnection(events=events, cursor=cursor)
|
||||
provider = RecordingProvider(connection=connection, events=events)
|
||||
builder = PostgresReplayPlanBuilder(connection_provider=provider)
|
||||
return builder, cursor, connection, provider, events
|
||||
|
||||
|
||||
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
|
||||
builder, _, _, provider, _ = dependencies()
|
||||
|
||||
assert provider.calls == 0
|
||||
assert not hasattr(builder, "__dict__")
|
||||
assert isinstance(builder, ReplayPlanBuilderProtocol)
|
||||
|
||||
|
||||
def test_constructor_rejects_non_callable_provider() -> None:
|
||||
with pytest.raises(TypeError, match="connection_provider"):
|
||||
PostgresReplayPlanBuilder(
|
||||
connection_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_exact_request_is_validated_before_connection_borrow() -> None:
|
||||
class RequestSubclass(ReplayPlanRequest):
|
||||
pass
|
||||
|
||||
builder, _, _, provider, _ = dependencies()
|
||||
request = RequestSubclass(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(START, END),
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="request"):
|
||||
builder.create_plan(request)
|
||||
|
||||
assert provider.calls == 0
|
||||
|
||||
|
||||
def test_trade_snapshot_uses_one_half_open_parameterized_query() -> None:
|
||||
builder, cursor, connection, provider, events = dependencies()
|
||||
request = make_request(max_records=7, venue="tenant'value")
|
||||
|
||||
plan = builder.create_plan(request)
|
||||
|
||||
assert plan.is_empty is True
|
||||
assert provider.calls == 1
|
||||
assert connection.transaction_calls == 1
|
||||
assert connection.cursor_calls == 1
|
||||
assert len(cursor.calls) == 2
|
||||
setup_sql, setup_parameters = cursor.calls[0]
|
||||
sql, parameters = cursor.calls[1]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert normalized_sql(setup_sql) == (
|
||||
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY"
|
||||
)
|
||||
assert setup_parameters is None
|
||||
assert "FROM market_data.trades" in compact_sql
|
||||
assert "FROM market_data.quotes" not in compact_sql
|
||||
assert "FROM market_data.candle_revisions" not in compact_sql
|
||||
assert "executed_at >= %s" in compact_sql
|
||||
assert "executed_at < %s" in compact_sql
|
||||
assert "ORDER BY replay_at ASC, replay_sequence ASC" in compact_sql
|
||||
assert compact_sql.endswith("LIMIT %s")
|
||||
assert "tenant'value" not in sql
|
||||
assert parameters == (
|
||||
"tenant'value",
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
8,
|
||||
)
|
||||
assert events == [
|
||||
("provider",),
|
||||
("connection_enter",),
|
||||
("transaction",),
|
||||
("transaction_enter",),
|
||||
("cursor",),
|
||||
("cursor_enter",),
|
||||
(
|
||||
"execute",
|
||||
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY",
|
||||
None,
|
||||
),
|
||||
("execute", compact_sql, parameters),
|
||||
("fetchall",),
|
||||
("cursor_exit", None),
|
||||
("transaction_exit", None),
|
||||
("connection_exit", None),
|
||||
]
|
||||
|
||||
|
||||
def test_all_requested_types_use_static_union_and_own_time_axes() -> None:
|
||||
builder, cursor, _, _, _ = dependencies()
|
||||
request = make_request(
|
||||
data_types=(
|
||||
ReplayDataType.CANDLE_REVISION,
|
||||
ReplayDataType.TRADE,
|
||||
ReplayDataType.QUOTE,
|
||||
),
|
||||
candle_intervals=("1m", "5m"),
|
||||
max_records=20,
|
||||
)
|
||||
|
||||
builder.create_plan(request)
|
||||
|
||||
sql, parameters = cursor.calls[1]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert compact_sql.count("UNION ALL") == 2
|
||||
assert "executed_at >= %s" in compact_sql
|
||||
assert "received_at >= %s" in compact_sql
|
||||
assert "observed_at >= %s" in compact_sql
|
||||
assert "open_time >= %s" not in compact_sql
|
||||
assert "interval = ANY(%s)" in compact_sql
|
||||
assert parameters == (
|
||||
VENUE,
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
VENUE,
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
VENUE,
|
||||
[SYMBOL],
|
||||
START,
|
||||
END,
|
||||
["1m", "5m"],
|
||||
21,
|
||||
)
|
||||
|
||||
|
||||
def test_materializes_globally_ordered_canonical_events() -> None:
|
||||
event_time = START + timedelta(minutes=2)
|
||||
builder, _, _, _, _ = dependencies(
|
||||
[
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=10,
|
||||
),
|
||||
make_quote_row(
|
||||
replay_at=event_time,
|
||||
replay_sequence=11,
|
||||
),
|
||||
make_candle_row(
|
||||
replay_at=event_time,
|
||||
replay_sequence=12,
|
||||
),
|
||||
]
|
||||
)
|
||||
request = make_request(
|
||||
data_types=(
|
||||
ReplayDataType.TRADE,
|
||||
ReplayDataType.QUOTE,
|
||||
ReplayDataType.CANDLE_REVISION,
|
||||
),
|
||||
candle_intervals=("1m",),
|
||||
)
|
||||
|
||||
plan = builder.create_plan(request)
|
||||
|
||||
assert [event.replay_sequence for event in plan.events] == [10, 11, 12]
|
||||
assert type(plan.events[0].payload) is Trade
|
||||
assert type(plan.events[1].payload) is Quote
|
||||
candle_payload = plan.events[2].payload
|
||||
assert type(candle_payload) is Candle
|
||||
assert isinstance(candle_payload, Candle)
|
||||
assert plan.events[2].replay_at == event_time
|
||||
assert candle_payload.open_time < START
|
||||
assert plan.events[2].candle_is_final is True
|
||||
|
||||
|
||||
def test_candle_snapshot_filters_by_observed_time_not_open_time() -> None:
|
||||
builder, cursor, _, _, _ = dependencies(
|
||||
[make_candle_row(open_time=START - timedelta(days=1))]
|
||||
)
|
||||
request = make_request(
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
candle_intervals=("1m",),
|
||||
)
|
||||
|
||||
plan = builder.create_plan(request)
|
||||
|
||||
sql, _ = cursor.calls[1]
|
||||
compact_sql = normalized_sql(sql)
|
||||
assert "observed_at >= %s" in compact_sql
|
||||
assert "observed_at < %s" in compact_sql
|
||||
assert "open_time >= %s" not in compact_sql
|
||||
assert len(plan.events) == 1
|
||||
|
||||
|
||||
def test_limit_plus_one_raises_without_partial_plan() -> None:
|
||||
builder, _, _, _, events = dependencies(
|
||||
[
|
||||
make_trade_row(replay_sequence=10),
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=11,
|
||||
trade_id=101,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(ReplayPlanLimitExceededError, match="max_records"):
|
||||
builder.create_plan(make_request(max_records=1))
|
||||
|
||||
assert events[-3:] == [
|
||||
("cursor_exit", None),
|
||||
("transaction_exit", None),
|
||||
("connection_exit", None),
|
||||
]
|
||||
|
||||
|
||||
def test_more_than_limit_plus_one_is_backend_integrity_error() -> None:
|
||||
builder, _, _, _, _ = dependencies(
|
||||
[
|
||||
make_trade_row(replay_sequence=10),
|
||||
make_trade_row(replay_sequence=11),
|
||||
make_trade_row(replay_sequence=12),
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
|
||||
builder.create_plan(make_request(max_records=1))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rows", (None, 1, "rows", b"rows"))
|
||||
def test_rejects_invalid_rows_container(rows: object) -> None:
|
||||
builder, _, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"row",
|
||||
(
|
||||
(),
|
||||
("trade",),
|
||||
("unknown",) + make_trade_row()[1:],
|
||||
(1,) + make_trade_row()[1:],
|
||||
),
|
||||
)
|
||||
def test_rejects_invalid_snapshot_row_shape_or_type(
|
||||
row: tuple[object, ...],
|
||||
) -> None:
|
||||
builder, _, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"row, replay_request",
|
||||
(
|
||||
(
|
||||
make_trade_row(price=Decimal("0")),
|
||||
make_request(),
|
||||
),
|
||||
(
|
||||
make_quote_row()[:-1] + (2,),
|
||||
make_request(data_types=(ReplayDataType.QUOTE,)),
|
||||
),
|
||||
(
|
||||
make_candle_row(is_final=1),
|
||||
make_request(
|
||||
data_types=(ReplayDataType.CANDLE_REVISION,),
|
||||
candle_intervals=("1m",),
|
||||
),
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_shared_mappers_reject_corrupt_canonical_values(
|
||||
row: tuple[object, ...],
|
||||
replay_request: ReplayPlanRequest,
|
||||
) -> None:
|
||||
builder, _, _, _, _ = dependencies([row])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError):
|
||||
builder.create_plan(replay_request)
|
||||
|
||||
|
||||
def test_rejects_replay_time_that_differs_from_payload_time() -> None:
|
||||
row = list(make_trade_row())
|
||||
row[1] = START + timedelta(minutes=2)
|
||||
builder, _, _, _, _ = dependencies([tuple(row)])
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="event"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rows",
|
||||
(
|
||||
(
|
||||
make_trade_row(replay_sequence=10),
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=10,
|
||||
trade_id=101,
|
||||
),
|
||||
),
|
||||
(
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=2),
|
||||
replay_sequence=11,
|
||||
),
|
||||
make_trade_row(
|
||||
replay_at=START + timedelta(minutes=1),
|
||||
replay_sequence=10,
|
||||
trade_id=101,
|
||||
),
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_rejects_duplicate_sequence_or_unordered_rows(
|
||||
rows: tuple[tuple[object, ...], ...],
|
||||
) -> None:
|
||||
builder, _, _, _, _ = dependencies(rows)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="snapshot"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
def test_rejects_event_outside_request_scope() -> None:
|
||||
builder, _, _, _, _ = dependencies(
|
||||
[make_trade_row(venue="another")]
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataAccessIntegrityError, match="snapshot"):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure_point", ("provider", "execute", "fetchall"))
|
||||
def test_backend_error_is_wrapped_with_original_cause(
|
||||
failure_point: str,
|
||||
) -> None:
|
||||
builder, cursor, _, provider, events = dependencies()
|
||||
backend_error = RuntimeError("backend failed")
|
||||
|
||||
if failure_point == "provider":
|
||||
provider.error = backend_error
|
||||
elif failure_point == "execute":
|
||||
cursor.execute_errors = [None, backend_error]
|
||||
else:
|
||||
cursor.fetchall_error = backend_error
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as raised:
|
||||
builder.create_plan(make_request())
|
||||
|
||||
assert raised.value.__cause__ is backend_error
|
||||
|
||||
if failure_point != "provider":
|
||||
assert events[-3:] == [
|
||||
("cursor_exit", RuntimeError),
|
||||
("transaction_exit", RuntimeError),
|
||||
("connection_exit", RuntimeError),
|
||||
]
|
||||
|
||||
|
||||
def test_transaction_setup_error_is_wrapped_and_query_is_not_executed() -> None:
|
||||
builder, cursor, _, _, _ = dependencies()
|
||||
setup_error = RuntimeError("cannot configure transaction")
|
||||
cursor.execute_errors = [setup_error]
|
||||
|
||||
with pytest.raises(MarketDataAccessOperationError) as raised:
|
||||
builder.create_plan(make_request())
|
||||
|
||||
assert raised.value.__cause__ is setup_error
|
||||
assert len(cursor.calls) == 1
|
||||
|
||||
|
||||
def test_base_exception_is_not_swallowed_and_contexts_are_closed() -> None:
|
||||
builder, cursor, _, _, events = dependencies()
|
||||
cursor.fetchall_error = KeyboardInterrupt()
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
builder.create_plan(make_request())
|
||||
|
||||
assert events[-3:] == [
|
||||
("cursor_exit", KeyboardInterrupt),
|
||||
("transaction_exit", KeyboardInterrupt),
|
||||
("connection_exit", KeyboardInterrupt),
|
||||
]
|
||||
807
app/tests/unit/market_data/replay/test_replay_session.py
Normal file
807
app/tests/unit/market_data/replay/test_replay_session.py
Normal file
@@ -0,0 +1,807 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
import src.market_data.replay as replay_package
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.replay import (
|
||||
MarketDataReplayValidationError,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanRequest,
|
||||
ReplaySession,
|
||||
ReplaySessionProtocol,
|
||||
ReplaySessionState,
|
||||
ReplaySessionStateError,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
FIRST_TIME = START + timedelta(minutes=1)
|
||||
SECOND_TIME = START + timedelta(minutes=2)
|
||||
THIRD_TIME = START + timedelta(minutes=3)
|
||||
END = START + timedelta(hours=1)
|
||||
|
||||
|
||||
def make_trade_event(
|
||||
*,
|
||||
trade_id: int,
|
||||
replay_at: datetime,
|
||||
replay_sequence: int,
|
||||
) -> ReplayEvent:
|
||||
trade = Trade(
|
||||
symbol=SYMBOL,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=replay_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source="dzengi_websocket_trade",
|
||||
)
|
||||
return ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=replay_sequence,
|
||||
payload=trade,
|
||||
)
|
||||
|
||||
|
||||
def make_plan(*, empty: bool = False) -> ReplayPlan:
|
||||
request = ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=START,
|
||||
end_time=END,
|
||||
),
|
||||
max_records=100,
|
||||
)
|
||||
events = () if empty else (
|
||||
make_trade_event(
|
||||
trade_id=1,
|
||||
replay_at=FIRST_TIME,
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=2,
|
||||
replay_at=FIRST_TIME,
|
||||
replay_sequence=2,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=3,
|
||||
replay_at=SECOND_TIME,
|
||||
replay_sequence=3,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=4,
|
||||
replay_at=THIRD_TIME,
|
||||
replay_sequence=4,
|
||||
),
|
||||
)
|
||||
return ReplayPlan(request=request, events=events)
|
||||
|
||||
|
||||
class RecordingClock:
|
||||
def __init__(self, now: datetime = START) -> None:
|
||||
self._now = now
|
||||
self.advance_calls: list[datetime] = []
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return self._now
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self.advance_calls.append(instant)
|
||||
self._now = instant
|
||||
|
||||
|
||||
class RecordingConsumer:
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
self.clock = clock
|
||||
self.events: list[ReplayEvent] = []
|
||||
self.observed_times: list[datetime] = []
|
||||
self.tasks: list[asyncio.Task[Any] | None] = []
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.events.append(event)
|
||||
self.observed_times.append(self.clock.now)
|
||||
self.tasks.append(asyncio.current_task())
|
||||
|
||||
|
||||
def make_session(
|
||||
*,
|
||||
plan: ReplayPlan | None = None,
|
||||
clock: RecordingClock | None = None,
|
||||
consumer: RecordingConsumer | None = None,
|
||||
) -> tuple[ReplaySession, ReplayPlan, RecordingClock, RecordingConsumer]:
|
||||
resolved_plan = make_plan() if plan is None else plan
|
||||
resolved_clock = (
|
||||
RecordingClock(resolved_plan.request.time_range.start_time)
|
||||
if clock is None
|
||||
else clock
|
||||
)
|
||||
resolved_consumer = (
|
||||
RecordingConsumer(resolved_clock)
|
||||
if consumer is None
|
||||
else consumer
|
||||
)
|
||||
session = ReplaySession(
|
||||
plan=resolved_plan,
|
||||
clock=resolved_clock,
|
||||
consumer=resolved_consumer,
|
||||
)
|
||||
return (
|
||||
session,
|
||||
resolved_plan,
|
||||
resolved_clock,
|
||||
resolved_consumer,
|
||||
)
|
||||
|
||||
|
||||
def test_matches_protocol_uses_slots_and_preserves_dependencies() -> None:
|
||||
session, plan, clock, _ = make_session()
|
||||
|
||||
assert isinstance(session, ReplaySessionProtocol)
|
||||
assert not hasattr(session, "__dict__")
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert session.plan is plan
|
||||
assert session.clock is clock
|
||||
assert replay_package.ReplaySession is ReplaySession
|
||||
assert not hasattr(replay_package, "ReplayEngine")
|
||||
|
||||
|
||||
def test_construction_does_not_advance_clock_or_call_consumer() -> None:
|
||||
session, _, clock, consumer = make_session()
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert clock.advance_calls == []
|
||||
assert consumer.events == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "plan"))
|
||||
def test_rejects_non_plan(invalid: object) -> None:
|
||||
clock = RecordingClock()
|
||||
|
||||
with pytest.raises(TypeError, match="plan"):
|
||||
ReplaySession(
|
||||
plan=invalid, # type: ignore[arg-type]
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_plan_subclass() -> None:
|
||||
class ReplayPlanSubclass(ReplayPlan):
|
||||
pass
|
||||
|
||||
plan = make_plan()
|
||||
subclass = ReplayPlanSubclass(
|
||||
request=plan.request,
|
||||
events=plan.events,
|
||||
)
|
||||
clock = RecordingClock()
|
||||
|
||||
with pytest.raises(TypeError, match="plan"):
|
||||
ReplaySession(
|
||||
plan=subclass,
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_clock_without_protocol() -> None:
|
||||
class InvalidClock:
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return START
|
||||
|
||||
with pytest.raises(TypeError, match="clock"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=InvalidClock(), # type: ignore[arg-type]
|
||||
consumer=RecordingConsumer(RecordingClock()),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_clock_class_before_any_runtime_action() -> None:
|
||||
class ClockClass:
|
||||
advance_calls: list[datetime] = []
|
||||
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return START
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self.advance_calls.append(instant)
|
||||
|
||||
consumer_clock = RecordingClock()
|
||||
consumer = RecordingConsumer(consumer_clock)
|
||||
|
||||
with pytest.raises(TypeError, match="clock"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=ClockClass, # type: ignore[arg-type]
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
assert ClockClass.advance_calls == []
|
||||
assert consumer.events == []
|
||||
|
||||
|
||||
def test_rejects_asynchronous_clock_advance() -> None:
|
||||
class AsyncClock:
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return START
|
||||
|
||||
async def advance_to(self, instant: datetime) -> None:
|
||||
return None
|
||||
|
||||
clock = AsyncClock()
|
||||
|
||||
with pytest.raises(TypeError, match="synchronous"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock, # type: ignore[arg-type]
|
||||
consumer=RecordingConsumer(RecordingClock()),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_consumer_without_protocol() -> None:
|
||||
class InvalidConsumer:
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=RecordingClock(),
|
||||
consumer=InvalidConsumer(), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_consumer_class_before_advancing_clock() -> None:
|
||||
clock = RecordingClock()
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
assert clock.now == START
|
||||
assert clock.advance_calls == []
|
||||
|
||||
|
||||
def test_rejects_synchronous_consumer() -> None:
|
||||
class SyncConsumer:
|
||||
def consume(self, event: ReplayEvent) -> None:
|
||||
return None
|
||||
|
||||
with pytest.raises(TypeError, match="asynchronous"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=RecordingClock(),
|
||||
consumer=SyncConsumer(), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"now",
|
||||
(
|
||||
START.replace(tzinfo=None),
|
||||
START.astimezone(timezone(timedelta(hours=3))),
|
||||
START.replace(fold=1),
|
||||
),
|
||||
)
|
||||
def test_rejects_non_canonical_clock_time(now: datetime) -> None:
|
||||
clock = RecordingClock(now)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
assert clock.advance_calls == []
|
||||
|
||||
|
||||
def test_rejects_datetime_subclass_from_clock() -> None:
|
||||
class CompatibleDatetime(datetime):
|
||||
pass
|
||||
|
||||
now = CompatibleDatetime(
|
||||
2026,
|
||||
8,
|
||||
2,
|
||||
12,
|
||||
0,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
clock = RecordingClock(now)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_non_datetime_clock_time() -> None:
|
||||
class InvalidNowClock:
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
return object() # type: ignore[return-value]
|
||||
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
return None
|
||||
|
||||
clock = InvalidNowClock()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(RecordingClock()),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"now",
|
||||
(START - timedelta(microseconds=1), START + timedelta(microseconds=1)),
|
||||
)
|
||||
def test_rejects_clock_outside_exact_start_without_reset(
|
||||
now: datetime,
|
||||
) -> None:
|
||||
clock = RecordingClock(now)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="start"):
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
assert clock.now is now
|
||||
assert clock.advance_calls == []
|
||||
|
||||
|
||||
def test_clock_now_error_is_not_swallowed() -> None:
|
||||
expected = RuntimeError("clock now failed")
|
||||
|
||||
class BrokenClock(RecordingClock):
|
||||
@property
|
||||
def now(self) -> datetime:
|
||||
raise expected
|
||||
|
||||
clock = BrokenClock()
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
ReplaySession(
|
||||
plan=make_plan(),
|
||||
clock=clock,
|
||||
consumer=RecordingConsumer(clock),
|
||||
)
|
||||
|
||||
assert captured.value is expected
|
||||
|
||||
|
||||
def test_run_preserves_order_identity_and_advances_before_consumer() -> None:
|
||||
async def scenario() -> None:
|
||||
session, plan, clock, consumer = make_session()
|
||||
caller_task = asyncio.current_task()
|
||||
|
||||
result = await session.run()
|
||||
|
||||
assert result is None
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
assert all(
|
||||
actual is expected
|
||||
for actual, expected in zip(
|
||||
consumer.events,
|
||||
plan.events,
|
||||
strict=True,
|
||||
)
|
||||
)
|
||||
assert clock.advance_calls == [
|
||||
event.replay_at for event in plan.events
|
||||
]
|
||||
assert consumer.observed_times == clock.advance_calls
|
||||
assert all(task is caller_task for task in consumer.tasks)
|
||||
assert clock.now == plan.events[-1].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_empty_plan_completes_without_dependency_calls() -> None:
|
||||
async def scenario() -> None:
|
||||
plan = make_plan(empty=True)
|
||||
session, _, clock, consumer = make_session(plan=plan)
|
||||
initial_time = clock.now
|
||||
|
||||
result = await session.run()
|
||||
|
||||
assert result is None
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert clock.now is initial_time
|
||||
assert clock.advance_calls == []
|
||||
assert consumer.events == []
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_delivery_is_strictly_sequential() -> None:
|
||||
class YieldingConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.active = 0
|
||||
self.maximum_active = 0
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.active += 1
|
||||
self.maximum_active = max(self.maximum_active, self.active)
|
||||
try:
|
||||
await asyncio.sleep(0)
|
||||
await super().consume(event)
|
||||
finally:
|
||||
self.active -= 1
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = YieldingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
await session.run()
|
||||
|
||||
assert consumer.maximum_active == 1
|
||||
assert tuple(consumer.events) == plan.events
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_clock_error_fails_without_delivering_current_or_suffix() -> None:
|
||||
expected = RuntimeError("clock failed")
|
||||
|
||||
class BrokenClock(RecordingClock):
|
||||
def advance_to(self, instant: datetime) -> None:
|
||||
self.advance_calls.append(instant)
|
||||
if len(self.advance_calls) == 3:
|
||||
raise expected
|
||||
self._now = instant
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = BrokenClock()
|
||||
consumer = RecordingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert tuple(consumer.events) == plan.events[:2]
|
||||
assert clock.advance_calls == [
|
||||
event.replay_at for event in plan.events[:3]
|
||||
]
|
||||
assert clock.now == plan.events[1].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_consumer_error_fails_without_retry_rollback_or_suffix() -> None:
|
||||
expected = RuntimeError("consumer failed")
|
||||
|
||||
class BrokenConsumer(RecordingConsumer):
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
if len(self.events) == 3:
|
||||
raise expected
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = BrokenConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert tuple(consumer.events) == plan.events[:3]
|
||||
assert clock.now == plan.events[2].replay_at
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert tuple(consumer.events) == plan.events[:3]
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_non_exception_base_error_preserves_identity_and_failed_state() -> None:
|
||||
class ReplaySignal(BaseException):
|
||||
pass
|
||||
|
||||
expected = ReplaySignal("stop")
|
||||
|
||||
class BrokenConsumer(RecordingConsumer):
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
raise expected
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=BrokenConsumer(clock),
|
||||
)
|
||||
|
||||
with pytest.raises(ReplaySignal) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert clock.now == plan.events[0].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_consumer_cancellation_preserves_identity_and_cancelled_state() -> None:
|
||||
expected = asyncio.CancelledError("consumer cancelled")
|
||||
|
||||
class CancellingConsumer(RecordingConsumer):
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
raise expected
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = CancellingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError) as captured:
|
||||
await session.run()
|
||||
|
||||
assert captured.value is expected
|
||||
assert session.state is ReplaySessionState.CANCELLED
|
||||
assert consumer.events == [plan.events[0]]
|
||||
assert clock.now == plan.events[0].replay_at
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.CANCELLED
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_external_cancellation_during_consumer_stays_cancelled() -> None:
|
||||
class BlockingConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.entered = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
self.entered.set()
|
||||
await self.release.wait()
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = BlockingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
task = asyncio.create_task(session.run())
|
||||
await consumer.entered.wait()
|
||||
|
||||
assert session.state is ReplaySessionState.RUNNING
|
||||
task.cancel()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert session.state is ReplaySessionState.CANCELLED
|
||||
assert consumer.events == [plan.events[0]]
|
||||
assert clock.now == plan.events[0].replay_at
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_concurrent_caller_is_rejected_without_damaging_first() -> None:
|
||||
class BlockingConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.entered = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
if len(self.events) == 1:
|
||||
self.entered.set()
|
||||
await self.release.wait()
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = BlockingConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
first = asyncio.create_task(session.run())
|
||||
await consumer.entered.wait()
|
||||
|
||||
assert first.done() is False
|
||||
assert session.state is ReplaySessionState.RUNNING
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert first.done() is False
|
||||
assert session.state is ReplaySessionState.RUNNING
|
||||
consumer.release.set()
|
||||
await first
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_caught_reentrant_run_does_not_damage_outer_run() -> None:
|
||||
class ReentrantConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.session: ReplaySession | None = None
|
||||
self.reentrant_errors: list[ReplaySessionStateError] = []
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
assert self.session is not None
|
||||
try:
|
||||
await self.session.run()
|
||||
except ReplaySessionStateError as error:
|
||||
self.reentrant_errors.append(error)
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = ReentrantConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
consumer.session = session
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
assert len(consumer.reentrant_errors) == len(plan.events)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_uncaught_reentrant_run_fails_outer_run() -> None:
|
||||
class ReentrantConsumer(RecordingConsumer):
|
||||
def __init__(self, clock: RecordingClock) -> None:
|
||||
super().__init__(clock)
|
||||
self.session: ReplaySession | None = None
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
await super().consume(event)
|
||||
assert self.session is not None
|
||||
await self.session.run()
|
||||
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
clock = RecordingClock()
|
||||
consumer = ReentrantConsumer(clock)
|
||||
session, *_ = make_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
consumer.session = session
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.FAILED
|
||||
assert consumer.events == [plan.events[0]]
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_repeated_run_after_completion_is_rejected() -> None:
|
||||
async def scenario() -> None:
|
||||
session, plan, _, consumer = make_session()
|
||||
await session.run()
|
||||
|
||||
with pytest.raises(ReplaySessionStateError):
|
||||
await session.run()
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert tuple(consumer.events) == plan.events
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_two_sessions_share_immutable_plan_but_not_state_or_clock() -> None:
|
||||
async def scenario() -> None:
|
||||
plan = make_plan()
|
||||
first, _, first_clock, first_consumer = make_session(plan=plan)
|
||||
second, _, second_clock, second_consumer = make_session(plan=plan)
|
||||
|
||||
await first.run()
|
||||
|
||||
assert first.state is ReplaySessionState.COMPLETED
|
||||
assert second.state is ReplaySessionState.CREATED
|
||||
assert first.plan is second.plan is plan
|
||||
assert first_clock is not second_clock
|
||||
assert second_clock.now == START
|
||||
assert second_consumer.events == []
|
||||
|
||||
await second.run()
|
||||
|
||||
assert tuple(first_consumer.events) == plan.events
|
||||
assert tuple(second_consumer.events) == plan.events
|
||||
assert second.state is ReplaySessionState.COMPLETED
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_lifecycle_and_hidden_task_extensions_are_absent() -> None:
|
||||
session, *_ = make_session()
|
||||
|
||||
assert not hasattr(session, "start")
|
||||
assert not hasattr(session, "stop")
|
||||
assert not hasattr(session, "close")
|
||||
assert not hasattr(session, "reset")
|
||||
assert not hasattr(session, "pause")
|
||||
assert not hasattr(session, "resume")
|
||||
883
app/tests/unit/market_data/replay/test_replay_session_factory.py
Normal file
883
app/tests/unit/market_data/replay/test_replay_session_factory.py
Normal file
@@ -0,0 +1,883 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
|
||||
import src.market_data.replay as replay_package
|
||||
import src.market_data.replay.replay_session_factory as factory_module
|
||||
from src.market_data.access import HistoricalTimeRange
|
||||
from src.market_data.acquisition.models.trade import (
|
||||
Trade,
|
||||
TradeAggressorSide,
|
||||
)
|
||||
from src.market_data.replay import (
|
||||
DeterministicReplayClock,
|
||||
MarketDataClockProtocol,
|
||||
MarketDataReplayValidationError,
|
||||
ReplayConsumerFactoryProtocol,
|
||||
ReplayConsumerProtocol,
|
||||
ReplayClockProtocol,
|
||||
ReplayDataType,
|
||||
ReplayEvent,
|
||||
ReplayPlan,
|
||||
ReplayPlanBuilderProtocol,
|
||||
ReplayPlanRequest,
|
||||
ReplaySession,
|
||||
ReplaySessionFactory,
|
||||
ReplaySessionState,
|
||||
)
|
||||
|
||||
|
||||
VENUE = "dzengi"
|
||||
SYMBOL = "BTC/USD_LEVERAGE"
|
||||
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
|
||||
FIRST_TIME = START + timedelta(minutes=1)
|
||||
SECOND_TIME = START + timedelta(minutes=2)
|
||||
END = START + timedelta(hours=1)
|
||||
|
||||
|
||||
def make_request() -> ReplayPlanRequest:
|
||||
return ReplayPlanRequest(
|
||||
venue=VENUE,
|
||||
symbols=(SYMBOL,),
|
||||
data_types=(ReplayDataType.TRADE,),
|
||||
time_range=HistoricalTimeRange(
|
||||
start_time=START,
|
||||
end_time=END,
|
||||
),
|
||||
max_records=100,
|
||||
)
|
||||
|
||||
|
||||
def make_trade_event(
|
||||
*,
|
||||
trade_id: int,
|
||||
replay_at: datetime,
|
||||
replay_sequence: int,
|
||||
) -> ReplayEvent:
|
||||
return ReplayEvent(
|
||||
venue=VENUE,
|
||||
replay_at=replay_at,
|
||||
replay_sequence=replay_sequence,
|
||||
payload=Trade(
|
||||
symbol=SYMBOL,
|
||||
trade_id=trade_id,
|
||||
price=Decimal("65000"),
|
||||
quantity=Decimal("0.001"),
|
||||
executed_at=replay_at,
|
||||
aggressor_side=TradeAggressorSide.BUY,
|
||||
source="dzengi_websocket_trade",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def make_plan(
|
||||
*,
|
||||
request: ReplayPlanRequest | None = None,
|
||||
empty: bool = False,
|
||||
) -> ReplayPlan:
|
||||
resolved_request = make_request() if request is None else request
|
||||
events = () if empty else (
|
||||
make_trade_event(
|
||||
trade_id=1,
|
||||
replay_at=FIRST_TIME,
|
||||
replay_sequence=1,
|
||||
),
|
||||
make_trade_event(
|
||||
trade_id=2,
|
||||
replay_at=SECOND_TIME,
|
||||
replay_sequence=2,
|
||||
),
|
||||
)
|
||||
return ReplayPlan(
|
||||
request=resolved_request,
|
||||
events=events,
|
||||
)
|
||||
|
||||
|
||||
class RecordingConsumer:
|
||||
def __init__(self, clock: MarketDataClockProtocol) -> None:
|
||||
self.clock = clock
|
||||
self.events: list[ReplayEvent] = []
|
||||
self.observed_times: list[datetime] = []
|
||||
self.start_calls = 0
|
||||
self.stop_calls = 0
|
||||
self.close_calls = 0
|
||||
|
||||
async def consume(self, event: ReplayEvent) -> None:
|
||||
self.events.append(event)
|
||||
self.observed_times.append(self.clock.now)
|
||||
|
||||
def start(self) -> None:
|
||||
self.start_calls += 1
|
||||
|
||||
def stop(self) -> None:
|
||||
self.stop_calls += 1
|
||||
|
||||
def close(self) -> None:
|
||||
self.close_calls += 1
|
||||
|
||||
|
||||
class RecordingPlanBuilder:
|
||||
def __init__(
|
||||
self,
|
||||
plan: ReplayPlan,
|
||||
*,
|
||||
actions: list[str] | None = None,
|
||||
) -> None:
|
||||
self.plan = plan
|
||||
self.actions = actions
|
||||
self.requests: list[ReplayPlanRequest] = []
|
||||
self.thread_ids: list[int] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
if self.actions is not None:
|
||||
self.actions.append("builder")
|
||||
self.requests.append(request)
|
||||
self.thread_ids.append(threading.get_ident())
|
||||
return self.plan
|
||||
|
||||
|
||||
class RecordingConsumerFactory:
|
||||
def __init__(self, *, actions: list[str] | None = None) -> None:
|
||||
self.actions = actions
|
||||
self.plans: list[ReplayPlan] = []
|
||||
self.clocks: list[MarketDataClockProtocol] = []
|
||||
self.consumers: list[RecordingConsumer] = []
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
if self.actions is not None:
|
||||
self.actions.append("consumer")
|
||||
consumer = RecordingConsumer(clock)
|
||||
self.plans.append(plan)
|
||||
self.clocks.append(clock)
|
||||
self.consumers.append(consumer)
|
||||
return consumer
|
||||
|
||||
|
||||
def create_factory(
|
||||
*,
|
||||
plan: ReplayPlan | None = None,
|
||||
) -> tuple[
|
||||
ReplaySessionFactory,
|
||||
ReplayPlan,
|
||||
RecordingPlanBuilder,
|
||||
RecordingConsumerFactory,
|
||||
]:
|
||||
resolved_plan = make_plan() if plan is None else plan
|
||||
plan_builder = RecordingPlanBuilder(resolved_plan)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
return (
|
||||
ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
),
|
||||
resolved_plan,
|
||||
plan_builder,
|
||||
consumer_factory,
|
||||
)
|
||||
|
||||
|
||||
def test_matches_protocols_uses_slots_and_is_exported() -> None:
|
||||
factory, _, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
assert isinstance(plan_builder, ReplayPlanBuilderProtocol)
|
||||
assert isinstance(consumer_factory, ReplayConsumerFactoryProtocol)
|
||||
assert not hasattr(factory, "__dict__")
|
||||
assert replay_package.ReplaySessionFactory is ReplaySessionFactory
|
||||
assert (
|
||||
replay_package.ReplayConsumerFactoryProtocol
|
||||
is ReplayConsumerFactoryProtocol
|
||||
)
|
||||
|
||||
|
||||
def test_constructor_only_preserves_dependencies() -> None:
|
||||
plan = make_plan()
|
||||
plan_builder = RecordingPlanBuilder(plan)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
assert isinstance(factory, ReplaySessionFactory)
|
||||
assert plan_builder.requests == []
|
||||
assert consumer_factory.plans == []
|
||||
assert consumer_factory.clocks == []
|
||||
assert consumer_factory.consumers == []
|
||||
|
||||
|
||||
def test_constructor_does_not_create_clock_session_or_task(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class ForbiddenClock:
|
||||
def __init__(self, initial_time: datetime) -> None:
|
||||
raise AssertionError("Clock must not be created")
|
||||
|
||||
def forbidden_session(**kwargs: object) -> ReplaySession:
|
||||
raise AssertionError("Session must not be created")
|
||||
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"DeterministicReplayClock",
|
||||
ForbiddenClock,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"ReplaySession",
|
||||
forbidden_session,
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
tasks_before = set(asyncio.all_tasks())
|
||||
|
||||
factory, _, _, _ = create_factory()
|
||||
|
||||
assert isinstance(factory, ReplaySessionFactory)
|
||||
assert set(asyncio.all_tasks()) == tasks_before
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "builder"))
|
||||
def test_rejects_invalid_plan_builder(invalid: object) -> None:
|
||||
with pytest.raises(TypeError, match="plan_builder"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=cast(ReplayPlanBuilderProtocol, invalid),
|
||||
consumer_factory=RecordingConsumerFactory(),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_plan_builder_class() -> None:
|
||||
class PlanBuilderClass:
|
||||
def create_plan(
|
||||
self,
|
||||
request: ReplayPlanRequest,
|
||||
) -> ReplayPlan:
|
||||
return make_plan(request=request)
|
||||
|
||||
with pytest.raises(TypeError, match="plan_builder"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=cast(
|
||||
ReplayPlanBuilderProtocol,
|
||||
PlanBuilderClass,
|
||||
),
|
||||
consumer_factory=RecordingConsumerFactory(),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_asynchronous_plan_builder() -> None:
|
||||
class AsyncPlanBuilder:
|
||||
async def create_plan(
|
||||
self,
|
||||
request: ReplayPlanRequest,
|
||||
) -> ReplayPlan:
|
||||
return make_plan(request=request)
|
||||
|
||||
with pytest.raises(TypeError, match="synchronous"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=cast(
|
||||
ReplayPlanBuilderProtocol,
|
||||
AsyncPlanBuilder(),
|
||||
),
|
||||
consumer_factory=RecordingConsumerFactory(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "factory"))
|
||||
def test_rejects_invalid_consumer_factory(invalid: object) -> None:
|
||||
with pytest.raises(TypeError, match="consumer_factory"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan()),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
invalid,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_consumer_factory_class() -> None:
|
||||
class ConsumerFactoryClass:
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
return RecordingConsumer(clock)
|
||||
|
||||
with pytest.raises(TypeError, match="consumer_factory"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan()),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
ConsumerFactoryClass,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_asynchronous_consumer_factory() -> None:
|
||||
class AsyncConsumerFactory:
|
||||
async def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
return RecordingConsumer(clock)
|
||||
|
||||
with pytest.raises(TypeError, match="synchronous"):
|
||||
ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan()),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
AsyncConsumerFactory(),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_prepare_is_synchronous_and_run_remains_asynchronous() -> None:
|
||||
assert not inspect.iscoroutinefunction(
|
||||
ReplaySessionFactory.prepare_session
|
||||
)
|
||||
assert inspect.iscoroutinefunction(ReplaySession.run)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "request"))
|
||||
def test_rejects_invalid_request_before_dependencies(
|
||||
invalid: object,
|
||||
) -> None:
|
||||
factory, _, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="request"):
|
||||
factory.prepare_session(
|
||||
cast(ReplayPlanRequest, invalid)
|
||||
)
|
||||
|
||||
assert plan_builder.requests == []
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_rejects_request_subclass_before_dependencies() -> None:
|
||||
class ReplayPlanRequestSubclass(ReplayPlanRequest):
|
||||
pass
|
||||
|
||||
request = make_request()
|
||||
subclass = ReplayPlanRequestSubclass(
|
||||
venue=request.venue,
|
||||
symbols=request.symbols,
|
||||
data_types=request.data_types,
|
||||
time_range=request.time_range,
|
||||
candle_intervals=request.candle_intervals,
|
||||
max_records=request.max_records,
|
||||
)
|
||||
factory, _, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="request"):
|
||||
factory.prepare_session(subclass)
|
||||
|
||||
assert plan_builder.requests == []
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_builds_in_order_and_preserves_all_identities(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
actions: list[str] = []
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
plan_builder = RecordingPlanBuilder(plan, actions=actions)
|
||||
consumer_factory = RecordingConsumerFactory(actions=actions)
|
||||
real_session = ReplaySession
|
||||
|
||||
class OrderedClock(DeterministicReplayClock):
|
||||
def __init__(self, initial_time: datetime) -> None:
|
||||
actions.append("clock")
|
||||
super().__init__(initial_time)
|
||||
|
||||
def ordered_session(
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: ReplayClockProtocol,
|
||||
consumer: ReplayConsumerProtocol,
|
||||
) -> ReplaySession:
|
||||
actions.append("session")
|
||||
return real_session(
|
||||
plan=plan,
|
||||
clock=clock,
|
||||
consumer=consumer,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"DeterministicReplayClock",
|
||||
OrderedClock,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
factory_module,
|
||||
"ReplaySession",
|
||||
ordered_session,
|
||||
)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
session = factory.prepare_session(request)
|
||||
clock = consumer_factory.clocks[0]
|
||||
|
||||
assert actions == ["builder", "clock", "consumer", "session"]
|
||||
assert plan_builder.requests == [request]
|
||||
assert plan_builder.requests[0] is request
|
||||
assert consumer_factory.plans == [plan]
|
||||
assert consumer_factory.plans[0] is plan
|
||||
assert session.plan is plan
|
||||
assert session.clock is clock
|
||||
assert isinstance(clock, OrderedClock)
|
||||
assert clock.now == START
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
|
||||
|
||||
def test_prepare_runs_builder_in_caller_thread() -> None:
|
||||
factory, plan, plan_builder, consumer_factory = create_factory()
|
||||
caller_thread_id = threading.get_ident()
|
||||
|
||||
session = factory.prepare_session(plan.request)
|
||||
|
||||
assert isinstance(session, ReplaySession)
|
||||
assert plan_builder.thread_ids == [caller_thread_id]
|
||||
assert consumer_factory.consumers[0].events == []
|
||||
|
||||
|
||||
def test_prepare_does_not_start_consumer_or_session_lifecycle() -> None:
|
||||
factory, plan, _, consumer_factory = create_factory()
|
||||
|
||||
async def scenario() -> None:
|
||||
tasks_before = set(asyncio.all_tasks())
|
||||
|
||||
session = factory.prepare_session(plan.request)
|
||||
consumer = consumer_factory.consumers[0]
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert consumer.events == []
|
||||
assert consumer.start_calls == 0
|
||||
assert consumer.stop_calls == 0
|
||||
assert consumer.close_calls == 0
|
||||
assert set(asyncio.all_tasks()) == tasks_before
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_consumer_observes_start_before_run_and_event_time_during_run(
|
||||
) -> None:
|
||||
factory, plan, _, consumer_factory = create_factory()
|
||||
session = factory.prepare_session(plan.request)
|
||||
consumer = consumer_factory.consumers[0]
|
||||
|
||||
assert consumer.clock is session.clock
|
||||
assert consumer.clock.now == START
|
||||
assert consumer.events == []
|
||||
|
||||
asyncio.run(session.run())
|
||||
|
||||
assert tuple(consumer.events) == plan.events
|
||||
assert consumer.observed_times == [FIRST_TIME, SECOND_TIME]
|
||||
assert session.clock.now == SECOND_TIME
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
|
||||
|
||||
def test_empty_plan_builds_created_no_op_session() -> None:
|
||||
request = make_request()
|
||||
plan = make_plan(request=request, empty=True)
|
||||
factory, _, _, consumer_factory = create_factory(plan=plan)
|
||||
|
||||
session = factory.prepare_session(request)
|
||||
consumer = consumer_factory.consumers[0]
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert session.clock.now == START
|
||||
|
||||
asyncio.run(session.run())
|
||||
|
||||
assert session.state is ReplaySessionState.COMPLETED
|
||||
assert session.clock.now == START
|
||||
assert consumer.events == []
|
||||
|
||||
|
||||
def test_repeated_prepare_creates_fresh_dependency_graph() -> None:
|
||||
factory, plan, plan_builder, consumer_factory = create_factory()
|
||||
|
||||
first = factory.prepare_session(plan.request)
|
||||
second = factory.prepare_session(plan.request)
|
||||
|
||||
assert first is not second
|
||||
assert first.plan is plan
|
||||
assert second.plan is plan
|
||||
assert first.clock is not second.clock
|
||||
assert consumer_factory.consumers[0] is not (
|
||||
consumer_factory.consumers[1]
|
||||
)
|
||||
assert consumer_factory.clocks == [first.clock, second.clock]
|
||||
assert plan_builder.requests == [plan.request, plan.request]
|
||||
|
||||
asyncio.run(first.run())
|
||||
|
||||
assert first.state is ReplaySessionState.COMPLETED
|
||||
assert second.state is ReplaySessionState.CREATED
|
||||
assert second.clock.now == START
|
||||
assert consumer_factory.consumers[1].events == []
|
||||
|
||||
|
||||
class FatalPreparationError(BaseException):
|
||||
pass
|
||||
|
||||
|
||||
class RaisingPlanBuilder:
|
||||
def __init__(self, error: BaseException) -> None:
|
||||
self.error = error
|
||||
self.requests: list[ReplayPlanRequest] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
self.requests.append(request)
|
||||
raise self.error
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expected",
|
||||
(
|
||||
RuntimeError("builder failed"),
|
||||
FatalPreparationError("builder fatal"),
|
||||
asyncio.CancelledError("builder cancelled"),
|
||||
),
|
||||
)
|
||||
def test_builder_error_is_not_wrapped_and_stops_preparation(
|
||||
expected: BaseException,
|
||||
) -> None:
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
plan_builder = RaisingPlanBuilder(expected)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
request = make_request()
|
||||
|
||||
with pytest.raises(type(expected)) as captured:
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert captured.value is expected
|
||||
assert plan_builder.requests == [request]
|
||||
assert consumer_factory.plans == []
|
||||
assert consumer_factory.clocks == []
|
||||
|
||||
|
||||
class InvalidResultPlanBuilder:
|
||||
def __init__(self, result: object) -> None:
|
||||
self.result = result
|
||||
self.requests: list[ReplayPlanRequest] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
self.requests.append(request)
|
||||
return cast(ReplayPlan, self.result)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "plan"))
|
||||
def test_rejects_invalid_builder_result_before_consumer(
|
||||
invalid: object,
|
||||
) -> None:
|
||||
plan_builder = InvalidResultPlanBuilder(invalid)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
request = make_request()
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="ReplayPlan"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert plan_builder.requests == [request]
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_rejects_plan_subclass_before_consumer() -> None:
|
||||
class ReplayPlanSubclass(ReplayPlan):
|
||||
pass
|
||||
|
||||
request = make_request()
|
||||
valid = make_plan(request=request)
|
||||
subclass = ReplayPlanSubclass(
|
||||
request=valid.request,
|
||||
events=valid.events,
|
||||
)
|
||||
plan_builder = InvalidResultPlanBuilder(subclass)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="exact"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
def test_rejects_equal_plan_with_different_request_identity() -> None:
|
||||
request = make_request()
|
||||
copied_request = make_request()
|
||||
assert copied_request == request
|
||||
assert copied_request is not request
|
||||
plan = make_plan(request=copied_request)
|
||||
plan_builder = RecordingPlanBuilder(plan)
|
||||
consumer_factory = RecordingConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataReplayValidationError, match="identity"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert plan_builder.requests == [request]
|
||||
assert consumer_factory.plans == []
|
||||
|
||||
|
||||
class RaisingConsumerFactory:
|
||||
def __init__(self, error: BaseException) -> None:
|
||||
self.error = error
|
||||
self.plans: list[ReplayPlan] = []
|
||||
self.clocks: list[MarketDataClockProtocol] = []
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
self.plans.append(plan)
|
||||
self.clocks.append(clock)
|
||||
raise self.error
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expected",
|
||||
(
|
||||
RuntimeError("consumer factory failed"),
|
||||
FatalPreparationError("consumer factory fatal"),
|
||||
asyncio.CancelledError("consumer factory cancelled"),
|
||||
),
|
||||
)
|
||||
def test_consumer_factory_error_is_not_wrapped_or_retried(
|
||||
expected: BaseException,
|
||||
) -> None:
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
consumer_factory = RaisingConsumerFactory(expected)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(plan),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(type(expected)) as captured:
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert captured.value is expected
|
||||
assert consumer_factory.plans == [plan]
|
||||
assert len(consumer_factory.clocks) == 1
|
||||
assert consumer_factory.clocks[0].now == START
|
||||
|
||||
|
||||
class InvalidConsumerFactory:
|
||||
def __init__(self, result: object) -> None:
|
||||
self.result = result
|
||||
self.calls = 0
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
self.calls += 1
|
||||
return cast(ReplayConsumerProtocol, self.result)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("invalid", (None, object(), "consumer"))
|
||||
def test_rejects_invalid_consumer_result(invalid: object) -> None:
|
||||
request = make_request()
|
||||
consumer_factory = InvalidConsumerFactory(invalid)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
assert consumer_factory.calls == 1
|
||||
|
||||
|
||||
def test_rejects_consumer_class_result() -> None:
|
||||
request = make_request()
|
||||
consumer_factory = InvalidConsumerFactory(RecordingConsumer)
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="consumer"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
|
||||
def test_rejects_synchronous_consumer_result() -> None:
|
||||
class SynchronousConsumer:
|
||||
def consume(self, event: ReplayEvent) -> None:
|
||||
return None
|
||||
|
||||
request = make_request()
|
||||
consumer_factory = InvalidConsumerFactory(SynchronousConsumer())
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="asynchronous"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
|
||||
def test_keyword_incompatible_consumer_factory_error_is_not_wrapped() -> None:
|
||||
class PositionalOnlyConsumerFactory:
|
||||
def create_consumer(
|
||||
self,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
/,
|
||||
) -> ReplayConsumerProtocol:
|
||||
return RecordingConsumer(clock)
|
||||
|
||||
request = make_request()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
|
||||
consumer_factory=cast(
|
||||
ReplayConsumerFactoryProtocol,
|
||||
PositionalOnlyConsumerFactory(),
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match="keyword"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
|
||||
def test_failed_preparation_does_not_poison_next_call() -> None:
|
||||
class RecoveringConsumerFactory(RecordingConsumerFactory):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.attempts = 0
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
self.attempts += 1
|
||||
if self.attempts == 1:
|
||||
raise RuntimeError("first attempt failed")
|
||||
return super().create_consumer(plan=plan, clock=clock)
|
||||
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
consumer_factory = RecoveringConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=RecordingPlanBuilder(plan),
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="first attempt failed"):
|
||||
factory.prepare_session(request)
|
||||
|
||||
session = factory.prepare_session(request)
|
||||
|
||||
assert session.state is ReplaySessionState.CREATED
|
||||
assert consumer_factory.attempts == 2
|
||||
assert len(consumer_factory.consumers) == 1
|
||||
|
||||
|
||||
def test_two_concurrent_callers_receive_independent_graphs() -> None:
|
||||
request = make_request()
|
||||
plan = make_plan(request=request)
|
||||
builder_barrier = threading.Barrier(2)
|
||||
consumer_barrier = threading.Barrier(2)
|
||||
lock = threading.Lock()
|
||||
|
||||
class ConcurrentPlanBuilder:
|
||||
def __init__(self) -> None:
|
||||
self.callers: list[int] = []
|
||||
|
||||
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
|
||||
with lock:
|
||||
self.callers.append(threading.get_ident())
|
||||
builder_barrier.wait(timeout=5)
|
||||
return plan
|
||||
|
||||
class ConcurrentConsumerFactory:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[
|
||||
tuple[int, MarketDataClockProtocol, RecordingConsumer]
|
||||
] = []
|
||||
|
||||
def create_consumer(
|
||||
self,
|
||||
*,
|
||||
plan: ReplayPlan,
|
||||
clock: MarketDataClockProtocol,
|
||||
) -> ReplayConsumerProtocol:
|
||||
consumer = RecordingConsumer(clock)
|
||||
with lock:
|
||||
self.calls.append(
|
||||
(threading.get_ident(), clock, consumer)
|
||||
)
|
||||
consumer_barrier.wait(timeout=5)
|
||||
return consumer
|
||||
|
||||
plan_builder = ConcurrentPlanBuilder()
|
||||
consumer_factory = ConcurrentConsumerFactory()
|
||||
factory = ReplaySessionFactory(
|
||||
plan_builder=plan_builder,
|
||||
consumer_factory=consumer_factory,
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
futures = [
|
||||
executor.submit(factory.prepare_session, request)
|
||||
for _ in range(2)
|
||||
]
|
||||
sessions = [future.result(timeout=10) for future in futures]
|
||||
|
||||
assert len(set(plan_builder.callers)) == 2
|
||||
assert len({call[0] for call in consumer_factory.calls}) == 2
|
||||
assert sessions[0] is not sessions[1]
|
||||
assert sessions[0].clock is not sessions[1].clock
|
||||
assert consumer_factory.calls[0][2] is not (
|
||||
consumer_factory.calls[1][2]
|
||||
)
|
||||
assert {id(call[1]) for call in consumer_factory.calls} == {
|
||||
id(session.clock) for session in sessions
|
||||
}
|
||||
assert all(
|
||||
session.state is ReplaySessionState.CREATED
|
||||
for session in sessions
|
||||
)
|
||||
@@ -38,6 +38,7 @@ TradeRow = dict[str, Any]
|
||||
@dataclass
|
||||
class TransactionalTradeDatabase:
|
||||
rows: dict[TradeKey, TradeRow] = field(default_factory=dict)
|
||||
next_replay_sequence: int = 1
|
||||
|
||||
|
||||
class TransactionalCursor:
|
||||
@@ -47,6 +48,7 @@ class TransactionalCursor:
|
||||
) -> None:
|
||||
self._connection = connection
|
||||
self._fetchone_result: object = None
|
||||
self._fetchall_result: list[tuple[int]] = []
|
||||
|
||||
def __enter__(self) -> TransactionalCursor:
|
||||
return self
|
||||
@@ -78,6 +80,12 @@ class TransactionalCursor:
|
||||
self._insert(parameters)
|
||||
return
|
||||
|
||||
if normalized.startswith(
|
||||
"SELECT nextval('market_data.replay_sequence'::regclass)"
|
||||
):
|
||||
self._allocate_replay_sequences(parameters)
|
||||
return
|
||||
|
||||
if normalized.startswith("SELECT price, quantity"):
|
||||
self._select(parameters)
|
||||
return
|
||||
@@ -91,21 +99,45 @@ class TransactionalCursor:
|
||||
def fetchone(self) -> object:
|
||||
return self._fetchone_result
|
||||
|
||||
def fetchall(self) -> list[tuple[int]]:
|
||||
return list(self._fetchall_result)
|
||||
|
||||
def _insert(self, parameters: tuple[Any, ...]) -> None:
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
canonical_schema_version,
|
||||
) = parameters
|
||||
if len(parameters) == 12:
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
canonical_schema_version,
|
||||
) = parameters
|
||||
replay_sequence = self._next_replay_sequence()
|
||||
elif len(parameters) == 13:
|
||||
(
|
||||
venue,
|
||||
symbol,
|
||||
trade_id,
|
||||
executed_at,
|
||||
price,
|
||||
quantity,
|
||||
aggressor_side,
|
||||
source,
|
||||
first_observed_at,
|
||||
last_observed_at,
|
||||
observation_sources,
|
||||
replay_sequence,
|
||||
canonical_schema_version,
|
||||
) = parameters
|
||||
else:
|
||||
raise AssertionError("Unexpected Trade INSERT parameters")
|
||||
|
||||
key = (venue, symbol, trade_id, executed_at)
|
||||
working_rows = self._connection.working_rows
|
||||
|
||||
@@ -121,10 +153,26 @@ class TransactionalCursor:
|
||||
"first_observed_at": first_observed_at,
|
||||
"last_observed_at": last_observed_at,
|
||||
"observation_sources": list(observation_sources),
|
||||
"replay_sequence": replay_sequence,
|
||||
"canonical_schema_version": canonical_schema_version,
|
||||
}
|
||||
self._fetchone_result = (1,)
|
||||
|
||||
def _allocate_replay_sequences(
|
||||
self,
|
||||
parameters: tuple[Any, ...],
|
||||
) -> None:
|
||||
(count,) = parameters
|
||||
self._fetchall_result = [
|
||||
(self._next_replay_sequence(),)
|
||||
for _ in range(count)
|
||||
]
|
||||
|
||||
def _next_replay_sequence(self) -> int:
|
||||
value = self._connection.database.next_replay_sequence
|
||||
self._connection.database.next_replay_sequence += 1
|
||||
return value
|
||||
|
||||
def _select(self, parameters: tuple[Any, ...]) -> None:
|
||||
key = parameters
|
||||
row = self._connection.working_rows.get(key)
|
||||
@@ -181,6 +229,10 @@ class TransactionalConnection:
|
||||
def cursor(self) -> TransactionalCursor:
|
||||
return TransactionalCursor(self)
|
||||
|
||||
@property
|
||||
def database(self) -> TransactionalTradeDatabase:
|
||||
return self._database
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingConnectionProvider:
|
||||
@@ -261,6 +313,7 @@ def test_store_trade_inserts_canonical_payload_and_provenance() -> None:
|
||||
"first_observed_at": OBSERVED_AT,
|
||||
"last_observed_at": OBSERVED_AT,
|
||||
"observation_sources": ["dzengi_websocket_trade"],
|
||||
"replay_sequence": 1,
|
||||
"canonical_schema_version": 1,
|
||||
}
|
||||
|
||||
@@ -459,6 +512,122 @@ def test_batch_uses_stable_identity_order() -> None:
|
||||
assert inserted_symbols == ("A", "B", "C")
|
||||
|
||||
|
||||
def test_batch_allocates_replay_sequence_in_input_order() -> None:
|
||||
repository, database, connection, _ = _repository()
|
||||
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
_trade(symbol="C", trade_id=3),
|
||||
_trade(symbol="A", trade_id=1),
|
||||
_trade(symbol="B", trade_id=2),
|
||||
),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
inserted = tuple(
|
||||
(parameters[1], parameters[11])
|
||||
for statement, parameters in connection.calls
|
||||
if statement.startswith("INSERT INTO market_data.trades")
|
||||
)
|
||||
|
||||
assert inserted == (
|
||||
("A", 2),
|
||||
("B", 3),
|
||||
("C", 1),
|
||||
)
|
||||
assert {
|
||||
key[1]: row["replay_sequence"]
|
||||
for key, row in database.rows.items()
|
||||
} == {
|
||||
"A": 2,
|
||||
"B": 3,
|
||||
"C": 1,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("first_trade_id", "second_trade_id"),
|
||||
(
|
||||
(SIGNED_TRADE_ID_MAX, SIGNED_TRADE_ID_MIN),
|
||||
(-1, 0),
|
||||
),
|
||||
)
|
||||
def test_batch_replay_sequence_preserves_signed_rollover_input_order(
|
||||
first_trade_id: int,
|
||||
second_trade_id: int,
|
||||
) -> None:
|
||||
repository, database, _, _ = _repository()
|
||||
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
_trade(trade_id=first_trade_id),
|
||||
_trade(trade_id=second_trade_id),
|
||||
),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
assert database.rows[
|
||||
(VENUE, SYMBOL, first_trade_id, EXECUTED_AT)
|
||||
]["replay_sequence"] == 1
|
||||
assert database.rows[
|
||||
(VENUE, SYMBOL, second_trade_id, EXECUTED_AT)
|
||||
]["replay_sequence"] == 2
|
||||
|
||||
|
||||
def test_batch_duplicate_and_provenance_keep_original_replay_sequence() -> None:
|
||||
repository, database, _, _ = _repository()
|
||||
original = _trade()
|
||||
repository.store_trade(
|
||||
venue=VENUE,
|
||||
trade=original,
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
replace(original, source="dzengi"),
|
||||
original,
|
||||
),
|
||||
observed_at=OBSERVED_AT + timedelta(seconds=1),
|
||||
)
|
||||
|
||||
assert _only_row(database)["replay_sequence"] == 1
|
||||
assert database.next_replay_sequence == 4
|
||||
|
||||
|
||||
def test_failed_batch_keeps_consumed_replay_sequence_gap() -> None:
|
||||
repository, database, _, _ = _repository()
|
||||
existing = _trade(symbol="B", trade_id=2)
|
||||
repository.store_trade(
|
||||
venue=VENUE,
|
||||
trade=existing,
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
with pytest.raises(MarketDataStorageConflictError):
|
||||
repository.store_trades(
|
||||
venue=VENUE,
|
||||
trades=(
|
||||
_trade(symbol="A", trade_id=1),
|
||||
replace(existing, price=Decimal("999")),
|
||||
),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
repository.store_trade(
|
||||
venue=VENUE,
|
||||
trade=_trade(symbol="C", trade_id=3),
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
assert database.rows[
|
||||
(VENUE, "C", 3, EXECUTED_AT)
|
||||
]["replay_sequence"] == 4
|
||||
|
||||
|
||||
def test_batch_conflict_rolls_back_preceding_insert() -> None:
|
||||
repository, database, connection, _ = _repository()
|
||||
existing = _trade(symbol="B", trade_id=2)
|
||||
|
||||
@@ -7,6 +7,7 @@ import pytest
|
||||
|
||||
from src.storage.exceptions import StorageMigrationError
|
||||
from src.storage.migrations import (
|
||||
MARKET_DATA_PARTITION_ADVISORY_LOCK_ID,
|
||||
STORAGE_MIGRATION_ADVISORY_LOCK_ID,
|
||||
STORAGE_MIGRATIONS,
|
||||
StorageMigration,
|
||||
@@ -106,6 +107,7 @@ def test_default_migrations_have_stable_order_and_names() -> None:
|
||||
(6, "add_quote_and_candle_observation_sources"),
|
||||
(7, "create_market_data_partition_registry"),
|
||||
(8, "create_trade_stream_checkpoints"),
|
||||
(9, "add_global_market_data_replay_sequence"),
|
||||
)
|
||||
|
||||
|
||||
@@ -158,12 +160,73 @@ def test_default_schema_defines_partitions_identities_and_constraints() -> None:
|
||||
assert "CHECK (checkpoint_schema_version > 0)" in sql
|
||||
|
||||
|
||||
def test_replay_sequence_migration_has_atomic_global_order_contract() -> None:
|
||||
migration = STORAGE_MIGRATIONS[-1]
|
||||
statements = tuple(
|
||||
" ".join(statement.split())
|
||||
for statement in migration.statements
|
||||
)
|
||||
sql = "\n".join(statements)
|
||||
|
||||
assert migration.version == 9
|
||||
assert migration.name == "add_global_market_data_replay_sequence"
|
||||
assert statements[:4] == (
|
||||
(
|
||||
"SELECT pg_advisory_xact_lock("
|
||||
f"{MARKET_DATA_PARTITION_ADVISORY_LOCK_ID}"
|
||||
")"
|
||||
),
|
||||
"LOCK TABLE market_data.trades IN ACCESS EXCLUSIVE MODE",
|
||||
"LOCK TABLE market_data.quotes IN ACCESS EXCLUSIVE MODE",
|
||||
(
|
||||
"LOCK TABLE market_data.candle_revisions "
|
||||
"IN ACCESS EXCLUSIVE MODE"
|
||||
),
|
||||
)
|
||||
assert "CREATE SEQUENCE market_data.replay_sequence AS BIGINT" in sql
|
||||
assert "MINVALUE 1" in sql
|
||||
assert "CACHE 1" in sql
|
||||
assert "NO CYCLE" in sql
|
||||
assert "OWNED BY NONE" in sql
|
||||
assert sql.count("ADD COLUMN replay_sequence BIGINT") == 3
|
||||
assert "CREATE TABLE market_data.replay_sequence" not in sql
|
||||
assert "CREATE TEMPORARY TABLE market_data_replay_sequence_backfill" in sql
|
||||
assert "ON COMMIT DROP" in sql
|
||||
assert "ROW_NUMBER() OVER" in sql
|
||||
assert "UNION ALL" in sql
|
||||
assert "event_time, data_type_rank, venue COLLATE \"C\"" in sql
|
||||
assert "symbol COLLATE \"C\"" in sql
|
||||
assert "candle_interval COLLATE \"C\"" in sql
|
||||
assert "ctid" not in sql.lower()
|
||||
assert sql.count("SET replay_sequence = backfill.replay_sequence") == 3
|
||||
assert "target.trade_id = backfill.trade_id" in sql
|
||||
assert "target.received_at = backfill.received_at" in sql
|
||||
assert "target.interval = backfill.interval" in sql
|
||||
assert "target.open_time = backfill.open_time" in sql
|
||||
assert "target.observed_at = backfill.observed_at" in sql
|
||||
assert "SELECT pg_catalog.setval(" in sql
|
||||
assert "EXISTS ( SELECT 1 FROM market_data_replay_sequence_backfill )" in sql
|
||||
assert sql.count("SET DEFAULT nextval(") == 3
|
||||
assert sql.count("ALTER COLUMN replay_sequence SET NOT NULL") == 3
|
||||
assert sql.count("CHECK (replay_sequence > 0)") == 3
|
||||
assert "CREATE FUNCTION market_data.reject_replay_sequence_change()" in sql
|
||||
assert sql.count("BEFORE UPDATE OF replay_sequence") == 3
|
||||
assert "CREATE INDEX trades_history_keyset_idx" in sql
|
||||
assert "CREATE INDEX quotes_history_keyset_idx" in sql
|
||||
assert "CREATE INDEX candle_revisions_history_keyset_idx" in sql
|
||||
assert "CREATE INDEX candle_revisions_replay_keyset_idx" in sql
|
||||
assert "executed_at, replay_sequence" in sql
|
||||
assert "received_at, replay_sequence" in sql
|
||||
assert "interval, open_time, replay_sequence" in sql
|
||||
assert "interval, observed_at, replay_sequence" in sql
|
||||
|
||||
|
||||
def test_run_locks_and_applies_every_pending_migration_in_order() -> None:
|
||||
runner, cursor, connection, provider = _runner()
|
||||
|
||||
result = runner.run()
|
||||
|
||||
assert result == (1, 2, 3, 4, 5, 6, 7, 8)
|
||||
assert result == (1, 2, 3, 4, 5, 6, 7, 8, 9)
|
||||
assert provider.calls == 1
|
||||
assert connection.entered == 1
|
||||
assert connection.exited == 1
|
||||
@@ -180,7 +243,7 @@ def test_run_locks_and_applies_every_pending_migration_in_order() -> None:
|
||||
)
|
||||
and isinstance(parameters, tuple)
|
||||
)
|
||||
assert inserted_versions == (1, 2, 3, 4, 5, 6, 7, 8)
|
||||
assert inserted_versions == (1, 2, 3, 4, 5, 6, 7, 8, 9)
|
||||
|
||||
|
||||
def test_run_skips_already_applied_migrations() -> None:
|
||||
@@ -209,7 +272,7 @@ def test_run_applies_only_migrations_after_existing_prefix() -> None:
|
||||
|
||||
result = runner.run()
|
||||
|
||||
assert result == (3, 4, 5, 6, 7, 8)
|
||||
assert result == (3, 4, 5, 6, 7, 8, 9)
|
||||
inserted_versions = tuple(
|
||||
parameters[0]
|
||||
for statement, parameters in cursor.calls
|
||||
@@ -218,7 +281,7 @@ def test_run_applies_only_migrations_after_existing_prefix() -> None:
|
||||
)
|
||||
and isinstance(parameters, tuple)
|
||||
)
|
||||
assert inserted_versions == (3, 4, 5, 6, 7, 8)
|
||||
assert inserted_versions == (3, 4, 5, 6, 7, 8, 9)
|
||||
|
||||
|
||||
def test_run_rejects_unknown_applied_version() -> None:
|
||||
|
||||
Reference in New Issue
Block a user