Build 060.29: implement Market Data Access and Replay

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

View File

@@ -1,4 +1,6 @@
{
"python.defaultInterpreterPath": "app/.venv/bin/python",
"python-envs.defaultEnvManager": "ms-python.python:system"
"python-envs.defaultEnvManager": "ms-python.python:system",
"python.analysis.typeCheckingMode": "standard",
"python.analysis.diagnosticMode": "workspace"
}

4
app/requirements-dev.txt Normal file
View File

@@ -0,0 +1,4 @@
-r requirements.txt
pytest==9.1.1
pyright[nodejs]==1.1.411

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

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

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

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

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

View File

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

View 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

View 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

View 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

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

View 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:
...

View 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

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

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

View 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

View 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

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View 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

View File

@@ -0,0 +1,55 @@
from __future__ import annotations
from src.market_data.access import (
CandleRevisionHistoryPage,
CandleRevisionHistoryQuery,
CandleRevisionHistoryReaderProtocol,
MarketDataAccessError,
MarketDataAccessIntegrityError,
MarketDataAccessOperationError,
MarketDataAccessValidationError,
MarketDataCursorError,
MarketDataHistoricalAccessProtocol,
QuoteHistoryPage,
QuoteHistoryQuery,
QuoteHistoryReaderProtocol,
TradeHistoryPage,
TradeHistoryQuery,
TradeHistoryReaderProtocol,
)
class RecordingHistoricalAccess:
def query_trades(
self,
query: TradeHistoryQuery,
) -> TradeHistoryPage:
return TradeHistoryPage(query=query, items=())
def query_quotes(
self,
query: QuoteHistoryQuery,
) -> QuoteHistoryPage:
return QuoteHistoryPage(query=query, items=())
def query_candle_revisions(
self,
query: CandleRevisionHistoryQuery,
) -> CandleRevisionHistoryPage:
return CandleRevisionHistoryPage(query=query, items=())
def test_history_reader_protocols_are_runtime_checkable() -> None:
access = RecordingHistoricalAccess()
assert isinstance(access, TradeHistoryReaderProtocol)
assert isinstance(access, QuoteHistoryReaderProtocol)
assert isinstance(access, CandleRevisionHistoryReaderProtocol)
assert isinstance(access, MarketDataHistoricalAccessProtocol)
def test_access_error_hierarchy_is_specialized() -> None:
assert issubclass(MarketDataAccessValidationError, MarketDataAccessError)
assert issubclass(MarketDataCursorError, MarketDataAccessValidationError)
assert issubclass(MarketDataAccessIntegrityError, MarketDataAccessError)
assert issubclass(MarketDataAccessOperationError, MarketDataAccessError)

View File

@@ -0,0 +1,882 @@
from __future__ import annotations
from dataclasses import FrozenInstanceError, replace
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.market_data.access import (
HISTORY_PAGE_LIMIT_MAX,
CandleRevisionHistoryCursor,
CandleRevisionHistoryPage,
CandleRevisionHistoryQuery,
CandleRevisionHistoryRecord,
HistoricalTimeRange,
QuoteHistoryCursor,
QuoteHistoryPage,
QuoteHistoryQuery,
QuoteHistoryRecord,
TradeHistoryCursor,
TradeHistoryPage,
TradeHistoryQuery,
TradeHistoryRecord,
)
from src.market_data.acquisition.models.candle import Candle
from src.market_data.acquisition.models.quote import Quote
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
START = datetime(2026, 8, 2, 10, 0, tzinfo=timezone.utc)
END = START + timedelta(hours=1)
SOURCE = "dzengi_websocket_trade"
def make_trade(
*,
trade_id: int = 100,
executed_at: datetime = START + timedelta(minutes=1),
symbol: str = SYMBOL,
source: str = SOURCE,
) -> Trade:
return Trade(
symbol=symbol,
trade_id=trade_id,
price=Decimal("65000.25"),
quantity=Decimal("0.001"),
executed_at=executed_at,
aggressor_side=TradeAggressorSide.BUY,
source=source,
)
def make_quote(
*,
received_at: datetime = START + timedelta(minutes=2),
) -> Quote:
return Quote(
symbol=SYMBOL,
last_price=Decimal("65000"),
bid_price=Decimal("64999"),
ask_price=Decimal("65001"),
exchange_timestamp=received_at - timedelta(milliseconds=1),
received_at=received_at,
source="dzengi_rest_quote",
)
def make_candle(
*,
open_time: datetime = START,
interval: str = "1m",
) -> Candle:
return Candle(
symbol=SYMBOL,
interval=interval,
open_time=open_time,
open_price=Decimal("64900"),
high_price=Decimal("65100"),
low_price=Decimal("64800"),
close_price=Decimal("65000"),
volume=Decimal("12.5"),
source="dzengi_rest_candle",
)
def make_range() -> HistoricalTimeRange:
return HistoricalTimeRange(start_time=START, end_time=END)
def make_trade_query(
*,
cursor: TradeHistoryCursor | None = None,
limit: int = 500,
) -> TradeHistoryQuery:
return TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
limit=limit,
cursor=cursor,
)
def make_trade_record(
*,
trade_id: int = 100,
executed_at: datetime = START + timedelta(minutes=1),
replay_sequence: int = 10,
) -> TradeHistoryRecord:
return TradeHistoryRecord(
venue=VENUE,
trade=make_trade(
trade_id=trade_id,
executed_at=executed_at,
),
first_observed_at=executed_at + timedelta(seconds=1),
last_observed_at=executed_at + timedelta(seconds=2),
observation_sources=(SOURCE,),
replay_sequence=replay_sequence,
)
def make_quote_record(
*,
received_at: datetime = START + timedelta(minutes=2),
replay_sequence: int = 20,
) -> QuoteHistoryRecord:
quote = make_quote(received_at=received_at)
return QuoteHistoryRecord(
venue=VENUE,
quote=quote,
observation_sources=(quote.source,),
replay_sequence=replay_sequence,
)
def make_candle_record(
*,
open_time: datetime = START,
observed_at: datetime = START + timedelta(seconds=30),
replay_sequence: int = 30,
interval: str = "1m",
) -> CandleRevisionHistoryRecord:
candle = make_candle(open_time=open_time, interval=interval)
return CandleRevisionHistoryRecord(
venue=VENUE,
candle=candle,
observed_at=observed_at,
is_final=False,
observation_sources=(candle.source,),
replay_sequence=replay_sequence,
)
def test_time_range_is_half_open_and_normalized_to_utc() -> None:
offset = timezone(timedelta(hours=3))
time_range = HistoricalTimeRange(
start_time=START.astimezone(offset),
end_time=END.astimezone(offset),
)
assert time_range.start_time == START
assert time_range.start_time.tzinfo is timezone.utc
assert time_range.end_time == END
assert time_range.contains(START)
assert time_range.contains(END - timedelta(microseconds=1))
assert not time_range.contains(END)
@pytest.mark.parametrize(
("start_time", "end_time", "error_type"),
(
(datetime(2026, 8, 2, 10, 0), END, ValueError),
(START, datetime(2026, 8, 2, 11, 0), ValueError),
(START, START, ValueError),
(END, START, ValueError),
("2026-08-02", END, TypeError),
),
)
def test_time_range_rejects_invalid_boundaries(
start_time: Any,
end_time: Any,
error_type: type[Exception],
) -> None:
with pytest.raises(error_type):
HistoricalTimeRange(
start_time=start_time,
end_time=end_time,
)
def test_history_records_preserve_exact_canonical_payloads() -> None:
trade = make_trade()
quote = make_quote()
candle = make_candle()
trade_record = TradeHistoryRecord(
venue=" dzengi ",
trade=trade,
first_observed_at=trade.executed_at,
last_observed_at=trade.executed_at,
observation_sources=(trade.source,),
replay_sequence=1,
)
quote_record = QuoteHistoryRecord(
venue=VENUE,
quote=quote,
observation_sources=(quote.source,),
replay_sequence=2,
)
candle_record = CandleRevisionHistoryRecord(
venue=VENUE,
candle=candle,
observed_at=candle.open_time,
is_final=True,
observation_sources=(candle.source,),
replay_sequence=3,
)
assert trade_record.venue == VENUE
assert trade_record.trade is trade
assert quote_record.quote is quote
assert candle_record.candle is candle
assert trade_record.event_time == trade.executed_at
assert quote_record.event_time == quote.received_at
assert candle_record.event_time == candle.open_time
assert candle_record.replay_at == candle.open_time
@pytest.mark.parametrize("invalid_sequence", (True, 0, -1, 1.5, "1"))
def test_records_reject_invalid_replay_sequence(
invalid_sequence: Any,
) -> None:
with pytest.raises((TypeError, ValueError), match="replay_sequence"):
make_trade_record(replay_sequence=invalid_sequence)
@pytest.mark.parametrize(
"observation_sources",
(
[],
(),
("",),
(SOURCE, SOURCE),
("another_source",),
),
)
def test_trade_record_rejects_invalid_provenance(
observation_sources: Any,
) -> None:
trade = make_trade()
with pytest.raises((TypeError, ValueError)):
TradeHistoryRecord(
venue=VENUE,
trade=trade,
first_observed_at=trade.executed_at,
last_observed_at=trade.executed_at,
observation_sources=observation_sources,
replay_sequence=1,
)
def test_trade_record_rejects_reversed_observation_times() -> None:
trade = make_trade()
with pytest.raises(ValueError, match="last_observed_at"):
TradeHistoryRecord(
venue=VENUE,
trade=trade,
first_observed_at=trade.executed_at + timedelta(seconds=1),
last_observed_at=trade.executed_at,
observation_sources=(trade.source,),
replay_sequence=1,
)
def test_candle_record_uses_open_time_for_history_and_observed_for_replay() -> None:
record = make_candle_record()
assert record.event_time == START
assert record.replay_at == START + timedelta(seconds=30)
assert record.interval == "1m"
def test_candle_record_rejects_observation_before_open_time() -> None:
with pytest.raises(ValueError, match="observed_at"):
make_candle_record(
observed_at=START - timedelta(microseconds=1),
)
@pytest.mark.parametrize("is_final", (0, 1, None, "true"))
def test_candle_record_requires_exact_boolean(is_final: Any) -> None:
candle = make_candle()
with pytest.raises(TypeError, match="is_final"):
CandleRevisionHistoryRecord(
venue=VENUE,
candle=candle,
observed_at=candle.open_time,
is_final=is_final,
observation_sources=(candle.source,),
replay_sequence=1,
)
def test_trade_cursor_normalizes_scope_and_time() -> None:
offset = timezone(timedelta(hours=3))
cursor = TradeHistoryCursor(
venue=" dzengi ",
symbol=" btc/usd_leverage ",
time_range=make_range(),
executed_at=(START + timedelta(minutes=1)).astimezone(offset),
replay_sequence=5,
)
assert cursor.venue == VENUE
assert cursor.symbol == SYMBOL
assert cursor.executed_at.tzinfo is timezone.utc
@pytest.mark.parametrize(
"cursor_time",
(
START - timedelta(microseconds=1),
END,
END + timedelta(microseconds=1),
),
)
def test_cursor_position_must_belong_to_query_range(
cursor_time: datetime,
) -> None:
with pytest.raises(ValueError, match="query range"):
TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=cursor_time,
replay_sequence=1,
)
def test_cursor_rejects_unknown_version() -> None:
with pytest.raises(ValueError, match="version"):
TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=START,
replay_sequence=1,
version=2,
)
@pytest.mark.parametrize("limit", (1, HISTORY_PAGE_LIMIT_MAX))
def test_trade_query_accepts_limit_boundaries(limit: int) -> None:
query = TradeHistoryQuery(
venue=" dzengi ",
symbol="btc/usd_leverage",
time_range=make_range(),
limit=limit,
)
assert query.venue == VENUE
assert query.symbol == SYMBOL
assert query.limit == limit
@pytest.mark.parametrize("limit", (True, 0, -1, 1.5, 1001))
def test_query_rejects_invalid_limit(limit: Any) -> None:
with pytest.raises((TypeError, ValueError), match="limit"):
TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
limit=limit,
)
def test_query_accepts_cursor_when_only_page_limit_changes() -> None:
time_range = make_range()
cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=time_range,
executed_at=START,
replay_sequence=1,
)
query = TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=time_range,
limit=17,
cursor=cursor,
)
assert query.cursor is cursor
def test_query_rejects_cursor_from_another_scope() -> None:
cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=START,
replay_sequence=1,
)
with pytest.raises(ValueError, match="scope"):
TradeHistoryQuery(
venue=VENUE,
symbol="ETH/USD_LEVERAGE",
time_range=make_range(),
cursor=cursor,
)
def test_query_rejects_cursor_of_another_data_type() -> None:
quote_cursor = QuoteHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
received_at=START,
replay_sequence=1,
)
with pytest.raises(TypeError, match="TradeHistoryCursor"):
TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
cursor=quote_cursor, # type: ignore[arg-type]
)
def test_candle_query_preserves_interval_case() -> None:
query = CandleRevisionHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
interval=" 1M ",
time_range=make_range(),
)
assert query.interval == "1M"
def test_empty_page_is_valid_without_cursor() -> None:
page = TradeHistoryPage(query=make_trade_query(), items=())
assert page.items == ()
assert page.has_more is False
def test_empty_page_rejects_cursor() -> None:
cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=START,
replay_sequence=1,
)
with pytest.raises(ValueError, match="empty page"):
TradeHistoryPage(
query=make_trade_query(),
items=(),
next_cursor=cursor,
)
def test_trade_page_accepts_signed_rollover_at_equal_timestamp() -> None:
rollover_time = START + timedelta(minutes=1)
first = make_trade_record(
trade_id=2_147_483_647,
executed_at=rollover_time,
replay_sequence=10,
)
second = make_trade_record(
trade_id=-2_147_483_648,
executed_at=rollover_time,
replay_sequence=11,
)
cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=rollover_time,
replay_sequence=11,
)
page = TradeHistoryPage(
query=make_trade_query(),
items=(first, second),
next_cursor=cursor,
)
assert page.items == (first, second)
assert page.has_more is True
def test_trade_page_accepts_negative_one_to_zero_boundary() -> None:
event_time = START + timedelta(minutes=1)
page = TradeHistoryPage(
query=make_trade_query(),
items=(
make_trade_record(
trade_id=-1,
executed_at=event_time,
replay_sequence=20,
),
make_trade_record(
trade_id=0,
executed_at=event_time,
replay_sequence=21,
),
),
)
assert [item.trade.trade_id for item in page.items] == [-1, 0]
def test_page_rejects_reverse_or_duplicate_order_key() -> None:
first = make_trade_record(replay_sequence=2)
second = make_trade_record(replay_sequence=1)
with pytest.raises(ValueError, match="strictly ordered"):
TradeHistoryPage(
query=make_trade_query(),
items=(first, second),
)
def test_page_rejects_cursor_not_pointing_to_last_item() -> None:
item = make_trade_record(replay_sequence=10)
cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=item.event_time,
replay_sequence=9,
)
with pytest.raises(ValueError, match="last page item"):
TradeHistoryPage(
query=make_trade_query(),
items=(item,),
next_cursor=cursor,
)
@pytest.mark.parametrize(
"page",
(
QuoteHistoryPage(
query=QuoteHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
),
items=(make_quote_record(),),
),
CandleRevisionHistoryPage(
query=CandleRevisionHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
interval="1m",
time_range=make_range(),
),
items=(make_candle_record(),),
),
),
)
def test_quote_and_candle_pages_are_typed_and_immutable(page: Any) -> None:
assert not hasattr(page, "__dict__")
with pytest.raises(FrozenInstanceError):
setattr(page, "items", ())
def test_cursor_and_query_classes_use_slots_and_are_frozen() -> None:
query = QuoteHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
)
assert not hasattr(query, "__dict__")
with pytest.raises(FrozenInstanceError):
setattr(query, "venue", "other")
def test_cursor_window_mismatch_is_rejected() -> None:
cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=START,
replay_sequence=1,
)
shifted = HistoricalTimeRange(
start_time=START - timedelta(minutes=1),
end_time=END,
)
with pytest.raises(ValueError, match="scope"):
TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=shifted,
cursor=cursor,
)
def test_candle_cursor_interval_mismatch_is_rejected() -> None:
cursor = CandleRevisionHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
interval="1m",
time_range=make_range(),
open_time=START,
replay_sequence=1,
)
with pytest.raises(ValueError, match="interval"):
CandleRevisionHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
interval="5m",
time_range=make_range(),
cursor=cursor,
)
def test_page_rejects_mixed_query_scope() -> None:
first = make_trade_record(replay_sequence=1)
second = replace(
make_trade_record(
executed_at=START + timedelta(minutes=2),
replay_sequence=2,
),
venue="other",
)
with pytest.raises(ValueError, match="one query scope"):
TradeHistoryPage(
query=make_trade_query(),
items=(first, second),
)
def test_terminal_page_rejects_item_outside_query_range() -> None:
item = make_trade_record(
executed_at=START - timedelta(microseconds=1),
)
with pytest.raises(ValueError, match="query range"):
TradeHistoryPage(
query=make_trade_query(),
items=(item,),
)
def test_page_rejects_more_items_than_query_limit() -> None:
first = make_trade_record(replay_sequence=1)
second = make_trade_record(
executed_at=START + timedelta(minutes=2),
replay_sequence=2,
)
with pytest.raises(ValueError, match="query limit"):
TradeHistoryPage(
query=make_trade_query(limit=1),
items=(first, second),
)
def test_page_items_must_follow_incoming_cursor() -> None:
cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=START + timedelta(minutes=1),
replay_sequence=10,
)
with pytest.raises(ValueError, match="follow query cursor"):
TradeHistoryPage(
query=make_trade_query(cursor=cursor),
items=(make_trade_record(replay_sequence=10),),
)
def test_records_reject_canonical_payload_subclasses() -> None:
class TradeSubclass(Trade):
pass
class QuoteSubclass(Quote):
pass
class CandleSubclass(Candle):
pass
trade = make_trade()
quote = make_quote()
candle = make_candle()
trade_subclass = TradeSubclass(
symbol=trade.symbol,
trade_id=trade.trade_id,
price=trade.price,
quantity=trade.quantity,
executed_at=trade.executed_at,
aggressor_side=trade.aggressor_side,
source=trade.source,
)
quote_subclass = QuoteSubclass(
symbol=quote.symbol,
last_price=quote.last_price,
bid_price=quote.bid_price,
ask_price=quote.ask_price,
exchange_timestamp=quote.exchange_timestamp,
received_at=quote.received_at,
source=quote.source,
)
candle_subclass = CandleSubclass(
symbol=candle.symbol,
interval=candle.interval,
open_time=candle.open_time,
open_price=candle.open_price,
high_price=candle.high_price,
low_price=candle.low_price,
close_price=candle.close_price,
volume=candle.volume,
source=candle.source,
)
with pytest.raises(TypeError, match="Canonical Trade"):
TradeHistoryRecord(
venue=VENUE,
trade=trade_subclass,
first_observed_at=trade.executed_at,
last_observed_at=trade.executed_at,
observation_sources=(trade.source,),
replay_sequence=1,
)
with pytest.raises(TypeError, match="Canonical Quote"):
QuoteHistoryRecord(
venue=VENUE,
quote=quote_subclass,
observation_sources=(quote.source,),
replay_sequence=1,
)
with pytest.raises(TypeError, match="Canonical Candle"):
CandleRevisionHistoryRecord(
venue=VENUE,
candle=candle_subclass,
observed_at=candle.open_time,
is_final=False,
observation_sources=(candle.source,),
replay_sequence=1,
)
def test_cursor_and_query_reject_time_range_subclass() -> None:
class HistoricalTimeRangeSubclass(HistoricalTimeRange):
pass
time_range = HistoricalTimeRangeSubclass(START, END)
with pytest.raises(TypeError, match="HistoricalTimeRange"):
TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=time_range,
executed_at=START,
replay_sequence=1,
)
with pytest.raises(TypeError, match="HistoricalTimeRange"):
TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=time_range,
)
def test_query_rejects_cursor_subclass() -> None:
class TradeHistoryCursorSubclass(TradeHistoryCursor):
pass
cursor = TradeHistoryCursorSubclass(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=START,
replay_sequence=1,
)
with pytest.raises(TypeError, match="TradeHistoryCursor"):
TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
cursor=cursor,
)
def test_page_rejects_query_record_and_cursor_subclasses() -> None:
class TradeHistoryQuerySubclass(TradeHistoryQuery):
pass
class TradeHistoryRecordSubclass(TradeHistoryRecord):
@property
def order_key(self) -> tuple[datetime, int]:
return (END, 1)
class TradeHistoryCursorSubclass(TradeHistoryCursor):
pass
item = make_trade_record(replay_sequence=10)
with pytest.raises(TypeError, match="TradeHistoryQuery"):
TradeHistoryPage(
query=TradeHistoryQuerySubclass(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
),
items=(item,),
)
item_subclass = TradeHistoryRecordSubclass(
venue=item.venue,
trade=item.trade,
first_observed_at=item.first_observed_at,
last_observed_at=item.last_observed_at,
observation_sources=item.observation_sources,
replay_sequence=item.replay_sequence,
)
with pytest.raises(TypeError, match="TradeHistoryRecord"):
TradeHistoryPage(
query=make_trade_query(),
items=(item_subclass,),
)
cursor_subclass = TradeHistoryCursorSubclass(
venue=VENUE,
symbol=SYMBOL,
time_range=make_range(),
executed_at=item.event_time,
replay_sequence=item.replay_sequence,
)
with pytest.raises(TypeError, match="TradeHistoryCursor"):
TradeHistoryPage(
query=make_trade_query(),
items=(item,),
next_cursor=cursor_subclass,
)
def test_page_rejects_items_tuple_subclass() -> None:
class ItemsTupleSubclass(tuple):
pass
with pytest.raises(TypeError, match="items must be a tuple"):
TradeHistoryPage(
query=make_trade_query(),
items=ItemsTupleSubclass((make_trade_record(),)),
)

View File

@@ -0,0 +1,195 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
import pytest
from src.market_data.access.contracts import (
MarketDataHistoricalAccessProtocol,
)
from src.market_data.access.market_data_historical_access import (
MarketDataHistoricalAccess,
)
from src.market_data.access.models import (
CandleRevisionHistoryPage,
CandleRevisionHistoryQuery,
HistoricalTimeRange,
QuoteHistoryPage,
QuoteHistoryQuery,
TradeHistoryPage,
TradeHistoryQuery,
)
NOW = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
TIME_RANGE = HistoricalTimeRange(
start_time=NOW,
end_time=NOW + timedelta(hours=1),
)
class RecordingTradeReader:
def __init__(self, page: TradeHistoryPage) -> None:
self.page = page
self.queries: list[TradeHistoryQuery] = []
def query_trades(self, query: TradeHistoryQuery) -> TradeHistoryPage:
self.queries.append(query)
return self.page
class RecordingQuoteReader:
def __init__(self, page: QuoteHistoryPage) -> None:
self.page = page
self.queries: list[QuoteHistoryQuery] = []
def query_quotes(self, query: QuoteHistoryQuery) -> QuoteHistoryPage:
self.queries.append(query)
return self.page
class RecordingCandleReader:
def __init__(self, page: CandleRevisionHistoryPage) -> None:
self.page = page
self.queries: list[CandleRevisionHistoryQuery] = []
def query_candle_revisions(
self,
query: CandleRevisionHistoryQuery,
) -> CandleRevisionHistoryPage:
self.queries.append(query)
return self.page
class BrokenTradeReader:
def __init__(self, error: RuntimeError) -> None:
self.error = error
def query_trades(self, query: TradeHistoryQuery) -> TradeHistoryPage:
del query
raise self.error
def make_dependencies() -> tuple[
MarketDataHistoricalAccess,
RecordingTradeReader,
RecordingQuoteReader,
RecordingCandleReader,
]:
trade_query = TradeHistoryQuery(
venue="dzengi",
symbol="BTC/USD_LEVERAGE",
time_range=TIME_RANGE,
)
quote_query = QuoteHistoryQuery(
venue="dzengi",
symbol="BTC/USD_LEVERAGE",
time_range=TIME_RANGE,
)
candle_query = CandleRevisionHistoryQuery(
venue="dzengi",
symbol="BTC/USD_LEVERAGE",
interval="1m",
time_range=TIME_RANGE,
)
trade_reader = RecordingTradeReader(
TradeHistoryPage(query=trade_query, items=()),
)
quote_reader = RecordingQuoteReader(
QuoteHistoryPage(query=quote_query, items=()),
)
candle_reader = RecordingCandleReader(
CandleRevisionHistoryPage(query=candle_query, items=()),
)
access = MarketDataHistoricalAccess(
trade_reader=trade_reader,
quote_reader=quote_reader,
candle_revision_reader=candle_reader,
)
return access, trade_reader, quote_reader, candle_reader
def test_implements_combined_protocol_and_uses_slots() -> None:
access, *_ = make_dependencies()
assert isinstance(access, MarketDataHistoricalAccessProtocol)
assert not hasattr(access, "__dict__")
@pytest.mark.parametrize(
"dependency_name",
(
"trade_reader",
"quote_reader",
"candle_revision_reader",
),
)
def test_rejects_dependency_without_required_protocol(
dependency_name: str,
) -> None:
_, trade_reader, quote_reader, candle_reader = make_dependencies()
dependencies: dict[str, Any] = {
"trade_reader": trade_reader,
"quote_reader": quote_reader,
"candle_revision_reader": candle_reader,
}
dependencies[dependency_name] = object()
with pytest.raises(TypeError, match=dependency_name):
MarketDataHistoricalAccess(**dependencies)
def test_delegates_each_query_without_rebuilding_page() -> None:
access, trade_reader, quote_reader, candle_reader = make_dependencies()
trade_query = TradeHistoryQuery(
venue="dzengi",
symbol="BTC/USD_LEVERAGE",
time_range=TIME_RANGE,
limit=10,
)
quote_query = QuoteHistoryQuery(
venue="dzengi",
symbol="BTC/USD_LEVERAGE",
time_range=TIME_RANGE,
limit=20,
)
candle_query = CandleRevisionHistoryQuery(
venue="dzengi",
symbol="BTC/USD_LEVERAGE",
interval="1m",
time_range=TIME_RANGE,
limit=30,
)
trade_page = access.query_trades(trade_query)
quote_page = access.query_quotes(quote_query)
candle_page = access.query_candle_revisions(candle_query)
assert trade_page is trade_reader.page
assert quote_page is quote_reader.page
assert candle_page is candle_reader.page
assert trade_reader.queries == [trade_query]
assert quote_reader.queries == [quote_query]
assert candle_reader.queries == [candle_query]
def test_does_not_swallow_reader_error() -> None:
_, _, quote_reader, candle_reader = make_dependencies()
expected = RuntimeError("reader failed")
broken_reader = BrokenTradeReader(expected)
access = MarketDataHistoricalAccess(
trade_reader=broken_reader,
quote_reader=quote_reader,
candle_revision_reader=candle_reader,
)
query = TradeHistoryQuery(
venue="dzengi",
symbol="BTC/USD_LEVERAGE",
time_range=TIME_RANGE,
)
with pytest.raises(RuntimeError) as captured:
access.query_trades(query)
assert captured.value is expected

View File

@@ -0,0 +1,620 @@
from __future__ import annotations
from contextlib import nullcontext
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.market_data.access.contracts import (
CandleRevisionHistoryReaderProtocol,
)
from src.market_data.access.exceptions import (
MarketDataAccessIntegrityError,
MarketDataAccessOperationError,
MarketDataAccessValidationError,
)
from src.market_data.access.models import (
CandleRevisionHistoryCursor,
CandleRevisionHistoryQuery,
HistoricalTimeRange,
)
from src.market_data.access.postgres_candle_revision_history_repository import (
PostgresCandleRevisionHistoryRepository,
)
from src.market_data.acquisition.models.candle import Candle
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
INTERVAL = "1m"
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
END = START + timedelta(hours=1)
SOURCE = "dzengi_websocket_candle"
_DEFAULT = object()
class RecordingCursor:
def __init__(self, rows: object = ()) -> None:
self.rows = rows
self.calls: list[tuple[str, tuple[object, ...]]] = []
self.enter_calls = 0
self.exit_exception_types: list[type[BaseException] | None] = []
self.execute_error: BaseException | None = None
self.fetchall_error: BaseException | None = None
self.exit_error: BaseException | None = None
def __enter__(self) -> RecordingCursor:
self.enter_calls += 1
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self.exit_exception_types.append(exception_type)
if self.exit_error is not None:
raise self.exit_error
return None
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
self.calls.append((sql, parameters))
if self.execute_error is not None:
raise self.execute_error
def fetchall(self) -> object:
if self.fetchall_error is not None:
raise self.fetchall_error
return self.rows
class RecordingConnection:
def __init__(self, cursor: RecordingCursor) -> None:
self._cursor = cursor
self.enter_calls = 0
self.cursor_calls = 0
self.exit_exception_types: list[type[BaseException] | None] = []
self.cursor_error: BaseException | None = None
self.exit_error: BaseException | None = None
def __enter__(self) -> RecordingConnection:
self.enter_calls += 1
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self.exit_exception_types.append(exception_type)
if self.exit_error is not None:
raise self.exit_error
return None
def cursor(self) -> RecordingCursor:
self.cursor_calls += 1
if self.cursor_error is not None:
raise self.cursor_error
return self._cursor
@dataclass
class RecordingProvider:
connection: RecordingConnection
calls: int = 0
error: BaseException | None = None
def __call__(self) -> RecordingConnection:
self.calls += 1
if self.error is not None:
raise self.error
return self.connection
def make_query(
*,
interval: str = INTERVAL,
limit: int = 3,
cursor: CandleRevisionHistoryCursor | None = None,
) -> CandleRevisionHistoryQuery:
return CandleRevisionHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
interval=interval,
time_range=HistoricalTimeRange(START, END),
limit=limit,
cursor=cursor,
)
def make_row(
*,
venue: object = VENUE,
symbol: object = SYMBOL,
interval: object = INTERVAL,
open_time: object = START + timedelta(minutes=1),
observed_at: object = _DEFAULT,
open_price: object = Decimal("64000"),
high_price: object = Decimal("64200"),
low_price: object = Decimal("63900"),
close_price: object = Decimal("64150"),
volume: object = Decimal("1.25"),
is_final: object = False,
source: object = SOURCE,
observation_sources: object = _DEFAULT,
replay_sequence: object = 10,
canonical_schema_version: object = 1,
) -> tuple[object, ...]:
resolved_observed_at = (
START + timedelta(minutes=1, seconds=10)
if observed_at is _DEFAULT
else observed_at
)
resolved_sources = (
[SOURCE]
if observation_sources is _DEFAULT
else observation_sources
)
return (
venue,
symbol,
interval,
open_time,
resolved_observed_at,
open_price,
high_price,
low_price,
close_price,
volume,
is_final,
source,
resolved_sources,
replay_sequence,
canonical_schema_version,
)
def dependencies(
rows: object = (),
) -> tuple[
PostgresCandleRevisionHistoryRepository,
RecordingCursor,
RecordingConnection,
RecordingProvider,
]:
cursor = RecordingCursor(rows)
connection = RecordingConnection(cursor)
provider = RecordingProvider(connection)
repository = PostgresCandleRevisionHistoryRepository(
connection_provider=provider,
)
return repository, cursor, connection, provider
def normalized_sql(sql: str) -> str:
return " ".join(sql.split())
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
repository, _, _, provider = dependencies()
assert provider.calls == 0
assert not hasattr(repository, "__dict__")
assert isinstance(repository, CandleRevisionHistoryReaderProtocol)
def test_constructor_rejects_non_callable_provider() -> None:
with pytest.raises(TypeError, match="connection_provider"):
PostgresCandleRevisionHistoryRepository(
connection_provider=None, # type: ignore[arg-type]
)
def test_exact_query_is_validated_before_connection_borrow() -> None:
class QuerySubclass(CandleRevisionHistoryQuery):
pass
repository, _, _, provider = dependencies()
query = QuerySubclass(
venue=VENUE,
symbol=SYMBOL,
interval=INTERVAL,
time_range=HistoricalTimeRange(START, END),
)
with pytest.raises(MarketDataAccessValidationError, match="query"):
repository.query_candle_revisions(query)
assert provider.calls == 0
def test_first_page_uses_open_time_half_open_order_and_limit_plus_one() -> None:
first_open = START + timedelta(minutes=1)
second_open = START + timedelta(minutes=2)
repository, cursor, connection, provider = dependencies(
[
make_row(open_time=first_open, replay_sequence=10),
make_row(
open_time=second_open,
observed_at=second_open + timedelta(seconds=10),
is_final=True,
replay_sequence=11,
),
]
)
query = make_query(limit=3)
page = repository.query_candle_revisions(query)
sql, parameters = cursor.calls[0]
compact_sql = normalized_sql(sql)
assert "open_time >= %s" in compact_sql
assert "open_time < %s" in compact_sql
assert "observed_at >= %s" not in compact_sql
assert "(open_time, replay_sequence) >" not in compact_sql
assert "ORDER BY open_time ASC, replay_sequence ASC" in compact_sql
assert compact_sql.endswith("LIMIT %s")
assert parameters == (VENUE, SYMBOL, INTERVAL, START, END, 4)
assert [item.replay_sequence for item in page.items] == [10, 11]
assert type(page.items[0].candle) is Candle
assert page.items[0].event_time is first_open
assert page.items[1].is_final is True
assert page.next_cursor is None
assert provider.calls == 1
assert connection.enter_calls == 1
assert connection.cursor_calls == 1
assert cursor.enter_calls == 1
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [None]
def test_interval_is_case_sensitive_and_not_normalized() -> None:
repository, cursor, _, _ = dependencies([])
query = make_query(interval="1M")
repository.query_candle_revisions(query)
assert cursor.calls[0][1][2] == "1M"
def test_keyset_query_uses_exact_open_time_cursor_tuple() -> None:
cursor_time = START + timedelta(minutes=5)
incoming = CandleRevisionHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
interval=INTERVAL,
time_range=HistoricalTimeRange(START, END),
open_time=cursor_time,
replay_sequence=25,
)
repository, cursor, _, _ = dependencies(
[
make_row(
open_time=cursor_time,
observed_at=cursor_time + timedelta(seconds=10),
replay_sequence=26,
)
]
)
repository.query_candle_revisions(make_query(limit=2, cursor=incoming))
sql, parameters = cursor.calls[0]
assert (
"(open_time, replay_sequence) > (%s, %s)"
in normalized_sql(sql)
)
assert parameters == (
VENUE,
SYMBOL,
INTERVAL,
START,
END,
cursor_time,
25,
3,
)
def test_limit_plus_one_creates_cursor_from_last_returned_item() -> None:
open_times = tuple(
START + timedelta(minutes=index) for index in (1, 2, 3)
)
repository, _, _, _ = dependencies(
[
make_row(
open_time=open_time,
observed_at=open_time + timedelta(seconds=10),
replay_sequence=10 + index,
)
for index, open_time in enumerate(open_times)
]
)
query = make_query(limit=2)
page = repository.query_candle_revisions(query)
assert [item.replay_sequence for item in page.items] == [10, 11]
assert page.next_cursor is not None
assert page.next_cursor.venue == VENUE
assert page.next_cursor.symbol == SYMBOL
assert page.next_cursor.interval == INTERVAL
assert page.next_cursor.time_range is query.time_range
assert page.next_cursor.open_time == open_times[1]
assert page.next_cursor.replay_sequence == 11
def test_empty_result_is_valid_without_cursor() -> None:
repository, _, _, _ = dependencies([])
page = repository.query_candle_revisions(make_query())
assert page.items == ()
assert page.next_cursor is None
assert page.has_more is False
def test_history_range_uses_open_time_not_observed_at() -> None:
open_time = START + timedelta(minutes=1)
observed_at = END + timedelta(minutes=10)
repository, _, _, _ = dependencies(
[make_row(open_time=open_time, observed_at=observed_at)]
)
page = repository.query_candle_revisions(make_query())
assert page.items[0].event_time == open_time
assert page.items[0].replay_at == observed_at
@pytest.mark.parametrize(
("overrides", "error_match"),
(
({"venue": "other"}, "scope"),
({"symbol": "ETH/USD_LEVERAGE"}, "scope"),
({"interval": "1M"}, "scope"),
({"open_time": END, "observed_at": END}, "range"),
),
)
def test_rows_outside_exact_scope_or_range_are_rejected(
overrides: dict[str, object],
error_match: str,
) -> None:
repository, _, _, _ = dependencies([make_row(**overrides)])
with pytest.raises(MarketDataAccessIntegrityError, match=error_match):
repository.query_candle_revisions(make_query())
@pytest.mark.parametrize(
"overrides",
(
{"venue": " dzengi"},
{"symbol": "btc/usd_leverage"},
{"interval": " 1m"},
{"open_time": START.replace(tzinfo=None)},
{"observed_at": START.replace(tzinfo=None)},
{"observed_at": START},
{"open_price": Decimal("0")},
{"high_price": Decimal("NaN")},
{"low_price": Decimal("65000")},
{"close_price": Decimal("65000")},
{"volume": Decimal("-0.01")},
{"is_final": 1},
{"source": " source"},
{"observation_sources": ()},
{"observation_sources": []},
{"observation_sources": [SOURCE, SOURCE]},
{"observation_sources": ["recovery", SOURCE]},
{"replay_sequence": True},
{"replay_sequence": 0},
{"canonical_schema_version": 2},
),
)
def test_corrupt_stored_values_raise_integrity_error(
overrides: dict[str, object],
) -> None:
repository, cursor, connection, _ = dependencies(
[make_row(**overrides)]
)
with pytest.raises(
MarketDataAccessIntegrityError,
match="invalid Canonical Candle",
):
repository.query_candle_revisions(make_query())
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
assert connection.exit_exception_types == [
MarketDataAccessIntegrityError
]
@pytest.mark.parametrize("row", ((), [object()] * 15, (object(),) * 14))
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
repository, _, _, _ = dependencies([row])
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
repository.query_candle_revisions(make_query())
def test_invalid_rows_collection_and_excess_rows_are_rejected() -> None:
repository, _, _, _ = dependencies(iter(()))
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
repository.query_candle_revisions(make_query())
repository, _, _, _ = dependencies(
[
make_row(
open_time=START + timedelta(minutes=index + 1),
observed_at=START + timedelta(minutes=index + 1, seconds=1),
replay_sequence=index + 1,
)
for index in range(3)
]
)
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
repository.query_candle_revisions(make_query(limit=1))
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
earlier = START + timedelta(minutes=1)
later = START + timedelta(minutes=2)
unordered, _, _, _ = dependencies(
[
make_row(
open_time=later,
observed_at=later,
replay_sequence=10,
),
make_row(
open_time=earlier,
observed_at=earlier,
replay_sequence=11,
),
]
)
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
unordered.query_candle_revisions(make_query())
duplicate, _, _, _ = dependencies(
[
make_row(
open_time=earlier,
observed_at=earlier,
replay_sequence=10,
),
make_row(
open_time=later,
observed_at=later,
replay_sequence=10,
),
]
)
with pytest.raises(
MarketDataAccessIntegrityError,
match="duplicate replay_sequence",
):
duplicate.query_candle_revisions(make_query())
def test_rows_must_strictly_follow_incoming_cursor() -> None:
cursor_time = START + timedelta(minutes=1)
incoming = CandleRevisionHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
interval=INTERVAL,
time_range=HistoricalTimeRange(START, END),
open_time=cursor_time,
replay_sequence=10,
)
repository, _, _, _ = dependencies(
[
make_row(
open_time=cursor_time,
observed_at=cursor_time,
replay_sequence=10,
)
]
)
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
repository.query_candle_revisions(make_query(cursor=incoming))
def test_database_error_is_wrapped_after_context_cleanup() -> None:
repository, cursor, connection, _ = dependencies()
cursor.execute_error = RuntimeError("database failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_candle_revisions(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert cursor.exit_exception_types == [RuntimeError]
assert connection.exit_exception_types == [RuntimeError]
def test_provider_error_is_wrapped_without_entering_connection() -> None:
repository, _, connection, provider = dependencies()
provider.error = RuntimeError("provider failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_candle_revisions(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert provider.calls == 1
assert connection.enter_calls == 0
def test_base_exception_is_not_wrapped_and_reaches_cleanup() -> None:
repository, cursor, connection, _ = dependencies()
cursor.execute_error = KeyboardInterrupt()
with pytest.raises(KeyboardInterrupt):
repository.query_candle_revisions(make_query())
assert cursor.exit_exception_types == [KeyboardInterrupt]
assert connection.exit_exception_types == [KeyboardInterrupt]
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
repository, cursor, connection, _ = dependencies([])
cursor.exit_error = RuntimeError("cursor cleanup failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_candle_revisions(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [RuntimeError]
repository, cursor, connection, _ = dependencies([])
cursor.exit_error = KeyboardInterrupt()
with pytest.raises(KeyboardInterrupt):
repository.query_candle_revisions(make_query())
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [KeyboardInterrupt]
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
cursor = RecordingCursor([make_row()])
connection = RecordingConnection(cursor)
provider_calls = 0
def bound_provider() -> Any:
nonlocal provider_calls
provider_calls += 1
return nullcontext(connection)
repository = PostgresCandleRevisionHistoryRepository(
connection_provider=bound_provider,
)
page = repository.query_candle_revisions(make_query())
assert len(page.items) == 1
assert provider_calls == 1
assert connection.enter_calls == 0
assert connection.exit_exception_types == []
assert cursor.exit_exception_types == [None]

View File

@@ -0,0 +1,529 @@
from __future__ import annotations
from contextlib import nullcontext
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.market_data.access.contracts import QuoteHistoryReaderProtocol
from src.market_data.access.exceptions import (
MarketDataAccessIntegrityError,
MarketDataAccessOperationError,
MarketDataAccessValidationError,
)
from src.market_data.access.models import (
HistoricalTimeRange,
QuoteHistoryCursor,
QuoteHistoryQuery,
)
from src.market_data.access.postgres_quote_history_repository import (
PostgresQuoteHistoryRepository,
)
from src.market_data.acquisition.models.quote import Quote
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
END = START + timedelta(hours=1)
SOURCE = "dzengi_websocket_quote"
class RecordingCursor:
def __init__(self, rows: object = ()) -> None:
self.rows = rows
self.calls: list[tuple[str, tuple[object, ...]]] = []
self.enter_calls = 0
self.exit_exception_types: list[type[BaseException] | None] = []
self.execute_error: BaseException | None = None
self.fetchall_error: BaseException | None = None
self.exit_error: BaseException | None = None
def __enter__(self) -> RecordingCursor:
self.enter_calls += 1
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self.exit_exception_types.append(exception_type)
if self.exit_error is not None:
raise self.exit_error
return None
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
self.calls.append((sql, parameters))
if self.execute_error is not None:
raise self.execute_error
def fetchall(self) -> object:
if self.fetchall_error is not None:
raise self.fetchall_error
return self.rows
class RecordingConnection:
def __init__(self, cursor: RecordingCursor) -> None:
self._cursor = cursor
self.enter_calls = 0
self.cursor_calls = 0
self.exit_exception_types: list[type[BaseException] | None] = []
self.cursor_error: BaseException | None = None
self.exit_error: BaseException | None = None
def __enter__(self) -> RecordingConnection:
self.enter_calls += 1
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self.exit_exception_types.append(exception_type)
if self.exit_error is not None:
raise self.exit_error
return None
def cursor(self) -> RecordingCursor:
self.cursor_calls += 1
if self.cursor_error is not None:
raise self.cursor_error
return self._cursor
@dataclass
class RecordingProvider:
connection: RecordingConnection
calls: int = 0
error: BaseException | None = None
def __call__(self) -> RecordingConnection:
self.calls += 1
if self.error is not None:
raise self.error
return self.connection
def make_query(
*,
limit: int = 3,
cursor: QuoteHistoryCursor | None = None,
) -> QuoteHistoryQuery:
return QuoteHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
limit=limit,
cursor=cursor,
)
def make_row(
*,
venue: object = VENUE,
symbol: object = SYMBOL,
received_at: object = START + timedelta(minutes=1),
exchange_timestamp: object = START + timedelta(seconds=59),
last_price: object = Decimal("64159.45"),
bid_price: object = Decimal("64159.40"),
ask_price: object = Decimal("64159.50"),
source: object = SOURCE,
observation_sources: object = None,
replay_sequence: object = 10,
canonical_schema_version: object = 1,
) -> tuple[object, ...]:
sources = [SOURCE] if observation_sources is None else observation_sources
return (
venue,
symbol,
received_at,
exchange_timestamp,
last_price,
bid_price,
ask_price,
source,
sources,
replay_sequence,
canonical_schema_version,
)
def dependencies(
rows: object = (),
) -> tuple[
PostgresQuoteHistoryRepository,
RecordingCursor,
RecordingConnection,
RecordingProvider,
]:
cursor = RecordingCursor(rows)
connection = RecordingConnection(cursor)
provider = RecordingProvider(connection)
repository = PostgresQuoteHistoryRepository(
connection_provider=provider,
)
return repository, cursor, connection, provider
def normalized_sql(sql: str) -> str:
return " ".join(sql.split())
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
repository, _, _, provider = dependencies()
assert provider.calls == 0
assert not hasattr(repository, "__dict__")
assert isinstance(repository, QuoteHistoryReaderProtocol)
def test_constructor_rejects_non_callable_provider() -> None:
with pytest.raises(TypeError, match="connection_provider"):
PostgresQuoteHistoryRepository(
connection_provider=None, # type: ignore[arg-type]
)
def test_exact_query_is_validated_before_connection_borrow() -> None:
class QuoteHistoryQuerySubclass(QuoteHistoryQuery):
pass
repository, _, _, provider = dependencies()
query = QuoteHistoryQuerySubclass(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
)
with pytest.raises(MarketDataAccessValidationError, match="query"):
repository.query_quotes(query)
assert provider.calls == 0
def test_first_page_uses_received_time_half_open_order_and_limit() -> None:
first_time = START + timedelta(minutes=1)
second_time = START + timedelta(minutes=2)
repository, cursor, connection, provider = dependencies(
[
make_row(received_at=first_time, replay_sequence=10),
make_row(received_at=second_time, replay_sequence=11),
]
)
query = make_query(limit=3)
page = repository.query_quotes(query)
sql, parameters = cursor.calls[0]
compact_sql = normalized_sql(sql)
assert "received_at >= %s" in compact_sql
assert "received_at < %s" in compact_sql
assert "(received_at, replay_sequence) >" not in compact_sql
assert "ORDER BY received_at ASC, replay_sequence ASC" in compact_sql
assert compact_sql.endswith("LIMIT %s")
assert parameters == (VENUE, SYMBOL, START, END, 4)
assert [item.replay_sequence for item in page.items] == [10, 11]
assert type(page.items[0].quote) is Quote
assert page.items[0].quote.received_at is first_time
assert page.items[0].observation_sources == (SOURCE,)
assert page.next_cursor is None
assert provider.calls == 1
assert connection.enter_calls == 1
assert connection.cursor_calls == 1
assert cursor.enter_calls == 1
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [None]
def test_keyset_query_uses_exact_received_at_and_sequence_cursor() -> None:
cursor_position = START + timedelta(minutes=5)
incoming_cursor = QuoteHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
received_at=cursor_position,
replay_sequence=25,
)
repository, recording_cursor, _, _ = dependencies(
[make_row(received_at=cursor_position, replay_sequence=26)]
)
repository.query_quotes(make_query(limit=2, cursor=incoming_cursor))
sql, parameters = recording_cursor.calls[0]
assert (
"(received_at, replay_sequence) > (%s, %s)"
in normalized_sql(sql)
)
assert parameters == (
VENUE,
SYMBOL,
START,
END,
cursor_position,
25,
3,
)
def test_limit_plus_one_creates_cursor_from_last_returned_quote() -> None:
times = tuple(START + timedelta(minutes=index) for index in (1, 2, 3))
repository, _, _, _ = dependencies(
[
make_row(received_at=event_time, replay_sequence=10 + index)
for index, event_time in enumerate(times)
]
)
query = make_query(limit=2)
page = repository.query_quotes(query)
assert len(page.items) == 2
assert [item.replay_sequence for item in page.items] == [10, 11]
assert page.next_cursor is not None
assert page.next_cursor.venue == query.venue
assert page.next_cursor.symbol == query.symbol
assert page.next_cursor.time_range is query.time_range
assert page.next_cursor.received_at == times[1]
assert page.next_cursor.replay_sequence == 11
def test_empty_result_and_exact_start_are_valid() -> None:
empty_repository, _, _, _ = dependencies([])
empty_page = empty_repository.query_quotes(make_query())
assert empty_page.items == ()
assert empty_page.next_cursor is None
assert empty_page.has_more is False
start_repository, _, _, _ = dependencies(
[make_row(received_at=START)]
)
start_page = start_repository.query_quotes(make_query())
assert start_page.items[0].event_time == START
def test_none_exchange_timestamp_is_preserved() -> None:
repository, _, _, _ = dependencies(
[make_row(exchange_timestamp=None)]
)
page = repository.query_quotes(make_query())
assert page.items[0].quote.exchange_timestamp is None
def test_backend_cannot_return_more_than_limit_plus_one() -> None:
repository, _, _, _ = dependencies(
[
make_row(
received_at=START + timedelta(minutes=index + 1),
replay_sequence=10 + index,
)
for index in range(3)
]
)
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
repository.query_quotes(make_query(limit=1))
@pytest.mark.parametrize(
"overrides",
(
{"venue": "other"},
{"venue": " dzengi "},
{"symbol": "btc/usd_leverage"},
{"received_at": END},
{"received_at": START.replace(tzinfo=None)},
{"exchange_timestamp": START.replace(tzinfo=None)},
{"last_price": Decimal("0")},
{"last_price": Decimal("NaN")},
{"last_price": 1},
{"bid_price": Decimal("64160")},
{"ask_price": Decimal("0")},
{"source": " source "},
{"observation_sources": ()},
{"observation_sources": []},
{"observation_sources": [SOURCE, SOURCE]},
{"observation_sources": ["recovery", SOURCE]},
{"replay_sequence": True},
{"replay_sequence": 0},
{"canonical_schema_version": 2},
),
)
def test_corrupt_or_out_of_scope_values_raise_integrity_error(
overrides: dict[str, object],
) -> None:
repository, cursor, connection, _ = dependencies(
[make_row(**overrides)]
)
with pytest.raises(MarketDataAccessIntegrityError):
repository.query_quotes(make_query())
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
assert connection.exit_exception_types == [MarketDataAccessIntegrityError]
@pytest.mark.parametrize(
"row",
((), [object()] * 11, (object(),) * 10),
)
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
repository, _, _, _ = dependencies([row])
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
repository.query_quotes(make_query())
@pytest.mark.parametrize("rows", (None, "rows", object()))
def test_invalid_rows_collection_raises_integrity_error(rows: object) -> None:
repository, _, _, _ = dependencies(rows)
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
repository.query_quotes(make_query())
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
later = START + timedelta(minutes=2)
earlier = START + timedelta(minutes=1)
unordered_repository, _, _, _ = dependencies(
[
make_row(received_at=later, replay_sequence=10),
make_row(received_at=earlier, replay_sequence=11),
]
)
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
unordered_repository.query_quotes(make_query())
duplicate_repository, _, _, _ = dependencies(
[
make_row(received_at=earlier, replay_sequence=10),
make_row(received_at=later, replay_sequence=10),
]
)
with pytest.raises(
MarketDataAccessIntegrityError,
match="duplicate replay_sequence",
):
duplicate_repository.query_quotes(make_query())
def test_rows_must_strictly_follow_incoming_cursor() -> None:
cursor_time = START + timedelta(minutes=1)
incoming = QuoteHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
received_at=cursor_time,
replay_sequence=10,
)
repository, _, _, _ = dependencies(
[make_row(received_at=cursor_time, replay_sequence=10)]
)
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
repository.query_quotes(make_query(cursor=incoming))
def test_database_error_is_wrapped_after_context_cleanup() -> None:
repository, cursor, connection, _ = dependencies()
cursor.execute_error = RuntimeError("database failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_quotes(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert cursor.exit_exception_types == [RuntimeError]
assert connection.exit_exception_types == [RuntimeError]
def test_provider_error_is_wrapped_without_entering_connection() -> None:
repository, _, connection, provider = dependencies()
provider.error = RuntimeError("provider failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_quotes(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert provider.calls == 1
assert connection.enter_calls == 0
def test_keyboard_interrupt_is_not_wrapped_and_reaches_cleanup() -> None:
repository, cursor, connection, _ = dependencies()
cursor.execute_error = KeyboardInterrupt()
with pytest.raises(KeyboardInterrupt):
repository.query_quotes(make_query())
assert cursor.exit_exception_types == [KeyboardInterrupt]
assert connection.exit_exception_types == [KeyboardInterrupt]
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
repository, cursor, connection, _ = dependencies([])
cursor.exit_error = RuntimeError("cursor cleanup failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_quotes(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [RuntimeError]
repository, cursor, connection, _ = dependencies([])
cursor.exit_error = KeyboardInterrupt()
with pytest.raises(KeyboardInterrupt):
repository.query_quotes(make_query())
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [KeyboardInterrupt]
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
cursor = RecordingCursor([make_row()])
connection = RecordingConnection(cursor)
provider_calls = 0
def bound_provider() -> Any:
nonlocal provider_calls
provider_calls += 1
return nullcontext(connection)
repository = PostgresQuoteHistoryRepository(
connection_provider=bound_provider,
)
page = repository.query_quotes(make_query())
assert len(page.items) == 1
assert provider_calls == 1
assert connection.enter_calls == 0
assert connection.exit_exception_types == []
assert cursor.exit_exception_types == [None]

View File

@@ -0,0 +1,622 @@
from __future__ import annotations
from contextlib import nullcontext
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.market_data.access import (
HistoricalTimeRange,
MarketDataAccessIntegrityError,
MarketDataAccessOperationError,
MarketDataAccessValidationError,
PostgresTradeHistoryRepository,
TradeHistoryCursor,
TradeHistoryQuery,
TradeHistoryReaderProtocol,
)
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.acquisition.trade_id_sequence import (
SIGNED_TRADE_ID_MAX,
SIGNED_TRADE_ID_MIN,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
END = START + timedelta(hours=1)
SOURCE = "dzengi_websocket_trade"
class RecordingCursor:
def __init__(self, rows: object = ()) -> None:
self.rows = rows
self.calls: list[tuple[str, tuple[object, ...]]] = []
self.enter_calls = 0
self.exit_exception_types: list[type[BaseException] | None] = []
self.execute_error: BaseException | None = None
self.fetchall_error: BaseException | None = None
self.exit_error: BaseException | None = None
def __enter__(self) -> RecordingCursor:
self.enter_calls += 1
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self.exit_exception_types.append(exception_type)
if self.exit_error is not None:
raise self.exit_error
return None
def execute(self, sql: str, parameters: tuple[object, ...]) -> None:
self.calls.append((sql, parameters))
if self.execute_error is not None:
raise self.execute_error
def fetchall(self) -> object:
if self.fetchall_error is not None:
raise self.fetchall_error
return self.rows
class RecordingConnection:
def __init__(self, cursor: RecordingCursor) -> None:
self._cursor = cursor
self.enter_calls = 0
self.cursor_calls = 0
self.exit_exception_types: list[type[BaseException] | None] = []
self.cursor_error: BaseException | None = None
self.exit_error: BaseException | None = None
def __enter__(self) -> RecordingConnection:
self.enter_calls += 1
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self.exit_exception_types.append(exception_type)
if self.exit_error is not None:
raise self.exit_error
return None
def cursor(self) -> RecordingCursor:
self.cursor_calls += 1
if self.cursor_error is not None:
raise self.cursor_error
return self._cursor
@dataclass
class RecordingProvider:
connection: RecordingConnection
calls: int = 0
error: BaseException | None = None
def __call__(self) -> RecordingConnection:
self.calls += 1
if self.error is not None:
raise self.error
return self.connection
def make_query(
*,
limit: int = 3,
cursor: TradeHistoryCursor | None = None,
) -> TradeHistoryQuery:
return TradeHistoryQuery(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
limit=limit,
cursor=cursor,
)
def make_row(
*,
venue: object = VENUE,
symbol: object = SYMBOL,
trade_id: object = 100,
executed_at: object = START + timedelta(minutes=1),
price: object = Decimal("64159.45"),
quantity: object = Decimal("0.125"),
aggressor_side: object = "buy",
source: object = SOURCE,
first_observed_at: object = START + timedelta(minutes=1, seconds=1),
last_observed_at: object = START + timedelta(minutes=1, seconds=2),
observation_sources: object = None,
replay_sequence: object = 10,
canonical_schema_version: object = 1,
) -> tuple[object, ...]:
sources = [SOURCE] if observation_sources is None else observation_sources
return (
venue,
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
first_observed_at,
last_observed_at,
sources,
replay_sequence,
canonical_schema_version,
)
def dependencies(
rows: object = (),
) -> tuple[
PostgresTradeHistoryRepository,
RecordingCursor,
RecordingConnection,
RecordingProvider,
]:
cursor = RecordingCursor(rows)
connection = RecordingConnection(cursor)
provider = RecordingProvider(connection)
repository = PostgresTradeHistoryRepository(
connection_provider=provider,
)
return repository, cursor, connection, provider
def normalized_sql(sql: str) -> str:
return " ".join(sql.split())
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
repository, _, _, provider = dependencies()
assert provider.calls == 0
assert not hasattr(repository, "__dict__")
assert isinstance(repository, TradeHistoryReaderProtocol)
def test_constructor_rejects_non_callable_provider() -> None:
with pytest.raises(TypeError, match="connection_provider"):
PostgresTradeHistoryRepository(
connection_provider=None, # type: ignore[arg-type]
)
def test_exact_query_is_validated_before_connection_borrow() -> None:
class TradeHistoryQuerySubclass(TradeHistoryQuery):
pass
repository, _, _, provider = dependencies()
query = TradeHistoryQuerySubclass(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
)
with pytest.raises(MarketDataAccessValidationError, match="query"):
repository.query_trades(query)
assert provider.calls == 0
def test_first_page_uses_half_open_ordered_limit_plus_one_query() -> None:
first_time = START + timedelta(minutes=1)
second_time = START + timedelta(minutes=2)
repository, cursor, connection, provider = dependencies(
[
make_row(executed_at=first_time, replay_sequence=10),
make_row(
trade_id=101,
executed_at=second_time,
first_observed_at=second_time + timedelta(seconds=1),
last_observed_at=second_time + timedelta(seconds=2),
replay_sequence=11,
),
]
)
query = make_query(limit=3)
page = repository.query_trades(query)
sql, parameters = cursor.calls[0]
compact_sql = normalized_sql(sql)
assert "executed_at >= %s" in compact_sql
assert "executed_at < %s" in compact_sql
assert "(executed_at, replay_sequence) >" not in compact_sql
assert "ORDER BY executed_at ASC, replay_sequence ASC" in compact_sql
assert compact_sql.endswith("LIMIT %s")
assert parameters == (VENUE, SYMBOL, START, END, 4)
assert [item.replay_sequence for item in page.items] == [10, 11]
assert type(page.items[0].trade) is Trade
assert page.items[0].trade.executed_at is first_time
assert page.items[0].trade.aggressor_side is TradeAggressorSide.BUY
assert page.items[0].observation_sources == (SOURCE,)
assert page.next_cursor is None
assert provider.calls == 1
assert connection.enter_calls == 1
assert connection.cursor_calls == 1
assert cursor.enter_calls == 1
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [None]
def test_keyset_query_uses_exact_cursor_tuple() -> None:
cursor_position = START + timedelta(minutes=5)
incoming_cursor = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
executed_at=cursor_position,
replay_sequence=25,
)
row_time = cursor_position
repository, recording_cursor, _, _ = dependencies(
[
make_row(
executed_at=row_time,
first_observed_at=row_time + timedelta(seconds=1),
last_observed_at=row_time + timedelta(seconds=2),
replay_sequence=26,
)
]
)
query = make_query(limit=2, cursor=incoming_cursor)
repository.query_trades(query)
sql, parameters = recording_cursor.calls[0]
assert (
"(executed_at, replay_sequence) > (%s, %s)"
in normalized_sql(sql)
)
assert parameters == (
VENUE,
SYMBOL,
START,
END,
cursor_position,
25,
3,
)
def test_trade_id_rollover_does_not_participate_in_history_order() -> None:
event_time = START + timedelta(minutes=1)
repository, _, _, _ = dependencies(
[
make_row(
trade_id=SIGNED_TRADE_ID_MAX,
executed_at=event_time,
replay_sequence=10,
),
make_row(
trade_id=SIGNED_TRADE_ID_MIN,
executed_at=event_time,
replay_sequence=11,
),
make_row(
trade_id=-1,
executed_at=event_time,
replay_sequence=12,
),
make_row(
trade_id=0,
executed_at=event_time,
replay_sequence=13,
),
]
)
page = repository.query_trades(make_query(limit=4))
assert [item.trade.trade_id for item in page.items] == [
SIGNED_TRADE_ID_MAX,
SIGNED_TRADE_ID_MIN,
-1,
0,
]
def test_limit_plus_one_creates_cursor_from_last_returned_item() -> None:
times = tuple(START + timedelta(minutes=index) for index in (1, 2, 3))
repository, _, _, _ = dependencies(
[
make_row(
trade_id=100 + index,
executed_at=event_time,
first_observed_at=event_time + timedelta(seconds=1),
last_observed_at=event_time + timedelta(seconds=2),
replay_sequence=10 + index,
)
for index, event_time in enumerate(times)
]
)
query = make_query(limit=2)
page = repository.query_trades(query)
assert len(page.items) == 2
assert [item.replay_sequence for item in page.items] == [10, 11]
assert page.next_cursor is not None
assert page.next_cursor.venue == query.venue
assert page.next_cursor.symbol == query.symbol
assert page.next_cursor.time_range is query.time_range
assert page.next_cursor.executed_at == times[1]
assert page.next_cursor.replay_sequence == 11
def test_empty_result_is_valid_without_cursor() -> None:
repository, _, _, _ = dependencies([])
page = repository.query_trades(make_query())
assert page.items == ()
assert page.next_cursor is None
assert page.has_more is False
def test_half_open_range_includes_exact_start() -> None:
repository, _, _, _ = dependencies(
[
make_row(
executed_at=START,
first_observed_at=START,
last_observed_at=START,
)
]
)
page = repository.query_trades(make_query())
assert page.items[0].event_time == START
def test_backend_cannot_return_more_than_limit_plus_one() -> None:
query = make_query(limit=1)
rows = [
make_row(
trade_id=100 + index,
executed_at=START + timedelta(minutes=index + 1),
first_observed_at=START + timedelta(minutes=index + 1, seconds=1),
last_observed_at=START + timedelta(minutes=index + 1, seconds=2),
replay_sequence=10 + index,
)
for index in range(3)
]
repository, _, _, _ = dependencies(rows)
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
repository.query_trades(query)
@pytest.mark.parametrize(
("overrides", "error_match"),
(
({"venue": "other"}, "invalid Canonical Trade"),
({"symbol": "btc/usd_leverage"}, "invalid Canonical Trade"),
({"trade_id": True}, "invalid Canonical Trade"),
({"trade_id": SIGNED_TRADE_ID_MAX + 1}, "invalid Canonical Trade"),
({"executed_at": END}, "invalid Canonical Trade"),
({"price": Decimal("0")}, "invalid Canonical Trade"),
({"price": Decimal("NaN")}, "invalid Canonical Trade"),
({"quantity": Decimal("0")}, "invalid Canonical Trade"),
({"aggressor_side": "hold"}, "invalid Canonical Trade"),
({"source": " source "}, "invalid Canonical Trade"),
(
{"first_observed_at": START.replace(tzinfo=None)},
"invalid Canonical Trade",
),
(
{
"first_observed_at": START + timedelta(minutes=3),
"last_observed_at": START + timedelta(minutes=2),
},
"invalid Canonical Trade",
),
({"observation_sources": ()}, "invalid Canonical Trade"),
({"observation_sources": []}, "invalid Canonical Trade"),
(
{"observation_sources": [SOURCE, SOURCE]},
"invalid Canonical Trade",
),
(
{"observation_sources": ["recovery", SOURCE]},
"invalid Canonical Trade",
),
({"replay_sequence": True}, "invalid Canonical Trade"),
({"replay_sequence": 0}, "invalid Canonical Trade"),
({"canonical_schema_version": 2}, "invalid Canonical Trade"),
),
)
def test_corrupt_stored_values_raise_integrity_error(
overrides: dict[str, object],
error_match: str,
) -> None:
repository, cursor, connection, _ = dependencies(
[make_row(**overrides)]
)
with pytest.raises(MarketDataAccessIntegrityError, match=error_match):
repository.query_trades(make_query())
assert cursor.exit_exception_types == [MarketDataAccessIntegrityError]
assert connection.exit_exception_types == [MarketDataAccessIntegrityError]
@pytest.mark.parametrize("row", ((), [object()] * 13, (object(),) * 12))
def test_invalid_row_shape_raises_integrity_error(row: object) -> None:
repository, _, _, _ = dependencies([row])
with pytest.raises(MarketDataAccessIntegrityError, match="row"):
repository.query_trades(make_query())
def test_unordered_rows_and_duplicate_global_sequence_are_rejected() -> None:
later = START + timedelta(minutes=2)
earlier = START + timedelta(minutes=1)
unordered_repository, _, _, _ = dependencies(
[
make_row(
executed_at=later,
first_observed_at=later + timedelta(seconds=1),
last_observed_at=later + timedelta(seconds=2),
replay_sequence=10,
),
make_row(
executed_at=earlier,
first_observed_at=earlier + timedelta(seconds=1),
last_observed_at=earlier + timedelta(seconds=2),
replay_sequence=11,
),
]
)
with pytest.raises(MarketDataAccessIntegrityError, match="ordered"):
unordered_repository.query_trades(make_query())
duplicate_repository, _, _, _ = dependencies(
[
make_row(
executed_at=earlier,
first_observed_at=earlier + timedelta(seconds=1),
last_observed_at=earlier + timedelta(seconds=2),
replay_sequence=10,
),
make_row(
trade_id=101,
executed_at=later,
first_observed_at=later + timedelta(seconds=1),
last_observed_at=later + timedelta(seconds=2),
replay_sequence=10,
),
]
)
with pytest.raises(
MarketDataAccessIntegrityError,
match="duplicate replay_sequence",
):
duplicate_repository.query_trades(make_query())
def test_rows_must_strictly_follow_incoming_cursor() -> None:
cursor_time = START + timedelta(minutes=1)
incoming = TradeHistoryCursor(
venue=VENUE,
symbol=SYMBOL,
time_range=HistoricalTimeRange(START, END),
executed_at=cursor_time,
replay_sequence=10,
)
repository, _, _, _ = dependencies(
[make_row(executed_at=cursor_time, replay_sequence=10)]
)
with pytest.raises(MarketDataAccessIntegrityError, match="cursor"):
repository.query_trades(make_query(cursor=incoming))
def test_database_error_is_wrapped_after_context_cleanup() -> None:
repository, cursor, connection, _ = dependencies()
cursor.execute_error = RuntimeError("database failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_trades(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert cursor.exit_exception_types == [RuntimeError]
assert connection.exit_exception_types == [RuntimeError]
def test_provider_error_is_wrapped_without_entering_connection() -> None:
repository, _, connection, provider = dependencies()
provider.error = RuntimeError("provider failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_trades(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert provider.calls == 1
assert connection.enter_calls == 0
def test_keyboard_interrupt_is_not_wrapped_and_reaches_cleanup() -> None:
repository, cursor, connection, _ = dependencies()
cursor.execute_error = KeyboardInterrupt()
with pytest.raises(KeyboardInterrupt):
repository.query_trades(make_query())
assert cursor.exit_exception_types == [KeyboardInterrupt]
assert connection.exit_exception_types == [KeyboardInterrupt]
def test_cleanup_error_obeys_exception_and_base_exception_contract() -> None:
repository, cursor, connection, _ = dependencies([])
cursor.exit_error = RuntimeError("cursor cleanup failed")
with pytest.raises(MarketDataAccessOperationError) as error_info:
repository.query_trades(make_query())
assert isinstance(error_info.value.__cause__, RuntimeError)
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [RuntimeError]
repository, cursor, connection, _ = dependencies([])
cursor.exit_error = KeyboardInterrupt()
with pytest.raises(KeyboardInterrupt):
repository.query_trades(make_query())
assert cursor.exit_exception_types == [None]
assert connection.exit_exception_types == [KeyboardInterrupt]
def test_bound_provider_can_reuse_caller_owned_connection() -> None:
cursor = RecordingCursor([make_row()])
connection = RecordingConnection(cursor)
provider_calls = 0
def bound_provider() -> Any:
nonlocal provider_calls
provider_calls += 1
return nullcontext(connection)
repository = PostgresTradeHistoryRepository(
connection_provider=bound_provider,
)
page = repository.query_trades(make_query())
assert len(page.items) == 1
assert provider_calls == 1
assert connection.enter_calls == 0
assert connection.exit_exception_types == []
assert cursor.exit_exception_types == [None]

View File

@@ -0,0 +1,28 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
import pytest
from src.market_data.access import HistoricalTimeRange
from src.market_data.replay import (
ReplayDataType,
ReplayPlan,
ReplayPlanRequest,
)
@pytest.fixture
def empty_replay_plan() -> ReplayPlan:
start = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
request = ReplayPlanRequest(
venue="dzengi",
symbols=("BTC/USD_LEVERAGE",),
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRange(
start_time=start,
end_time=start + timedelta(hours=1),
),
max_records=100,
)
return ReplayPlan(request=request, events=())

View File

@@ -0,0 +1,214 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
import inspect
import pytest
from src.market_data.replay.contracts import (
MarketDataClockProtocol,
ReplayClockProtocol,
)
from src.market_data.replay.deterministic_replay_clock import (
DeterministicReplayClock,
)
from src.market_data.replay.exceptions import ReplayClockError
START = datetime(
2026,
8,
2,
12,
0,
0,
123456,
tzinfo=timezone.utc,
)
def test_matches_protocols_uses_slots_and_is_synchronous() -> None:
clock = DeterministicReplayClock(START)
assert isinstance(clock, MarketDataClockProtocol)
assert isinstance(clock, ReplayClockProtocol)
assert not hasattr(clock, "__dict__")
assert inspect.iscoroutinefunction(clock.advance_to) is False
def test_starts_at_exact_canonical_utc_time() -> None:
clock = DeterministicReplayClock(START)
assert clock.now == START
assert type(clock.now) is datetime
assert clock.now.tzinfo is timezone.utc
def test_normalizes_non_utc_initial_time_without_losing_precision() -> None:
offset = timezone(timedelta(hours=3, minutes=30))
initial_time = datetime(
2026,
8,
2,
15,
30,
0,
654321,
tzinfo=offset,
)
clock = DeterministicReplayClock(initial_time)
assert clock.now == datetime(
2026,
8,
2,
12,
0,
0,
654321,
tzinfo=timezone.utc,
)
assert type(clock.now) is datetime
def test_accepts_datetime_subclass_but_stores_base_datetime() -> None:
class CompatibleDatetime(datetime):
pass
initial_time = CompatibleDatetime(
2026,
8,
2,
12,
0,
tzinfo=timezone.utc,
)
clock = DeterministicReplayClock(initial_time)
assert clock.now == initial_time
assert type(clock.now) is datetime
@pytest.mark.parametrize("invalid", (None, "2026-08-02", 1, object()))
def test_rejects_non_datetime_initial_value(invalid: object) -> None:
with pytest.raises(TypeError, match="initial_time"):
DeterministicReplayClock(invalid) # type: ignore[arg-type]
def test_rejects_naive_initial_time() -> None:
with pytest.raises(ValueError, match="timezone"):
DeterministicReplayClock(START.replace(tzinfo=None))
def test_advances_forward_with_microsecond_precision() -> None:
clock = DeterministicReplayClock(START)
expected = START + timedelta(microseconds=1)
result = clock.advance_to(expected)
assert result is None
assert clock.now == expected
def test_allows_repeated_equal_time_for_distinct_replay_events() -> None:
clock = DeterministicReplayClock(START)
clock.advance_to(START)
clock.advance_to(START)
assert clock.now == START
def test_equal_instant_with_another_offset_is_idempotent() -> None:
clock = DeterministicReplayClock(START)
equal_instant = START.astimezone(timezone(timedelta(hours=-4)))
clock.advance_to(equal_instant)
assert clock.now == START
assert clock.now.tzinfo is timezone.utc
def test_equal_utc_instant_with_another_fold_is_complete_no_op() -> None:
clock = DeterministicReplayClock(START)
previous = clock.now
clock.advance_to(START.replace(fold=1))
assert clock.now is previous
assert clock.now.fold == 0
def test_accepts_datetime_subclass_during_advance() -> None:
class CompatibleDatetime(datetime):
pass
clock = DeterministicReplayClock(START)
later = CompatibleDatetime(
2026,
8,
2,
12,
1,
tzinfo=timezone.utc,
)
clock.advance_to(later)
assert clock.now == later
assert type(clock.now) is datetime
def test_rejects_backward_transition_without_changing_state() -> None:
clock = DeterministicReplayClock(START)
later = START + timedelta(minutes=1)
clock.advance_to(later)
with pytest.raises(ReplayClockError, match="backwards"):
clock.advance_to(START)
assert clock.now == later
@pytest.mark.parametrize("invalid", (None, "later", 1, object()))
def test_invalid_advance_type_does_not_change_state(invalid: object) -> None:
clock = DeterministicReplayClock(START)
with pytest.raises(TypeError, match="instant"):
clock.advance_to(invalid) # type: ignore[arg-type]
assert clock.now == START
def test_naive_advance_does_not_change_state() -> None:
clock = DeterministicReplayClock(START)
with pytest.raises(ValueError, match="timezone"):
clock.advance_to(START.replace(tzinfo=None))
assert clock.now == START
def test_two_clocks_have_independent_state() -> None:
first = DeterministicReplayClock(START)
second = DeterministicReplayClock(START)
first.advance_to(START + timedelta(hours=1))
assert first.now == START + timedelta(hours=1)
assert second.now == START
def test_now_is_read_only_and_lifecycle_extensions_are_absent() -> None:
clock = DeterministicReplayClock(START)
with pytest.raises(AttributeError):
setattr(clock, "now", START + timedelta(hours=1))
assert not hasattr(clock, "reset")
assert not hasattr(clock, "advance_by")
assert not hasattr(clock, "start")
assert not hasattr(clock, "stop")
assert clock.now == START

View File

@@ -0,0 +1,94 @@
from __future__ import annotations
from datetime import datetime
from src.market_data.replay import (
MarketDataClockProtocol,
MarketDataReplayError,
MarketDataReplayValidationError,
ReplayClockError,
ReplayClockProtocol,
ReplayConsumerProtocol,
ReplayEvent,
ReplayPlan,
ReplayPlanBuilderProtocol,
ReplayPlanLimitExceededError,
ReplayPlanRequest,
ReplaySessionProtocol,
ReplaySessionState,
ReplaySessionStateError,
)
class FakeClock:
def __init__(self, now: datetime) -> None:
self._now = now
@property
def now(self) -> datetime:
return self._now
def advance_to(self, instant: datetime) -> None:
self._now = instant
class RecordingConsumer:
def __init__(self) -> None:
self.events: list[ReplayEvent] = []
async def consume(self, event: ReplayEvent) -> None:
self.events.append(event)
class RecordingPlanBuilder:
def __init__(self, plan: ReplayPlan) -> None:
self.plan = plan
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
return self.plan
class RecordingSession:
def __init__(self, plan: ReplayPlan, clock: FakeClock) -> None:
self._plan = plan
self._clock = clock
@property
def state(self) -> ReplaySessionState:
return ReplaySessionState.CREATED
@property
def plan(self) -> ReplayPlan:
return self._plan
@property
def clock(self) -> MarketDataClockProtocol:
return self._clock
async def run(self) -> None:
return None
def test_replay_protocols_are_runtime_checkable(
empty_replay_plan: ReplayPlan,
) -> None:
clock = FakeClock(empty_replay_plan.request.time_range.start_time)
assert isinstance(clock, MarketDataClockProtocol)
assert isinstance(clock, ReplayClockProtocol)
assert isinstance(RecordingConsumer(), ReplayConsumerProtocol)
assert isinstance(
RecordingPlanBuilder(empty_replay_plan),
ReplayPlanBuilderProtocol,
)
assert isinstance(
RecordingSession(empty_replay_plan, clock),
ReplaySessionProtocol,
)
def test_replay_error_hierarchy_is_specialized() -> None:
assert issubclass(MarketDataReplayValidationError, MarketDataReplayError)
assert issubclass(ReplayPlanLimitExceededError, MarketDataReplayError)
assert issubclass(ReplayClockError, MarketDataReplayError)
assert issubclass(ReplaySessionStateError, MarketDataReplayError)

View File

@@ -0,0 +1,659 @@
from __future__ import annotations
from dataclasses import FrozenInstanceError
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.market_data.access import HistoricalTimeRange
from src.market_data.acquisition.models.candle import Candle
from src.market_data.acquisition.models.quote import Quote
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.replay import (
REPLAY_PLAN_MAX_RECORDS_LIMIT,
ReplayDataType,
ReplayEvent,
ReplayPlan,
ReplayPlanRequest,
ReplaySessionState,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
END = START + timedelta(hours=1)
def make_trade(
*,
trade_id: int = 100,
executed_at: datetime = START + timedelta(minutes=1),
symbol: str = SYMBOL,
) -> Trade:
return Trade(
symbol=symbol,
trade_id=trade_id,
price=Decimal("65000"),
quantity=Decimal("0.001"),
executed_at=executed_at,
aggressor_side=TradeAggressorSide.BUY,
source="dzengi_websocket_trade",
)
def make_quote(
*,
received_at: datetime = START + timedelta(minutes=2),
) -> Quote:
return Quote(
symbol=SYMBOL,
last_price=Decimal("65000"),
bid_price=Decimal("64999"),
ask_price=Decimal("65001"),
exchange_timestamp=received_at - timedelta(milliseconds=1),
received_at=received_at,
source="dzengi_rest_quote",
)
def make_candle(*, interval: str = "1m") -> Candle:
return Candle(
symbol=SYMBOL,
interval=interval,
open_time=START,
open_price=Decimal("64900"),
high_price=Decimal("65100"),
low_price=Decimal("64800"),
close_price=Decimal("65000"),
volume=Decimal("10"),
source="dzengi_rest_candle",
)
def make_request(
*,
data_types: tuple[ReplayDataType, ...] = (ReplayDataType.TRADE,),
candle_intervals: tuple[str, ...] = (),
max_records: int = 100,
) -> ReplayPlanRequest:
return ReplayPlanRequest(
venue=VENUE,
symbols=(SYMBOL,),
data_types=data_types,
time_range=HistoricalTimeRange(
start_time=START,
end_time=END,
),
candle_intervals=candle_intervals,
max_records=max_records,
)
def make_trade_event(
*,
trade_id: int = 100,
replay_at: datetime = START + timedelta(minutes=1),
replay_sequence: int = 1,
venue: str = VENUE,
symbol: str = SYMBOL,
) -> ReplayEvent:
trade = make_trade(
trade_id=trade_id,
executed_at=replay_at,
symbol=symbol,
)
return ReplayEvent(
venue=venue,
replay_at=replay_at,
replay_sequence=replay_sequence,
payload=trade,
)
def trade_id_from_event(event: ReplayEvent) -> int:
payload = event.payload
if not isinstance(payload, Trade):
raise AssertionError("Replay event must contain Trade payload.")
return payload.trade_id
def test_trade_event_preserves_payload_identity_and_exact_event_time() -> None:
trade = make_trade()
event = ReplayEvent(
venue=" dzengi ",
replay_at=trade.executed_at,
replay_sequence=10,
payload=trade,
)
assert event.venue == VENUE
assert event.payload is trade
assert event.symbol == SYMBOL
assert event.data_type is ReplayDataType.TRADE
assert event.order_key == (trade.executed_at, 10)
assert event.candle_is_final is None
def test_quote_event_uses_received_time_not_exchange_time() -> None:
quote = make_quote()
event = ReplayEvent(
venue=VENUE,
replay_at=quote.received_at,
replay_sequence=11,
payload=quote,
)
assert event.data_type is ReplayDataType.QUOTE
assert event.replay_at == quote.received_at
assert event.replay_at != quote.exchange_timestamp
def test_candle_event_uses_observation_time_and_final_metadata() -> None:
candle = make_candle()
observed_at = candle.open_time + timedelta(seconds=30)
event = ReplayEvent(
venue=VENUE,
replay_at=observed_at,
replay_sequence=12,
payload=candle,
candle_is_final=False,
)
assert event.data_type is ReplayDataType.CANDLE_REVISION
assert event.replay_at == observed_at
assert event.candle_is_final is False
def test_event_normalizes_aware_time_to_utc() -> None:
offset = timezone(timedelta(hours=3))
trade = make_trade()
event = ReplayEvent(
venue=VENUE,
replay_at=trade.executed_at.astimezone(offset),
replay_sequence=1,
payload=trade,
)
assert event.replay_at.tzinfo is timezone.utc
assert event.replay_at == trade.executed_at
@pytest.mark.parametrize("invalid_sequence", (True, 0, -1, 1.5, "1"))
def test_event_rejects_invalid_sequence(invalid_sequence: Any) -> None:
trade = make_trade()
with pytest.raises((TypeError, ValueError), match="replay_sequence"):
ReplayEvent(
venue=VENUE,
replay_at=trade.executed_at,
replay_sequence=invalid_sequence,
payload=trade,
)
def test_trade_event_rejects_mismatched_replay_time() -> None:
trade = make_trade()
with pytest.raises(ValueError, match="event time"):
ReplayEvent(
venue=VENUE,
replay_at=trade.executed_at + timedelta(microseconds=1),
replay_sequence=1,
payload=trade,
)
def test_candle_event_rejects_time_before_open() -> None:
candle = make_candle()
with pytest.raises(ValueError, match="open_time"):
ReplayEvent(
venue=VENUE,
replay_at=candle.open_time - timedelta(microseconds=1),
replay_sequence=1,
payload=candle,
candle_is_final=False,
)
@pytest.mark.parametrize("candle_is_final", (None, 0, 1, "true"))
def test_candle_event_requires_exact_final_boolean(
candle_is_final: Any,
) -> None:
candle = make_candle()
with pytest.raises(TypeError, match="candle_is_final"):
ReplayEvent(
venue=VENUE,
replay_at=candle.open_time,
replay_sequence=1,
payload=candle,
candle_is_final=candle_is_final,
)
def test_trade_event_rejects_candle_metadata() -> None:
trade = make_trade()
with pytest.raises(ValueError, match="must be None"):
ReplayEvent(
venue=VENUE,
replay_at=trade.executed_at,
replay_sequence=1,
payload=trade,
candle_is_final=False,
)
def test_request_normalizes_symbols_and_preserves_interval_case() -> None:
request = ReplayPlanRequest(
venue=" dzengi ",
symbols=(" btc/usd_leverage ",),
data_types=(ReplayDataType.CANDLE_REVISION,),
time_range=HistoricalTimeRange(START, END),
candle_intervals=(" 1M ",),
)
assert request.venue == VENUE
assert request.symbols == (SYMBOL,)
assert request.candle_intervals == ("1M",)
@pytest.mark.parametrize(
"symbols",
(
[],
(),
("",),
("BTC/USD_LEVERAGE", "btc/usd_leverage"),
),
)
def test_request_rejects_invalid_symbols(symbols: Any) -> None:
with pytest.raises((TypeError, ValueError)):
ReplayPlanRequest(
venue=VENUE,
symbols=symbols,
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRange(START, END),
)
def test_request_requires_intervals_only_for_candles() -> None:
with pytest.raises(ValueError, match="required"):
make_request(data_types=(ReplayDataType.CANDLE_REVISION,))
with pytest.raises(ValueError, match="require Candle"):
make_request(candle_intervals=("1m",))
@pytest.mark.parametrize(
"max_records",
(True, 0, -1, 1.5, REPLAY_PLAN_MAX_RECORDS_LIMIT + 1),
)
def test_request_rejects_invalid_max_records(max_records: Any) -> None:
with pytest.raises((TypeError, ValueError), match="max_records"):
make_request(max_records=max_records)
def test_empty_plan_is_valid_and_preserves_request_scope() -> None:
request = make_request()
plan = ReplayPlan(request=request, events=())
assert plan.request is request
assert plan.events == ()
assert plan.is_empty is True
assert len(plan) == 0
def test_plan_preserves_event_and_payload_identity() -> None:
request = make_request()
event = make_trade_event()
plan = ReplayPlan(request=request, events=(event,))
assert plan.events[0] is event
assert plan.events[0].payload is event.payload
def test_plan_accepts_rollover_order_by_time_and_sequence() -> None:
same_time = START + timedelta(minutes=1)
first = make_trade_event(
trade_id=2_147_483_647,
replay_at=same_time,
replay_sequence=10,
)
second = make_trade_event(
trade_id=-2_147_483_648,
replay_at=same_time,
replay_sequence=11,
)
plan = ReplayPlan(
request=make_request(),
events=(first, second),
)
assert [trade_id_from_event(event) for event in plan.events] == [
2_147_483_647,
-2_147_483_648,
]
def test_plan_accepts_negative_one_to_zero_at_equal_time() -> None:
same_time = START + timedelta(minutes=1)
plan = ReplayPlan(
request=make_request(),
events=(
make_trade_event(
trade_id=-1,
replay_at=same_time,
replay_sequence=20,
),
make_trade_event(
trade_id=0,
replay_at=same_time,
replay_sequence=21,
),
),
)
assert [trade_id_from_event(event) for event in plan.events] == [-1, 0]
def test_plan_rejects_reverse_order() -> None:
with pytest.raises(ValueError, match="strictly ordered"):
ReplayPlan(
request=make_request(),
events=(
make_trade_event(replay_sequence=2),
make_trade_event(replay_sequence=1),
),
)
def test_plan_rejects_reverse_event_time_with_increasing_sequence() -> None:
with pytest.raises(ValueError, match="strictly ordered"):
ReplayPlan(
request=make_request(),
events=(
make_trade_event(
replay_at=START + timedelta(minutes=2),
replay_sequence=1,
),
make_trade_event(
replay_at=START + timedelta(minutes=1),
replay_sequence=2,
),
),
)
def test_plan_rejects_duplicate_global_sequence() -> None:
with pytest.raises(ValueError, match="globally unique"):
ReplayPlan(
request=make_request(),
events=(
make_trade_event(
replay_at=START + timedelta(minutes=1),
replay_sequence=1,
),
make_trade_event(
replay_at=START + timedelta(minutes=2),
replay_sequence=1,
),
),
)
def test_plan_rejects_event_outside_time_range() -> None:
event = make_trade_event(replay_at=END)
with pytest.raises(ValueError, match="outside Replay request"):
ReplayPlan(request=make_request(), events=(event,))
def test_plan_rejects_data_type_outside_request() -> None:
quote = make_quote()
event = ReplayEvent(
venue=VENUE,
replay_at=quote.received_at,
replay_sequence=1,
payload=quote,
)
with pytest.raises(ValueError, match="data type"):
ReplayPlan(request=make_request(), events=(event,))
def test_plan_rejects_event_from_another_venue() -> None:
event = make_trade_event(venue="other")
with pytest.raises(ValueError, match="event venue"):
ReplayPlan(request=make_request(), events=(event,))
def test_plan_rejects_event_from_another_symbol() -> None:
event = make_trade_event(symbol="ETH/USD_LEVERAGE")
with pytest.raises(ValueError, match="event symbol"):
ReplayPlan(request=make_request(), events=(event,))
def test_plan_rejects_candle_interval_outside_request() -> None:
candle = make_candle(interval="5m")
event = ReplayEvent(
venue=VENUE,
replay_at=candle.open_time,
replay_sequence=1,
payload=candle,
candle_is_final=False,
)
request = make_request(
data_types=(ReplayDataType.CANDLE_REVISION,),
candle_intervals=("1m",),
)
with pytest.raises(ValueError, match="Candle interval"):
ReplayPlan(request=request, events=(event,))
def test_plan_rejects_more_events_than_request_limit() -> None:
event = make_trade_event()
with pytest.raises(ValueError, match="max_records"):
ReplayPlan(
request=make_request(max_records=1),
events=(event, make_trade_event(replay_sequence=2)),
)
def test_replay_models_are_frozen_and_slotted() -> None:
event = make_trade_event()
request = make_request()
plan = ReplayPlan(request=request, events=(event,))
for model in (event, request, plan):
assert not hasattr(model, "__dict__")
with pytest.raises(FrozenInstanceError):
setattr(event, "venue", "other")
def test_session_states_are_explicit_and_complete() -> None:
assert tuple(ReplaySessionState) == (
ReplaySessionState.CREATED,
ReplaySessionState.RUNNING,
ReplaySessionState.COMPLETED,
ReplaySessionState.FAILED,
ReplaySessionState.CANCELLED,
)
def test_event_rejects_canonical_payload_subclasses() -> None:
class TradeSubclass(Trade):
pass
class QuoteSubclass(Quote):
pass
class CandleSubclass(Candle):
pass
trade = make_trade()
quote = make_quote()
candle = make_candle()
payloads = (
(
TradeSubclass(
symbol=trade.symbol,
trade_id=trade.trade_id,
price=trade.price,
quantity=trade.quantity,
executed_at=trade.executed_at,
aggressor_side=trade.aggressor_side,
source=trade.source,
),
trade.executed_at,
None,
),
(
QuoteSubclass(
symbol=quote.symbol,
last_price=quote.last_price,
bid_price=quote.bid_price,
ask_price=quote.ask_price,
exchange_timestamp=quote.exchange_timestamp,
received_at=quote.received_at,
source=quote.source,
),
quote.received_at,
None,
),
(
CandleSubclass(
symbol=candle.symbol,
interval=candle.interval,
open_time=candle.open_time,
open_price=candle.open_price,
high_price=candle.high_price,
low_price=candle.low_price,
close_price=candle.close_price,
volume=candle.volume,
source=candle.source,
),
candle.open_time,
False,
),
)
for payload, replay_at, candle_is_final in payloads:
with pytest.raises(TypeError, match="Canonical Trade, Quote or Candle"):
ReplayEvent(
venue=VENUE,
replay_at=replay_at,
replay_sequence=1,
payload=payload,
candle_is_final=candle_is_final,
)
def test_request_rejects_time_range_subclass() -> None:
class HistoricalTimeRangeSubclass(HistoricalTimeRange):
def contains(self, instant: datetime) -> bool:
return True
with pytest.raises(TypeError, match="HistoricalTimeRange"):
ReplayPlanRequest(
venue=VENUE,
symbols=(SYMBOL,),
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRangeSubclass(START, END),
)
def test_request_rejects_tuple_subclasses() -> None:
class TupleSubclass(tuple):
pass
with pytest.raises(TypeError, match="symbols must be a tuple"):
ReplayPlanRequest(
venue=VENUE,
symbols=TupleSubclass((SYMBOL,)),
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRange(START, END),
)
with pytest.raises(TypeError, match="data_types must be a tuple"):
ReplayPlanRequest(
venue=VENUE,
symbols=(SYMBOL,),
data_types=TupleSubclass((ReplayDataType.TRADE,)),
time_range=HistoricalTimeRange(START, END),
)
with pytest.raises(TypeError, match="candle_intervals must be a tuple"):
ReplayPlanRequest(
venue=VENUE,
symbols=(SYMBOL,),
data_types=(ReplayDataType.CANDLE_REVISION,),
time_range=HistoricalTimeRange(START, END),
candle_intervals=TupleSubclass(("1m",)),
)
def test_plan_rejects_request_and_event_subclasses() -> None:
class ReplayPlanRequestSubclass(ReplayPlanRequest):
pass
class ReplayEventSubclass(ReplayEvent):
@property
def symbol(self) -> str:
return SYMBOL
@property
def data_type(self) -> ReplayDataType:
return ReplayDataType.TRADE
@property
def order_key(self) -> tuple[datetime, int]:
return (START, 1)
request_subclass = ReplayPlanRequestSubclass(
venue=VENUE,
symbols=(SYMBOL,),
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRange(START, END),
)
with pytest.raises(TypeError, match="ReplayPlanRequest"):
ReplayPlan(request=request_subclass, events=())
event = make_trade_event()
event_subclass = ReplayEventSubclass(
venue=event.venue,
replay_at=event.replay_at,
replay_sequence=event.replay_sequence,
payload=event.payload,
)
with pytest.raises(TypeError, match="ReplayEvent"):
ReplayPlan(request=make_request(), events=(event_subclass,))
def test_plan_rejects_events_tuple_subclass() -> None:
class EventsTupleSubclass(tuple):
pass
with pytest.raises(TypeError, match="events must be a tuple"):
ReplayPlan(
request=make_request(),
events=EventsTupleSubclass((make_trade_event(),)),
)

View File

@@ -0,0 +1,699 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.market_data.access.exceptions import (
MarketDataAccessIntegrityError,
MarketDataAccessOperationError,
)
from src.market_data.access.models import HistoricalTimeRange
from src.market_data.acquisition.models.candle import Candle
from src.market_data.acquisition.models.quote import Quote
from src.market_data.acquisition.models.trade import Trade
from src.market_data.replay.contracts import ReplayPlanBuilderProtocol
from src.market_data.replay.exceptions import (
MarketDataReplayValidationError,
ReplayPlanLimitExceededError,
)
from src.market_data.replay.models import (
ReplayDataType,
ReplayPlanRequest,
)
from src.market_data.replay.postgres_replay_plan_builder import (
PostgresReplayPlanBuilder,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
END = START + timedelta(hours=1)
TRADE_SOURCE = "dzengi_websocket_trade"
QUOTE_SOURCE = "dzengi_rest_quote"
CANDLE_SOURCE = "dzengi_rest_candle"
class RecordingCursor:
def __init__(
self,
*,
events: list[tuple[object, ...]],
rows: object = (),
) -> None:
self._events = events
self.rows = rows
self.calls: list[tuple[str, object]] = []
self.execute_errors: list[BaseException | None] = []
self.fetchall_error: BaseException | None = None
def __enter__(self) -> RecordingCursor:
self._events.append(("cursor_enter",))
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self._events.append(("cursor_exit", exception_type))
return None
def execute(
self,
sql: str,
parameters: object = None,
) -> None:
self.calls.append((sql, parameters))
self._events.append(("execute", normalized_sql(sql), parameters))
error = (
self.execute_errors.pop(0)
if self.execute_errors
else None
)
if error is not None:
raise error
def fetchall(self) -> object:
self._events.append(("fetchall",))
if self.fetchall_error is not None:
raise self.fetchall_error
return self.rows
class RecordingTransaction:
def __init__(self, events: list[tuple[object, ...]]) -> None:
self._events = events
def __enter__(self) -> RecordingTransaction:
self._events.append(("transaction_enter",))
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self._events.append(("transaction_exit", exception_type))
return None
class RecordingConnection:
def __init__(
self,
*,
events: list[tuple[object, ...]],
cursor: RecordingCursor,
) -> None:
self._events = events
self._cursor = cursor
self.transaction_calls = 0
self.cursor_calls = 0
def __enter__(self) -> RecordingConnection:
self._events.append(("connection_enter",))
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self._events.append(("connection_exit", exception_type))
return None
def transaction(self) -> RecordingTransaction:
self.transaction_calls += 1
self._events.append(("transaction",))
return RecordingTransaction(self._events)
def cursor(self) -> RecordingCursor:
self.cursor_calls += 1
self._events.append(("cursor",))
return self._cursor
@dataclass
class RecordingProvider:
connection: RecordingConnection
events: list[tuple[object, ...]]
calls: int = 0
error: BaseException | None = None
def __call__(self) -> RecordingConnection:
self.calls += 1
self.events.append(("provider",))
if self.error is not None:
raise self.error
return self.connection
def normalized_sql(sql: str) -> str:
return " ".join(sql.split())
def make_request(
*,
data_types: tuple[ReplayDataType, ...] = (ReplayDataType.TRADE,),
candle_intervals: tuple[str, ...] = (),
max_records: int = 10,
venue: str = VENUE,
) -> ReplayPlanRequest:
return ReplayPlanRequest(
venue=venue,
symbols=(SYMBOL,),
data_types=data_types,
time_range=HistoricalTimeRange(START, END),
candle_intervals=candle_intervals,
max_records=max_records,
)
def make_trade_row(
*,
replay_at: object = START + timedelta(minutes=1),
replay_sequence: object = 10,
venue: object = VENUE,
trade_id: object = 100,
price: object = Decimal("65000"),
) -> tuple[object, ...]:
return (
"trade",
replay_at,
replay_sequence,
venue,
SYMBOL,
trade_id,
replay_at,
price,
Decimal("0.01"),
"buy",
TRADE_SOURCE,
START + timedelta(seconds=1),
START + timedelta(seconds=2),
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
[TRADE_SOURCE],
1,
)
def make_quote_row(
*,
replay_at: object = START + timedelta(minutes=2),
replay_sequence: object = 11,
) -> tuple[object, ...]:
return (
"quote",
replay_at,
replay_sequence,
VENUE,
SYMBOL,
None,
None,
None,
None,
None,
QUOTE_SOURCE,
None,
None,
replay_at,
START + timedelta(minutes=2) - timedelta(milliseconds=1),
Decimal("65001"),
Decimal("65000"),
Decimal("65002"),
None,
None,
None,
None,
None,
None,
None,
None,
None,
[QUOTE_SOURCE],
1,
)
def make_candle_row(
*,
replay_at: object = START + timedelta(minutes=3),
replay_sequence: object = 12,
open_time: object = START - timedelta(hours=1),
interval: object = "1m",
is_final: object = True,
) -> tuple[object, ...]:
return (
"candle_revision",
replay_at,
replay_sequence,
VENUE,
SYMBOL,
None,
None,
None,
None,
None,
CANDLE_SOURCE,
None,
None,
None,
None,
None,
None,
None,
interval,
open_time,
replay_at,
Decimal("64900"),
Decimal("65100"),
Decimal("64800"),
Decimal("65000"),
Decimal("10"),
is_final,
[CANDLE_SOURCE],
1,
)
def dependencies(
rows: object = (),
) -> tuple[
PostgresReplayPlanBuilder,
RecordingCursor,
RecordingConnection,
RecordingProvider,
list[tuple[object, ...]],
]:
events: list[tuple[object, ...]] = []
cursor = RecordingCursor(events=events, rows=rows)
connection = RecordingConnection(events=events, cursor=cursor)
provider = RecordingProvider(connection=connection, events=events)
builder = PostgresReplayPlanBuilder(connection_provider=provider)
return builder, cursor, connection, provider, events
def test_constructor_is_no_io_slotted_and_matches_protocol() -> None:
builder, _, _, provider, _ = dependencies()
assert provider.calls == 0
assert not hasattr(builder, "__dict__")
assert isinstance(builder, ReplayPlanBuilderProtocol)
def test_constructor_rejects_non_callable_provider() -> None:
with pytest.raises(TypeError, match="connection_provider"):
PostgresReplayPlanBuilder(
connection_provider=None, # type: ignore[arg-type]
)
def test_exact_request_is_validated_before_connection_borrow() -> None:
class RequestSubclass(ReplayPlanRequest):
pass
builder, _, _, provider, _ = dependencies()
request = RequestSubclass(
venue=VENUE,
symbols=(SYMBOL,),
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRange(START, END),
)
with pytest.raises(MarketDataReplayValidationError, match="request"):
builder.create_plan(request)
assert provider.calls == 0
def test_trade_snapshot_uses_one_half_open_parameterized_query() -> None:
builder, cursor, connection, provider, events = dependencies()
request = make_request(max_records=7, venue="tenant'value")
plan = builder.create_plan(request)
assert plan.is_empty is True
assert provider.calls == 1
assert connection.transaction_calls == 1
assert connection.cursor_calls == 1
assert len(cursor.calls) == 2
setup_sql, setup_parameters = cursor.calls[0]
sql, parameters = cursor.calls[1]
compact_sql = normalized_sql(sql)
assert normalized_sql(setup_sql) == (
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY"
)
assert setup_parameters is None
assert "FROM market_data.trades" in compact_sql
assert "FROM market_data.quotes" not in compact_sql
assert "FROM market_data.candle_revisions" not in compact_sql
assert "executed_at >= %s" in compact_sql
assert "executed_at < %s" in compact_sql
assert "ORDER BY replay_at ASC, replay_sequence ASC" in compact_sql
assert compact_sql.endswith("LIMIT %s")
assert "tenant'value" not in sql
assert parameters == (
"tenant'value",
[SYMBOL],
START,
END,
8,
)
assert events == [
("provider",),
("connection_enter",),
("transaction",),
("transaction_enter",),
("cursor",),
("cursor_enter",),
(
"execute",
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ, READ ONLY",
None,
),
("execute", compact_sql, parameters),
("fetchall",),
("cursor_exit", None),
("transaction_exit", None),
("connection_exit", None),
]
def test_all_requested_types_use_static_union_and_own_time_axes() -> None:
builder, cursor, _, _, _ = dependencies()
request = make_request(
data_types=(
ReplayDataType.CANDLE_REVISION,
ReplayDataType.TRADE,
ReplayDataType.QUOTE,
),
candle_intervals=("1m", "5m"),
max_records=20,
)
builder.create_plan(request)
sql, parameters = cursor.calls[1]
compact_sql = normalized_sql(sql)
assert compact_sql.count("UNION ALL") == 2
assert "executed_at >= %s" in compact_sql
assert "received_at >= %s" in compact_sql
assert "observed_at >= %s" in compact_sql
assert "open_time >= %s" not in compact_sql
assert "interval = ANY(%s)" in compact_sql
assert parameters == (
VENUE,
[SYMBOL],
START,
END,
VENUE,
[SYMBOL],
START,
END,
VENUE,
[SYMBOL],
START,
END,
["1m", "5m"],
21,
)
def test_materializes_globally_ordered_canonical_events() -> None:
event_time = START + timedelta(minutes=2)
builder, _, _, _, _ = dependencies(
[
make_trade_row(
replay_at=START + timedelta(minutes=1),
replay_sequence=10,
),
make_quote_row(
replay_at=event_time,
replay_sequence=11,
),
make_candle_row(
replay_at=event_time,
replay_sequence=12,
),
]
)
request = make_request(
data_types=(
ReplayDataType.TRADE,
ReplayDataType.QUOTE,
ReplayDataType.CANDLE_REVISION,
),
candle_intervals=("1m",),
)
plan = builder.create_plan(request)
assert [event.replay_sequence for event in plan.events] == [10, 11, 12]
assert type(plan.events[0].payload) is Trade
assert type(plan.events[1].payload) is Quote
candle_payload = plan.events[2].payload
assert type(candle_payload) is Candle
assert isinstance(candle_payload, Candle)
assert plan.events[2].replay_at == event_time
assert candle_payload.open_time < START
assert plan.events[2].candle_is_final is True
def test_candle_snapshot_filters_by_observed_time_not_open_time() -> None:
builder, cursor, _, _, _ = dependencies(
[make_candle_row(open_time=START - timedelta(days=1))]
)
request = make_request(
data_types=(ReplayDataType.CANDLE_REVISION,),
candle_intervals=("1m",),
)
plan = builder.create_plan(request)
sql, _ = cursor.calls[1]
compact_sql = normalized_sql(sql)
assert "observed_at >= %s" in compact_sql
assert "observed_at < %s" in compact_sql
assert "open_time >= %s" not in compact_sql
assert len(plan.events) == 1
def test_limit_plus_one_raises_without_partial_plan() -> None:
builder, _, _, _, events = dependencies(
[
make_trade_row(replay_sequence=10),
make_trade_row(
replay_at=START + timedelta(minutes=2),
replay_sequence=11,
trade_id=101,
),
]
)
with pytest.raises(ReplayPlanLimitExceededError, match="max_records"):
builder.create_plan(make_request(max_records=1))
assert events[-3:] == [
("cursor_exit", None),
("transaction_exit", None),
("connection_exit", None),
]
def test_more_than_limit_plus_one_is_backend_integrity_error() -> None:
builder, _, _, _, _ = dependencies(
[
make_trade_row(replay_sequence=10),
make_trade_row(replay_sequence=11),
make_trade_row(replay_sequence=12),
]
)
with pytest.raises(MarketDataAccessIntegrityError, match="limit"):
builder.create_plan(make_request(max_records=1))
@pytest.mark.parametrize("rows", (None, 1, "rows", b"rows"))
def test_rejects_invalid_rows_container(rows: object) -> None:
builder, _, _, _, _ = dependencies(rows)
with pytest.raises(MarketDataAccessIntegrityError, match="rows"):
builder.create_plan(make_request())
@pytest.mark.parametrize(
"row",
(
(),
("trade",),
("unknown",) + make_trade_row()[1:],
(1,) + make_trade_row()[1:],
),
)
def test_rejects_invalid_snapshot_row_shape_or_type(
row: tuple[object, ...],
) -> None:
builder, _, _, _, _ = dependencies([row])
with pytest.raises(MarketDataAccessIntegrityError):
builder.create_plan(make_request())
@pytest.mark.parametrize(
"row, replay_request",
(
(
make_trade_row(price=Decimal("0")),
make_request(),
),
(
make_quote_row()[:-1] + (2,),
make_request(data_types=(ReplayDataType.QUOTE,)),
),
(
make_candle_row(is_final=1),
make_request(
data_types=(ReplayDataType.CANDLE_REVISION,),
candle_intervals=("1m",),
),
),
),
)
def test_shared_mappers_reject_corrupt_canonical_values(
row: tuple[object, ...],
replay_request: ReplayPlanRequest,
) -> None:
builder, _, _, _, _ = dependencies([row])
with pytest.raises(MarketDataAccessIntegrityError):
builder.create_plan(replay_request)
def test_rejects_replay_time_that_differs_from_payload_time() -> None:
row = list(make_trade_row())
row[1] = START + timedelta(minutes=2)
builder, _, _, _, _ = dependencies([tuple(row)])
with pytest.raises(MarketDataAccessIntegrityError, match="event"):
builder.create_plan(make_request())
@pytest.mark.parametrize(
"rows",
(
(
make_trade_row(replay_sequence=10),
make_trade_row(
replay_at=START + timedelta(minutes=2),
replay_sequence=10,
trade_id=101,
),
),
(
make_trade_row(
replay_at=START + timedelta(minutes=2),
replay_sequence=11,
),
make_trade_row(
replay_at=START + timedelta(minutes=1),
replay_sequence=10,
trade_id=101,
),
),
),
)
def test_rejects_duplicate_sequence_or_unordered_rows(
rows: tuple[tuple[object, ...], ...],
) -> None:
builder, _, _, _, _ = dependencies(rows)
with pytest.raises(MarketDataAccessIntegrityError, match="snapshot"):
builder.create_plan(make_request())
def test_rejects_event_outside_request_scope() -> None:
builder, _, _, _, _ = dependencies(
[make_trade_row(venue="another")]
)
with pytest.raises(MarketDataAccessIntegrityError, match="snapshot"):
builder.create_plan(make_request())
@pytest.mark.parametrize("failure_point", ("provider", "execute", "fetchall"))
def test_backend_error_is_wrapped_with_original_cause(
failure_point: str,
) -> None:
builder, cursor, _, provider, events = dependencies()
backend_error = RuntimeError("backend failed")
if failure_point == "provider":
provider.error = backend_error
elif failure_point == "execute":
cursor.execute_errors = [None, backend_error]
else:
cursor.fetchall_error = backend_error
with pytest.raises(MarketDataAccessOperationError) as raised:
builder.create_plan(make_request())
assert raised.value.__cause__ is backend_error
if failure_point != "provider":
assert events[-3:] == [
("cursor_exit", RuntimeError),
("transaction_exit", RuntimeError),
("connection_exit", RuntimeError),
]
def test_transaction_setup_error_is_wrapped_and_query_is_not_executed() -> None:
builder, cursor, _, _, _ = dependencies()
setup_error = RuntimeError("cannot configure transaction")
cursor.execute_errors = [setup_error]
with pytest.raises(MarketDataAccessOperationError) as raised:
builder.create_plan(make_request())
assert raised.value.__cause__ is setup_error
assert len(cursor.calls) == 1
def test_base_exception_is_not_swallowed_and_contexts_are_closed() -> None:
builder, cursor, _, _, events = dependencies()
cursor.fetchall_error = KeyboardInterrupt()
with pytest.raises(KeyboardInterrupt):
builder.create_plan(make_request())
assert events[-3:] == [
("cursor_exit", KeyboardInterrupt),
("transaction_exit", KeyboardInterrupt),
("connection_exit", KeyboardInterrupt),
]

View File

@@ -0,0 +1,807 @@
from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
import pytest
import src.market_data.replay as replay_package
from src.market_data.access import HistoricalTimeRange
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.replay import (
MarketDataReplayValidationError,
ReplayDataType,
ReplayEvent,
ReplayPlan,
ReplayPlanRequest,
ReplaySession,
ReplaySessionProtocol,
ReplaySessionState,
ReplaySessionStateError,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
FIRST_TIME = START + timedelta(minutes=1)
SECOND_TIME = START + timedelta(minutes=2)
THIRD_TIME = START + timedelta(minutes=3)
END = START + timedelta(hours=1)
def make_trade_event(
*,
trade_id: int,
replay_at: datetime,
replay_sequence: int,
) -> ReplayEvent:
trade = Trade(
symbol=SYMBOL,
trade_id=trade_id,
price=Decimal("65000"),
quantity=Decimal("0.001"),
executed_at=replay_at,
aggressor_side=TradeAggressorSide.BUY,
source="dzengi_websocket_trade",
)
return ReplayEvent(
venue=VENUE,
replay_at=replay_at,
replay_sequence=replay_sequence,
payload=trade,
)
def make_plan(*, empty: bool = False) -> ReplayPlan:
request = ReplayPlanRequest(
venue=VENUE,
symbols=(SYMBOL,),
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRange(
start_time=START,
end_time=END,
),
max_records=100,
)
events = () if empty else (
make_trade_event(
trade_id=1,
replay_at=FIRST_TIME,
replay_sequence=1,
),
make_trade_event(
trade_id=2,
replay_at=FIRST_TIME,
replay_sequence=2,
),
make_trade_event(
trade_id=3,
replay_at=SECOND_TIME,
replay_sequence=3,
),
make_trade_event(
trade_id=4,
replay_at=THIRD_TIME,
replay_sequence=4,
),
)
return ReplayPlan(request=request, events=events)
class RecordingClock:
def __init__(self, now: datetime = START) -> None:
self._now = now
self.advance_calls: list[datetime] = []
@property
def now(self) -> datetime:
return self._now
def advance_to(self, instant: datetime) -> None:
self.advance_calls.append(instant)
self._now = instant
class RecordingConsumer:
def __init__(self, clock: RecordingClock) -> None:
self.clock = clock
self.events: list[ReplayEvent] = []
self.observed_times: list[datetime] = []
self.tasks: list[asyncio.Task[Any] | None] = []
async def consume(self, event: ReplayEvent) -> None:
self.events.append(event)
self.observed_times.append(self.clock.now)
self.tasks.append(asyncio.current_task())
def make_session(
*,
plan: ReplayPlan | None = None,
clock: RecordingClock | None = None,
consumer: RecordingConsumer | None = None,
) -> tuple[ReplaySession, ReplayPlan, RecordingClock, RecordingConsumer]:
resolved_plan = make_plan() if plan is None else plan
resolved_clock = (
RecordingClock(resolved_plan.request.time_range.start_time)
if clock is None
else clock
)
resolved_consumer = (
RecordingConsumer(resolved_clock)
if consumer is None
else consumer
)
session = ReplaySession(
plan=resolved_plan,
clock=resolved_clock,
consumer=resolved_consumer,
)
return (
session,
resolved_plan,
resolved_clock,
resolved_consumer,
)
def test_matches_protocol_uses_slots_and_preserves_dependencies() -> None:
session, plan, clock, _ = make_session()
assert isinstance(session, ReplaySessionProtocol)
assert not hasattr(session, "__dict__")
assert session.state is ReplaySessionState.CREATED
assert session.plan is plan
assert session.clock is clock
assert replay_package.ReplaySession is ReplaySession
assert not hasattr(replay_package, "ReplayEngine")
def test_construction_does_not_advance_clock_or_call_consumer() -> None:
session, _, clock, consumer = make_session()
assert session.state is ReplaySessionState.CREATED
assert clock.advance_calls == []
assert consumer.events == []
@pytest.mark.parametrize("invalid", (None, object(), "plan"))
def test_rejects_non_plan(invalid: object) -> None:
clock = RecordingClock()
with pytest.raises(TypeError, match="plan"):
ReplaySession(
plan=invalid, # type: ignore[arg-type]
clock=clock,
consumer=RecordingConsumer(clock),
)
def test_rejects_plan_subclass() -> None:
class ReplayPlanSubclass(ReplayPlan):
pass
plan = make_plan()
subclass = ReplayPlanSubclass(
request=plan.request,
events=plan.events,
)
clock = RecordingClock()
with pytest.raises(TypeError, match="plan"):
ReplaySession(
plan=subclass,
clock=clock,
consumer=RecordingConsumer(clock),
)
def test_rejects_clock_without_protocol() -> None:
class InvalidClock:
@property
def now(self) -> datetime:
return START
with pytest.raises(TypeError, match="clock"):
ReplaySession(
plan=make_plan(),
clock=InvalidClock(), # type: ignore[arg-type]
consumer=RecordingConsumer(RecordingClock()),
)
def test_rejects_clock_class_before_any_runtime_action() -> None:
class ClockClass:
advance_calls: list[datetime] = []
@property
def now(self) -> datetime:
return START
def advance_to(self, instant: datetime) -> None:
self.advance_calls.append(instant)
consumer_clock = RecordingClock()
consumer = RecordingConsumer(consumer_clock)
with pytest.raises(TypeError, match="clock"):
ReplaySession(
plan=make_plan(),
clock=ClockClass, # type: ignore[arg-type]
consumer=consumer,
)
assert ClockClass.advance_calls == []
assert consumer.events == []
def test_rejects_asynchronous_clock_advance() -> None:
class AsyncClock:
@property
def now(self) -> datetime:
return START
async def advance_to(self, instant: datetime) -> None:
return None
clock = AsyncClock()
with pytest.raises(TypeError, match="synchronous"):
ReplaySession(
plan=make_plan(),
clock=clock, # type: ignore[arg-type]
consumer=RecordingConsumer(RecordingClock()),
)
def test_rejects_consumer_without_protocol() -> None:
class InvalidConsumer:
pass
with pytest.raises(TypeError, match="consumer"):
ReplaySession(
plan=make_plan(),
clock=RecordingClock(),
consumer=InvalidConsumer(), # type: ignore[arg-type]
)
def test_rejects_consumer_class_before_advancing_clock() -> None:
clock = RecordingClock()
with pytest.raises(TypeError, match="consumer"):
ReplaySession(
plan=make_plan(),
clock=clock,
consumer=RecordingConsumer, # type: ignore[arg-type]
)
assert clock.now == START
assert clock.advance_calls == []
def test_rejects_synchronous_consumer() -> None:
class SyncConsumer:
def consume(self, event: ReplayEvent) -> None:
return None
with pytest.raises(TypeError, match="asynchronous"):
ReplaySession(
plan=make_plan(),
clock=RecordingClock(),
consumer=SyncConsumer(), # type: ignore[arg-type]
)
@pytest.mark.parametrize(
"now",
(
START.replace(tzinfo=None),
START.astimezone(timezone(timedelta(hours=3))),
START.replace(fold=1),
),
)
def test_rejects_non_canonical_clock_time(now: datetime) -> None:
clock = RecordingClock(now)
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
ReplaySession(
plan=make_plan(),
clock=clock,
consumer=RecordingConsumer(clock),
)
assert clock.advance_calls == []
def test_rejects_datetime_subclass_from_clock() -> None:
class CompatibleDatetime(datetime):
pass
now = CompatibleDatetime(
2026,
8,
2,
12,
0,
tzinfo=timezone.utc,
)
clock = RecordingClock(now)
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
ReplaySession(
plan=make_plan(),
clock=clock,
consumer=RecordingConsumer(clock),
)
def test_rejects_non_datetime_clock_time() -> None:
class InvalidNowClock:
@property
def now(self) -> datetime:
return object() # type: ignore[return-value]
def advance_to(self, instant: datetime) -> None:
return None
clock = InvalidNowClock()
with pytest.raises(MarketDataReplayValidationError, match="canonical"):
ReplaySession(
plan=make_plan(),
clock=clock,
consumer=RecordingConsumer(RecordingClock()),
)
@pytest.mark.parametrize(
"now",
(START - timedelta(microseconds=1), START + timedelta(microseconds=1)),
)
def test_rejects_clock_outside_exact_start_without_reset(
now: datetime,
) -> None:
clock = RecordingClock(now)
with pytest.raises(MarketDataReplayValidationError, match="start"):
ReplaySession(
plan=make_plan(),
clock=clock,
consumer=RecordingConsumer(clock),
)
assert clock.now is now
assert clock.advance_calls == []
def test_clock_now_error_is_not_swallowed() -> None:
expected = RuntimeError("clock now failed")
class BrokenClock(RecordingClock):
@property
def now(self) -> datetime:
raise expected
clock = BrokenClock()
with pytest.raises(RuntimeError) as captured:
ReplaySession(
plan=make_plan(),
clock=clock,
consumer=RecordingConsumer(clock),
)
assert captured.value is expected
def test_run_preserves_order_identity_and_advances_before_consumer() -> None:
async def scenario() -> None:
session, plan, clock, consumer = make_session()
caller_task = asyncio.current_task()
result = await session.run()
assert result is None
assert session.state is ReplaySessionState.COMPLETED
assert tuple(consumer.events) == plan.events
assert all(
actual is expected
for actual, expected in zip(
consumer.events,
plan.events,
strict=True,
)
)
assert clock.advance_calls == [
event.replay_at for event in plan.events
]
assert consumer.observed_times == clock.advance_calls
assert all(task is caller_task for task in consumer.tasks)
assert clock.now == plan.events[-1].replay_at
asyncio.run(scenario())
def test_empty_plan_completes_without_dependency_calls() -> None:
async def scenario() -> None:
plan = make_plan(empty=True)
session, _, clock, consumer = make_session(plan=plan)
initial_time = clock.now
result = await session.run()
assert result is None
assert session.state is ReplaySessionState.COMPLETED
assert clock.now is initial_time
assert clock.advance_calls == []
assert consumer.events == []
asyncio.run(scenario())
def test_delivery_is_strictly_sequential() -> None:
class YieldingConsumer(RecordingConsumer):
def __init__(self, clock: RecordingClock) -> None:
super().__init__(clock)
self.active = 0
self.maximum_active = 0
async def consume(self, event: ReplayEvent) -> None:
self.active += 1
self.maximum_active = max(self.maximum_active, self.active)
try:
await asyncio.sleep(0)
await super().consume(event)
finally:
self.active -= 1
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
consumer = YieldingConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
await session.run()
assert consumer.maximum_active == 1
assert tuple(consumer.events) == plan.events
asyncio.run(scenario())
def test_clock_error_fails_without_delivering_current_or_suffix() -> None:
expected = RuntimeError("clock failed")
class BrokenClock(RecordingClock):
def advance_to(self, instant: datetime) -> None:
self.advance_calls.append(instant)
if len(self.advance_calls) == 3:
raise expected
self._now = instant
async def scenario() -> None:
plan = make_plan()
clock = BrokenClock()
consumer = RecordingConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
with pytest.raises(RuntimeError) as captured:
await session.run()
assert captured.value is expected
assert session.state is ReplaySessionState.FAILED
assert tuple(consumer.events) == plan.events[:2]
assert clock.advance_calls == [
event.replay_at for event in plan.events[:3]
]
assert clock.now == plan.events[1].replay_at
asyncio.run(scenario())
def test_consumer_error_fails_without_retry_rollback_or_suffix() -> None:
expected = RuntimeError("consumer failed")
class BrokenConsumer(RecordingConsumer):
async def consume(self, event: ReplayEvent) -> None:
await super().consume(event)
if len(self.events) == 3:
raise expected
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
consumer = BrokenConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
with pytest.raises(RuntimeError) as captured:
await session.run()
assert captured.value is expected
assert session.state is ReplaySessionState.FAILED
assert tuple(consumer.events) == plan.events[:3]
assert clock.now == plan.events[2].replay_at
with pytest.raises(ReplaySessionStateError):
await session.run()
assert tuple(consumer.events) == plan.events[:3]
assert session.state is ReplaySessionState.FAILED
asyncio.run(scenario())
def test_non_exception_base_error_preserves_identity_and_failed_state() -> None:
class ReplaySignal(BaseException):
pass
expected = ReplaySignal("stop")
class BrokenConsumer(RecordingConsumer):
async def consume(self, event: ReplayEvent) -> None:
raise expected
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=BrokenConsumer(clock),
)
with pytest.raises(ReplaySignal) as captured:
await session.run()
assert captured.value is expected
assert session.state is ReplaySessionState.FAILED
assert clock.now == plan.events[0].replay_at
asyncio.run(scenario())
def test_consumer_cancellation_preserves_identity_and_cancelled_state() -> None:
expected = asyncio.CancelledError("consumer cancelled")
class CancellingConsumer(RecordingConsumer):
async def consume(self, event: ReplayEvent) -> None:
await super().consume(event)
raise expected
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
consumer = CancellingConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
with pytest.raises(asyncio.CancelledError) as captured:
await session.run()
assert captured.value is expected
assert session.state is ReplaySessionState.CANCELLED
assert consumer.events == [plan.events[0]]
assert clock.now == plan.events[0].replay_at
with pytest.raises(ReplaySessionStateError):
await session.run()
assert session.state is ReplaySessionState.CANCELLED
asyncio.run(scenario())
def test_external_cancellation_during_consumer_stays_cancelled() -> None:
class BlockingConsumer(RecordingConsumer):
def __init__(self, clock: RecordingClock) -> None:
super().__init__(clock)
self.entered = asyncio.Event()
self.release = asyncio.Event()
async def consume(self, event: ReplayEvent) -> None:
await super().consume(event)
self.entered.set()
await self.release.wait()
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
consumer = BlockingConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
task = asyncio.create_task(session.run())
await consumer.entered.wait()
assert session.state is ReplaySessionState.RUNNING
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert session.state is ReplaySessionState.CANCELLED
assert consumer.events == [plan.events[0]]
assert clock.now == plan.events[0].replay_at
asyncio.run(scenario())
def test_concurrent_caller_is_rejected_without_damaging_first() -> None:
class BlockingConsumer(RecordingConsumer):
def __init__(self, clock: RecordingClock) -> None:
super().__init__(clock)
self.entered = asyncio.Event()
self.release = asyncio.Event()
async def consume(self, event: ReplayEvent) -> None:
await super().consume(event)
if len(self.events) == 1:
self.entered.set()
await self.release.wait()
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
consumer = BlockingConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
first = asyncio.create_task(session.run())
await consumer.entered.wait()
assert first.done() is False
assert session.state is ReplaySessionState.RUNNING
with pytest.raises(ReplaySessionStateError):
await session.run()
assert first.done() is False
assert session.state is ReplaySessionState.RUNNING
consumer.release.set()
await first
assert session.state is ReplaySessionState.COMPLETED
assert tuple(consumer.events) == plan.events
asyncio.run(scenario())
def test_caught_reentrant_run_does_not_damage_outer_run() -> None:
class ReentrantConsumer(RecordingConsumer):
def __init__(self, clock: RecordingClock) -> None:
super().__init__(clock)
self.session: ReplaySession | None = None
self.reentrant_errors: list[ReplaySessionStateError] = []
async def consume(self, event: ReplayEvent) -> None:
await super().consume(event)
assert self.session is not None
try:
await self.session.run()
except ReplaySessionStateError as error:
self.reentrant_errors.append(error)
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
consumer = ReentrantConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
consumer.session = session
await session.run()
assert session.state is ReplaySessionState.COMPLETED
assert tuple(consumer.events) == plan.events
assert len(consumer.reentrant_errors) == len(plan.events)
asyncio.run(scenario())
def test_uncaught_reentrant_run_fails_outer_run() -> None:
class ReentrantConsumer(RecordingConsumer):
def __init__(self, clock: RecordingClock) -> None:
super().__init__(clock)
self.session: ReplaySession | None = None
async def consume(self, event: ReplayEvent) -> None:
await super().consume(event)
assert self.session is not None
await self.session.run()
async def scenario() -> None:
plan = make_plan()
clock = RecordingClock()
consumer = ReentrantConsumer(clock)
session, *_ = make_session(
plan=plan,
clock=clock,
consumer=consumer,
)
consumer.session = session
with pytest.raises(ReplaySessionStateError):
await session.run()
assert session.state is ReplaySessionState.FAILED
assert consumer.events == [plan.events[0]]
asyncio.run(scenario())
def test_repeated_run_after_completion_is_rejected() -> None:
async def scenario() -> None:
session, plan, _, consumer = make_session()
await session.run()
with pytest.raises(ReplaySessionStateError):
await session.run()
assert session.state is ReplaySessionState.COMPLETED
assert tuple(consumer.events) == plan.events
asyncio.run(scenario())
def test_two_sessions_share_immutable_plan_but_not_state_or_clock() -> None:
async def scenario() -> None:
plan = make_plan()
first, _, first_clock, first_consumer = make_session(plan=plan)
second, _, second_clock, second_consumer = make_session(plan=plan)
await first.run()
assert first.state is ReplaySessionState.COMPLETED
assert second.state is ReplaySessionState.CREATED
assert first.plan is second.plan is plan
assert first_clock is not second_clock
assert second_clock.now == START
assert second_consumer.events == []
await second.run()
assert tuple(first_consumer.events) == plan.events
assert tuple(second_consumer.events) == plan.events
assert second.state is ReplaySessionState.COMPLETED
asyncio.run(scenario())
def test_lifecycle_and_hidden_task_extensions_are_absent() -> None:
session, *_ = make_session()
assert not hasattr(session, "start")
assert not hasattr(session, "stop")
assert not hasattr(session, "close")
assert not hasattr(session, "reset")
assert not hasattr(session, "pause")
assert not hasattr(session, "resume")

View File

@@ -0,0 +1,883 @@
from __future__ import annotations
import asyncio
import inspect
import threading
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import cast
import pytest
import src.market_data.replay as replay_package
import src.market_data.replay.replay_session_factory as factory_module
from src.market_data.access import HistoricalTimeRange
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.replay import (
DeterministicReplayClock,
MarketDataClockProtocol,
MarketDataReplayValidationError,
ReplayConsumerFactoryProtocol,
ReplayConsumerProtocol,
ReplayClockProtocol,
ReplayDataType,
ReplayEvent,
ReplayPlan,
ReplayPlanBuilderProtocol,
ReplayPlanRequest,
ReplaySession,
ReplaySessionFactory,
ReplaySessionState,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
START = datetime(2026, 8, 2, 12, 0, tzinfo=timezone.utc)
FIRST_TIME = START + timedelta(minutes=1)
SECOND_TIME = START + timedelta(minutes=2)
END = START + timedelta(hours=1)
def make_request() -> ReplayPlanRequest:
return ReplayPlanRequest(
venue=VENUE,
symbols=(SYMBOL,),
data_types=(ReplayDataType.TRADE,),
time_range=HistoricalTimeRange(
start_time=START,
end_time=END,
),
max_records=100,
)
def make_trade_event(
*,
trade_id: int,
replay_at: datetime,
replay_sequence: int,
) -> ReplayEvent:
return ReplayEvent(
venue=VENUE,
replay_at=replay_at,
replay_sequence=replay_sequence,
payload=Trade(
symbol=SYMBOL,
trade_id=trade_id,
price=Decimal("65000"),
quantity=Decimal("0.001"),
executed_at=replay_at,
aggressor_side=TradeAggressorSide.BUY,
source="dzengi_websocket_trade",
),
)
def make_plan(
*,
request: ReplayPlanRequest | None = None,
empty: bool = False,
) -> ReplayPlan:
resolved_request = make_request() if request is None else request
events = () if empty else (
make_trade_event(
trade_id=1,
replay_at=FIRST_TIME,
replay_sequence=1,
),
make_trade_event(
trade_id=2,
replay_at=SECOND_TIME,
replay_sequence=2,
),
)
return ReplayPlan(
request=resolved_request,
events=events,
)
class RecordingConsumer:
def __init__(self, clock: MarketDataClockProtocol) -> None:
self.clock = clock
self.events: list[ReplayEvent] = []
self.observed_times: list[datetime] = []
self.start_calls = 0
self.stop_calls = 0
self.close_calls = 0
async def consume(self, event: ReplayEvent) -> None:
self.events.append(event)
self.observed_times.append(self.clock.now)
def start(self) -> None:
self.start_calls += 1
def stop(self) -> None:
self.stop_calls += 1
def close(self) -> None:
self.close_calls += 1
class RecordingPlanBuilder:
def __init__(
self,
plan: ReplayPlan,
*,
actions: list[str] | None = None,
) -> None:
self.plan = plan
self.actions = actions
self.requests: list[ReplayPlanRequest] = []
self.thread_ids: list[int] = []
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
if self.actions is not None:
self.actions.append("builder")
self.requests.append(request)
self.thread_ids.append(threading.get_ident())
return self.plan
class RecordingConsumerFactory:
def __init__(self, *, actions: list[str] | None = None) -> None:
self.actions = actions
self.plans: list[ReplayPlan] = []
self.clocks: list[MarketDataClockProtocol] = []
self.consumers: list[RecordingConsumer] = []
def create_consumer(
self,
*,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
) -> ReplayConsumerProtocol:
if self.actions is not None:
self.actions.append("consumer")
consumer = RecordingConsumer(clock)
self.plans.append(plan)
self.clocks.append(clock)
self.consumers.append(consumer)
return consumer
def create_factory(
*,
plan: ReplayPlan | None = None,
) -> tuple[
ReplaySessionFactory,
ReplayPlan,
RecordingPlanBuilder,
RecordingConsumerFactory,
]:
resolved_plan = make_plan() if plan is None else plan
plan_builder = RecordingPlanBuilder(resolved_plan)
consumer_factory = RecordingConsumerFactory()
return (
ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
),
resolved_plan,
plan_builder,
consumer_factory,
)
def test_matches_protocols_uses_slots_and_is_exported() -> None:
factory, _, plan_builder, consumer_factory = create_factory()
assert isinstance(plan_builder, ReplayPlanBuilderProtocol)
assert isinstance(consumer_factory, ReplayConsumerFactoryProtocol)
assert not hasattr(factory, "__dict__")
assert replay_package.ReplaySessionFactory is ReplaySessionFactory
assert (
replay_package.ReplayConsumerFactoryProtocol
is ReplayConsumerFactoryProtocol
)
def test_constructor_only_preserves_dependencies() -> None:
plan = make_plan()
plan_builder = RecordingPlanBuilder(plan)
consumer_factory = RecordingConsumerFactory()
factory = ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
)
assert isinstance(factory, ReplaySessionFactory)
assert plan_builder.requests == []
assert consumer_factory.plans == []
assert consumer_factory.clocks == []
assert consumer_factory.consumers == []
def test_constructor_does_not_create_clock_session_or_task(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class ForbiddenClock:
def __init__(self, initial_time: datetime) -> None:
raise AssertionError("Clock must not be created")
def forbidden_session(**kwargs: object) -> ReplaySession:
raise AssertionError("Session must not be created")
monkeypatch.setattr(
factory_module,
"DeterministicReplayClock",
ForbiddenClock,
)
monkeypatch.setattr(
factory_module,
"ReplaySession",
forbidden_session,
)
async def scenario() -> None:
tasks_before = set(asyncio.all_tasks())
factory, _, _, _ = create_factory()
assert isinstance(factory, ReplaySessionFactory)
assert set(asyncio.all_tasks()) == tasks_before
asyncio.run(scenario())
@pytest.mark.parametrize("invalid", (None, object(), "builder"))
def test_rejects_invalid_plan_builder(invalid: object) -> None:
with pytest.raises(TypeError, match="plan_builder"):
ReplaySessionFactory(
plan_builder=cast(ReplayPlanBuilderProtocol, invalid),
consumer_factory=RecordingConsumerFactory(),
)
def test_rejects_plan_builder_class() -> None:
class PlanBuilderClass:
def create_plan(
self,
request: ReplayPlanRequest,
) -> ReplayPlan:
return make_plan(request=request)
with pytest.raises(TypeError, match="plan_builder"):
ReplaySessionFactory(
plan_builder=cast(
ReplayPlanBuilderProtocol,
PlanBuilderClass,
),
consumer_factory=RecordingConsumerFactory(),
)
def test_rejects_asynchronous_plan_builder() -> None:
class AsyncPlanBuilder:
async def create_plan(
self,
request: ReplayPlanRequest,
) -> ReplayPlan:
return make_plan(request=request)
with pytest.raises(TypeError, match="synchronous"):
ReplaySessionFactory(
plan_builder=cast(
ReplayPlanBuilderProtocol,
AsyncPlanBuilder(),
),
consumer_factory=RecordingConsumerFactory(),
)
@pytest.mark.parametrize("invalid", (None, object(), "factory"))
def test_rejects_invalid_consumer_factory(invalid: object) -> None:
with pytest.raises(TypeError, match="consumer_factory"):
ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(make_plan()),
consumer_factory=cast(
ReplayConsumerFactoryProtocol,
invalid,
),
)
def test_rejects_consumer_factory_class() -> None:
class ConsumerFactoryClass:
def create_consumer(
self,
*,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
) -> ReplayConsumerProtocol:
return RecordingConsumer(clock)
with pytest.raises(TypeError, match="consumer_factory"):
ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(make_plan()),
consumer_factory=cast(
ReplayConsumerFactoryProtocol,
ConsumerFactoryClass,
),
)
def test_rejects_asynchronous_consumer_factory() -> None:
class AsyncConsumerFactory:
async def create_consumer(
self,
*,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
) -> ReplayConsumerProtocol:
return RecordingConsumer(clock)
with pytest.raises(TypeError, match="synchronous"):
ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(make_plan()),
consumer_factory=cast(
ReplayConsumerFactoryProtocol,
AsyncConsumerFactory(),
),
)
def test_prepare_is_synchronous_and_run_remains_asynchronous() -> None:
assert not inspect.iscoroutinefunction(
ReplaySessionFactory.prepare_session
)
assert inspect.iscoroutinefunction(ReplaySession.run)
@pytest.mark.parametrize("invalid", (None, object(), "request"))
def test_rejects_invalid_request_before_dependencies(
invalid: object,
) -> None:
factory, _, plan_builder, consumer_factory = create_factory()
with pytest.raises(MarketDataReplayValidationError, match="request"):
factory.prepare_session(
cast(ReplayPlanRequest, invalid)
)
assert plan_builder.requests == []
assert consumer_factory.plans == []
def test_rejects_request_subclass_before_dependencies() -> None:
class ReplayPlanRequestSubclass(ReplayPlanRequest):
pass
request = make_request()
subclass = ReplayPlanRequestSubclass(
venue=request.venue,
symbols=request.symbols,
data_types=request.data_types,
time_range=request.time_range,
candle_intervals=request.candle_intervals,
max_records=request.max_records,
)
factory, _, plan_builder, consumer_factory = create_factory()
with pytest.raises(MarketDataReplayValidationError, match="request"):
factory.prepare_session(subclass)
assert plan_builder.requests == []
assert consumer_factory.plans == []
def test_builds_in_order_and_preserves_all_identities(
monkeypatch: pytest.MonkeyPatch,
) -> None:
actions: list[str] = []
request = make_request()
plan = make_plan(request=request)
plan_builder = RecordingPlanBuilder(plan, actions=actions)
consumer_factory = RecordingConsumerFactory(actions=actions)
real_session = ReplaySession
class OrderedClock(DeterministicReplayClock):
def __init__(self, initial_time: datetime) -> None:
actions.append("clock")
super().__init__(initial_time)
def ordered_session(
*,
plan: ReplayPlan,
clock: ReplayClockProtocol,
consumer: ReplayConsumerProtocol,
) -> ReplaySession:
actions.append("session")
return real_session(
plan=plan,
clock=clock,
consumer=consumer,
)
monkeypatch.setattr(
factory_module,
"DeterministicReplayClock",
OrderedClock,
)
monkeypatch.setattr(
factory_module,
"ReplaySession",
ordered_session,
)
factory = ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
)
session = factory.prepare_session(request)
clock = consumer_factory.clocks[0]
assert actions == ["builder", "clock", "consumer", "session"]
assert plan_builder.requests == [request]
assert plan_builder.requests[0] is request
assert consumer_factory.plans == [plan]
assert consumer_factory.plans[0] is plan
assert session.plan is plan
assert session.clock is clock
assert isinstance(clock, OrderedClock)
assert clock.now == START
assert session.state is ReplaySessionState.CREATED
def test_prepare_runs_builder_in_caller_thread() -> None:
factory, plan, plan_builder, consumer_factory = create_factory()
caller_thread_id = threading.get_ident()
session = factory.prepare_session(plan.request)
assert isinstance(session, ReplaySession)
assert plan_builder.thread_ids == [caller_thread_id]
assert consumer_factory.consumers[0].events == []
def test_prepare_does_not_start_consumer_or_session_lifecycle() -> None:
factory, plan, _, consumer_factory = create_factory()
async def scenario() -> None:
tasks_before = set(asyncio.all_tasks())
session = factory.prepare_session(plan.request)
consumer = consumer_factory.consumers[0]
assert session.state is ReplaySessionState.CREATED
assert consumer.events == []
assert consumer.start_calls == 0
assert consumer.stop_calls == 0
assert consumer.close_calls == 0
assert set(asyncio.all_tasks()) == tasks_before
asyncio.run(scenario())
def test_consumer_observes_start_before_run_and_event_time_during_run(
) -> None:
factory, plan, _, consumer_factory = create_factory()
session = factory.prepare_session(plan.request)
consumer = consumer_factory.consumers[0]
assert consumer.clock is session.clock
assert consumer.clock.now == START
assert consumer.events == []
asyncio.run(session.run())
assert tuple(consumer.events) == plan.events
assert consumer.observed_times == [FIRST_TIME, SECOND_TIME]
assert session.clock.now == SECOND_TIME
assert session.state is ReplaySessionState.COMPLETED
def test_empty_plan_builds_created_no_op_session() -> None:
request = make_request()
plan = make_plan(request=request, empty=True)
factory, _, _, consumer_factory = create_factory(plan=plan)
session = factory.prepare_session(request)
consumer = consumer_factory.consumers[0]
assert session.state is ReplaySessionState.CREATED
assert session.clock.now == START
asyncio.run(session.run())
assert session.state is ReplaySessionState.COMPLETED
assert session.clock.now == START
assert consumer.events == []
def test_repeated_prepare_creates_fresh_dependency_graph() -> None:
factory, plan, plan_builder, consumer_factory = create_factory()
first = factory.prepare_session(plan.request)
second = factory.prepare_session(plan.request)
assert first is not second
assert first.plan is plan
assert second.plan is plan
assert first.clock is not second.clock
assert consumer_factory.consumers[0] is not (
consumer_factory.consumers[1]
)
assert consumer_factory.clocks == [first.clock, second.clock]
assert plan_builder.requests == [plan.request, plan.request]
asyncio.run(first.run())
assert first.state is ReplaySessionState.COMPLETED
assert second.state is ReplaySessionState.CREATED
assert second.clock.now == START
assert consumer_factory.consumers[1].events == []
class FatalPreparationError(BaseException):
pass
class RaisingPlanBuilder:
def __init__(self, error: BaseException) -> None:
self.error = error
self.requests: list[ReplayPlanRequest] = []
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
self.requests.append(request)
raise self.error
@pytest.mark.parametrize(
"expected",
(
RuntimeError("builder failed"),
FatalPreparationError("builder fatal"),
asyncio.CancelledError("builder cancelled"),
),
)
def test_builder_error_is_not_wrapped_and_stops_preparation(
expected: BaseException,
) -> None:
consumer_factory = RecordingConsumerFactory()
plan_builder = RaisingPlanBuilder(expected)
factory = ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
)
request = make_request()
with pytest.raises(type(expected)) as captured:
factory.prepare_session(request)
assert captured.value is expected
assert plan_builder.requests == [request]
assert consumer_factory.plans == []
assert consumer_factory.clocks == []
class InvalidResultPlanBuilder:
def __init__(self, result: object) -> None:
self.result = result
self.requests: list[ReplayPlanRequest] = []
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
self.requests.append(request)
return cast(ReplayPlan, self.result)
@pytest.mark.parametrize("invalid", (None, object(), "plan"))
def test_rejects_invalid_builder_result_before_consumer(
invalid: object,
) -> None:
plan_builder = InvalidResultPlanBuilder(invalid)
consumer_factory = RecordingConsumerFactory()
factory = ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
)
request = make_request()
with pytest.raises(MarketDataReplayValidationError, match="ReplayPlan"):
factory.prepare_session(request)
assert plan_builder.requests == [request]
assert consumer_factory.plans == []
def test_rejects_plan_subclass_before_consumer() -> None:
class ReplayPlanSubclass(ReplayPlan):
pass
request = make_request()
valid = make_plan(request=request)
subclass = ReplayPlanSubclass(
request=valid.request,
events=valid.events,
)
plan_builder = InvalidResultPlanBuilder(subclass)
consumer_factory = RecordingConsumerFactory()
factory = ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
)
with pytest.raises(MarketDataReplayValidationError, match="exact"):
factory.prepare_session(request)
assert consumer_factory.plans == []
def test_rejects_equal_plan_with_different_request_identity() -> None:
request = make_request()
copied_request = make_request()
assert copied_request == request
assert copied_request is not request
plan = make_plan(request=copied_request)
plan_builder = RecordingPlanBuilder(plan)
consumer_factory = RecordingConsumerFactory()
factory = ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
)
with pytest.raises(MarketDataReplayValidationError, match="identity"):
factory.prepare_session(request)
assert plan_builder.requests == [request]
assert consumer_factory.plans == []
class RaisingConsumerFactory:
def __init__(self, error: BaseException) -> None:
self.error = error
self.plans: list[ReplayPlan] = []
self.clocks: list[MarketDataClockProtocol] = []
def create_consumer(
self,
*,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
) -> ReplayConsumerProtocol:
self.plans.append(plan)
self.clocks.append(clock)
raise self.error
@pytest.mark.parametrize(
"expected",
(
RuntimeError("consumer factory failed"),
FatalPreparationError("consumer factory fatal"),
asyncio.CancelledError("consumer factory cancelled"),
),
)
def test_consumer_factory_error_is_not_wrapped_or_retried(
expected: BaseException,
) -> None:
request = make_request()
plan = make_plan(request=request)
consumer_factory = RaisingConsumerFactory(expected)
factory = ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(plan),
consumer_factory=consumer_factory,
)
with pytest.raises(type(expected)) as captured:
factory.prepare_session(request)
assert captured.value is expected
assert consumer_factory.plans == [plan]
assert len(consumer_factory.clocks) == 1
assert consumer_factory.clocks[0].now == START
class InvalidConsumerFactory:
def __init__(self, result: object) -> None:
self.result = result
self.calls = 0
def create_consumer(
self,
*,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
) -> ReplayConsumerProtocol:
self.calls += 1
return cast(ReplayConsumerProtocol, self.result)
@pytest.mark.parametrize("invalid", (None, object(), "consumer"))
def test_rejects_invalid_consumer_result(invalid: object) -> None:
request = make_request()
consumer_factory = InvalidConsumerFactory(invalid)
factory = ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
consumer_factory=consumer_factory,
)
with pytest.raises(TypeError, match="consumer"):
factory.prepare_session(request)
assert consumer_factory.calls == 1
def test_rejects_consumer_class_result() -> None:
request = make_request()
consumer_factory = InvalidConsumerFactory(RecordingConsumer)
factory = ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
consumer_factory=consumer_factory,
)
with pytest.raises(TypeError, match="consumer"):
factory.prepare_session(request)
def test_rejects_synchronous_consumer_result() -> None:
class SynchronousConsumer:
def consume(self, event: ReplayEvent) -> None:
return None
request = make_request()
consumer_factory = InvalidConsumerFactory(SynchronousConsumer())
factory = ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
consumer_factory=consumer_factory,
)
with pytest.raises(TypeError, match="asynchronous"):
factory.prepare_session(request)
def test_keyword_incompatible_consumer_factory_error_is_not_wrapped() -> None:
class PositionalOnlyConsumerFactory:
def create_consumer(
self,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
/,
) -> ReplayConsumerProtocol:
return RecordingConsumer(clock)
request = make_request()
factory = ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(make_plan(request=request)),
consumer_factory=cast(
ReplayConsumerFactoryProtocol,
PositionalOnlyConsumerFactory(),
),
)
with pytest.raises(TypeError, match="keyword"):
factory.prepare_session(request)
def test_failed_preparation_does_not_poison_next_call() -> None:
class RecoveringConsumerFactory(RecordingConsumerFactory):
def __init__(self) -> None:
super().__init__()
self.attempts = 0
def create_consumer(
self,
*,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
) -> ReplayConsumerProtocol:
self.attempts += 1
if self.attempts == 1:
raise RuntimeError("first attempt failed")
return super().create_consumer(plan=plan, clock=clock)
request = make_request()
plan = make_plan(request=request)
consumer_factory = RecoveringConsumerFactory()
factory = ReplaySessionFactory(
plan_builder=RecordingPlanBuilder(plan),
consumer_factory=consumer_factory,
)
with pytest.raises(RuntimeError, match="first attempt failed"):
factory.prepare_session(request)
session = factory.prepare_session(request)
assert session.state is ReplaySessionState.CREATED
assert consumer_factory.attempts == 2
assert len(consumer_factory.consumers) == 1
def test_two_concurrent_callers_receive_independent_graphs() -> None:
request = make_request()
plan = make_plan(request=request)
builder_barrier = threading.Barrier(2)
consumer_barrier = threading.Barrier(2)
lock = threading.Lock()
class ConcurrentPlanBuilder:
def __init__(self) -> None:
self.callers: list[int] = []
def create_plan(self, request: ReplayPlanRequest) -> ReplayPlan:
with lock:
self.callers.append(threading.get_ident())
builder_barrier.wait(timeout=5)
return plan
class ConcurrentConsumerFactory:
def __init__(self) -> None:
self.calls: list[
tuple[int, MarketDataClockProtocol, RecordingConsumer]
] = []
def create_consumer(
self,
*,
plan: ReplayPlan,
clock: MarketDataClockProtocol,
) -> ReplayConsumerProtocol:
consumer = RecordingConsumer(clock)
with lock:
self.calls.append(
(threading.get_ident(), clock, consumer)
)
consumer_barrier.wait(timeout=5)
return consumer
plan_builder = ConcurrentPlanBuilder()
consumer_factory = ConcurrentConsumerFactory()
factory = ReplaySessionFactory(
plan_builder=plan_builder,
consumer_factory=consumer_factory,
)
with ThreadPoolExecutor(max_workers=2) as executor:
futures = [
executor.submit(factory.prepare_session, request)
for _ in range(2)
]
sessions = [future.result(timeout=10) for future in futures]
assert len(set(plan_builder.callers)) == 2
assert len({call[0] for call in consumer_factory.calls}) == 2
assert sessions[0] is not sessions[1]
assert sessions[0].clock is not sessions[1].clock
assert consumer_factory.calls[0][2] is not (
consumer_factory.calls[1][2]
)
assert {id(call[1]) for call in consumer_factory.calls} == {
id(session.clock) for session in sessions
}
assert all(
session.state is ReplaySessionState.CREATED
for session in sessions
)

View File

@@ -38,6 +38,7 @@ TradeRow = dict[str, Any]
@dataclass
class TransactionalTradeDatabase:
rows: dict[TradeKey, TradeRow] = field(default_factory=dict)
next_replay_sequence: int = 1
class TransactionalCursor:
@@ -47,6 +48,7 @@ class TransactionalCursor:
) -> None:
self._connection = connection
self._fetchone_result: object = None
self._fetchall_result: list[tuple[int]] = []
def __enter__(self) -> TransactionalCursor:
return self
@@ -78,6 +80,12 @@ class TransactionalCursor:
self._insert(parameters)
return
if normalized.startswith(
"SELECT nextval('market_data.replay_sequence'::regclass)"
):
self._allocate_replay_sequences(parameters)
return
if normalized.startswith("SELECT price, quantity"):
self._select(parameters)
return
@@ -91,21 +99,45 @@ class TransactionalCursor:
def fetchone(self) -> object:
return self._fetchone_result
def fetchall(self) -> list[tuple[int]]:
return list(self._fetchall_result)
def _insert(self, parameters: tuple[Any, ...]) -> None:
(
venue,
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
first_observed_at,
last_observed_at,
observation_sources,
canonical_schema_version,
) = parameters
if len(parameters) == 12:
(
venue,
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
first_observed_at,
last_observed_at,
observation_sources,
canonical_schema_version,
) = parameters
replay_sequence = self._next_replay_sequence()
elif len(parameters) == 13:
(
venue,
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
first_observed_at,
last_observed_at,
observation_sources,
replay_sequence,
canonical_schema_version,
) = parameters
else:
raise AssertionError("Unexpected Trade INSERT parameters")
key = (venue, symbol, trade_id, executed_at)
working_rows = self._connection.working_rows
@@ -121,10 +153,26 @@ class TransactionalCursor:
"first_observed_at": first_observed_at,
"last_observed_at": last_observed_at,
"observation_sources": list(observation_sources),
"replay_sequence": replay_sequence,
"canonical_schema_version": canonical_schema_version,
}
self._fetchone_result = (1,)
def _allocate_replay_sequences(
self,
parameters: tuple[Any, ...],
) -> None:
(count,) = parameters
self._fetchall_result = [
(self._next_replay_sequence(),)
for _ in range(count)
]
def _next_replay_sequence(self) -> int:
value = self._connection.database.next_replay_sequence
self._connection.database.next_replay_sequence += 1
return value
def _select(self, parameters: tuple[Any, ...]) -> None:
key = parameters
row = self._connection.working_rows.get(key)
@@ -181,6 +229,10 @@ class TransactionalConnection:
def cursor(self) -> TransactionalCursor:
return TransactionalCursor(self)
@property
def database(self) -> TransactionalTradeDatabase:
return self._database
@dataclass
class RecordingConnectionProvider:
@@ -261,6 +313,7 @@ def test_store_trade_inserts_canonical_payload_and_provenance() -> None:
"first_observed_at": OBSERVED_AT,
"last_observed_at": OBSERVED_AT,
"observation_sources": ["dzengi_websocket_trade"],
"replay_sequence": 1,
"canonical_schema_version": 1,
}
@@ -459,6 +512,122 @@ def test_batch_uses_stable_identity_order() -> None:
assert inserted_symbols == ("A", "B", "C")
def test_batch_allocates_replay_sequence_in_input_order() -> None:
repository, database, connection, _ = _repository()
repository.store_trades(
venue=VENUE,
trades=(
_trade(symbol="C", trade_id=3),
_trade(symbol="A", trade_id=1),
_trade(symbol="B", trade_id=2),
),
observed_at=OBSERVED_AT,
)
inserted = tuple(
(parameters[1], parameters[11])
for statement, parameters in connection.calls
if statement.startswith("INSERT INTO market_data.trades")
)
assert inserted == (
("A", 2),
("B", 3),
("C", 1),
)
assert {
key[1]: row["replay_sequence"]
for key, row in database.rows.items()
} == {
"A": 2,
"B": 3,
"C": 1,
}
@pytest.mark.parametrize(
("first_trade_id", "second_trade_id"),
(
(SIGNED_TRADE_ID_MAX, SIGNED_TRADE_ID_MIN),
(-1, 0),
),
)
def test_batch_replay_sequence_preserves_signed_rollover_input_order(
first_trade_id: int,
second_trade_id: int,
) -> None:
repository, database, _, _ = _repository()
repository.store_trades(
venue=VENUE,
trades=(
_trade(trade_id=first_trade_id),
_trade(trade_id=second_trade_id),
),
observed_at=OBSERVED_AT,
)
assert database.rows[
(VENUE, SYMBOL, first_trade_id, EXECUTED_AT)
]["replay_sequence"] == 1
assert database.rows[
(VENUE, SYMBOL, second_trade_id, EXECUTED_AT)
]["replay_sequence"] == 2
def test_batch_duplicate_and_provenance_keep_original_replay_sequence() -> None:
repository, database, _, _ = _repository()
original = _trade()
repository.store_trade(
venue=VENUE,
trade=original,
observed_at=OBSERVED_AT,
)
repository.store_trades(
venue=VENUE,
trades=(
replace(original, source="dzengi"),
original,
),
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
assert _only_row(database)["replay_sequence"] == 1
assert database.next_replay_sequence == 4
def test_failed_batch_keeps_consumed_replay_sequence_gap() -> None:
repository, database, _, _ = _repository()
existing = _trade(symbol="B", trade_id=2)
repository.store_trade(
venue=VENUE,
trade=existing,
observed_at=OBSERVED_AT,
)
with pytest.raises(MarketDataStorageConflictError):
repository.store_trades(
venue=VENUE,
trades=(
_trade(symbol="A", trade_id=1),
replace(existing, price=Decimal("999")),
),
observed_at=OBSERVED_AT,
)
repository.store_trade(
venue=VENUE,
trade=_trade(symbol="C", trade_id=3),
observed_at=OBSERVED_AT,
)
assert database.rows[
(VENUE, "C", 3, EXECUTED_AT)
]["replay_sequence"] == 4
def test_batch_conflict_rolls_back_preceding_insert() -> None:
repository, database, connection, _ = _repository()
existing = _trade(symbol="B", trade_id=2)

View File

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

View File

@@ -0,0 +1,236 @@
# Build 060.29 — Market Data Access and Replay
**Engineering Migration Report**
---
## Контроль документа
| Свойство | Значение |
|---|---|
| Build | 060.29 |
| Статус | Completed |
| Подсистема | Market Data / Historical Access and Replay |
| Компонент | Canonical History and Deterministic Replay |
| Дата завершения | 2026-08-02 |
| Версия | 1.0 |
---
## Связанные документы
- `build_060_29_architecture.md` — архитектура, решения и подробные
результаты подэтапов 060.29.0060.29.8;
- `build_060_28.md` — Persistent Checkpoint and Startup Recovery;
- `dzentra_target_architecture.md` — целевая архитектура Dzentra;
- `master-roadmap.md` — дальнейшая последовательность Build.
---
## 1. Назначение Build
Build 060.29 добавил безопасное чтение сохранённых Canonical Market
Data и их детерминированное воспроизведение.
Итоговая цепочка:
```text
PostgreSQL Canonical Market Data
Historical Access
bounded immutable ReplayPlan
ReplaySession + Deterministic Clock
явно переданный Canonical consumer
```
Historical Access и Replay являются read-side подсистемами. Они не
заменяют Acquisition Runtime, не продвигают operational checkpoint и не
записывают воспроизводимые события обратно в Market Data Storage.
---
## 2. Завершённые подэтапы
| Подэтап | Название | Статус |
|---|---|---|
| 060.29.0 | Architecture, Boundaries and Ordering Policy | Accepted |
| 060.29.1 | Historical Access and Replay Contracts | Accepted |
| 060.29.2 | Global Replay Sequence Migration and PostgreSQL Trade History | Accepted |
| 060.29.3 | Quote/Candle History and Snapshot Plan Builder | Accepted |
| 060.29.4 | Deterministic Replay Clock | Accepted |
| 060.29.5 | Replay Session and Engine | Accepted |
| 060.29.6 | Consumer and Composition Integration | Accepted |
| 060.29.7 | PostgreSQL Replay and Failure Verification | Accepted |
| 060.29.8 | Final Regression and Acceptance | Accepted |
Каждый подэтап проходил отдельный read-only review. Findings
исправлялись и закрывались regression-тестами до принятия.
---
## 3. Реализованная архитектура
### 3.1. Глобальная последовательность Replay
Migration 9 добавила один shared PostgreSQL sequence для Trades, Quotes
и Candle revisions. Неизменяемый `replay_sequence` задаёт общий
детерминированный порядок для событий с одинаковым временем и
сохраняется при duplicate или provenance update.
Существующие writers совместимы с sequence для одиночных, пакетных и
конкурентных записей. Default, существующие месячные и будущие partitions
получают одинаковые constraints, trigger и keyset indexes.
Migration выполняется атомарно и использует единый порядок блокировок с
Partition Manager и writers.
### 3.2. Historical Access
DB-neutral read-контракты отделены от write-only Storage API. Отдельные
PostgreSQL readers возвращают типизированные неизменяемые страницы:
- Trades — по `executed_at`;
- Quotes — по `received_at`;
- Candle revisions — по `open_time`.
Запросы используют полуоткрытый временной диапазон и forward keyset
pagination. Cursor связан с точным query scope. Строки PostgreSQL строго
проверяются перед созданием Canonical `Trade`, `Quote` или `Candle`.
### 3.3. Bounded Replay snapshot
Replay Plan Builder выполняет один статический snapshot-запрос в
короткой транзакции `READ ONLY REPEATABLE READ`. Все выбранные события
материализуются до возврата `ReplayPlan`; cursor, transaction и
connection освобождаются до первого вызова consumer.
Порядок Replay равен `(replay_at, replay_sequence)`. Для Candle временем
Replay является `observed_at`, поэтому ревизия не появляется раньше
момента, когда она стала известна системе. Превышение `max_records`
завершается явной ошибкой без частичного plan.
### 3.4. Clock, Session и Composition
`DeterministicReplayClock` работает только с aware UTC-временем,
разрешает равное время и запрещает движение назад. Он не читает wall
clock и не имеет reset.
`ReplaySession` является необратимой one-shot сущностью. Она
последовательно продвигает Clock и ожидает один consumer-вызов для
каждого события. Ошибки и cancellation сохраняют явное terminal state и
не скрываются.
`ReplaySessionFactory` синхронно подготавливает новый граф
`Plan → Clock → consumer → Session` для каждого вызова. Default consumer,
автоматический startup, скрытая задача и отдельный `ReplayEngine` не
добавлены. Caller отдельно владеет `session.run()`.
### 3.5. Статическая проверка типов
Pyright закреплён как обязательный gate в режиме `standard`, совпадающем
с Pylance проекта. Канонический запуск — `scripts/check_python_types.sh`;
допустимый результат — только `0 errors, 0 warnings`. Тот же gate входит
в обычную offline pytest-регрессию.
---
## 4. Real PostgreSQL verification
Безопасный opt-in harness требует отдельную локальную базу с именем
`dzentra_test_*`, явный флаг и отдельный DSN.
Настоящий PostgreSQL 16 подтвердил:
- атомарный детерминированный backfill migration 9;
- общий sequence и совместимость всех writers/partitions;
- keyset pagination Trades, Quotes и Candle revisions;
- единый mixed-type Replay order;
- устойчивый `REPEATABLE READ` snapshot при concurrent commit;
- освобождение PostgreSQL resources до playback;
- limit, integrity, backend, consumer и cancellation paths;
- независимые графы двух одновременно готовящихся Replay Session;
- полный путь `Loopback Runtime → Storage → History → Replay`;
- отсутствие изменения durable данных и operational checkpoint.
---
## 5. Финальные результаты
Итоговая приёмка выполнена 2026-08-02:
```text
Pyright mandatory gate: 0 errors, 0 warnings
Pyright pytest gate: 1 passed
Expanded Access/Replay/Storage unit: 545 passed
PostgreSQL Replay target repeated: 10 × 12 passed
Full PostgreSQL Storage integration: 77 passed
Full integration with PostgreSQL: 90 passed
PostgreSQL suite without opt-in: 77 skipped
Fixed stress target: 3 passed, 1 deselected
Full offline regression: 2681 passed, 95 deselected
pip check: clean
Compileall через временный pycache: clean
Tracked/untracked whitespace checks: clean
Три независимых read-only review: clean
```
В integration и stress наборах `ResourceWarning` считался ошибкой.
Production-код на финальном подэтапе 060.29.8 не изменялся: матрица не
выявила реального дефекта.
---
## 6. Эксплуатационное предупреждение migration 9
Migration 9 является блокирующей. Её длительность пропорциональна уже
накопленному объёму Market Data, потому что существующие строки получают
глобальный `replay_sequence`, после чего создаются constraints и
индексы.
Для текущей небольшой базы выбранный атомарный вариант разумен. Перед
применением к большой production-базе обязательны резервная копия,
замер длительности на сопоставимом объёме и отдельное maintenance
window. Online staged migration в Build 060.29 не реализована и требует
отдельного архитектурного решения.
---
## 7. Границы Build
Build 060.29 намеренно не реализует:
- initial historical backfill с биржи;
- гарантию полной биржевой истории без completeness metadata;
- точное воспроизведение порядка исходных сетевых пакетов;
- production consumers для persistent Quotes и Candles;
- автоматический запуск Replay из Bootstrap;
- default/no-op Replay consumer;
- Backtesting, Analytics API или торговую симуляцию;
- online migration 9 для большой production-базы.
Каждая Historical page отражает committed-состояние на время своего
запроса. Последовательная pagination охватывает все фактически
сохранённые строки диапазона только при неизменном dataset между
страницами; единого межстраничного snapshot нет. Период доступной истории
определяется моментом включения persistent storage и Retention Policy, а
не самим Replay API.
Постороннее пользовательское изменение `.gitignore` не относится к
Build 060.29 и не должно включаться в его staging.
---
## 8. Итог
Build 060.29 завершён и принят.
Dzentra получила отдельный Historical Access, общий детерминированный
порядок сохранённых Canonical Market Data и caller-owned Replay без
скрытого lifecycle. Сохранённую рыночную историю теперь можно безопасно
читать и воспроизводить одинаковыми Canonical объектами.
Следующий этап — Build 060.30, итоговый аудит и документальное закрытие
ветки Trades Feed.

File diff suppressed because it is too large Load Diff

View File

@@ -6,10 +6,10 @@
|---|---|
| Тип | Master Delivery Roadmap |
| Статус | Active |
| Версия | 2.2 |
| Дата актуализации | 2026-08-01 |
| Текущий завершённый Build | 060.28 |
| Текущий Build | 060.29 — Planned |
| Версия | 2.3 |
| Дата актуализации | 2026-08-02 |
| Текущий завершённый Build | 060.29 |
| Текущий Build | 060.30 — Planned |
---
@@ -42,33 +42,33 @@ test evidence находятся в документах конкретных Bu
```text
Market Data Acquisition
Persistent Checkpoint and Startup Recovery завершён
Market Data Access and Replay завершён
Build 060.28
Build 060.29
Completed → следующий Build 060.29
Completed → следующий Build 060.30
```
Build 060.28 завершён и принят:
Build 060.29 завершён и принят:
- persistent checkpoint подтверждается Canonical Trade history;
- Trade и checkpoint продвигаются одной PostgreSQL transaction;
- при запуске восстанавливаются checkpoint и deduplication tail;
- Startup Recovery завершается до buffered Live processing;
- Bootstrap и cancellation сохраняют строгий lifecycle pool/Runtime;
- restart/failure, concurrency и финальная регрессия приняты.
- Historical Access отделён от write-only Storage API;
- migration 9 добавила общий immutable Replay sequence;
- Trades, Quotes и Candle revisions читаются устойчивыми keyset pages;
- Replay строится как bounded `REPEATABLE READ` snapshot;
- Clock, Session и Composition остаются детерминированными и caller-owned;
- PostgreSQL failure paths и финальная регрессия приняты.
Подробности:
```text
docs/migrations/build_060_28.md
docs/migrations/build_060_29.md
```
---
# Активная программа — Market Data Acquisition
## Завершённая ветка Trades Feed
## Завершённые Build ветки Trades Feed
| Build | Результат | Статус |
|---|---|---|
@@ -83,6 +83,7 @@ docs/migrations/build_060_28.md
| 060.26 | Integration and Regression | Completed |
| 060.27 | Persistent Market Data Storage | Completed |
| 060.28 | Persistent Checkpoint and Startup Recovery | Completed |
| 060.29 | Market Data Access and Replay | Completed |
## Build 060.26 — Integration and Regression
@@ -145,7 +146,7 @@ REST и только затем продолжает buffered Live processing.
### Build 060.29 — Market Data Access and Replay
**Статус:** Planned
**Статус:** Completed
Назначение:
@@ -159,6 +160,12 @@ REST и только затем продолжает buffered Live processing.
Build. Они принадлежат Market Data Processing и Feature Engineering и
получат отдельный scope после появления устойчивого Storage/Replay.
Результат: сохранённые Canonical Trades, Quotes и Candle revisions
доступны через отдельный Historical Access и могут воспроизводиться в
детерминированном global order через caller-owned Replay Session.
Подробный итог: `docs/migrations/build_060_29.md`.
### Build 060.30 — Market Data Acquisition Final Documentation
**Статус:** Planned
@@ -1739,6 +1746,6 @@ read-only архитектурного анализа. Старый ориент
Актуальная контрольная точка:
```text
Завершён: Build 060.28Persistent Checkpoint and Startup Recovery
Следующий: Build 060.29 — Market Data Access and Replay
Завершён: Build 060.29Market Data Access and Replay
Следующий: Build 060.30 — Market Data Acquisition Final Documentation
```

35
pyrightconfig.json Normal file
View File

@@ -0,0 +1,35 @@
{
"include": [
"app/src/bootstrap",
"app/src/core",
"app/src/integrations",
"app/src/market_data",
"app/src/runtime_events",
"app/src/storage",
"app/tests/integration/market_data",
"app/tests/live/market_data",
"app/tests/static",
"app/tests/stress/market_data",
"app/tests/support",
"app/tests/unit/bootstrap",
"app/tests/unit/core",
"app/tests/unit/integrations",
"app/tests/unit/market_data",
"app/tests/unit/storage",
"app/tests/unit/test_live_trade_stream_support.py",
"app/tests/unit/test_postgres_market_data_support.py",
"app/tests/unit/test_trade_stream_runtime_support.py"
],
"venvPath": "app",
"venv": ".venv",
"pythonVersion": "3.12",
"typeCheckingMode": "standard",
"executionEnvironments": [
{
"root": "app",
"extraPaths": [
"app"
]
}
]
}

8
scripts/check_python_types.sh Executable file
View File

@@ -0,0 +1,8 @@
#!/usr/bin/env bash
set -euo pipefail
project_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
cd "${project_root}"
app/.venv/bin/python -m pyright --project pyrightconfig.json --warnings