Build 060.28: implement Persistent Checkpoint and Startup Recovery

This commit is contained in:
2026-08-01 20:55:32 +03:00
parent 58e5a12a4d
commit 8c485e32b1
63 changed files with 12430 additions and 92 deletions

View File

@@ -33,3 +33,5 @@ TRADE_STREAM_CLOSE_TIMEOUT_SECONDS=10
TRADE_STREAM_HEARTBEAT_TIMEOUT_SECONDS=30
TRADE_STREAM_SCHEDULER_INTERVAL_SECONDS=5
TRADE_STREAM_RECOVERY_WINDOW_MS=3599999
TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS=10
TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY=10000

View File

@@ -52,6 +52,11 @@ def create_app() -> ApplicationComposition:
if market_data_storage is not None
else None
),
checkpoint_storage=(
market_data_storage.trade_repository
if market_data_storage is not None
else None
),
)
bot = Bot(

View File

@@ -38,12 +38,16 @@ from src.market_data.acquisition.runtime.websocket_subscription_manager import (
from src.market_data.acquisition.trade_stream_runtime_composition import (
build_trade_stream_runtime_composition,
)
from src.market_data.storage.contracts import (
TradeCheckpointStorageProtocol,
)
def build_trade_stream_production_runtime(
settings: Settings,
*,
trade_observation_sink: TradeObservationSinkProtocol | None = None,
checkpoint_storage: TradeCheckpointStorageProtocol | None = None,
) -> TradeStreamProductionRuntime | None:
"""
Собрать Production Trade Stream Runtime без запуска lifecycle.
@@ -53,10 +57,31 @@ def build_trade_stream_production_runtime(
Runtime создаются один раз и затем передаются по identity.
"""
trade_stream = settings.trade_stream
storage_enabled = settings.market_data_storage.enabled
if not trade_stream.enabled:
if storage_enabled:
raise RuntimeError(
"Trade Stream must be enabled when Market Data "
"Storage is enabled."
)
return None
if storage_enabled and (
trade_observation_sink is None
or checkpoint_storage is None
):
raise RuntimeError(
"Enabled Market Data Storage requires both Trade "
"observation sink and checkpoint storage."
)
if not storage_enabled and checkpoint_storage is not None:
raise RuntimeError(
"Checkpoint storage requires enabled Market Data Storage."
)
headers = {
"Origin": settings.exchange_base_url.rstrip("/"),
"Content-Type": "application/json",
@@ -105,6 +130,12 @@ def build_trade_stream_production_runtime(
),
trade_observation_sink=trade_observation_sink,
max_recovery_window_ms=trade_stream.recovery_window_ms,
checkpoint_storage=checkpoint_storage,
checkpoint_venue=(
settings.exchange_name
if checkpoint_storage is not None
else None
),
)
return TradeStreamProductionRuntime(
@@ -128,4 +159,13 @@ def build_trade_stream_production_runtime(
runtime_supervisor=composition.runtime_supervisor,
runtime_scheduler=composition.runtime_scheduler,
symbols=trade_stream.symbols,
startup_recovery_coordinator=(
composition.runtime_startup_recovery_coordinator
),
subscription_ack_timeout_seconds=(
trade_stream.subscription_ack_timeout_seconds
),
startup_market_buffer_capacity=(
trade_stream.startup_market_buffer_capacity
),
)

View File

@@ -9,6 +9,8 @@ from pathlib import Path
from dotenv import load_dotenv
from src.market_data.acquisition.symbols import normalize_symbol
# корень проекта
BASE_DIR = Path(__file__).resolve().parents[2]
@@ -33,6 +35,8 @@ class TradeStreamSettings:
heartbeat_timeout_seconds: float
scheduler_interval_seconds: float
recovery_window_ms: int
subscription_ack_timeout_seconds: float = 10.0
startup_market_buffer_capacity: int = 10_000
@dataclass(frozen=True, slots=True)
@@ -200,7 +204,7 @@ def _parse_trade_stream_symbols(
return tuple(
sorted(
{
symbol.strip()
normalize_symbol(symbol)
for symbol in raw_symbols
}
)
@@ -228,6 +232,8 @@ def _load_trade_stream_settings(
heartbeat_timeout_seconds=30.0,
scheduler_interval_seconds=5.0,
recovery_window_ms=3_599_999,
subscription_ack_timeout_seconds=10.0,
startup_market_buffer_capacity=10_000,
)
websocket_url = os.getenv(
@@ -300,6 +306,22 @@ def _load_trade_stream_settings(
name="TRADE_STREAM_RECOVERY_WINDOW_MS",
default=3_599_999,
),
subscription_ack_timeout_seconds=_parse_positive_float(
os.getenv(
"TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS",
"10",
),
name="TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS",
default=10.0,
),
startup_market_buffer_capacity=_parse_positive_int(
os.getenv(
"TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY",
"10000",
),
name="TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY",
default=10_000,
),
)

View File

@@ -0,0 +1,11 @@
"""Восстановление persistent Trade Stream checkpoint."""
from src.market_data.acquisition.checkpoint.trade_stream_state_hydrator import (
TradeStreamStateHydrator,
TradeStreamStateHydratorProtocol,
)
__all__ = (
"TradeStreamStateHydrator",
"TradeStreamStateHydratorProtocol",
)

View File

@@ -0,0 +1,327 @@
from __future__ import annotations
from typing import Protocol, runtime_checkable
from src.market_data.acquisition.consistency.trade_stream_exceptions import (
TradeConsistencyError,
TradeOrderingError,
)
from src.market_data.acquisition.consistency.trade_stream_state import (
DEFAULT_DEDUPLICATION_WINDOW_SIZE,
TradeStreamState,
)
from src.market_data.acquisition.consistency.trade_stream_state_store_exceptions import (
TradeStreamStateStoreInitializationError,
)
from src.market_data.acquisition.consistency.trade_stream_state_store_protocol import (
TradeStreamStateStoreProtocol,
)
from src.market_data.acquisition.models.trade import Trade
from src.market_data.acquisition.symbols import normalize_symbol
from src.market_data.storage.contracts import (
PersistentTradeCheckpoint,
TradeCheckpointStorageProtocol,
)
from src.market_data.storage.exceptions import (
MarketDataCheckpointIntegrityError,
)
@runtime_checkable
class TradeStreamStateHydratorProtocol(Protocol):
"""Контракт восстановления Consistency state до запуска Runtime."""
def hydrate(
self,
*,
symbols: tuple[str, ...],
) -> tuple[TradeStreamState, ...]:
"""Восстановить и атомарно опубликовать состояния symbols."""
...
class TradeStreamStateHydrator:
"""Восстанавливает TradeStreamState из persistent checkpoint."""
__slots__ = (
"_checkpoint_storage",
"_state_store",
"_venue",
"_deduplication_window_size",
)
def __init__(
self,
*,
checkpoint_storage: TradeCheckpointStorageProtocol,
state_store: TradeStreamStateStoreProtocol,
venue: str,
deduplication_window_size: int = (
DEFAULT_DEDUPLICATION_WINDOW_SIZE
),
) -> None:
if not isinstance(
checkpoint_storage,
TradeCheckpointStorageProtocol,
):
raise TypeError(
"checkpoint_storage must implement "
"TradeCheckpointStorageProtocol"
)
if not isinstance(state_store, TradeStreamStateStoreProtocol):
raise TypeError(
"state_store must implement TradeStreamStateStoreProtocol"
)
if not isinstance(venue, str):
raise TypeError("venue must be a string")
normalized_venue = venue.strip().lower()
if not normalized_venue:
raise ValueError("venue must not be empty")
if (
isinstance(deduplication_window_size, bool)
or not isinstance(deduplication_window_size, int)
):
raise TypeError(
"deduplication_window_size must be an integer"
)
if deduplication_window_size <= 0:
raise ValueError(
"deduplication_window_size must be positive"
)
self._checkpoint_storage = checkpoint_storage
self._state_store = state_store
self._venue = normalized_venue
self._deduplication_window_size = deduplication_window_size
def hydrate(
self,
*,
symbols: tuple[str, ...],
) -> tuple[TradeStreamState, ...]:
normalized_symbols = self._normalize_symbols(symbols)
if not self._state_store.is_empty():
raise TradeStreamStateStoreInitializationError(
"Trade Stream State Store уже содержит состояния."
)
states = tuple(
self._prepare_state(symbol)
for symbol in normalized_symbols
)
self._state_store.initialize(states)
return states
def _prepare_state(
self,
symbol: str,
) -> TradeStreamState:
checkpoint = self._checkpoint_storage.load_checkpoint(
venue=self._venue,
symbol=symbol,
)
if checkpoint is not None:
return self._state_from_checkpoint(
symbol=symbol,
checkpoint=checkpoint,
)
latest_tail = self._checkpoint_storage.load_latest_trade_tail(
venue=self._venue,
symbol=symbol,
limit=self._deduplication_window_size,
)
if not isinstance(latest_tail, tuple):
raise MarketDataCheckpointIntegrityError(
"Checkpoint storage returned non-tuple latest Trade tail."
)
if not latest_tail:
concurrent_checkpoint = (
self._checkpoint_storage.load_checkpoint(
venue=self._venue,
symbol=symbol,
)
)
if concurrent_checkpoint is not None:
return self._state_from_checkpoint(
symbol=symbol,
checkpoint=concurrent_checkpoint,
)
return TradeStreamState.from_history(
symbol=symbol,
trades=(),
deduplication_window_size=(
self._deduplication_window_size
),
)
candidate_state = self._state_from_tail(
symbol=symbol,
trades=latest_tail,
)
candidate = candidate_state.last_trade
if candidate is None:
raise MarketDataCheckpointIntegrityError(
"Latest Trade tail unexpectedly produced empty state."
)
adopted = (
self._checkpoint_storage.adopt_existing_trade_as_checkpoint(
venue=self._venue,
trade=candidate,
)
)
self._validate_checkpoint(
checkpoint=adopted,
symbol=symbol,
expected_trade=candidate,
)
return self._state_from_checkpoint(
symbol=symbol,
checkpoint=adopted,
)
def _state_from_checkpoint(
self,
*,
symbol: str,
checkpoint: PersistentTradeCheckpoint,
) -> TradeStreamState:
self._validate_checkpoint(
checkpoint=checkpoint,
symbol=symbol,
expected_trade=None,
)
tail = self._checkpoint_storage.load_checkpoint_tail(
venue=self._venue,
checkpoint=checkpoint,
limit=self._deduplication_window_size,
)
state = self._state_from_tail(
symbol=symbol,
trades=tail,
)
if state.last_trade is None or not self._same_market_trade(
state.last_trade,
checkpoint.trade,
):
raise MarketDataCheckpointIntegrityError(
"Hydrated tail does not end at persistent checkpoint."
)
return state
def _state_from_tail(
self,
*,
symbol: str,
trades: object,
) -> TradeStreamState:
if not isinstance(trades, tuple):
raise MarketDataCheckpointIntegrityError(
"Checkpoint storage returned non-tuple Trade tail."
)
try:
return TradeStreamState.from_history(
symbol=symbol,
trades=trades,
deduplication_window_size=(
self._deduplication_window_size
),
)
except (
TypeError,
ValueError,
TradeConsistencyError,
TradeOrderingError,
) as error:
raise MarketDataCheckpointIntegrityError(
"Persistent Trade tail cannot hydrate Consistency state."
) from error
def _validate_checkpoint(
self,
*,
checkpoint: object,
symbol: str,
expected_trade: Trade | None,
) -> None:
if not isinstance(checkpoint, PersistentTradeCheckpoint):
raise MarketDataCheckpointIntegrityError(
"Checkpoint storage returned invalid checkpoint."
)
if checkpoint.venue.strip().lower() != self._venue:
raise MarketDataCheckpointIntegrityError(
"Persistent checkpoint venue does not match hydrator."
)
if checkpoint.trade.symbol.strip().upper() != symbol:
raise MarketDataCheckpointIntegrityError(
"Persistent checkpoint symbol does not match request."
)
if expected_trade is not None and not self._same_market_trade(
checkpoint.trade,
expected_trade,
):
raise MarketDataCheckpointIntegrityError(
"Adopted checkpoint does not match candidate Trade."
)
@staticmethod
def _same_market_trade(
first: Trade,
second: Trade,
) -> bool:
return (
first.symbol == second.symbol
and first.trade_id == second.trade_id
and first.executed_at == second.executed_at
and first.price == second.price
and first.quantity == second.quantity
and first.aggressor_side is second.aggressor_side
)
@staticmethod
def _normalize_symbols(
symbols: tuple[str, ...],
) -> tuple[str, ...]:
if not isinstance(symbols, tuple):
raise TypeError("symbols must be a tuple")
if not symbols:
raise ValueError("symbols must not be empty")
normalized: list[str] = []
for symbol in symbols:
if not isinstance(symbol, str):
raise TypeError("symbols must contain only strings")
value = normalize_symbol(symbol)
if not value:
raise ValueError("symbols must not contain empty values")
if value in normalized:
raise ValueError("symbols must not contain duplicates")
normalized.append(value)
return tuple(normalized)

View File

@@ -9,9 +9,18 @@ from src.market_data.acquisition.models.trade import Trade
class TradeObservationSinkProtocol(Protocol):
"""Надёжный приёмник одного наблюдения канонической сделки."""
def persist(
def persist_accepted(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
"""Сохранить новую Trade и атомарно продвинуть checkpoint."""
...
def persist_duplicate(
self,
trade: Trade,
) -> None:
"""Надёжно сохранить наблюдение до продвижения контрольной точки."""
"""Сохранить provenance дубликата без продвижения checkpoint."""
...

View File

@@ -60,7 +60,12 @@ class TradeStreamConsistencyController(
return state.accept(
trade,
before_checkpoint=(
sink.persist
sink.persist_accepted
if sink is not None
else None
),
on_duplicate=(
sink.persist_duplicate
if sink is not None
else None
),

View File

@@ -5,6 +5,7 @@ from __future__ import annotations
from collections import deque
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Protocol, Self
from src.market_data.acquisition.consistency.trade_stream_exceptions import (
TradeConsistencyError,
@@ -17,6 +18,19 @@ from src.market_data.acquisition.trade_id_sequence import (
)
DEFAULT_DEDUPLICATION_WINDOW_SIZE = 10_000
TradeDuplicateCallback = Callable[[Trade], None]
class TradeCheckpointCallback(Protocol):
"""Callback атомарной записи перед in-memory checkpoint."""
def __call__(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
...
@dataclass(slots=True)
@@ -49,11 +63,47 @@ class TradeStreamState:
)
self._trades = {}
@classmethod
def from_history(
cls,
*,
symbol: str,
trades: tuple[Trade, ...],
deduplication_window_size: int = (
DEFAULT_DEDUPLICATION_WINDOW_SIZE
),
) -> Self:
"""Строго восстановить временное состояние из durable tail."""
if not isinstance(trades, tuple):
raise TypeError("trades must be a tuple")
if len(trades) > deduplication_window_size:
raise ValueError(
"trades must fit deduplication_window_size"
)
state = cls(
symbol=symbol,
deduplication_window_size=deduplication_window_size,
)
for trade in trades:
if not isinstance(trade, Trade):
raise TypeError("trades must contain only Trade")
if state.accept(trade) is None:
raise TradeConsistencyError(
"Hydration history contains a duplicate Trade."
)
return state
def accept(
self,
trade: Trade,
*,
before_checkpoint: Callable[[Trade], None] | None = None,
before_checkpoint: TradeCheckpointCallback | None = None,
on_duplicate: TradeDuplicateCallback | None = None,
) -> Trade | None:
"""
Проверить сделку на согласованность.
@@ -71,6 +121,9 @@ class TradeStreamState:
):
raise TypeError("before_checkpoint must be callable")
if on_duplicate is not None and not callable(on_duplicate):
raise TypeError("on_duplicate must be callable")
if trade.symbol != self.symbol:
raise ValueError(
f"Unexpected symbol: {trade.symbol!r}"
@@ -93,9 +146,9 @@ class TradeStreamState:
previous,
trade,
):
self._run_before_checkpoint(
self._run_duplicate_observation(
trade,
callback=before_checkpoint,
callback=on_duplicate,
)
return None
@@ -108,9 +161,9 @@ class TradeStreamState:
previous,
trade,
):
self._run_before_checkpoint(
self._run_duplicate_observation(
trade,
callback=before_checkpoint,
callback=on_duplicate,
)
return None
@@ -118,6 +171,7 @@ class TradeStreamState:
self._run_before_checkpoint(
trade,
expected_trade=self.last_trade,
callback=before_checkpoint,
)
@@ -132,7 +186,20 @@ class TradeStreamState:
def _run_before_checkpoint(
trade: Trade,
*,
callback: Callable[[Trade], None] | None,
expected_trade: Trade | None,
callback: TradeCheckpointCallback | None,
) -> None:
if callback is not None:
callback(
trade,
expected_trade=expected_trade,
)
@staticmethod
def _run_duplicate_observation(
trade: Trade,
*,
callback: TradeDuplicateCallback | None,
) -> None:
if callback is not None:
callback(trade)

View File

@@ -19,6 +19,7 @@ from src.market_data.acquisition.consistency.trade_stream_state import (
)
from src.market_data.acquisition.consistency.trade_stream_state_store_exceptions import (
TradeStreamStateNotFoundError,
TradeStreamStateStoreInitializationError,
)
from src.market_data.acquisition.consistency.trade_stream_state_store_protocol import (
TradeStreamStateStoreProtocol,
@@ -38,6 +39,43 @@ class TradeStreamStateStore(TradeStreamStateStoreProtocol):
Создаёт пустое хранилище состояний.
"""
self._states: dict[str, TradeStreamState] = {}
self._initialized = False
def initialize(
self,
states: tuple[TradeStreamState, ...],
) -> None:
"""Однократно опубликовать полностью подготовленный набор."""
if not isinstance(states, tuple):
raise TypeError("states must be a tuple")
if self._initialized or self._states:
raise TradeStreamStateStoreInitializationError(
"Trade Stream State Store уже содержит состояния."
)
prepared: dict[str, TradeStreamState] = {}
for state in states:
if not isinstance(state, TradeStreamState):
raise TypeError(
"states must contain only TradeStreamState"
)
if state.symbol in prepared:
raise TradeStreamStateStoreInitializationError(
f"Состояние для символа {state.symbol!r} "
"передано повторно."
)
prepared[state.symbol] = state
self._states.update(prepared)
self._initialized = True
def is_empty(self) -> bool:
"""Проверить, что хранилище не содержит состояний."""
return not self._states and not self._initialized
def get_or_create(
self,
@@ -126,4 +164,5 @@ class TradeStreamStateStore(TradeStreamStateStoreProtocol):
"""
Полностью очищает хранилище состояний.
"""
self._states.clear()
self._states.clear()
self._initialized = False

View File

@@ -25,4 +25,10 @@ class TradeStreamStateNotFoundError(
):
"""
Состояние торгового инструмента отсутствует.
"""
"""
class TradeStreamStateStoreInitializationError(
TradeStreamStateStoreError,
):
"""Хранилище состояний нельзя безопасно инициализировать."""

View File

@@ -29,6 +29,17 @@ class TradeStreamStateStoreProtocol(Protocol):
Протокол хранилища состояний Trade Stream Consistency.
"""
def initialize(
self,
states: tuple[TradeStreamState, ...],
) -> None:
"""Однократно и атомарно опубликовать готовые состояния."""
...
def is_empty(self) -> bool:
"""Проверить, что ни одно состояние ещё не опубликовано."""
...
def get_or_create(
self,
symbol: str,
@@ -104,4 +115,4 @@ class TradeStreamStateStoreProtocol(Protocol):
"""
Полностью очищает хранилище состояний.
"""
...
...

View File

@@ -166,6 +166,20 @@ class WebSocketControlMessageError(MarketDataAcquisitionError):
pass
# Истекло время ожидания подтверждения стартовой WebSocket-подписки.
class WebSocketSubscriptionAckTimeoutError(
MarketDataAcquisitionError,
):
pass
# Переполнен буфер market-сообщений до подтверждения подписки.
class WebSocketStartupMarketBufferOverflowError(
MarketDataAcquisitionError,
):
pass
# Ошибка низкоуровневой WebSocket-инфраструктуры Acquisition Runtime.
class WebSocketTransportError(MarketDataAcquisitionError):
pass

View File

@@ -18,6 +18,7 @@ from src.market_data.acquisition.runtime.reconnect import (
from src.market_data.acquisition.runtime.runtime_recovery_protocol import (
RuntimeRecoveryProtocol,
)
from src.market_data.acquisition.symbols import normalize_symbol
RuntimeUnixTimeMillisecondsClock = Callable[[], int]
@@ -354,7 +355,7 @@ class RuntimeReconnectRecoveryCoordinator:
"symbols must contain only strings"
)
normalized_symbol = symbol.strip()
normalized_symbol = normalize_symbol(symbol)
if normalized_symbol:
normalized_symbols.add(normalized_symbol)

View File

@@ -0,0 +1,295 @@
from __future__ import annotations
import asyncio
from typing import Protocol, TypeVar, runtime_checkable
from src.market_data.acquisition.checkpoint.trade_stream_state_hydrator import (
TradeStreamStateHydratorProtocol,
)
from src.market_data.acquisition.consistency.trade_stream_state import (
TradeStreamState,
)
from src.market_data.acquisition.recovery.trade_recovery_result import (
TradeRecoveryResult,
)
from src.market_data.acquisition.runtime.live_processing_gate import (
RuntimeLiveProcessingGateProtocol,
)
from src.market_data.acquisition.runtime.runtime_reconnect_recovery_coordinator import (
RuntimeUnixTimeMillisecondsClock,
system_unix_time_ms,
)
from src.market_data.acquisition.runtime.runtime_recovery_protocol import (
RuntimeRecoveryProtocol,
)
from src.market_data.acquisition.symbols import normalize_symbol
_WorkerResult = TypeVar("_WorkerResult")
@runtime_checkable
class RuntimeStartupRecoveryProtocol(Protocol):
"""Контракт восстановления состояния при запуске Trade Stream."""
@property
def live_processing_gate(
self,
) -> RuntimeLiveProcessingGateProtocol:
"""Вернуть общий gate Startup Recovery и Live processing."""
...
@property
def symbols(self) -> tuple[str, ...]:
"""Вернуть нормализованные символы Startup Recovery."""
...
async def hydrate_once(
self,
) -> tuple[TradeStreamState, ...]:
"""Один раз восстановить Consistency state из хранилища."""
...
async def recover_after_ack(
self,
) -> tuple[TradeRecoveryResult, ...]:
"""Выполнить REST Recovery после подтверждения подписки."""
...
class RuntimeStartupRecoveryCoordinator:
"""
Координирует blocking-операции Startup Hydration и Recovery.
Владельцем общей границы subscribe → ACK → Recovery → buffered Live
остаётся Production Runtime. Координатор хранит тот же gate для
проверки единства dependency graph, но самостоятельно его не
захватывает.
"""
__slots__ = (
"_state_hydrator",
"_recovery_coordinator",
"_live_processing_gate",
"_symbols",
"_clock",
"_hydration_lock",
"_hydration_completed",
"_hydrated_states",
"_hydration_task",
"_recovery_task",
)
def __init__(
self,
*,
state_hydrator: TradeStreamStateHydratorProtocol,
recovery_coordinator: RuntimeRecoveryProtocol,
live_processing_gate: RuntimeLiveProcessingGateProtocol,
symbols: tuple[str, ...],
clock: RuntimeUnixTimeMillisecondsClock = system_unix_time_ms,
) -> None:
if not callable(clock):
raise TypeError("clock must be callable")
self._state_hydrator = state_hydrator
self._recovery_coordinator = recovery_coordinator
self._live_processing_gate = live_processing_gate
self._symbols = self._normalize_symbols(symbols)
self._clock = clock
self._hydration_lock = asyncio.Lock()
self._hydration_completed = False
self._hydrated_states: tuple[TradeStreamState, ...] = ()
self._hydration_task: (
asyncio.Task[tuple[TradeStreamState, ...]] | None
) = None
self._recovery_task: (
asyncio.Task[tuple[TradeRecoveryResult, ...]] | None
) = None
@property
def live_processing_gate(
self,
) -> RuntimeLiveProcessingGateProtocol:
"""Вернуть общий gate Startup Recovery и Live processing."""
return self._live_processing_gate
@property
def symbols(self) -> tuple[str, ...]:
"""Вернуть нормализованные символы Startup Recovery."""
return self._symbols
async def hydrate_once(
self,
) -> tuple[TradeStreamState, ...]:
"""
Восстановить Consistency state не более одного успешного раза.
При отмене вызывающего coroutine уже запущенный blocking worker
завершается до возврата управления. Если worker успел успешно
восстановить state, результат сохраняется для следующего вызова.
"""
if self._hydration_completed:
return self._hydrated_states
async with self._hydration_lock:
if self._hydration_completed:
return self._hydrated_states
hydration_task = asyncio.create_task(
asyncio.to_thread(
self._hydrate_symbols,
),
name="trade-stream-state-hydration",
)
self._hydration_task = hydration_task
try:
states = await asyncio.shield(
hydration_task,
)
except asyncio.CancelledError as cancellation:
try:
states = await self._wait_for_worker_completion(
hydration_task,
)
except BaseException as hydration_error:
cancellation.add_note(
"Trade Stream state hydration also failed: "
f"{type(hydration_error).__name__}."
)
else:
self._store_hydration_result(states)
raise
else:
self._store_hydration_result(states)
return states
finally:
if self._hydration_task is hydration_task:
self._hydration_task = None
async def recover_after_ack(
self,
) -> tuple[TradeRecoveryResult, ...]:
"""
Последовательно восстановить symbols по одной временной границе.
Вызывающий Production Runtime обязан удерживать общий live gate
от отправки подписки до завершения buffered Live processing.
"""
recovery_end_time = self._clock()
self._validate_recovery_end_time(recovery_end_time)
recovery_task = asyncio.create_task(
asyncio.to_thread(
self._recover_symbols,
recovery_end_time,
),
name="trade-stream-startup-recovery",
)
self._recovery_task = recovery_task
try:
return await asyncio.shield(
recovery_task,
)
except asyncio.CancelledError as cancellation:
try:
await self._wait_for_worker_completion(
recovery_task,
)
except BaseException as recovery_error:
cancellation.add_note(
"Trade Stream Startup Recovery also failed: "
f"{type(recovery_error).__name__}."
)
raise
finally:
if self._recovery_task is recovery_task:
self._recovery_task = None
def _hydrate_symbols(
self,
) -> tuple[TradeStreamState, ...]:
return self._state_hydrator.hydrate(
symbols=self._symbols,
)
def _recover_symbols(
self,
recovery_end_time: int,
) -> tuple[TradeRecoveryResult, ...]:
return tuple(
self._recovery_coordinator.recover(
symbol=symbol,
recovery_end_time=recovery_end_time,
)
for symbol in self._symbols
)
def _store_hydration_result(
self,
states: tuple[TradeStreamState, ...],
) -> None:
self._hydrated_states = states
self._hydration_completed = True
@staticmethod
async def _wait_for_worker_completion(
worker_task: asyncio.Task[_WorkerResult],
) -> _WorkerResult:
while not worker_task.done():
try:
await asyncio.shield(
worker_task,
)
except asyncio.CancelledError:
continue
except BaseException:
break
return worker_task.result()
@staticmethod
def _normalize_symbols(
symbols: tuple[str, ...],
) -> tuple[str, ...]:
if not isinstance(symbols, tuple):
raise TypeError("symbols must be a tuple")
normalized_symbols: set[str] = set()
for symbol in symbols:
if not isinstance(symbol, str):
raise TypeError(
"symbols must contain only strings"
)
normalized_symbol = normalize_symbol(symbol)
if normalized_symbol:
normalized_symbols.add(normalized_symbol)
if not normalized_symbols:
raise ValueError(
"symbols must contain at least one non-empty symbol"
)
return tuple(sorted(normalized_symbols))
@staticmethod
def _validate_recovery_end_time(
recovery_end_time: int,
) -> None:
if isinstance(recovery_end_time, bool) or not isinstance(
recovery_end_time,
int,
):
raise TypeError("clock must return an integer")
if recovery_end_time < 0:
raise ValueError(
"clock must not return a negative value"
)

View File

@@ -2,12 +2,16 @@ from __future__ import annotations
import asyncio
import json
import math
from collections import deque
from enum import Enum
from typing import Protocol, runtime_checkable
from uuid import uuid4
from src.market_data.acquisition.exceptions import (
WebSocketMessageDecodeError,
WebSocketStartupMarketBufferOverflowError,
WebSocketSubscriptionAckTimeoutError,
WebSocketTransportError,
)
from src.market_data.acquisition.models.trade import Trade
@@ -17,6 +21,9 @@ from src.market_data.acquisition.runtime.live_processing_gate import (
from src.market_data.acquisition.runtime.runtime_reconnect_recovery_coordinator import (
RuntimeReconnectRecoveryProtocol,
)
from src.market_data.acquisition.runtime.runtime_startup_recovery_coordinator import (
RuntimeStartupRecoveryProtocol,
)
from src.market_data.acquisition.runtime.scheduler import (
RuntimeSchedulerProtocol,
)
@@ -47,6 +54,7 @@ from src.market_data.acquisition.runtime.websocket_protocol import (
from src.market_data.acquisition.trade_stream_acquisition_protocol import (
TradeStreamAcquisitionServiceProtocol,
)
from src.market_data.acquisition.symbols import normalize_symbol
class TradeStreamProductionRuntimeState(str, Enum):
@@ -116,6 +124,9 @@ class TradeStreamProductionRuntime:
"_control_message_handler",
"_live_processing_gate",
"_reconnect_recovery_coordinator",
"_startup_recovery_coordinator",
"_subscription_ack_timeout_seconds",
"_startup_market_buffer_capacity",
"_runtime_supervisor",
"_runtime_scheduler",
"_symbols",
@@ -148,6 +159,11 @@ class TradeStreamProductionRuntime:
runtime_supervisor: RuntimeSupervisorProtocol,
runtime_scheduler: RuntimeSchedulerProtocol,
symbols: tuple[str, ...],
startup_recovery_coordinator: (
RuntimeStartupRecoveryProtocol | None
) = None,
subscription_ack_timeout_seconds: float = 10.0,
startup_market_buffer_capacity: int = 10_000,
) -> None:
self._session = session
self._transport = transport
@@ -174,6 +190,63 @@ class TradeStreamProductionRuntime:
"must use the same symbols."
)
if startup_recovery_coordinator is not None:
if (
startup_recovery_coordinator.live_processing_gate
is not live_processing_gate
):
raise ValueError(
"Production Runtime and startup recovery "
"coordinator must share one live processing gate."
)
if (
startup_recovery_coordinator.symbols
!= normalized_symbols
):
raise ValueError(
"Production Runtime and startup recovery "
"coordinator must use the same symbols."
)
if (
isinstance(subscription_ack_timeout_seconds, bool)
or not isinstance(
subscription_ack_timeout_seconds,
(int, float),
)
):
raise TypeError(
"subscription_ack_timeout_seconds must be a number"
)
normalized_ack_timeout = float(
subscription_ack_timeout_seconds,
)
if (
not math.isfinite(normalized_ack_timeout)
or normalized_ack_timeout <= 0.0
):
raise ValueError(
"subscription_ack_timeout_seconds must be finite "
"and greater than zero"
)
if (
isinstance(startup_market_buffer_capacity, bool)
or not isinstance(startup_market_buffer_capacity, int)
):
raise TypeError(
"startup_market_buffer_capacity must be an integer"
)
if startup_market_buffer_capacity <= 0:
raise ValueError(
"startup_market_buffer_capacity must be greater "
"than zero"
)
if runtime_scheduler.runtime_supervisor is not runtime_supervisor:
raise ValueError(
"Production Runtime and Scheduler must share one "
@@ -190,6 +263,15 @@ class TradeStreamProductionRuntime:
self._reconnect_recovery_coordinator = (
reconnect_recovery_coordinator
)
self._startup_recovery_coordinator = (
startup_recovery_coordinator
)
self._subscription_ack_timeout_seconds = (
normalized_ack_timeout
)
self._startup_market_buffer_capacity = (
startup_market_buffer_capacity
)
self._runtime_supervisor = runtime_supervisor
self._runtime_scheduler = runtime_scheduler
self._symbols = normalized_symbols
@@ -386,6 +468,19 @@ class TradeStreamProductionRuntime:
raise
async def _start_sequence(self) -> None:
startup_recovery_coordinator = (
self._startup_recovery_coordinator
)
if startup_recovery_coordinator is not None:
try:
await startup_recovery_coordinator.hydrate_once()
except asyncio.CancelledError:
raise
except Exception as error:
self._live_processing_gate.fail(error)
raise
try:
await self._session.start()
except asyncio.CancelledError:
@@ -396,6 +491,10 @@ class TradeStreamProductionRuntime:
reason=str(error),
)
)
if startup_recovery_coordinator is not None:
self._live_processing_gate.fail(error)
raise
self._session_started = True
@@ -410,11 +509,110 @@ class TradeStreamProductionRuntime:
correlation_id = str(uuid4())
self._subscription_correlation_id = correlation_id
await self._trade_stream_service.subscribe(
self._symbols,
correlation_id=correlation_id,
if startup_recovery_coordinator is None:
await self._trade_stream_service.subscribe(
self._symbols,
correlation_id=correlation_id,
)
return
try:
async with self._live_processing_gate:
await self._trade_stream_service.subscribe(
self._symbols,
correlation_id=correlation_id,
)
buffered_documents = await self._wait_for_startup_ack(
correlation_id=correlation_id,
)
await startup_recovery_coordinator.recover_after_ack()
for document in buffered_documents:
await self._process_market_document(document)
except asyncio.CancelledError:
raise
except Exception as error:
self._live_processing_gate.fail(error)
raise
async def _wait_for_startup_ack(
self,
*,
correlation_id: str,
) -> deque[object]:
"""
Дождаться ACK одним WebSocket consumer и сохранить ранний market FIFO.
Timeout запускается только после успешной отправки подписки.
После принятого ACK этот метод больше не читает транспорт.
"""
buffered_documents: deque[object] = deque()
ack_timeout = asyncio.timeout(
self._subscription_ack_timeout_seconds,
)
try:
async with ack_timeout:
while True:
raw_message = await self._transport.receive()
transport_message = self._wrap_transport_message(
raw_message,
)
await self._event_publisher.publish(
MessageReceivedEvent(
message=transport_message,
)
)
document = self._decode_message(raw_message)
message_kind = self._message_classifier.classify(
document,
)
if (
message_kind
is WebSocketInboundMessageKind.CONTROL
):
self._control_message_handler.handle(
document,
expected_correlation_id=correlation_id,
)
return buffered_documents
if (
message_kind
is not WebSocketInboundMessageKind.MARKET
):
raise TypeError(
"message classifier must return "
"WebSocketInboundMessageKind"
)
if (
len(buffered_documents)
>= self._startup_market_buffer_capacity
):
raise (
WebSocketStartupMarketBufferOverflowError(
"Переполнен буфер market-сообщений "
"до подтверждения WebSocket-подписки."
)
)
buffered_documents.append(document)
except TimeoutError as error:
if not ack_timeout.expired():
raise
raise WebSocketSubscriptionAckTimeoutError(
"Истекло время ожидания подтверждения стартовой "
"WebSocket-подписки."
) from error
async def _receive_loop(self) -> None:
while not self._stop_requested.is_set():
observed_generation = (
@@ -810,7 +1008,7 @@ class TradeStreamProductionRuntime:
"symbols must contain only strings"
)
normalized_symbol = symbol.strip()
normalized_symbol = normalize_symbol(symbol)
if normalized_symbol:
normalized_symbols.add(normalized_symbol)

View File

@@ -12,6 +12,9 @@ from typing import Protocol
from src.market_data.acquisition.adapters.dzengi.rest import (
DzengiTradesDocumentSource,
)
from src.market_data.acquisition.checkpoint.trade_stream_state_hydrator import (
TradeStreamStateHydrator,
)
from src.market_data.acquisition.consistency.trade_stream_consistency_controller import (
TradeStreamConsistencyController,
)
@@ -51,6 +54,9 @@ from src.market_data.acquisition.runtime.runtime_liveness_probe import (
from src.market_data.acquisition.runtime.runtime_recovery_coordinator import (
RuntimeRecoveryCoordinator,
)
from src.market_data.acquisition.runtime.runtime_startup_recovery_coordinator import (
RuntimeStartupRecoveryCoordinator,
)
from src.market_data.acquisition.runtime.scheduler import (
RuntimeScheduler,
RuntimeSleep,
@@ -70,6 +76,9 @@ from src.market_data.acquisition.trade_stream_acquisition_service import (
from src.market_data.acquisition.trade_stream_message_adapter_protocol import (
TradeStreamMessageAdapterProtocol,
)
from src.market_data.storage.contracts import (
TradeCheckpointStorageProtocol,
)
class _TradeStreamRuntimeTransportProtocol(
@@ -99,6 +108,10 @@ class TradeStreamRuntimeComposition:
recovery_controller: TradeRecoveryController
recovery_window_planner: TradeRecoveryWindowPlanner
runtime_recovery_coordinator: RuntimeRecoveryCoordinator
state_hydrator: TradeStreamStateHydrator | None
runtime_startup_recovery_coordinator: (
RuntimeStartupRecoveryCoordinator | None
)
acquisition_runtime_service: AcquisitionRuntimeService
trade_stream_acquisition_service: TradeStreamAcquisitionService
@@ -132,6 +145,8 @@ def build_trade_stream_runtime_composition(
system_unix_time_ms
),
scheduler_sleep: RuntimeSleep = asyncio.sleep,
checkpoint_storage: TradeCheckpointStorageProtocol | None = None,
checkpoint_venue: str | None = None,
) -> TradeStreamRuntimeComposition:
"""
Построить изолированный граф зависимостей Trade Stream Runtime.
@@ -145,6 +160,12 @@ def build_trade_stream_runtime_composition(
Factory не создаёт production transport, не читает Settings,
не запускает lifecycle и не создаёт фоновые asyncio-задачи.
"""
if (checkpoint_storage is None) != (checkpoint_venue is None):
raise ValueError(
"checkpoint_storage и checkpoint_venue должны быть "
"переданы вместе"
)
state_store = TradeStreamStateStore()
consistency_controller = TradeStreamConsistencyController(
@@ -198,6 +219,29 @@ def build_trade_stream_runtime_composition(
)
)
state_hydrator: TradeStreamStateHydrator | None = None
runtime_startup_recovery_coordinator: (
RuntimeStartupRecoveryCoordinator | None
) = None
if checkpoint_storage is not None and checkpoint_venue is not None:
state_hydrator = TradeStreamStateHydrator(
checkpoint_storage=checkpoint_storage,
state_store=state_store,
venue=checkpoint_venue,
)
runtime_startup_recovery_coordinator = (
RuntimeStartupRecoveryCoordinator(
state_hydrator=state_hydrator,
recovery_coordinator=runtime_recovery_coordinator,
live_processing_gate=live_processing_gate,
symbols=(
runtime_reconnect_recovery_coordinator.symbols
),
clock=recovery_end_time_clock,
)
)
heartbeat_monitor = HeartbeatMonitor(
event_publisher=event_publisher,
timeout_seconds=heartbeat_timeout_seconds,
@@ -226,6 +270,10 @@ def build_trade_stream_runtime_composition(
recovery_controller=recovery_controller,
recovery_window_planner=recovery_window_planner,
runtime_recovery_coordinator=runtime_recovery_coordinator,
state_hydrator=state_hydrator,
runtime_startup_recovery_coordinator=(
runtime_startup_recovery_coordinator
),
acquisition_runtime_service=acquisition_runtime_service,
trade_stream_acquisition_service=trade_stream_acquisition_service,
reconnect_coordinator=reconnect_coordinator,

View File

@@ -5,10 +5,14 @@ from src.market_data.storage.contracts import (
MarketDataBatchWriteResult,
MarketDataWriteResult,
MarketDataWriteStatus,
PersistentTradeCheckpoint,
QuoteStorageProtocol,
TradeCheckpointStorageProtocol,
TradeStorageProtocol,
)
from src.market_data.storage.exceptions import (
MarketDataCheckpointConflictError,
MarketDataCheckpointIntegrityError,
MarketDataStorageConfigurationError,
MarketDataStorageConflictError,
MarketDataStorageError,
@@ -46,6 +50,8 @@ from src.market_data.storage.postgres_partitions import (
__all__ = (
"CandleStorageProtocol",
"MarketDataBatchWriteResult",
"MarketDataCheckpointConflictError",
"MarketDataCheckpointIntegrityError",
"MarketDataPartition",
"MarketDataPartitionResult",
"MarketDataPartitionType",
@@ -60,12 +66,14 @@ __all__ = (
"MarketDataStorageValidationError",
"MarketDataWriteResult",
"MarketDataWriteStatus",
"PersistentTradeCheckpoint",
"PostgresCandleRepository",
"PostgresMarketDataPartitionManager",
"PostgresMarketDataRetentionService",
"PostgresQuoteRepository",
"PostgresTradeRepository",
"QuoteStorageProtocol",
"TradeCheckpointStorageProtocol",
"TradeStorageProtocol",
"TradeObservationClock",
"TradeStorageObservationSink",

View File

@@ -8,6 +8,9 @@ from typing import Protocol, runtime_checkable
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.trade_id_sequence import (
validate_signed_trade_id,
)
class MarketDataWriteStatus(Enum):
@@ -61,6 +64,78 @@ class MarketDataBatchWriteResult:
)
@dataclass(frozen=True, slots=True)
class PersistentTradeCheckpoint:
"""
Подтверждённая persistent-точка состояния Trade Stream.
"""
venue: str
trade: Trade
revision: int
updated_at: datetime
checkpoint_schema_version: int = 1
def __post_init__(self) -> None:
if not isinstance(self.venue, str):
raise TypeError("venue must be a string")
if not self.venue.strip():
raise ValueError("venue must not be empty")
if not isinstance(self.trade, Trade):
raise TypeError("trade must be a Canonical Trade")
if not isinstance(self.trade.symbol, str):
raise TypeError("trade.symbol must be a string")
if not self.trade.symbol.strip():
raise ValueError("trade.symbol must not be empty")
validate_signed_trade_id(self.trade.trade_id)
if (
not isinstance(self.trade.executed_at, datetime)
or self.trade.executed_at.tzinfo is None
or self.trade.executed_at.utcoffset() is None
):
raise ValueError(
"trade.executed_at must contain timezone"
)
for field_name, value in (
("revision", self.revision),
(
"checkpoint_schema_version",
self.checkpoint_schema_version,
),
):
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 not isinstance(self.updated_at, datetime):
raise TypeError("updated_at must be a datetime")
if (
self.updated_at.tzinfo is None
or self.updated_at.utcoffset() is None
):
raise ValueError("updated_at must contain timezone")
@property
def identity(self) -> tuple[str, str, int, datetime]:
"""Вернуть durable identity checkpoint Trade."""
return (
self.venue,
self.trade.symbol,
self.trade.trade_id,
self.trade.executed_at,
)
@runtime_checkable
class TradeStorageProtocol(Protocol):
"""Граница постоянного хранения канонических сделок только для записи."""
@@ -84,6 +159,72 @@ class TradeStorageProtocol(Protocol):
...
@runtime_checkable
class TradeCheckpointStorageProtocol(Protocol):
"""
Узкая граница хранения и восстановления Trade checkpoint.
"""
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
"""
Загрузить checkpoint, подтверждённый Canonical Trade history.
"""
...
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
"""Загрузить bounded history до указанного checkpoint."""
...
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
"""
Загрузить bounded tail для первой инициализации checkpoint.
"""
...
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
"""
Создать первый checkpoint на существующей durable Trade.
Операция не должна повторно записывать наблюдение Trade либо
изменять её provenance и observation timestamps.
"""
...
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
"""
Атомарно сохранить принятую Trade и продвинуть checkpoint.
"""
...
@runtime_checkable
class QuoteStorageProtocol(Protocol):
"""Граница постоянного хранения канонических котировок для записи."""

View File

@@ -19,3 +19,15 @@ class MarketDataStorageConflictError(MarketDataStorageError):
class MarketDataStorageOperationError(MarketDataStorageError):
"""Ошибка операции постоянного хранилища с базой данных."""
class MarketDataCheckpointConflictError(
MarketDataStorageConflictError,
):
"""
Persistent checkpoint не совпал с ожидаемым состоянием writer.
"""
class MarketDataCheckpointIntegrityError(MarketDataStorageError):
"""Persistent checkpoint не подтверждается Canonical Trade history."""

View File

@@ -21,6 +21,32 @@ from src.market_data.storage.postgres_repository_support import (
MARKET_DATA_PARTITION_ADVISORY_LOCK_ID = 0x445A504152544E
_SCHEMA_NAME = "market_data"
_TRADE_CHECKPOINT_FK_NAME = "trade_stream_checkpoints_trade_fk"
def _drop_trade_checkpoint_foreign_key(*, cursor: Any) -> None:
cursor.execute(
sql.SQL("ALTER TABLE {} DROP CONSTRAINT {}").format(
sql.Identifier(_SCHEMA_NAME, "trade_stream_checkpoints"),
sql.Identifier(_TRADE_CHECKPOINT_FK_NAME),
)
)
def _restore_trade_checkpoint_foreign_key(*, cursor: Any) -> None:
cursor.execute(
sql.SQL(
"ALTER TABLE {} ADD CONSTRAINT {} "
"FOREIGN KEY (venue, symbol, trade_id, executed_at) "
"REFERENCES {} (venue, symbol, trade_id, executed_at) "
"ON UPDATE NO ACTION ON DELETE NO ACTION "
"DEFERRABLE INITIALLY DEFERRED"
).format(
sql.Identifier(_SCHEMA_NAME, "trade_stream_checkpoints"),
sql.Identifier(_TRADE_CHECKPOINT_FK_NAME),
sql.Identifier(_SCHEMA_NAME, "trades"),
)
)
class MarketDataPartitionType(Enum):
@@ -456,6 +482,9 @@ class PostgresMarketDataPartitionManager:
partition=partition,
spec=spec,
)
if partition.data_type is MarketDataPartitionType.TRADES:
_restore_trade_checkpoint_foreign_key(cursor=cursor)
except MarketDataStorageError:
raise
except Exception as error:
@@ -486,9 +515,20 @@ class PostgresMarketDataPartitionManager:
range_start = sql.Literal(partition.range_start)
range_end = sql.Literal(partition.range_end)
if partition.data_type is MarketDataPartitionType.TRADES:
cursor.execute(
sql.SQL(
"LOCK TABLE {} IN SHARE ROW EXCLUSIVE MODE"
).format(parent)
)
cursor.execute(
sql.SQL("LOCK TABLE {} IN ACCESS EXCLUSIVE MODE").format(default)
)
if partition.data_type is MarketDataPartitionType.TRADES:
_drop_trade_checkpoint_foreign_key(cursor=cursor)
cursor.execute(
sql.SQL("CREATE TABLE {} (LIKE {} INCLUDING ALL)").format(
child,
@@ -625,6 +665,14 @@ class PostgresMarketDataRetentionService:
)
)
maintains_trade_checkpoints = any(
data_type is MarketDataPartitionType.TRADES
for data_type, _ in configured
)
if maintains_trade_checkpoints:
_drop_trade_checkpoint_foreign_key(cursor=cursor)
entries = tuple(
self._apply_one(
cursor=cursor,
@@ -633,6 +681,9 @@ class PostgresMarketDataRetentionService:
)
for data_type, days in configured
)
if maintains_trade_checkpoints:
_restore_trade_checkpoint_foreign_key(cursor=cursor)
except MarketDataStorageError:
raise
except Exception as error:

View File

@@ -10,14 +10,18 @@ from src.market_data.acquisition.models.trade import (
TradeAggressorSide,
)
from src.market_data.acquisition.trade_id_sequence import (
is_trade_id_newer,
validate_signed_trade_id,
)
from src.market_data.storage.contracts import (
MarketDataBatchWriteResult,
MarketDataWriteResult,
MarketDataWriteStatus,
PersistentTradeCheckpoint,
)
from src.market_data.storage.exceptions import (
MarketDataCheckpointConflictError,
MarketDataCheckpointIntegrityError,
MarketDataStorageConflictError,
MarketDataStorageError,
MarketDataStorageOperationError,
@@ -33,6 +37,7 @@ from src.market_data.storage.postgres_repository_support import (
CANONICAL_TRADE_SCHEMA_VERSION = 1
CHECKPOINT_SCHEMA_VERSION = 1
_INSERT_TRADE_SQL = """
INSERT INTO market_data.trades (
@@ -84,6 +89,165 @@ WHERE venue = %s
AND executed_at = %s
"""
_SELECT_CHECKPOINT_SQL = """
SELECT
checkpoint.venue,
checkpoint.symbol,
checkpoint.trade_id,
checkpoint.executed_at,
checkpoint.revision,
checkpoint.updated_at,
checkpoint.checkpoint_schema_version,
trade_row.venue,
trade_row.symbol,
trade_row.trade_id,
trade_row.executed_at,
trade_row.price,
trade_row.quantity,
trade_row.aggressor_side,
trade_row.source,
trade_row.canonical_schema_version
FROM market_data.trade_stream_checkpoints AS checkpoint
LEFT JOIN market_data.trades AS trade_row
ON trade_row.venue = checkpoint.venue
AND trade_row.symbol = checkpoint.symbol
AND trade_row.trade_id = checkpoint.trade_id
AND trade_row.executed_at = checkpoint.executed_at
WHERE checkpoint.venue = %s
AND checkpoint.symbol = %s
"""
_SELECT_CHECKPOINT_FOR_UPDATE_SQL = (
_SELECT_CHECKPOINT_SQL
+ "FOR UPDATE OF checkpoint\n"
)
_SELECT_TRADE_FOR_ADOPTION_SQL = """
SELECT
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
canonical_schema_version
FROM market_data.trades
WHERE venue = %s
AND symbol = %s
AND trade_id = %s
AND executed_at = %s
FOR SHARE
"""
_INSERT_CHECKPOINT_SQL = """
INSERT INTO market_data.trade_stream_checkpoints (
venue,
symbol,
trade_id,
executed_at,
revision,
checkpoint_schema_version
)
VALUES (%s, %s, %s, %s, 1, %s)
ON CONFLICT (venue, symbol) DO NOTHING
RETURNING 1
"""
_UPDATE_CHECKPOINT_SQL = """
UPDATE market_data.trade_stream_checkpoints
SET trade_id = %s,
executed_at = %s,
revision = revision + 1,
updated_at = NOW(),
checkpoint_schema_version = %s
WHERE venue = %s
AND symbol = %s
AND revision = %s
RETURNING 1
"""
_SELECT_LATEST_TRADE_SQL = """
SELECT
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
canonical_schema_version
FROM market_data.trades
WHERE venue = %s
AND symbol = %s
ORDER BY first_observed_at DESC,
executed_at DESC,
trade_id DESC
LIMIT 1
"""
_SELECT_TRADE_TAIL_SQL = """
WITH reference_trade AS (
SELECT first_observed_at
FROM market_data.trades
WHERE venue = %s
AND symbol = %s
AND trade_id = %s
AND executed_at = %s
),
ranked_trade_cycles AS (
SELECT
trade_row.symbol,
trade_row.trade_id,
trade_row.executed_at,
trade_row.price,
trade_row.quantity,
trade_row.aggressor_side,
trade_row.source,
trade_row.canonical_schema_version,
trade_row.first_observed_at,
ROW_NUMBER() OVER (
PARTITION BY trade_row.trade_id
ORDER BY trade_row.first_observed_at DESC,
trade_row.executed_at DESC
) AS cycle_rank
FROM market_data.trades AS trade_row
CROSS JOIN reference_trade
WHERE trade_row.venue = %s
AND trade_row.symbol = %s
AND trade_row.first_observed_at
<= reference_trade.first_observed_at
),
current_trade_cycle AS (
SELECT
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
canonical_schema_version,
first_observed_at
FROM ranked_trade_cycles
WHERE cycle_rank = 1
)
SELECT
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
canonical_schema_version
FROM current_trade_cycle
ORDER BY first_observed_at DESC,
executed_at DESC,
trade_id DESC
LIMIT %s
"""
@dataclass(frozen=True, slots=True)
class _PreparedTrade:
@@ -106,7 +270,7 @@ class _PreparedTrade:
class PostgresTradeRepository:
"""Транзакционный модуль записи канонических сделок в PostgreSQL."""
"""Единый транзакционный модуль Trades и persistent checkpoint."""
__slots__ = ("_connection_provider",)
@@ -120,6 +284,284 @@ class PostgresTradeRepository:
self._connection_provider = connection_provider
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
normalized_venue = normalize_venue(venue)
normalized_symbol = normalize_non_empty_text(
symbol,
field_name="symbol",
).upper()
try:
with self._connection_provider() as connection:
with connection.cursor() as cursor:
return self._load_checkpoint(
cursor=cursor,
venue=normalized_venue,
symbol=normalized_symbol,
for_update=False,
)
except MarketDataStorageError:
raise
except Exception as error:
raise MarketDataStorageOperationError(
"Failed to load persistent Trade checkpoint."
) from error
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
normalized_venue = normalize_venue(venue)
normalized_limit = self._validate_tail_limit(limit)
if not isinstance(checkpoint, PersistentTradeCheckpoint):
raise MarketDataStorageValidationError(
"checkpoint must be PersistentTradeCheckpoint"
)
checkpoint_venue = normalize_venue(checkpoint.venue)
checkpoint_symbol = normalize_non_empty_text(
checkpoint.trade.symbol,
field_name="checkpoint.trade.symbol",
).upper()
if checkpoint_venue != normalized_venue:
raise MarketDataStorageValidationError(
"checkpoint venue does not match venue"
)
try:
with self._connection_provider() as connection:
with connection.cursor() as cursor:
trades = self._load_trade_tail(
cursor=cursor,
venue=normalized_venue,
symbol=checkpoint_symbol,
reference_trade_id=checkpoint.trade.trade_id,
reference_executed_at=normalize_aware_datetime(
checkpoint.trade.executed_at,
field_name="checkpoint.trade.executed_at",
),
limit=normalized_limit,
)
except MarketDataStorageError:
raise
except Exception as error:
raise MarketDataStorageOperationError(
"Failed to load persistent Trade checkpoint tail."
) from error
return self._validate_trade_tail(
trades=trades,
expected_last_identity=(
checkpoint_venue,
checkpoint_symbol,
checkpoint.trade.trade_id,
normalize_aware_datetime(
checkpoint.trade.executed_at,
field_name="checkpoint.trade.executed_at",
),
),
)
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
normalized_venue = normalize_venue(venue)
normalized_symbol = normalize_non_empty_text(
symbol,
field_name="symbol",
).upper()
normalized_limit = self._validate_tail_limit(limit)
try:
with self._connection_provider() as connection:
with connection.cursor() as cursor:
cursor.execute(
_SELECT_LATEST_TRADE_SQL,
(
normalized_venue,
normalized_symbol,
),
)
latest_row = cursor.fetchone()
if latest_row is None:
return ()
latest_trade = self._trade_from_row(latest_row)
trades = self._load_trade_tail(
cursor=cursor,
venue=normalized_venue,
symbol=normalized_symbol,
reference_trade_id=latest_trade.trade_id,
reference_executed_at=latest_trade.executed_at,
limit=normalized_limit,
)
except MarketDataStorageError:
raise
except Exception as error:
raise MarketDataStorageOperationError(
"Failed to load latest persistent Trade tail."
) from error
return self._validate_trade_tail(
trades=trades,
expected_last_identity=self._trade_identity(
venue=normalized_venue,
trade=latest_trade,
),
)
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
normalized_venue = normalize_venue(venue)
if not isinstance(trade, Trade):
raise MarketDataStorageValidationError(
"trade must be a Canonical Trade"
)
prepared = self._prepare_trade(
trade=trade,
observed_at=trade.executed_at,
)
candidate_identity = self._prepared_identity(
venue=normalized_venue,
trade=prepared,
)
try:
with self._connection_provider() as connection:
with connection.cursor() as cursor:
durable_trade = self._load_trade_for_adoption(
cursor=cursor,
venue=normalized_venue,
trade=prepared,
)
if durable_trade is None:
raise MarketDataCheckpointIntegrityError(
"Checkpoint candidate is absent from durable "
"Canonical Trade history."
)
if not self._prepared_trade_matches(
prepared=prepared,
stored=durable_trade,
):
raise MarketDataCheckpointConflictError(
"Checkpoint candidate conflicts with durable "
"Canonical Trade facts."
)
cursor.execute(
_INSERT_CHECKPOINT_SQL,
(
*candidate_identity,
CHECKPOINT_SCHEMA_VERSION,
),
)
inserted = cursor.fetchone() is not None
checkpoint = self._load_checkpoint(
cursor=cursor,
venue=normalized_venue,
symbol=prepared.symbol,
for_update=True,
)
if checkpoint is None:
raise MarketDataCheckpointIntegrityError(
"Adopted checkpoint disappeared after insert."
)
checkpoint_matches = (
self._normalized_checkpoint_identity(checkpoint)
== candidate_identity
and self._prepared_trade_matches(
prepared=prepared,
stored=checkpoint.trade,
)
and checkpoint.revision == 1
)
if not checkpoint_matches:
if inserted:
raise MarketDataCheckpointIntegrityError(
"Adopted checkpoint does not match durable "
"candidate Trade."
)
raise MarketDataCheckpointConflictError(
"Persistent checkpoint already points to another "
"Trade or revision."
)
except MarketDataStorageError:
raise
except Exception as error:
raise MarketDataStorageOperationError(
"Failed to adopt existing Trade as persistent checkpoint."
) from error
return checkpoint
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
normalized_venue = normalize_venue(venue)
prepared = self._prepare_trade(
trade=trade,
observed_at=observed_at,
)
expected_identity = self._prepare_expected_identity(
venue=normalized_venue,
expected_trade=expected_trade,
)
try:
with self._connection_provider() as connection:
with connection.cursor() as cursor:
self._store_prepared_trade(
cursor=cursor,
venue=normalized_venue,
trade=prepared,
)
checkpoint = self._advance_checkpoint(
cursor=cursor,
venue=normalized_venue,
expected_identity=expected_identity,
trade=prepared,
)
except MarketDataStorageError:
raise
except Exception as error:
raise MarketDataStorageOperationError(
"Failed to atomically store Trade and advance checkpoint."
) from error
return checkpoint
def store_trade(
self,
*,
@@ -219,6 +661,461 @@ class PostgresTradeRepository:
provenance_updated_count=provenance_updated_count,
)
def _advance_checkpoint(
self,
*,
cursor: Any,
venue: str,
expected_identity: tuple[str, str, int, datetime] | None,
trade: _PreparedTrade,
) -> PersistentTradeCheckpoint:
candidate_identity = self._prepared_identity(
venue=venue,
trade=trade,
)
if expected_identity is None:
cursor.execute(
_INSERT_CHECKPOINT_SQL,
(
*candidate_identity,
CHECKPOINT_SCHEMA_VERSION,
),
)
inserted = cursor.fetchone() is not None
current = self._load_checkpoint(
cursor=cursor,
venue=venue,
symbol=trade.symbol,
for_update=True,
)
if current is None:
raise MarketDataCheckpointIntegrityError(
"Checkpoint disappeared after first insert."
)
if inserted:
if (
self._normalized_checkpoint_identity(current)
!= candidate_identity
):
raise MarketDataCheckpointIntegrityError(
"First checkpoint does not match candidate Trade."
)
return current
else:
current = self._load_checkpoint(
cursor=cursor,
venue=venue,
symbol=trade.symbol,
for_update=True,
)
if current is None:
raise MarketDataCheckpointConflictError(
"Expected persistent checkpoint does not exist."
)
current_identity = self._normalized_checkpoint_identity(current)
if current_identity == candidate_identity:
return current
if expected_identity is None or current_identity != expected_identity:
raise MarketDataCheckpointConflictError(
"Persistent checkpoint differs from expected Trade."
)
try:
candidate_is_newer = is_trade_id_newer(
trade.trade_id,
current.trade.trade_id,
)
except (TypeError, ValueError) as error:
raise MarketDataCheckpointConflictError(
"Candidate Trade has ambiguous checkpoint ordering."
) from error
if not candidate_is_newer:
raise MarketDataCheckpointConflictError(
"Candidate Trade does not advance persistent checkpoint."
)
cursor.execute(
_UPDATE_CHECKPOINT_SQL,
(
trade.trade_id,
trade.executed_at,
CHECKPOINT_SCHEMA_VERSION,
venue,
trade.symbol,
current.revision,
),
)
if cursor.fetchone() is None:
raise MarketDataCheckpointConflictError(
"Persistent checkpoint changed during update."
)
updated = self._load_checkpoint(
cursor=cursor,
venue=venue,
symbol=trade.symbol,
for_update=True,
)
if (
updated is None
or self._normalized_checkpoint_identity(updated)
!= candidate_identity
or updated.revision != current.revision + 1
):
raise MarketDataCheckpointIntegrityError(
"Updated checkpoint does not match committed Trade."
)
return updated
def _load_checkpoint(
self,
*,
cursor: Any,
venue: str,
symbol: str,
for_update: bool,
) -> PersistentTradeCheckpoint | None:
cursor.execute(
(
_SELECT_CHECKPOINT_FOR_UPDATE_SQL
if for_update
else _SELECT_CHECKPOINT_SQL
),
(
venue,
symbol,
),
)
row = cursor.fetchone()
if row is None:
return None
return self._checkpoint_from_row(
row=row,
expected_venue=venue,
expected_symbol=symbol,
)
def _load_trade_for_adoption(
self,
*,
cursor: Any,
venue: str,
trade: _PreparedTrade,
) -> Trade | None:
cursor.execute(
_SELECT_TRADE_FOR_ADOPTION_SQL,
self._identity_parameters(
venue=venue,
trade=trade,
),
)
row = cursor.fetchone()
if row is None:
return None
return self._trade_from_row(row)
def _load_trade_tail(
self,
*,
cursor: Any,
venue: str,
symbol: str,
reference_trade_id: int,
reference_executed_at: datetime,
limit: int,
) -> tuple[Trade, ...]:
cursor.execute(
_SELECT_TRADE_TAIL_SQL,
(
venue,
symbol,
reference_trade_id,
reference_executed_at,
venue,
symbol,
limit,
),
)
descending = tuple(
self._trade_from_row(row)
for row in cursor.fetchall()
)
return tuple(reversed(descending))
def _checkpoint_from_row(
self,
*,
row: object,
expected_venue: str,
expected_symbol: str,
) -> PersistentTradeCheckpoint:
if not isinstance(row, tuple) or len(row) != 16:
raise MarketDataCheckpointIntegrityError(
"PostgreSQL returned invalid checkpoint row."
)
(
checkpoint_venue,
checkpoint_symbol,
checkpoint_trade_id,
checkpoint_executed_at,
revision,
updated_at,
checkpoint_schema_version,
trade_venue,
trade_symbol,
trade_id,
trade_executed_at,
price,
quantity,
aggressor_side,
source,
canonical_schema_version,
) = row
if (
checkpoint_venue != expected_venue
or checkpoint_symbol != expected_symbol
or trade_venue != checkpoint_venue
or trade_symbol != checkpoint_symbol
or trade_id != checkpoint_trade_id
or trade_executed_at != checkpoint_executed_at
or checkpoint_schema_version != CHECKPOINT_SCHEMA_VERSION
):
raise MarketDataCheckpointIntegrityError(
"Checkpoint is not confirmed by exact Canonical Trade."
)
trade = self._trade_from_row(
(
trade_symbol,
trade_id,
trade_executed_at,
price,
quantity,
aggressor_side,
source,
canonical_schema_version,
)
)
try:
return PersistentTradeCheckpoint(
venue=checkpoint_venue,
trade=trade,
revision=revision,
updated_at=updated_at,
checkpoint_schema_version=checkpoint_schema_version,
)
except (TypeError, ValueError) as error:
raise MarketDataCheckpointIntegrityError(
"PostgreSQL returned invalid checkpoint values."
) from error
@staticmethod
def _trade_from_row(row: object) -> Trade:
if not isinstance(row, tuple) or len(row) != 8:
raise MarketDataCheckpointIntegrityError(
"PostgreSQL returned invalid Canonical Trade row."
)
(
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
canonical_schema_version,
) = row
try:
normalized_symbol = normalize_non_empty_text(
symbol,
field_name="stored trade.symbol",
).upper()
validate_signed_trade_id(trade_id)
normalized_executed_at = normalize_aware_datetime(
executed_at,
field_name="stored trade.executed_at",
)
normalized_price = validate_decimal(
price,
field_name="stored trade.price",
)
normalized_quantity = validate_decimal(
quantity,
field_name="stored trade.quantity",
)
normalized_source = normalize_non_empty_text(
source,
field_name="stored trade.source",
)
normalized_side = TradeAggressorSide(aggressor_side)
except (TypeError, ValueError, MarketDataStorageError) as error:
raise MarketDataCheckpointIntegrityError(
"PostgreSQL returned invalid Canonical Trade values."
) from error
if canonical_schema_version != CANONICAL_TRADE_SCHEMA_VERSION:
raise MarketDataCheckpointIntegrityError(
"Unsupported Canonical Trade schema version."
)
return Trade(
symbol=normalized_symbol,
trade_id=trade_id,
price=normalized_price,
quantity=normalized_quantity,
executed_at=normalized_executed_at,
aggressor_side=normalized_side,
source=normalized_source,
)
def _prepare_expected_identity(
self,
*,
venue: str,
expected_trade: Trade | None,
) -> tuple[str, str, int, datetime] | None:
if expected_trade is None:
return None
if not isinstance(expected_trade, Trade):
raise MarketDataStorageValidationError(
"expected_trade must be a Canonical Trade or None"
)
prepared = self._prepare_trade(
trade=expected_trade,
observed_at=expected_trade.executed_at,
)
return self._prepared_identity(
venue=venue,
trade=prepared,
)
@staticmethod
def _prepared_identity(
*,
venue: str,
trade: _PreparedTrade,
) -> tuple[str, str, int, datetime]:
return (
venue,
trade.symbol,
trade.trade_id,
trade.executed_at,
)
@staticmethod
def _normalized_checkpoint_identity(
checkpoint: PersistentTradeCheckpoint,
) -> tuple[str, str, int, datetime]:
return (
checkpoint.venue.strip(),
checkpoint.trade.symbol.strip().upper(),
checkpoint.trade.trade_id,
checkpoint.trade.executed_at,
)
@staticmethod
def _trade_identity(
*,
venue: str,
trade: Trade,
) -> tuple[str, str, int, datetime]:
return (
venue,
trade.symbol,
trade.trade_id,
trade.executed_at,
)
@staticmethod
def _prepared_trade_matches(
*,
prepared: _PreparedTrade,
stored: Trade,
) -> bool:
return (
stored.symbol == prepared.symbol
and stored.trade_id == prepared.trade_id
and stored.executed_at == prepared.executed_at
and stored.price == prepared.price
and stored.quantity == prepared.quantity
and stored.aggressor_side.value == prepared.aggressor_side
and stored.source == prepared.source
)
@classmethod
def _validate_trade_tail(
cls,
*,
trades: tuple[Trade, ...],
expected_last_identity: tuple[str, str, int, datetime],
) -> tuple[Trade, ...]:
if (
not trades
or cls._trade_identity(
venue=expected_last_identity[0],
trade=trades[-1],
)
!= expected_last_identity
):
raise MarketDataCheckpointIntegrityError(
"Trade tail does not end at expected durable Trade."
)
for previous, current in zip(trades, trades[1:], strict=False):
try:
current_is_newer = is_trade_id_newer(
current.trade_id,
previous.trade_id,
)
except (TypeError, ValueError) as error:
raise MarketDataCheckpointIntegrityError(
"Trade tail contains ambiguous rollover ordering."
) from error
if not current_is_newer:
raise MarketDataCheckpointIntegrityError(
"Trade tail is not strictly rollover-ordered."
)
return trades
@staticmethod
def _validate_tail_limit(limit: int) -> int:
if isinstance(limit, bool) or not isinstance(limit, int):
raise MarketDataStorageValidationError(
"limit must be an integer"
)
if limit <= 0:
raise MarketDataStorageValidationError(
"limit must be positive"
)
return limit
def _store_prepared_trade(
self,
*,

View File

@@ -2,10 +2,13 @@ from __future__ import annotations
from collections.abc import Callable
from datetime import datetime, timezone
from typing import cast
from src.market_data.acquisition.models.trade import Trade
from src.market_data.storage.contracts import (
MarketDataWriteResult,
PersistentTradeCheckpoint,
TradeCheckpointStorageProtocol,
TradeStorageProtocol,
)
from src.market_data.storage.postgres_repository_support import (
@@ -25,6 +28,7 @@ class TradeStorageObservationSink:
"""Адаптер наблюдений канонических сделок к постоянному хранилищу."""
__slots__ = (
"_checkpoint_storage",
"_trade_storage",
"_venue",
"_clock",
@@ -42,14 +46,45 @@ class TradeStorageObservationSink:
"trade_storage must implement TradeStorageProtocol"
)
if not isinstance(trade_storage, TradeCheckpointStorageProtocol):
raise TypeError(
"trade_storage must implement "
"TradeCheckpointStorageProtocol"
)
if not callable(clock):
raise TypeError("clock must be callable")
self._trade_storage = trade_storage
self._checkpoint_storage = cast(
TradeCheckpointStorageProtocol,
trade_storage,
)
self._venue = normalize_venue(venue)
self._clock = clock
def persist(
def persist_accepted(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
checkpoint = (
self._checkpoint_storage.store_trade_and_advance_checkpoint(
venue=self._venue,
expected_trade=expected_trade,
trade=trade,
observed_at=self._clock(),
)
)
if not isinstance(checkpoint, PersistentTradeCheckpoint):
raise TypeError(
"trade_storage.store_trade_and_advance_checkpoint() "
"must return PersistentTradeCheckpoint"
)
def persist_duplicate(
self,
trade: Trade,
) -> None:

View File

@@ -292,6 +292,44 @@ STORAGE_MIGRATIONS = (
""",
),
),
StorageMigration(
version=8,
name="create_trade_stream_checkpoints",
statements=(
"""
CREATE TABLE market_data.trade_stream_checkpoints (
venue TEXT NOT NULL,
symbol TEXT NOT NULL,
trade_id INTEGER NOT NULL,
executed_at TIMESTAMPTZ NOT NULL,
revision BIGINT NOT NULL,
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
checkpoint_schema_version INTEGER NOT NULL DEFAULT 1,
PRIMARY KEY (venue, symbol),
CONSTRAINT trade_stream_checkpoints_trade_fk
FOREIGN KEY (
venue,
symbol,
trade_id,
executed_at
) REFERENCES market_data.trades (
venue,
symbol,
trade_id,
executed_at
)
ON UPDATE NO ACTION
ON DELETE NO ACTION
DEFERRABLE INITIALLY DEFERRED,
CHECK (trade_id BETWEEN -2147483648 AND 2147483647),
CHECK (BTRIM(venue) <> ''),
CHECK (BTRIM(symbol) <> ''),
CHECK (revision > 0),
CHECK (checkpoint_schema_version > 0)
)
""",
),
),
)

View File

@@ -214,6 +214,26 @@ class LoopbackTradeWebSocketServer:
),
)
async def send_ack(
self,
connection_index: int,
*,
correlation_id: str,
status: str = "OK",
destination: str = "trades.subscribe",
) -> None:
"""Отправить управляемый ответ на Trade subscription."""
await self.send_raw(
connection_index,
json.dumps(
{
"correlationId": correlation_id,
"destination": destination,
"status": status,
}
),
)
async def abort_connection(
self,
connection_index: int,

View File

@@ -3,10 +3,29 @@ from __future__ import annotations
import asyncio
import threading
import time
from datetime import datetime, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.core.config import Settings
from src.integrations.exchange.rest_client import ExchangeRestClient
from src.market_data.acquisition.adapters.dzengi.rest import (
DzengiTradesDocumentSource,
)
from src.market_data.acquisition.adapters.dzengi.websocket import (
DzengiUnifiedWebSocketAdapter,
)
from src.market_data.acquisition.adapters.dzengi.websocket_control_message_handler import (
DzengiWebSocketControlMessageHandler,
)
from src.market_data.acquisition.adapters.dzengi.websocket_inbound_message_classifier import (
DzengiWebSocketInboundMessageClassifier,
)
from src.market_data.acquisition.adapters.dzengi.websocket_transport import (
DzengiWebSocketTransport,
)
from tests.integration.market_data.acquisition.runtime import (
loopback_trade_exchange,
)
@@ -22,6 +41,13 @@ from src.market_data.acquisition.exceptions import (
TradeTransportError,
WebSocketMessageDecodeError,
)
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.acquisition.runtime.acquisition_runtime_event_publisher import (
AcquisitionRuntimeEventPublisher,
)
from src.market_data.acquisition.runtime.runtime_reconnect_recovery_coordinator import (
RuntimeReconnectRecoveryCoordinator,
)
@@ -29,10 +55,22 @@ from src.market_data.acquisition.runtime.trade_stream_production_runtime import
TradeStreamProductionRuntime,
TradeStreamProductionRuntimeState,
)
from src.market_data.acquisition.runtime.websocket_session import (
WebSocketSession,
)
from src.market_data.acquisition.runtime.websocket_subscription_manager import (
WebSocketSubscriptionManager,
)
from src.market_data.acquisition.trade_stream_runtime_composition import (
TradeStreamRuntimeComposition,
build_trade_stream_runtime_composition,
)
from src.market_data.storage.contracts import PersistentTradeCheckpoint
from tests.support.trade_stream_runtime import (
SYMBOL,
assert_no_owned_tasks,
build_runtime,
make_settings,
reconnect_coordinator_from,
run_scenario,
start_runtime,
@@ -44,6 +82,187 @@ from tests.support.trade_stream_runtime import (
pytestmark = pytest.mark.integration
class StaticCheckpointStorage:
"""Локальное checkpoint-хранилище для startup boundary."""
def __init__(
self,
*,
checkpoint: PersistentTradeCheckpoint,
tail: tuple[Trade, ...],
events: list[str],
) -> None:
self._checkpoint = checkpoint
self._tail = tail
self._events = events
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
assert venue == "dzengi"
assert symbol == SYMBOL
self._events.append("checkpoint.load")
return self._checkpoint
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
assert venue == "dzengi"
assert checkpoint is self._checkpoint
assert limit > 0
self._events.append("checkpoint.tail")
return self._tail
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
raise AssertionError("Checkpoint уже существует.")
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
raise AssertionError("Checkpoint уже существует.")
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
raise AssertionError("Startup Hydrator не выполняет запись.")
class RecordingTradeObservationSink:
"""Фиксирует точный порядок durable acceptance в тесте."""
def __init__(self) -> None:
self._lock = threading.Lock()
self._accepted_trade_ids: list[int] = []
@property
def accepted_trade_ids(self) -> tuple[int, ...]:
with self._lock:
return tuple(self._accepted_trade_ids)
def persist_accepted(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
del expected_trade
with self._lock:
self._accepted_trade_ids.append(trade.trade_id)
def persist_duplicate(self, trade: Trade) -> None:
del trade
def make_checkpoint_trade(*, timestamp_ms: int) -> Trade:
return Trade(
symbol=SYMBOL,
trade_id=100,
price=Decimal("64555.55"),
quantity=Decimal("0.002"),
executed_at=datetime.fromtimestamp(
timestamp_ms / 1_000,
tz=timezone.utc,
),
aggressor_side=TradeAggressorSide.BUY,
source="postgres_trade_history",
)
def build_startup_recovery_runtime(
*,
settings: Settings,
checkpoint_storage: StaticCheckpointStorage,
observation_sink: RecordingTradeObservationSink,
recovery_end_time_ms: int,
) -> tuple[TradeStreamProductionRuntime, TradeStreamRuntimeComposition]:
"""Собрать persistent startup-граф без использования Bootstrap."""
transport = DzengiWebSocketTransport(
url=settings.trade_stream.websocket_url,
headers={
"Origin": settings.exchange_base_url.rstrip("/"),
"Content-Type": "application/json",
},
open_timeout=settings.trade_stream.open_timeout_seconds,
ping_interval=None,
ping_timeout=None,
probe_timeout=settings.trade_stream.probe_timeout_seconds,
close_timeout=settings.trade_stream.close_timeout_seconds,
)
session = WebSocketSession(transport)
subscription_manager = WebSocketSubscriptionManager(
transport,
supports_unsubscribe=False,
)
event_publisher = AcquisitionRuntimeEventPublisher()
composition = build_trade_stream_runtime_composition(
session=session,
transport=transport,
subscription_manager=subscription_manager,
event_publisher=event_publisher,
message_adapter=DzengiUnifiedWebSocketAdapter(),
recovery_document_source=DzengiTradesDocumentSource(
ExchangeRestClient(settings=settings),
),
symbols=(SYMBOL,),
heartbeat_timeout_seconds=60.0,
scheduler_interval_seconds=60.0,
trade_observation_sink=observation_sink,
max_recovery_window_ms=3_599_999,
recovery_end_time_clock=lambda: recovery_end_time_ms,
checkpoint_storage=checkpoint_storage,
checkpoint_venue="dzengi",
)
startup_recovery = (
composition.runtime_startup_recovery_coordinator
)
assert startup_recovery is not None
runtime = TradeStreamProductionRuntime(
session=session,
transport=transport,
subscription_manager=subscription_manager,
event_publisher=event_publisher,
trade_stream_service=(
composition.trade_stream_acquisition_service
),
message_classifier=DzengiWebSocketInboundMessageClassifier(),
control_message_handler=DzengiWebSocketControlMessageHandler(),
live_processing_gate=composition.live_processing_gate,
reconnect_recovery_coordinator=(
composition.runtime_reconnect_recovery_coordinator
),
runtime_supervisor=composition.runtime_supervisor,
runtime_scheduler=composition.runtime_scheduler,
symbols=(SYMBOL,),
startup_recovery_coordinator=startup_recovery,
subscription_ack_timeout_seconds=1.0,
startup_market_buffer_capacity=10,
)
return runtime, composition
def make_recovered_trade(
*,
trade_id: int,
@@ -277,6 +496,218 @@ def test_real_loopback_websocket_updates_shared_checkpoint() -> None:
run_scenario(scenario())
def test_startup_recovery_precedes_buffered_live_market() -> None:
async def scenario() -> None:
events: list[str] = []
recovery_release = threading.Event()
base_time_ms = time.time_ns() // 1_000_000 - 10_000
checkpoint_trade = make_checkpoint_trade(
timestamp_ms=base_time_ms,
)
checkpoint = PersistentTradeCheckpoint(
venue="dzengi",
trade=checkpoint_trade,
revision=7,
updated_at=datetime.now(timezone.utc),
)
checkpoint_storage = StaticCheckpointStorage(
checkpoint=checkpoint,
tail=(checkpoint_trade,),
events=events,
)
observation_sink = RecordingTradeObservationSink()
websocket = LoopbackTradeWebSocketServer(
events=events,
auto_ack=False,
)
rest = LoopbackTradeRestServer(
responses=(
LoopbackHttpResponse(
body=[
make_recovered_trade(
trade_id=101,
timestamp_ms=base_time_ms + 100,
),
],
release=recovery_release,
),
),
events=events,
)
async with LoopbackTradeEnvironment(
websocket=websocket,
rest=rest,
) as environment:
settings = make_settings(
websocket_url=environment.websocket_url,
rest_base_url=rest.base_url,
)
runtime, composition = build_startup_recovery_runtime(
settings=settings,
checkpoint_storage=checkpoint_storage,
observation_sink=observation_sink,
recovery_end_time_ms=base_time_ms + 150,
)
runtime_task = asyncio.create_task(
runtime.run(),
name="trade-stream-runtime",
)
try:
await websocket.wait_for_subscriptions(1)
assert runtime.state is (
TradeStreamProductionRuntimeState.STARTING
)
assert composition.state_store.get(
SYMBOL,
).last_trade_id == 100
assert events.index("checkpoint.tail") < events.index(
"ws.connect:0"
)
await websocket.send_trade(
0,
symbol=SYMBOL,
trade_id=102,
timestamp_ms=base_time_ms + 200,
)
subscription = websocket.subscriptions[0]
await websocket.send_ack(
0,
correlation_id=subscription.correlation_id,
)
await rest.wait_for_requests(1)
assert runtime.state is (
TradeStreamProductionRuntimeState.STARTING
)
assert composition.state_store.get(
SYMBOL,
).last_trade_id == 100
assert observation_sink.accepted_trade_ids == ()
recovery_release.set()
await wait_until(
lambda: runtime.state
is TradeStreamProductionRuntimeState.RUNNING,
)
await wait_until(
lambda: composition.state_store.get(
SYMBOL,
).last_trade_id
== 102,
)
assert observation_sink.accepted_trade_ids == (
101,
102,
)
assert events.index("ws.subscribe:0") < events.index(
"rest.request"
)
assert runtime_task.done() is False
finally:
recovery_release.set()
await stop_runtime(runtime, runtime_task)
assert websocket.active_handler_count == 0
assert rest.thread_is_alive is False
await assert_no_owned_tasks()
run_scenario(scenario())
def test_startup_stop_waits_for_blocked_recovery_worker() -> None:
async def scenario() -> None:
events: list[str] = []
recovery_release = threading.Event()
base_time_ms = time.time_ns() // 1_000_000 - 10_000
checkpoint_trade = make_checkpoint_trade(
timestamp_ms=base_time_ms,
)
checkpoint = PersistentTradeCheckpoint(
venue="dzengi",
trade=checkpoint_trade,
revision=3,
updated_at=datetime.now(timezone.utc),
)
checkpoint_storage = StaticCheckpointStorage(
checkpoint=checkpoint,
tail=(checkpoint_trade,),
events=events,
)
observation_sink = RecordingTradeObservationSink()
websocket = LoopbackTradeWebSocketServer(events=events)
rest = LoopbackTradeRestServer(
responses=(
LoopbackHttpResponse(
body=[],
release=recovery_release,
),
),
events=events,
)
async with LoopbackTradeEnvironment(
websocket=websocket,
rest=rest,
) as environment:
settings = make_settings(
websocket_url=environment.websocket_url,
rest_base_url=rest.base_url,
)
runtime, composition = build_startup_recovery_runtime(
settings=settings,
checkpoint_storage=checkpoint_storage,
observation_sink=observation_sink,
recovery_end_time_ms=base_time_ms + 150,
)
runtime_task = asyncio.create_task(
runtime.run(),
name="trade-stream-runtime",
)
stop_task: asyncio.Task[None] | None = None
try:
await websocket.wait_for_subscriptions(1)
await rest.wait_for_requests(1)
stop_task = asyncio.create_task(runtime.stop())
await asyncio.sleep(0)
await asyncio.sleep(0)
assert stop_task.done() is False
assert runtime_task.done() is False
assert composition.live_processing_gate.locked is True
recovery_release.set()
await asyncio.wait_for(stop_task, timeout=3.0)
await asyncio.wait_for(runtime_task, timeout=3.0)
assert runtime.state is (
TradeStreamProductionRuntimeState.STOPPED
)
assert composition.live_processing_gate.locked is False
assert composition.live_processing_gate.failed is False
assert observation_sink.accepted_trade_ids == ()
finally:
recovery_release.set()
if stop_task is not None and not stop_task.done():
await stop_task
if not runtime_task.done():
await stop_runtime(runtime, runtime_task)
assert websocket.active_handler_count == 0
assert rest.thread_is_alive is False
await assert_no_owned_tasks()
run_scenario(scenario())
def test_reconnect_restores_then_recovers_before_buffered_live() -> None:
async def scenario() -> None:
events: list[str] = []

View File

@@ -14,9 +14,11 @@ from src.market_data.acquisition.models.trade import (
)
from src.market_data.storage import (
MARKET_DATA_PARTITION_ADVISORY_LOCK_ID,
MarketDataPartitionResult,
MarketDataPartitionType,
MarketDataRetentionPolicy,
MarketDataStorageOperationError,
PersistentTradeCheckpoint,
PostgresMarketDataPartitionManager,
PostgresMarketDataRetentionService,
PostgresQuoteRepository,
@@ -27,6 +29,7 @@ from tests.support.postgres_market_data import (
PostgresTestSettings,
connect_postgres_test_database,
wait_for_postgres_advisory_lock_waiters,
wait_for_postgres_relation_lock_waiters,
)
@@ -34,6 +37,7 @@ pytestmark = pytest.mark.integration
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
PARTITION_WRITER_BARRIER_LOCK_ID = 0x445A5041525457
def _trade(*, trade_id: int, executed_at: datetime) -> Trade:
@@ -60,6 +64,91 @@ def _quote(*, received_at: datetime) -> Quote:
)
def _insert_checkpoint(
pool: PostgresConnectionPool,
*,
trade: Trade,
revision: int = 1,
) -> None:
with pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
INSERT INTO market_data.trade_stream_checkpoints (
venue,
symbol,
trade_id,
executed_at,
revision
)
VALUES (%s, %s, %s, %s, %s)
""",
(
VENUE,
trade.symbol,
trade.trade_id,
trade.executed_at,
revision,
),
)
def _advance_checkpoint(
pool: PostgresConnectionPool,
*,
trade: Trade,
revision: int,
) -> None:
with pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
UPDATE market_data.trade_stream_checkpoints
SET trade_id = %s,
executed_at = %s,
revision = %s,
updated_at = NOW()
WHERE venue = %s
AND symbol = %s
""",
(
trade.trade_id,
trade.executed_at,
revision,
VENUE,
trade.symbol,
),
)
assert cursor.rowcount == 1
def _assert_checkpoint_foreign_key_is_deferred_no_action(
pool: PostgresConnectionPool,
) -> None:
with pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT constraint_row.confupdtype::text,
constraint_row.confdeltype::text,
constraint_row.condeferrable,
constraint_row.condeferred
FROM pg_catalog.pg_constraint AS constraint_row
JOIN pg_catalog.pg_class AS table_row
ON table_row.oid = constraint_row.conrelid
JOIN pg_catalog.pg_namespace AS namespace
ON namespace.oid = table_row.relnamespace
WHERE namespace.nspname = 'market_data'
AND table_row.relname = 'trade_stream_checkpoints'
AND constraint_row.conname =
'trade_stream_checkpoints_trade_fk'
"""
)
row = cursor.fetchone()
assert row == ("a", "a", True, True)
def test_real_partition_creation_moves_default_row_and_is_idempotent(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
@@ -99,6 +188,56 @@ def test_real_partition_creation_moves_default_row_and_is_idempotent(
assert relation == ("market_data.trades_2026_07",)
def test_checkpoint_trade_moves_from_default_to_month_partition(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
event_time = datetime(2026, 7, 15, tzinfo=timezone.utc)
trade = _trade(trade_id=1, executed_at=event_time)
repository = PostgresTradeRepository(
connection_provider=migrated_postgres_pool.connection,
)
manager = PostgresMarketDataPartitionManager(
connection_provider=migrated_postgres_pool.connection,
)
repository.store_trade(
venue=VENUE,
trade=trade,
observed_at=event_time,
)
_insert_checkpoint(
migrated_postgres_pool,
trade=trade,
)
result = manager.ensure_month_partition(
data_type=MarketDataPartitionType.TRADES,
month=event_time,
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT trade_row.tableoid::regclass::text,
checkpoint.revision
FROM market_data.trade_stream_checkpoints AS checkpoint
JOIN market_data.trades AS trade_row
ON trade_row.venue = checkpoint.venue
AND trade_row.symbol = checkpoint.symbol
AND trade_row.trade_id = checkpoint.trade_id
AND trade_row.executed_at = checkpoint.executed_at
"""
)
restored = cursor.fetchone()
assert result.created is True
assert result.moved_row_count == 1
assert restored == ("market_data.trades_2026_07", 1)
_assert_checkpoint_foreign_key_is_deferred_no_action(
migrated_postgres_pool
)
def test_two_real_partition_callers_create_one_partition(
migrated_postgres_pool: PostgresConnectionPool,
postgres_test_settings: PostgresTestSettings,
@@ -150,6 +289,164 @@ def test_two_real_partition_callers_create_one_partition(
assert sorted(results) == [(False, 0), (True, 0)]
def test_partition_creation_and_checkpoint_writer_use_one_lock_order(
migrated_postgres_pool: PostgresConnectionPool,
postgres_test_settings: PostgresTestSettings,
) -> None:
existing_month = datetime(2026, 7, 1, tzinfo=timezone.utc)
new_month = datetime(2026, 8, 1, tzinfo=timezone.utc)
previous_trade = _trade(
trade_id=1,
executed_at=datetime(2026, 7, 14, tzinfo=timezone.utc),
)
candidate_trade = _trade(
trade_id=2,
executed_at=datetime(2026, 7, 15, tzinfo=timezone.utc),
)
repository = PostgresTradeRepository(
connection_provider=migrated_postgres_pool.connection,
)
manager = PostgresMarketDataPartitionManager(
connection_provider=migrated_postgres_pool.connection,
)
manager.ensure_month_partition(
data_type=MarketDataPartitionType.TRADES,
month=existing_month,
)
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=previous_trade,
observed_at=previous_trade.executed_at,
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
f"""
CREATE FUNCTION market_data.block_partition_writer()
RETURNS trigger
LANGUAGE plpgsql
AS $$
BEGIN
PERFORM pg_advisory_xact_lock(
{PARTITION_WRITER_BARRIER_LOCK_ID}
);
RETURN NEW;
END;
$$
"""
)
cursor.execute(
"""
CREATE TRIGGER block_partition_writer
BEFORE INSERT ON market_data.trades_2026_07
FOR EACH ROW
EXECUTE FUNCTION market_data.block_partition_writer()
"""
)
caller_ids: set[int] = set()
caller_ids_lock = threading.Lock()
def store_and_advance_checkpoint() -> PersistentTradeCheckpoint:
with caller_ids_lock:
caller_ids.add(threading.get_ident())
return repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=previous_trade,
trade=candidate_trade,
observed_at=candidate_trade.executed_at,
)
def create_partition() -> MarketDataPartitionResult:
with caller_ids_lock:
caller_ids.add(threading.get_ident())
return manager.ensure_month_partition(
data_type=MarketDataPartitionType.TRADES,
month=new_month,
)
with connect_postgres_test_database(
postgres_test_settings,
) as control:
with control.cursor() as cursor:
cursor.execute(
"SELECT pg_advisory_lock(%s)",
(PARTITION_WRITER_BARRIER_LOCK_ID,),
)
with ThreadPoolExecutor(max_workers=2) as executor:
try:
writer_future = executor.submit(
store_and_advance_checkpoint
)
wait_for_postgres_advisory_lock_waiters(
control,
lock_id=PARTITION_WRITER_BARRIER_LOCK_ID,
expected_count=1,
)
manager_future = executor.submit(create_partition)
wait_for_postgres_relation_lock_waiters(
control,
relation_name="market_data.trades",
expected_count=1,
)
finally:
with control.cursor() as cursor:
cursor.execute(
"SELECT pg_advisory_unlock(%s)",
(PARTITION_WRITER_BARRIER_LOCK_ID,),
)
barrier_released = cursor.fetchone()
if barrier_released != (True,):
raise AssertionError("Writer barrier was not released.")
try:
checkpoint = writer_future.result(timeout=10.0)
partition_result = manager_future.result(timeout=10.0)
finally:
writer_future.cancel()
manager_future.cancel()
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT trade_row.tableoid::regclass::text,
checkpoint.trade_id,
checkpoint.revision
FROM market_data.trades AS trade_row
JOIN market_data.trade_stream_checkpoints AS checkpoint
ON checkpoint.venue = trade_row.venue
AND checkpoint.symbol = trade_row.symbol
AND checkpoint.trade_id = trade_row.trade_id
AND checkpoint.executed_at = trade_row.executed_at
WHERE trade_row.venue = %s
AND trade_row.symbol = %s
AND trade_row.trade_id = %s
""",
(
VENUE,
SYMBOL,
candidate_trade.trade_id,
),
)
persisted = cursor.fetchone()
assert len(caller_ids) == 2
assert checkpoint.trade == candidate_trade
assert checkpoint.revision == 2
assert partition_result.created is True
assert partition_result.moved_row_count == 0
assert persisted == ("market_data.trades_2026_07", 2, 2)
_assert_checkpoint_foreign_key_is_deferred_no_action(
migrated_postgres_pool
)
def test_real_retention_uses_exact_cutoff(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
@@ -189,6 +486,184 @@ def test_real_retention_uses_exact_cutoff(
assert remaining == (2, 3)
def test_retention_rolls_back_for_active_checkpoint_in_default(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
now = datetime(2026, 8, 15, 12, 0, tzinfo=timezone.utc)
old_time = now - timedelta(days=30)
trade = _trade(trade_id=1, executed_at=old_time)
repository = PostgresTradeRepository(
connection_provider=migrated_postgres_pool.connection,
)
service = PostgresMarketDataRetentionService(
connection_provider=migrated_postgres_pool.connection,
)
repository.store_trade(
venue=VENUE,
trade=trade,
observed_at=now,
)
_insert_checkpoint(
migrated_postgres_pool,
trade=trade,
)
with pytest.raises(MarketDataStorageOperationError):
service.apply(
policy=MarketDataRetentionPolicy(
enabled=True,
trade_days=10,
),
now=now,
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute("SELECT COUNT(*) FROM market_data.trades")
trade_count = cursor.fetchone()
cursor.execute(
"SELECT COUNT(*) "
"FROM market_data.trade_stream_checkpoints"
)
checkpoint_count = cursor.fetchone()
assert trade_count == (1,)
assert checkpoint_count == (1,)
_assert_checkpoint_foreign_key_is_deferred_no_action(
migrated_postgres_pool
)
def test_retention_rolls_back_for_active_checkpoint_in_partition(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
now = datetime(2026, 8, 15, 12, 0, tzinfo=timezone.utc)
old_time = datetime(2026, 6, 15, tzinfo=timezone.utc)
trade = _trade(trade_id=1, executed_at=old_time)
repository = PostgresTradeRepository(
connection_provider=migrated_postgres_pool.connection,
)
manager = PostgresMarketDataPartitionManager(
connection_provider=migrated_postgres_pool.connection,
)
service = PostgresMarketDataRetentionService(
connection_provider=migrated_postgres_pool.connection,
)
manager.ensure_month_partition(
data_type=MarketDataPartitionType.TRADES,
month=old_time,
)
repository.store_trade(
venue=VENUE,
trade=trade,
observed_at=now,
)
_insert_checkpoint(
migrated_postgres_pool,
trade=trade,
)
with pytest.raises(MarketDataStorageOperationError):
service.apply(
policy=MarketDataRetentionPolicy(
enabled=True,
trade_days=10,
),
now=now,
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute("SELECT COUNT(*) FROM market_data.trades")
trade_count = cursor.fetchone()
cursor.execute(
"""
SELECT COUNT(*)
FROM market_data.partition_registry
WHERE partition_name = 'trades_2026_06'
"""
)
registry_count = cursor.fetchone()
assert trade_count == (1,)
assert registry_count == (1,)
def test_retention_drops_old_partition_after_checkpoint_advances(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
now = datetime(2026, 8, 15, 12, 0, tzinfo=timezone.utc)
old_time = datetime(2026, 1, 15, tzinfo=timezone.utc)
current_time = datetime(2026, 7, 15, tzinfo=timezone.utc)
old_trade = _trade(trade_id=1, executed_at=old_time)
current_trade = _trade(trade_id=2, executed_at=current_time)
repository = PostgresTradeRepository(
connection_provider=migrated_postgres_pool.connection,
)
manager = PostgresMarketDataPartitionManager(
connection_provider=migrated_postgres_pool.connection,
)
service = PostgresMarketDataRetentionService(
connection_provider=migrated_postgres_pool.connection,
)
manager.ensure_month_partition(
data_type=MarketDataPartitionType.TRADES,
month=old_time,
)
repository.store_trade(
venue=VENUE,
trade=old_trade,
observed_at=now,
)
repository.store_trade(
venue=VENUE,
trade=current_trade,
observed_at=now,
)
_insert_checkpoint(
migrated_postgres_pool,
trade=old_trade,
)
_advance_checkpoint(
migrated_postgres_pool,
trade=current_trade,
revision=2,
)
result = service.apply(
policy=MarketDataRetentionPolicy(
enabled=True,
trade_days=90,
),
now=now,
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"SELECT trade_id FROM market_data.trades ORDER BY trade_id"
)
remaining_trade_ids = tuple(
row[0] for row in cursor.fetchall()
)
cursor.execute(
"""
SELECT trade_id, revision
FROM market_data.trade_stream_checkpoints
"""
)
checkpoint_row = cursor.fetchone()
trade_result = result.entries[0]
assert trade_result.dropped_partitions == ("trades_2026_01",)
assert trade_result.dropped_row_count == 1
assert remaining_trade_ids == (2,)
assert checkpoint_row == (2, 2)
_assert_checkpoint_foreign_key_is_deferred_no_action(
migrated_postgres_pool
)
def test_real_retention_failure_rolls_back_all_data_types(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:

View File

@@ -1,9 +1,11 @@
from __future__ import annotations
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime, timezone
import threading
import pytest
from psycopg.errors import ForeignKeyViolation
from src.storage.exceptions import StorageMigrationError
from src.storage.migrations import (
@@ -104,6 +106,51 @@ def test_real_migrations_create_expected_schema_and_are_idempotent(
partitioned_tables = tuple(
row[0] for row in cursor.fetchall()
)
cursor.execute(
"""
SELECT column_name, data_type, is_nullable
FROM information_schema.columns
WHERE table_schema = 'market_data'
AND table_name = 'trade_stream_checkpoints'
ORDER BY ordinal_position
"""
)
checkpoint_columns = tuple(cursor.fetchall())
cursor.execute(
"""
SELECT pg_get_constraintdef(constraint_row.oid)
FROM pg_catalog.pg_constraint AS constraint_row
JOIN pg_catalog.pg_class AS table_row
ON table_row.oid = constraint_row.conrelid
JOIN pg_catalog.pg_namespace AS namespace
ON namespace.oid = table_row.relnamespace
WHERE namespace.nspname = 'market_data'
AND table_row.relname = 'trade_stream_checkpoints'
ORDER BY constraint_row.contype,
constraint_row.conname
"""
)
checkpoint_constraints = tuple(
row[0] for row in cursor.fetchall()
)
cursor.execute(
"""
SELECT constraint_row.confupdtype::text,
constraint_row.confdeltype::text,
constraint_row.condeferrable,
constraint_row.condeferred
FROM pg_catalog.pg_constraint AS constraint_row
JOIN pg_catalog.pg_class AS table_row
ON table_row.oid = constraint_row.conrelid
JOIN pg_catalog.pg_namespace AS namespace
ON namespace.oid = table_row.relnamespace
WHERE namespace.nspname = 'market_data'
AND table_row.relname = 'trade_stream_checkpoints'
AND constraint_row.conname =
'trade_stream_checkpoints_trade_fk'
"""
)
checkpoint_foreign_key = cursor.fetchone()
finally:
pool.close()
@@ -121,11 +168,140 @@ def test_real_migrations_create_expected_schema_and_are_idempotent(
"quotes",
"trades",
)
assert checkpoint_columns == (
("venue", "text", "NO"),
("symbol", "text", "NO"),
("trade_id", "integer", "NO"),
("executed_at", "timestamp with time zone", "NO"),
("revision", "bigint", "NO"),
("updated_at", "timestamp with time zone", "NO"),
("checkpoint_schema_version", "integer", "NO"),
)
assert any(
definition
== "PRIMARY KEY (venue, symbol)"
for definition in checkpoint_constraints
)
assert any(
definition.startswith(
"FOREIGN KEY (venue, symbol, trade_id, executed_at) "
"REFERENCES market_data.trades"
)
and "DEFERRABLE INITIALLY DEFERRED" in definition
for definition in checkpoint_constraints
)
assert checkpoint_foreign_key == ("a", "a", True, True)
with connect_postgres_test_database(postgres_test_settings) as control:
assert count_other_test_connections(control) == 0
def test_checkpoint_foreign_key_rejects_orphan_and_trade_deletion(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
executed_at = datetime(2026, 8, 1, 10, 0, tzinfo=timezone.utc)
with pytest.raises(ForeignKeyViolation):
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
INSERT INTO market_data.trade_stream_checkpoints (
venue,
symbol,
trade_id,
executed_at,
revision
)
VALUES (%s, %s, %s, %s, %s)
""",
(
"DZENGI",
"BTC/USD_LEVERAGE",
123,
executed_at,
1,
),
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
INSERT INTO market_data.trades (
venue,
symbol,
trade_id,
executed_at,
price,
quantity,
aggressor_side,
source,
first_observed_at,
last_observed_at,
observation_sources,
canonical_schema_version
)
VALUES (
%s, %s, %s, %s, %s, %s,
%s, %s, %s, %s, %s, %s
)
""",
(
"DZENGI",
"BTC/USD_LEVERAGE",
123,
executed_at,
"65000.25",
"0.001",
"buy",
"dzengi_websocket_trade",
executed_at,
executed_at,
["dzengi_websocket_trade"],
1,
),
)
cursor.execute(
"""
INSERT INTO market_data.trade_stream_checkpoints (
venue,
symbol,
trade_id,
executed_at,
revision
)
VALUES (%s, %s, %s, %s, %s)
""",
(
"DZENGI",
"BTC/USD_LEVERAGE",
123,
executed_at,
1,
),
)
with pytest.raises(ForeignKeyViolation):
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
DELETE FROM market_data.trades
WHERE venue = %s
AND symbol = %s
AND trade_id = %s
AND executed_at = %s
""",
(
"DZENGI",
"BTC/USD_LEVERAGE",
123,
executed_at,
),
)
def test_two_real_migration_runners_apply_each_version_once(
postgres_test_settings: PostgresTestSettings,
) -> None:

View File

@@ -0,0 +1,707 @@
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.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 (
MarketDataCheckpointConflictError,
MarketDataCheckpointIntegrityError,
MarketDataStorageOperationError,
PostgresTradeRepository,
)
from src.storage.postgres_pool import PostgresConnectionPool
pytestmark = pytest.mark.integration
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
EVENT_TIME = datetime(2026, 8, 1, 10, 0, tzinfo=timezone.utc)
OBSERVED_AT = EVENT_TIME + timedelta(seconds=1)
def _trade(
*,
trade_id: int,
offset_ms: int = 0,
source: str = "dzengi_websocket_trade",
) -> Trade:
return Trade(
symbol=SYMBOL,
trade_id=trade_id,
price=Decimal("65000.25"),
quantity=Decimal("0.001"),
executed_at=EVENT_TIME + timedelta(milliseconds=offset_ms),
aggressor_side=TradeAggressorSide.BUY,
source=source,
)
def _repository(
pool: PostgresConnectionPool,
) -> PostgresTradeRepository:
return PostgresTradeRepository(
connection_provider=pool.connection,
)
def _stored_trade_metadata(
pool: PostgresConnectionPool,
*,
trade: Trade,
) -> tuple[object, ...] | None:
with pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT
source,
first_observed_at,
last_observed_at,
observation_sources
FROM market_data.trades
WHERE venue = %s
AND symbol = %s
AND trade_id = %s
AND executed_at = %s
""",
(
VENUE,
trade.symbol,
trade.trade_id,
trade.executed_at,
),
)
return cursor.fetchone()
def test_real_checkpoint_commit_load_and_rollover_tail(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=SIGNED_TRADE_ID_MAX)
second = _trade(
trade_id=SIGNED_TRADE_ID_MIN,
offset_ms=1,
)
third = _trade(
trade_id=SIGNED_TRADE_ID_MIN + 1,
offset_ms=2,
)
first_checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=first,
observed_at=OBSERVED_AT,
)
second_checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=first,
trade=second,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
third_checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=second,
trade=third,
observed_at=OBSERVED_AT + timedelta(seconds=2),
)
loaded = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
tail = repository.load_checkpoint_tail(
venue=VENUE,
checkpoint=third_checkpoint,
limit=3,
)
assert first_checkpoint.revision == 1
assert second_checkpoint.revision == 2
assert third_checkpoint.revision == 3
assert loaded == third_checkpoint
assert tail == (first, second, third)
def test_real_latest_tail_orders_minus_one_to_zero_without_checkpoint(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=-1)
second = _trade(trade_id=0, offset_ms=1)
repository.store_trade(
venue=VENUE,
trade=first,
observed_at=OBSERVED_AT,
)
repository.store_trade(
venue=VENUE,
trade=second,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
tail = repository.load_latest_trade_tail(
venue=VENUE,
symbol=SYMBOL,
limit=2,
)
assert tail == (first, second)
def test_real_latest_tail_uses_acceptance_order_not_exchange_time(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
first = _trade(
trade_id=SIGNED_TRADE_ID_MAX,
offset_ms=2,
)
second = _trade(
trade_id=SIGNED_TRADE_ID_MIN,
offset_ms=1,
)
repository.store_trade(
venue=VENUE,
trade=first,
observed_at=OBSERVED_AT,
)
repository.store_trade(
venue=VENUE,
trade=second,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
tail = repository.load_latest_trade_tail(
venue=VENUE,
symbol=SYMBOL,
limit=2,
)
assert first.executed_at > second.executed_at
assert tail == (first, second)
def test_real_checkpoint_tail_uses_only_latest_raw_id_cycle(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
previous_cycle = _trade(trade_id=100)
current_predecessor = _trade(trade_id=99, offset_ms=1)
current = _trade(trade_id=100, offset_ms=2)
repository.store_trade(
venue=VENUE,
trade=previous_cycle,
observed_at=OBSERVED_AT,
)
repository.store_trade(
venue=VENUE,
trade=current_predecessor,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=current,
observed_at=OBSERVED_AT + timedelta(seconds=2),
)
tail = repository.load_checkpoint_tail(
venue=VENUE,
checkpoint=checkpoint,
limit=2,
)
assert previous_cycle.trade_id == current.trade_id
assert previous_cycle.executed_at != current.executed_at
assert tail == (current_predecessor, current)
def test_real_adoption_preserves_trade_metadata_and_is_idempotent(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
trade = _trade(trade_id=100)
repository.store_trade(
venue=VENUE,
trade=trade,
observed_at=OBSERVED_AT,
)
before = _stored_trade_metadata(
migrated_postgres_pool,
trade=trade,
)
adopted = repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=trade,
)
repeated = repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=trade,
)
after = _stored_trade_metadata(
migrated_postgres_pool,
trade=trade,
)
assert adopted.revision == 1
assert repeated == adopted
assert before == after
def test_real_adoption_rejects_missing_and_conflicting_durable_trade(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
missing = _trade(trade_id=100)
with pytest.raises(MarketDataCheckpointIntegrityError):
repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=missing,
)
durable = _trade(trade_id=101, offset_ms=1)
repository.store_trade(
venue=VENUE,
trade=durable,
observed_at=OBSERVED_AT,
)
with pytest.raises(MarketDataCheckpointConflictError):
repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=replace(durable, price=Decimal("65000.26")),
)
assert repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
) is None
def test_two_real_identical_adoption_callers_converge(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
trade = _trade(trade_id=100)
repository.store_trade(
venue=VENUE,
trade=trade,
observed_at=OBSERVED_AT,
)
before = _stored_trade_metadata(
migrated_postgres_pool,
trade=trade,
)
start_barrier = threading.Barrier(2)
caller_ids: set[int] = set()
caller_ids_lock = threading.Lock()
def adopt() -> int:
with caller_ids_lock:
caller_ids.add(threading.get_ident())
start_barrier.wait(timeout=5.0)
checkpoint = repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=trade,
)
return checkpoint.revision
with ThreadPoolExecutor(max_workers=2) as executor:
revisions = tuple(executor.map(lambda _: adopt(), range(2)))
after = _stored_trade_metadata(
migrated_postgres_pool,
trade=trade,
)
assert len(caller_ids) == 2
assert revisions == (1, 1)
assert before == after
def test_two_real_different_adoption_callers_allow_one_winner(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
candidates = (
_trade(trade_id=100),
_trade(trade_id=101, offset_ms=1),
)
for index, candidate in enumerate(candidates):
repository.store_trade(
venue=VENUE,
trade=candidate,
observed_at=OBSERVED_AT + timedelta(seconds=index),
)
before = tuple(
_stored_trade_metadata(
migrated_postgres_pool,
trade=candidate,
)
for candidate in candidates
)
start_barrier = threading.Barrier(2)
caller_ids: set[int] = set()
caller_ids_lock = threading.Lock()
def adopt(candidate: Trade) -> tuple[str, int]:
with caller_ids_lock:
caller_ids.add(threading.get_ident())
start_barrier.wait(timeout=5.0)
try:
checkpoint = repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=candidate,
)
except MarketDataCheckpointConflictError:
return ("conflict", candidate.trade_id)
return ("committed", checkpoint.trade.trade_id)
with ThreadPoolExecutor(max_workers=2) as executor:
results = tuple(executor.map(adopt, candidates))
loaded = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
after = tuple(
_stored_trade_metadata(
migrated_postgres_pool,
trade=candidate,
)
for candidate in candidates
)
committed_ids = tuple(
trade_id
for status, trade_id in results
if status == "committed"
)
assert len(caller_ids) == 2
assert tuple(status for status, _ in results).count("committed") == 1
assert tuple(status for status, _ in results).count("conflict") == 1
assert loaded is not None
assert loaded.trade.trade_id == committed_ids[0]
assert loaded.revision == 1
assert before == after
def test_real_stale_writer_rolls_back_candidate_trade(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=100)
winner = _trade(trade_id=101, offset_ms=1)
stale_candidate = _trade(trade_id=102, offset_ms=2)
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=first,
observed_at=OBSERVED_AT,
)
winner_checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=first,
trade=winner,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
with pytest.raises(MarketDataCheckpointConflictError):
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=first,
trade=stale_candidate,
observed_at=OBSERVED_AT + timedelta(seconds=2),
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"SELECT trade_id FROM market_data.trades ORDER BY trade_id"
)
trade_ids = tuple(row[0] for row in cursor.fetchall())
loaded = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
assert trade_ids == (100, 101)
assert loaded == winner_checkpoint
def test_real_retry_after_commit_preserves_revision_and_updates_provenance(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=100)
current = _trade(trade_id=101, offset_ms=1)
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=first,
observed_at=OBSERVED_AT,
)
committed = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=first,
trade=current,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
repeated = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=first,
trade=replace(current, source="dzengi"),
observed_at=OBSERVED_AT + timedelta(seconds=2),
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT observation_sources
FROM market_data.trades
WHERE trade_id = 101
"""
)
sources = cursor.fetchone()
assert repeated == committed
assert repeated.revision == 2
assert sources is not None
assert set(sources[0]) == {"dzengi_websocket_trade", "dzengi"}
def test_two_real_first_writers_leave_only_winner_trade(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
candidates = (
_trade(trade_id=100),
_trade(trade_id=101, offset_ms=1),
)
start_barrier = threading.Barrier(2)
caller_ids: set[int] = set()
caller_ids_lock = threading.Lock()
def commit(trade: Trade) -> tuple[str, int]:
with caller_ids_lock:
caller_ids.add(threading.get_ident())
start_barrier.wait(timeout=5.0)
try:
checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=trade,
observed_at=OBSERVED_AT,
)
except MarketDataCheckpointConflictError:
return ("conflict", trade.trade_id)
return ("committed", checkpoint.trade.trade_id)
with ThreadPoolExecutor(max_workers=2) as executor:
results = tuple(executor.map(commit, candidates))
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute("SELECT trade_id FROM market_data.trades")
stored_trade_ids = tuple(row[0] for row in cursor.fetchall())
loaded = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
committed_ids = tuple(
trade_id
for status, trade_id in results
if status == "committed"
)
assert len(caller_ids) == 2
assert tuple(status for status, _ in results).count("committed") == 1
assert tuple(status for status, _ in results).count("conflict") == 1
assert stored_trade_ids == committed_ids
assert loaded is not None
assert loaded.trade.trade_id == committed_ids[0]
def test_two_real_identical_first_writers_converge_idempotently(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
candidates = (
_trade(trade_id=100),
_trade(
trade_id=100,
source="dzengi",
),
)
start_barrier = threading.Barrier(2)
caller_ids: set[int] = set()
caller_ids_lock = threading.Lock()
def commit(trade: Trade) -> int:
with caller_ids_lock:
caller_ids.add(threading.get_ident())
start_barrier.wait(timeout=5.0)
checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=trade,
observed_at=OBSERVED_AT,
)
return checkpoint.revision
with ThreadPoolExecutor(max_workers=2) as executor:
revisions = tuple(executor.map(commit, candidates))
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"SELECT observation_sources FROM market_data.trades"
)
rows = tuple(cursor.fetchall())
assert len(caller_ids) == 2
assert revisions == (1, 1)
assert len(rows) == 1
assert set(rows[0][0]) == {"dzengi_websocket_trade", "dzengi"}
def test_two_real_checkpoint_updates_allow_only_one_winner(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
previous = _trade(trade_id=100)
candidates = (
_trade(trade_id=101, offset_ms=1),
_trade(trade_id=102, offset_ms=2),
)
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=previous,
observed_at=OBSERVED_AT,
)
start_barrier = threading.Barrier(2)
caller_ids: set[int] = set()
caller_ids_lock = threading.Lock()
def commit(trade: Trade) -> tuple[str, int]:
with caller_ids_lock:
caller_ids.add(threading.get_ident())
start_barrier.wait(timeout=5.0)
try:
checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=previous,
trade=trade,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
except MarketDataCheckpointConflictError:
return ("conflict", trade.trade_id)
return ("committed", checkpoint.trade.trade_id)
with ThreadPoolExecutor(max_workers=2) as executor:
results = tuple(executor.map(commit, candidates))
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"SELECT trade_id FROM market_data.trades ORDER BY trade_id"
)
stored_trade_ids = tuple(row[0] for row in cursor.fetchall())
loaded = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
committed_ids = tuple(
trade_id
for status, trade_id in results
if status == "committed"
)
assert len(caller_ids) == 2
assert tuple(status for status, _ in results).count("committed") == 1
assert tuple(status for status, _ in results).count("conflict") == 1
assert stored_trade_ids == (100, committed_ids[0])
assert loaded is not None
assert loaded.trade.trade_id == committed_ids[0]
assert loaded.revision == 2
def test_real_checkpoint_failure_rolls_back_inserted_trade(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
trade = _trade(trade_id=100)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
CREATE FUNCTION market_data.reject_checkpoint_write()
RETURNS trigger
LANGUAGE plpgsql
AS $$
BEGIN
RAISE EXCEPTION 'checkpoint write rejected';
END;
$$
"""
)
cursor.execute(
"""
CREATE TRIGGER reject_checkpoint_write
BEFORE INSERT OR UPDATE
ON market_data.trade_stream_checkpoints
FOR EACH ROW
EXECUTE FUNCTION market_data.reject_checkpoint_write()
"""
)
with pytest.raises(MarketDataStorageOperationError):
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=trade,
observed_at=OBSERVED_AT,
)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute("SELECT COUNT(*) FROM market_data.trades")
trade_count = cursor.fetchone()
cursor.execute(
"SELECT COUNT(*) "
"FROM market_data.trade_stream_checkpoints"
)
checkpoint_count = cursor.fetchone()
assert trade_count == (0,)
assert checkpoint_count == (0,)

View File

@@ -0,0 +1,209 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from decimal import Decimal
import pytest
from src.market_data.acquisition.consistency.trade_stream_consistency_controller import (
TradeStreamConsistencyController,
)
from src.market_data.acquisition.consistency.trade_stream_state_store import (
TradeStreamStateStore,
)
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.storage import (
MarketDataCheckpointConflictError,
PostgresTradeRepository,
TradeStorageObservationSink,
)
from src.storage.postgres_pool import PostgresConnectionPool
pytestmark = pytest.mark.integration
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
EXECUTED_AT = datetime(2026, 8, 1, 12, 0, tzinfo=timezone.utc)
OBSERVED_AT = EXECUTED_AT + timedelta(seconds=1)
def _trade(
*,
trade_id: int,
offset_ms: int = 0,
source: str = "dzengi_websocket_trade",
) -> Trade:
return Trade(
symbol=SYMBOL,
trade_id=trade_id,
price=Decimal("65000.25"),
quantity=Decimal("0.001"),
executed_at=EXECUTED_AT + timedelta(milliseconds=offset_ms),
aggressor_side=TradeAggressorSide.BUY,
source=source,
)
def _repository(
pool: PostgresConnectionPool,
) -> PostgresTradeRepository:
return PostgresTradeRepository(
connection_provider=pool.connection,
)
def _controller(
pool: PostgresConnectionPool,
) -> tuple[
TradeStreamConsistencyController,
TradeStreamStateStore,
]:
repository = _repository(pool)
state_store = TradeStreamStateStore()
sink = TradeStorageObservationSink(
trade_storage=repository,
venue=VENUE,
clock=lambda: OBSERVED_AT,
)
return (
TradeStreamConsistencyController(
state_store=state_store,
trade_observation_sink=sink,
),
state_store,
)
def _stored_trade_ids(
pool: PostgresConnectionPool,
) -> tuple[int, ...]:
with pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT trade_id
FROM market_data.trades
WHERE venue = %s AND symbol = %s
ORDER BY trade_id
""",
(VENUE, SYMBOL),
)
return tuple(row[0] for row in cursor.fetchall())
def test_real_consistency_atomically_advances_durable_checkpoint(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
controller, state_store = _controller(migrated_postgres_pool)
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=100)
second = _trade(trade_id=101, offset_ms=1)
controller.accept(first)
result = controller.accept(second)
checkpoint = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
state = state_store.get(SYMBOL)
assert result is second
assert checkpoint is not None
assert checkpoint.trade == second
assert checkpoint.revision == 2
assert state.last_trade is second
assert state.last_trade_id == second.trade_id
assert _stored_trade_ids(migrated_postgres_pool) == (100, 101)
def test_real_duplicate_updates_provenance_without_checkpoint_advance(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
controller, state_store = _controller(migrated_postgres_pool)
repository = _repository(migrated_postgres_pool)
original = _trade(trade_id=200)
duplicate = _trade(
trade_id=200,
source="dzengi",
)
controller.accept(original)
before_duplicate = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
result = controller.accept(duplicate)
after_duplicate = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
state = state_store.get(SYMBOL)
with migrated_postgres_pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT observation_sources
FROM market_data.trades
WHERE venue = %s
AND symbol = %s
AND trade_id = %s
AND executed_at = %s
""",
(
VENUE,
SYMBOL,
original.trade_id,
original.executed_at,
),
)
provenance = cursor.fetchone()
assert result is None
assert before_duplicate is not None
assert after_duplicate is not None
assert after_duplicate == before_duplicate
assert after_duplicate.revision == 1
assert provenance == (
["dzengi_websocket_trade", "dzengi"],
)
assert state.last_trade is original
assert state.last_trade_id == original.trade_id
def test_real_checkpoint_conflict_rolls_back_trade_and_in_memory_state(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
controller, state_store = _controller(migrated_postgres_pool)
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=300)
stale_candidate = _trade(trade_id=301, offset_ms=1)
winner = _trade(trade_id=302, offset_ms=2)
controller.accept(first)
winner_checkpoint = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=first,
trade=winner,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
with pytest.raises(MarketDataCheckpointConflictError):
controller.accept(stale_candidate)
checkpoint = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
state = state_store.get(SYMBOL)
assert checkpoint == winner_checkpoint
assert _stored_trade_ids(migrated_postgres_pool) == (300, 302)
assert state.last_trade is first
assert state.last_trade_id == first.trade_id
assert stale_candidate.trade_id not in state._trades

View File

@@ -391,6 +391,7 @@ def test_application_owns_real_storage_before_and_after_runtime(
runtime = build_trade_stream_production_runtime(
settings,
trade_observation_sink=storage.trade_observation_sink,
checkpoint_storage=storage.trade_repository,
)
assert runtime is not None
dispatcher = ControlledDispatcher()

View File

@@ -0,0 +1,185 @@
from __future__ import annotations
import asyncio
import time
from datetime import datetime, timezone
from typing import Any
import pytest
from src.bootstrap.market_data_storage import build_market_data_storage
from src.market_data.acquisition.runtime.trade_stream_production_runtime import (
TradeStreamProductionRuntimeState,
)
from tests.integration.market_data.acquisition.runtime.loopback_trade_exchange import (
LoopbackTradeEnvironment,
LoopbackTradeRestServer,
LoopbackTradeWebSocketServer,
)
from tests.integration.market_data.storage.test_trade_stream_persistent_restart_integration import (
APPLICATION_TIMEOUT_SECONDS,
VENUE,
PersistentApplication,
_assert_no_pool_connections,
_build_application,
_canonical_trade,
_cleanup_application,
_settings_for_database,
)
from tests.support.postgres_market_data import (
PostgresTestSettings,
connect_postgres_test_database,
wait_for_postgres_relation_lock_waiters,
)
from tests.support.trade_stream_runtime import (
assert_no_owned_tasks,
run_scenario,
)
pytestmark = pytest.mark.integration
CHECKPOINT_RELATION_NAME = "market_data.trade_stream_checkpoints"
def test_cancellation_waits_for_blocked_hydration_before_pool_close(
postgres_test_settings: PostgresTestSettings,
) -> None:
async def scenario() -> None:
base_time_ms = time.time_ns() // 1_000_000 - 10_000
websocket = LoopbackTradeWebSocketServer(auto_ack=False)
rest = LoopbackTradeRestServer()
application: PersistentApplication | None = None
seed_storage = None
blocker = None
try:
async with LoopbackTradeEnvironment(
websocket=websocket,
rest=rest,
) as environment:
settings = _settings_for_database(
postgres=postgres_test_settings,
websocket_url=environment.websocket_url,
rest_base_url=rest.base_url,
)
seed_storage = build_market_data_storage(settings)
assert seed_storage is not None
await asyncio.to_thread(seed_storage.lifecycle.start)
await asyncio.to_thread(
seed_storage.trade_repository
.store_trade_and_advance_checkpoint,
venue=VENUE,
expected_trade=None,
trade=_canonical_trade(
trade_id=500,
timestamp_ms=base_time_ms,
),
observed_at=datetime.now(timezone.utc),
)
await asyncio.to_thread(seed_storage.lifecycle.stop)
assert seed_storage.connection_pool.is_open is False
await asyncio.to_thread(
_assert_no_pool_connections,
postgres_test_settings,
)
blocker = connect_postgres_test_database(
postgres_test_settings,
autocommit=False,
)
with blocker.cursor() as cursor:
cursor.execute(
"LOCK TABLE "
"market_data.trade_stream_checkpoints "
"IN ACCESS EXCLUSIVE MODE"
)
application = _build_application(
settings,
task_name="persistent-application-cancellation",
)
await application.dispatcher.started.wait()
await asyncio.to_thread(
wait_for_postgres_relation_lock_waiters,
blocker,
relation_name=CHECKPOINT_RELATION_NAME,
expected_count=1,
)
assert application.runtime.state is (
TradeStreamProductionRuntimeState.STARTING
)
assert application.storage.connection_pool.is_open is True
assert websocket.connection_count == 0
assert rest.request_count == 0
application.task.cancel()
await asyncio.wait_for(
application.dispatcher.cancelled.wait(),
timeout=APPLICATION_TIMEOUT_SECONDS,
)
assert application.task.done() is False
assert application.task.cancelling() == 1
assert application.storage.connection_pool.is_open is True
application.task.cancel()
await asyncio.sleep(0)
assert application.task.done() is False
assert application.task.cancelling() == 2
assert application.storage.connection_pool.is_open is True
await asyncio.to_thread(blocker.rollback)
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(
asyncio.shield(application.task),
timeout=APPLICATION_TIMEOUT_SECONDS,
)
runtime_graph: Any = application.runtime
assert application.task.cancelled() is True
assert application.runtime.state is (
TradeStreamProductionRuntimeState.STOPPED
)
assert runtime_graph._live_processing_gate.failed is False
assert runtime_graph._live_processing_gate.locked is False
assert application.dispatcher.cancelled.is_set()
assert application.storage.connection_pool.is_open is False
assert application.storage.lifecycle.started is False
assert application.bot_session.close_calls == 1
finally:
if blocker is not None:
try:
if not blocker.closed:
await asyncio.to_thread(blocker.rollback)
finally:
await asyncio.to_thread(blocker.close)
await _cleanup_application(application)
if (
seed_storage is not None
and seed_storage.connection_pool.is_open
):
await asyncio.to_thread(seed_storage.lifecycle.stop)
assert websocket.active_handler_count == 0
assert rest.thread_is_alive is False
await asyncio.to_thread(
_assert_no_pool_connections,
postgres_test_settings,
)
await assert_no_owned_tasks()
run_scenario(
scenario(),
timeout_seconds=60.0,
)

View File

@@ -0,0 +1,357 @@
from __future__ import annotations
import asyncio
import threading
import time
from datetime import datetime, timezone
from typing import Any
import pytest
from src.bootstrap.market_data_storage import (
MarketDataStorageBootstrapComposition,
build_market_data_storage,
)
from src.core.config import Settings
from src.integrations.exchange.exceptions import ExchangeResponseError
from src.market_data.acquisition.exceptions import TradeTransportError
from src.market_data.acquisition.runtime.trade_stream_production_runtime import (
TradeStreamProductionRuntimeState,
)
from src.market_data.storage import MarketDataCheckpointIntegrityError
from tests.integration.market_data.acquisition.runtime.loopback_trade_exchange import (
LoopbackHttpResponse,
LoopbackTradeEnvironment,
LoopbackTradeRestServer,
LoopbackTradeWebSocketServer,
)
from tests.integration.market_data.storage import (
test_trade_stream_persistent_restart_integration as restart_harness,
)
from tests.support.postgres_market_data import (
PostgresTestSettings,
connect_postgres_test_database,
)
from tests.support.trade_stream_runtime import (
SYMBOL,
assert_no_owned_tasks,
run_scenario,
state_store_from,
)
pytestmark = pytest.mark.integration
VENUE = "dzengi"
CHECKPOINT_FOREIGN_KEY = "trade_stream_checkpoints_trade_fk"
def _prepare_storage(
settings: Settings,
) -> MarketDataStorageBootstrapComposition:
storage = build_market_data_storage(settings)
assert storage is not None
storage.lifecycle.start()
return storage
def _stop_storage(
storage: MarketDataStorageBootstrapComposition | None,
) -> None:
if storage is not None and storage.connection_pool.is_open:
storage.lifecycle.stop()
def _insert_orphan_checkpoint(
postgres: PostgresTestSettings,
*,
trade_id: int,
executed_at: datetime,
) -> None:
"""Создать повреждённую точку только в одноразовой тестовой БД."""
with connect_postgres_test_database(postgres) as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT COUNT(*)
FROM pg_catalog.pg_constraint
WHERE conname = %s
AND conrelid = (
'market_data.trade_stream_checkpoints'::regclass
)
""",
(CHECKPOINT_FOREIGN_KEY,),
)
assert cursor.fetchone() == (1,)
cursor.execute(
"""
ALTER TABLE market_data.trade_stream_checkpoints
DROP CONSTRAINT trade_stream_checkpoints_trade_fk
"""
)
cursor.execute(
"""
INSERT INTO market_data.trade_stream_checkpoints (
venue,
symbol,
trade_id,
executed_at,
revision,
checkpoint_schema_version
)
VALUES (%s, %s, %s, %s, 1, 1)
""",
(
VENUE,
SYMBOL,
trade_id,
executed_at,
),
)
def _seed_checkpoint(
settings: Settings,
*,
trade_id: int,
timestamp_ms: int,
) -> None:
storage: MarketDataStorageBootstrapComposition | None = None
try:
storage = _prepare_storage(settings)
trade = restart_harness._canonical_trade(
trade_id=trade_id,
timestamp_ms=timestamp_ms,
)
storage.trade_repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=trade,
observed_at=datetime.now(timezone.utc),
)
finally:
_stop_storage(storage)
def test_orphan_checkpoint_fails_before_network_io(
postgres_test_settings: PostgresTestSettings,
) -> None:
async def scenario() -> None:
base_time_ms = time.time_ns() // 1_000_000 - 10_000
websocket = LoopbackTradeWebSocketServer(auto_ack=False)
rest = LoopbackTradeRestServer()
application: restart_harness.PersistentApplication | None = None
setup_storage: MarketDataStorageBootstrapComposition | None = None
try:
async with LoopbackTradeEnvironment(
websocket=websocket,
rest=rest,
) as environment:
settings = restart_harness._settings_for_database(
postgres=postgres_test_settings,
websocket_url=environment.websocket_url,
rest_base_url=rest.base_url,
)
setup_storage = await asyncio.to_thread(
_prepare_storage,
settings,
)
await asyncio.to_thread(
setup_storage.lifecycle.stop,
)
await asyncio.to_thread(
_insert_orphan_checkpoint,
postgres_test_settings,
trade_id=300,
executed_at=datetime.fromtimestamp(
base_time_ms / 1_000,
tz=timezone.utc,
),
)
application = restart_harness._build_application(
settings,
task_name="persistent-application-orphan-checkpoint",
)
with pytest.raises(
MarketDataCheckpointIntegrityError,
) as captured:
await asyncio.wait_for(
asyncio.shield(application.task),
timeout=restart_harness.APPLICATION_TIMEOUT_SECONDS,
)
assert type(captured.value) is (
MarketDataCheckpointIntegrityError
)
runtime_graph: Any = application.runtime
gate = runtime_graph._live_processing_gate
assert application.runtime.state is (
TradeStreamProductionRuntimeState.FAILED
)
assert gate.failed is True
assert gate.locked is False
assert gate._failure is captured.value
assert runtime_graph._session.is_connected is False
assert (
runtime_graph._subscription_manager.subscription_keys
== ()
)
assert application.dispatcher.cancelled.is_set()
assert state_store_from(application.runtime).is_empty() is True
assert websocket.connection_count == 0
assert websocket.subscriptions == ()
assert rest.request_count == 0
assert application.storage.connection_pool.is_open is False
assert application.storage.lifecycle.started is False
assert application.bot_session.close_calls == 1
await asyncio.to_thread(
restart_harness._assert_no_pool_connections,
postgres_test_settings,
)
finally:
await restart_harness._cleanup_application(application)
await asyncio.to_thread(_stop_storage, setup_storage)
assert websocket.active_handler_count == 0
assert rest.thread_is_alive is False
await assert_no_owned_tasks()
run_scenario(
scenario(),
timeout_seconds=60.0,
)
def test_rest_failure_rejects_buffered_live_and_preserves_checkpoint(
postgres_test_settings: PostgresTestSettings,
) -> None:
async def scenario() -> None:
recovery_release = threading.Event()
base_time_ms = time.time_ns() // 1_000_000 - 10_000
websocket = LoopbackTradeWebSocketServer(auto_ack=False)
rest = LoopbackTradeRestServer(
responses=(
LoopbackHttpResponse(
body={"error": "recovery unavailable"},
status=500,
release=recovery_release,
),
)
)
application: restart_harness.PersistentApplication | None = None
try:
async with LoopbackTradeEnvironment(
websocket=websocket,
rest=rest,
) as environment:
settings = restart_harness._settings_for_database(
postgres=postgres_test_settings,
websocket_url=environment.websocket_url,
rest_base_url=rest.base_url,
)
await asyncio.to_thread(
_seed_checkpoint,
settings,
trade_id=400,
timestamp_ms=base_time_ms,
)
seeded_snapshot = await asyncio.to_thread(
restart_harness._snapshot_database,
postgres_test_settings,
)
assert tuple(row[0] for row in seeded_snapshot.trades) == (
400,
)
assert seeded_snapshot.checkpoint == (400, 1, 1)
application = restart_harness._build_application(
settings,
task_name="persistent-application-rest-failure",
)
await application.dispatcher.started.wait()
await websocket.wait_for_subscriptions(1)
state_store = state_store_from(application.runtime)
assert state_store.get(SYMBOL).last_trade_id == 400
await websocket.send_trade(
0,
symbol=SYMBOL,
trade_id=402,
timestamp_ms=base_time_ms + 200,
)
subscription = websocket.subscriptions[0]
await websocket.send_ack(
0,
correlation_id=subscription.correlation_id,
)
await rest.wait_for_requests(1)
runtime_graph: Any = application.runtime
gate = runtime_graph._live_processing_gate
assert gate.locked is True
assert gate.failed is False
assert application.runtime.state is (
TradeStreamProductionRuntimeState.STARTING
)
assert state_store.get(SYMBOL).last_trade_id == 400
assert await asyncio.to_thread(
restart_harness._snapshot_database,
postgres_test_settings,
) == seeded_snapshot
recovery_release.set()
with pytest.raises(TradeTransportError) as captured:
await asyncio.wait_for(
asyncio.shield(application.task),
timeout=restart_harness.APPLICATION_TIMEOUT_SECONDS,
)
assert type(captured.value) is TradeTransportError
assert type(captured.value.__cause__) is ExchangeResponseError
failed_snapshot = await asyncio.to_thread(
restart_harness._snapshot_database,
postgres_test_settings,
)
assert failed_snapshot == seeded_snapshot
assert state_store.get(SYMBOL).last_trade_id == 400
assert application.runtime.state is (
TradeStreamProductionRuntimeState.FAILED
)
assert gate.failed is True
assert gate.locked is False
assert gate._failure is captured.value
assert runtime_graph._session.is_connected is False
assert (
runtime_graph._subscription_manager.subscription_keys
== ()
)
assert application.dispatcher.cancelled.is_set()
assert application.storage.connection_pool.is_open is False
assert application.storage.lifecycle.started is False
assert application.bot_session.close_calls == 1
await asyncio.to_thread(
restart_harness._assert_no_pool_connections,
postgres_test_settings,
)
finally:
recovery_release.set()
await restart_harness._cleanup_application(application)
assert websocket.active_handler_count == 0
assert rest.thread_is_alive is False
await assert_no_owned_tasks()
run_scenario(
scenario(),
timeout_seconds=60.0,
)

View File

@@ -0,0 +1,480 @@
from __future__ import annotations
import asyncio
import time
from datetime import datetime, timezone
from typing import Any
import psycopg
import pytest
from src.bootstrap.market_data_storage import (
MarketDataStorageBootstrapComposition,
build_market_data_storage,
)
from src.core.config import Settings
from src.market_data.acquisition.consistency.trade_observation_sink_protocol import (
TradeObservationSinkProtocol,
)
from src.market_data.acquisition.models.trade import Trade
from src.market_data.acquisition.runtime.trade_stream_production_runtime import (
TradeStreamProductionRuntimeState,
)
from src.market_data.storage.exceptions import (
MarketDataStorageOperationError,
)
from tests.integration.market_data.acquisition.runtime.loopback_trade_exchange import (
LoopbackHttpResponse,
LoopbackTradeEnvironment,
LoopbackTradeRestServer,
LoopbackTradeWebSocketServer,
wait_until,
)
from tests.integration.market_data.storage import (
test_trade_stream_persistent_restart_integration as restart_harness,
)
from tests.support.postgres_market_data import (
PostgresTestSettings,
connect_postgres_test_database,
)
from tests.support.trade_stream_runtime import (
SYMBOL,
assert_no_owned_tasks,
run_scenario,
state_store_from,
)
pytestmark = pytest.mark.integration
VENUE = "dzengi"
REJECT_TRIGGER = "reject_runtime_checkpoint_write"
REJECT_FUNCTION = "market_data.reject_runtime_checkpoint_write"
class PostCommitTestError(RuntimeError):
"""Искусственная ошибка после подтверждённой фиксации в PostgreSQL."""
class CommitThenFailObservationSink:
"""Сохранить принятую сделку и имитировать потерю ответа."""
__slots__ = (
"_accepted_trades",
"_delegate",
"_expected_error",
)
def __init__(
self,
*,
delegate: TradeObservationSinkProtocol,
expected_error: PostCommitTestError,
) -> None:
self._delegate = delegate
self._expected_error = expected_error
self._accepted_trades: list[Trade] = []
@property
def accepted_trades(self) -> tuple[Trade, ...]:
return tuple(self._accepted_trades)
def persist_accepted(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
self._delegate.persist_accepted(
trade,
expected_trade=expected_trade,
)
self._accepted_trades.append(trade)
raise self._expected_error
def persist_duplicate(
self,
trade: Trade,
) -> None:
self._delegate.persist_duplicate(trade)
def _seed_checkpoint(
settings: Settings,
*,
trade_id: int,
timestamp_ms: int,
) -> None:
storage = build_market_data_storage(settings)
assert storage is not None
try:
storage.lifecycle.start()
storage.trade_repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=restart_harness._canonical_trade(
trade_id=trade_id,
timestamp_ms=timestamp_ms,
),
observed_at=datetime.now(timezone.utc),
)
finally:
if storage.connection_pool.is_open:
storage.lifecycle.stop()
def _install_checkpoint_rejection_trigger(
postgres: PostgresTestSettings,
) -> None:
with connect_postgres_test_database(postgres) as connection:
with connection.cursor() as cursor:
cursor.execute(
f"""
CREATE FUNCTION {REJECT_FUNCTION}()
RETURNS trigger
LANGUAGE plpgsql
AS $$
BEGIN
RAISE EXCEPTION 'runtime checkpoint write rejected';
END;
$$
"""
)
cursor.execute(
f"""
CREATE TRIGGER {REJECT_TRIGGER}
BEFORE INSERT OR UPDATE
ON market_data.trade_stream_checkpoints
FOR EACH ROW
EXECUTE FUNCTION {REJECT_FUNCTION}()
"""
)
def _remove_checkpoint_rejection_trigger(
postgres: PostgresTestSettings,
) -> None:
with connect_postgres_test_database(postgres) as connection:
with connection.cursor() as cursor:
cursor.execute(
f"""
DROP TRIGGER IF EXISTS {REJECT_TRIGGER}
ON market_data.trade_stream_checkpoints
"""
)
cursor.execute(
f"DROP FUNCTION IF EXISTS {REJECT_FUNCTION}()"
)
def test_checkpoint_write_failure_rolls_back_trade_and_runtime_state(
postgres_test_settings: PostgresTestSettings,
) -> None:
async def scenario() -> None:
base_time_ms = time.time_ns() // 1_000_000 - 10_000
websocket = LoopbackTradeWebSocketServer(auto_ack=False)
rest = LoopbackTradeRestServer()
application: restart_harness.PersistentApplication | None = None
trigger_installed = False
try:
async with LoopbackTradeEnvironment(
websocket=websocket,
rest=rest,
) as environment:
settings = restart_harness._settings_for_database(
postgres=postgres_test_settings,
websocket_url=environment.websocket_url,
rest_base_url=rest.base_url,
)
await asyncio.to_thread(
_seed_checkpoint,
settings,
trade_id=300,
timestamp_ms=base_time_ms,
)
seeded_snapshot = await asyncio.to_thread(
restart_harness._snapshot_database,
postgres_test_settings,
)
await asyncio.to_thread(
_install_checkpoint_rejection_trigger,
postgres_test_settings,
)
trigger_installed = True
application = restart_harness._build_application(
settings,
task_name="persistent-application-before-commit-failure",
)
await application.dispatcher.started.wait()
await websocket.wait_for_subscriptions(1)
subscription = websocket.subscriptions[0]
await websocket.send_ack(
0,
correlation_id=subscription.correlation_id,
)
await wait_until(
lambda: application.runtime.state
is TradeStreamProductionRuntimeState.RUNNING,
)
state_store = state_store_from(application.runtime)
state_before_failure = state_store.get(SYMBOL)
assert state_before_failure.last_trade_id == 300
await websocket.send_trade(
0,
symbol=SYMBOL,
trade_id=301,
timestamp_ms=base_time_ms + 100,
)
with pytest.raises(
MarketDataStorageOperationError,
) as captured:
await asyncio.wait_for(
asyncio.shield(application.task),
timeout=restart_harness.APPLICATION_TIMEOUT_SECONDS,
)
assert type(captured.value) is (
MarketDataStorageOperationError
)
database_error = captured.value.__cause__
assert type(database_error) is psycopg.errors.RaiseException
assert database_error.sqlstate == "P0001"
failed_snapshot = await asyncio.to_thread(
restart_harness._snapshot_database,
postgres_test_settings,
)
runtime_graph: Any = application.runtime
gate = runtime_graph._live_processing_gate
assert failed_snapshot == seeded_snapshot
assert state_store.get(SYMBOL) is state_before_failure
assert state_before_failure.last_trade_id == 300
assert application.runtime.state is (
TradeStreamProductionRuntimeState.FAILED
)
assert gate.failed is True
assert gate.locked is False
assert gate._failure is captured.value
assert runtime_graph._session.is_connected is False
assert (
runtime_graph._subscription_manager.subscription_keys
== ()
)
assert application.dispatcher.cancelled.is_set()
assert application.storage.connection_pool.is_open is False
assert application.storage.lifecycle.started is False
assert application.bot_session.close_calls == 1
await asyncio.to_thread(
restart_harness._assert_no_pool_connections,
postgres_test_settings,
)
finally:
await restart_harness._cleanup_application(application)
if trigger_installed:
await asyncio.to_thread(
_remove_checkpoint_rejection_trigger,
postgres_test_settings,
)
assert websocket.active_handler_count == 0
assert rest.thread_is_alive is False
await assert_no_owned_tasks()
run_scenario(
scenario(),
timeout_seconds=60.0,
)
def test_post_commit_error_is_recovered_by_fresh_runtime_graph(
postgres_test_settings: PostgresTestSettings,
) -> None:
async def scenario() -> None:
base_time_ms = time.time_ns() // 1_000_000 - 10_000
expected_error = PostCommitTestError(
"caller did not observe committed checkpoint"
)
websocket = LoopbackTradeWebSocketServer(auto_ack=False)
rest = LoopbackTradeRestServer(
responses=(
LoopbackHttpResponse(
body=[
restart_harness._rest_trade(
trade_id=600,
timestamp_ms=base_time_ms,
),
]
),
)
)
first: restart_harness.PersistentApplication | None = None
second: restart_harness.PersistentApplication | None = None
failing_sink: CommitThenFailObservationSink | None = None
try:
async with LoopbackTradeEnvironment(
websocket=websocket,
rest=rest,
) as environment:
settings = restart_harness._settings_for_database(
postgres=postgres_test_settings,
websocket_url=environment.websocket_url,
rest_base_url=rest.base_url,
)
def failing_sink_factory(
storage: MarketDataStorageBootstrapComposition,
) -> TradeObservationSinkProtocol:
nonlocal failing_sink
failing_sink = CommitThenFailObservationSink(
delegate=storage.trade_observation_sink,
expected_error=expected_error,
)
return failing_sink
first = restart_harness._build_application(
settings,
task_name="persistent-application-post-commit-failure",
trade_observation_sink_factory=failing_sink_factory,
)
await first.dispatcher.started.wait()
await websocket.wait_for_subscriptions(1)
first_subscription = websocket.subscriptions[0]
await websocket.send_ack(
0,
correlation_id=first_subscription.correlation_id,
)
await wait_until(
lambda: first.runtime.state
is TradeStreamProductionRuntimeState.RUNNING,
)
first_store = state_store_from(first.runtime)
first_state = first_store.get(SYMBOL)
assert first_state.last_trade_id is None
await websocket.send_trade(
0,
symbol=SYMBOL,
trade_id=600,
timestamp_ms=base_time_ms,
)
with pytest.raises(PostCommitTestError) as captured:
await asyncio.wait_for(
asyncio.shield(first.task),
timeout=restart_harness.APPLICATION_TIMEOUT_SECONDS,
)
assert captured.value is expected_error
assert failing_sink is not None
assert len(failing_sink.accepted_trades) == 1
committed_object = failing_sink.accepted_trades[0]
committed_snapshot = await asyncio.to_thread(
restart_harness._snapshot_database,
postgres_test_settings,
)
assert tuple(row[0] for row in committed_snapshot.trades) == (
600,
)
assert committed_snapshot.checkpoint == (600, 1, 1)
assert first_state.last_trade_id is None
assert first.runtime.state is (
TradeStreamProductionRuntimeState.FAILED
)
first_graph: Any = first.runtime
assert first_graph._live_processing_gate._failure is (
expected_error
)
assert first_graph._session.is_connected is False
assert (
first_graph._subscription_manager.subscription_keys
== ()
)
assert first.dispatcher.cancelled.is_set()
assert first.storage.connection_pool.is_open is False
await wait_until(
lambda: websocket.active_handler_count == 0,
)
await assert_no_owned_tasks()
await asyncio.to_thread(
restart_harness._assert_no_pool_connections,
postgres_test_settings,
)
second = restart_harness._build_application(
settings,
task_name="persistent-application-post-commit-restart",
)
assert second.storage is not first.storage
assert second.storage.trade_repository is not (
first.storage.trade_repository
)
assert second.runtime is not first.runtime
await second.dispatcher.started.wait()
await websocket.wait_for_subscriptions(2)
second_store = state_store_from(second.runtime)
second_state = second_store.get(SYMBOL)
hydrated_object = second_state.last_trade
assert second_state is not first_state
assert hydrated_object == committed_object
assert hydrated_object is not committed_object
assert second_state.last_trade_id == 600
assert tuple(second_state._trade_window) == (600,)
second_subscription = websocket.subscriptions[1]
await websocket.send_ack(
1,
correlation_id=second_subscription.correlation_id,
)
await rest.wait_for_requests(1)
await wait_until(
lambda: second.runtime.state
is TradeStreamProductionRuntimeState.RUNNING,
)
recovered_snapshot = await asyncio.to_thread(
restart_harness._snapshot_database,
postgres_test_settings,
)
assert tuple(row[0] for row in recovered_snapshot.trades) == (
600,
)
assert recovered_snapshot.checkpoint == (600, 1, 1)
assert recovered_snapshot.trades[0][2] == [
"dzengi_websocket_trade",
"dzengi",
]
assert second_state.last_trade_id == 600
assert second_state.last_trade is hydrated_object
await restart_harness._stop_application(second)
await asyncio.to_thread(
restart_harness._assert_no_pool_connections,
postgres_test_settings,
)
finally:
await restart_harness._cleanup_application(second)
await restart_harness._cleanup_application(first)
assert websocket.active_handler_count == 0
assert rest.thread_is_alive is False
await assert_no_owned_tasks()
run_scenario(
scenario(),
timeout_seconds=60.0,
)

View File

@@ -0,0 +1,283 @@
from __future__ import annotations
from dataclasses import replace
from datetime import datetime, timedelta, timezone
from decimal import Decimal
import pytest
from src.market_data.acquisition.checkpoint import (
TradeStreamStateHydrator,
)
from src.market_data.acquisition.consistency.trade_stream_consistency_controller import (
TradeStreamConsistencyController,
)
from src.market_data.acquisition.consistency.trade_stream_state_store import (
TradeStreamStateStore,
)
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 (
MarketDataCheckpointIntegrityError,
PostgresTradeRepository,
TradeStorageObservationSink,
)
from src.storage.postgres_pool import PostgresConnectionPool
pytestmark = pytest.mark.integration
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
SECOND_SYMBOL = "ETH/USD_LEVERAGE"
EVENT_TIME = datetime(2026, 8, 1, 14, 0, tzinfo=timezone.utc)
OBSERVED_AT = EVENT_TIME + timedelta(seconds=1)
def _trade(
*,
trade_id: int,
offset_ms: int = 0,
symbol: str = SYMBOL,
source: str = "dzengi_websocket_trade",
) -> Trade:
return Trade(
symbol=symbol,
trade_id=trade_id,
price=Decimal("65000.25"),
quantity=Decimal("0.001"),
executed_at=EVENT_TIME + timedelta(milliseconds=offset_ms),
aggressor_side=TradeAggressorSide.BUY,
source=source,
)
def _repository(
pool: PostgresConnectionPool,
) -> PostgresTradeRepository:
return PostgresTradeRepository(
connection_provider=pool.connection,
)
def _trade_metadata(
pool: PostgresConnectionPool,
*,
symbol: str,
) -> tuple[tuple[object, ...], ...]:
with pool.connection() as connection:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT
trade_id,
executed_at,
source,
first_observed_at,
last_observed_at,
observation_sources
FROM market_data.trades
WHERE venue = %s AND symbol = %s
ORDER BY first_observed_at, executed_at, trade_id
""",
(VENUE, symbol),
)
return tuple(cursor.fetchall())
def test_real_hydration_adopts_history_without_rewriting_trades(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=100)
second = _trade(trade_id=101, offset_ms=1)
repository.store_trade(
venue=VENUE,
trade=first,
observed_at=OBSERVED_AT,
)
repository.store_trade(
venue=VENUE,
trade=second,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
metadata_before = _trade_metadata(
migrated_postgres_pool,
symbol=SYMBOL,
)
first_store = TradeStreamStateStore()
first_states = TradeStreamStateHydrator(
checkpoint_storage=repository,
state_store=first_store,
venue=VENUE,
deduplication_window_size=2,
).hydrate(symbols=(SYMBOL,))
first_checkpoint = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
assert len(first_states) == 1
assert tuple(first_states[0]._trade_window) == (100, 101)
assert first_states[0].last_trade == second
assert first_checkpoint is not None
assert first_checkpoint.trade == second
assert first_checkpoint.revision == 1
assert _trade_metadata(
migrated_postgres_pool,
symbol=SYMBOL,
) == metadata_before
restart_store = TradeStreamStateStore()
restart_state = TradeStreamStateHydrator(
checkpoint_storage=repository,
state_store=restart_store,
venue=VENUE,
deduplication_window_size=2,
).hydrate(symbols=(SYMBOL,))[0]
repeated_checkpoint = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
assert tuple(restart_state._trade_window) == (100, 101)
assert restart_state.accept(
replace(first, source="dzengi_recovery"),
) is None
assert repeated_checkpoint is not None
assert repeated_checkpoint == first_checkpoint
assert repeated_checkpoint.revision == 1
assert _trade_metadata(
migrated_postgres_pool,
symbol=SYMBOL,
) == metadata_before
controller = TradeStreamConsistencyController(
state_store=restart_store,
trade_observation_sink=TradeStorageObservationSink(
trade_storage=repository,
venue=VENUE,
clock=lambda: OBSERVED_AT + timedelta(seconds=2),
),
)
third = _trade(trade_id=102, offset_ms=2)
assert controller.accept(third) is third
advanced_checkpoint = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
assert advanced_checkpoint is not None
assert advanced_checkpoint.trade == third
assert advanced_checkpoint.revision == 2
assert restart_state.last_trade is third
@pytest.mark.parametrize(
("first_trade_id", "second_trade_id", "next_trade_id"),
(
(SIGNED_TRADE_ID_MAX, SIGNED_TRADE_ID_MIN, SIGNED_TRADE_ID_MIN + 1),
(-1, 0, 1),
),
)
def test_real_hydration_restores_rollover_aware_deduplication_tail(
migrated_postgres_pool: PostgresConnectionPool,
first_trade_id: int,
second_trade_id: int,
next_trade_id: int,
) -> None:
repository = _repository(migrated_postgres_pool)
first = _trade(trade_id=first_trade_id)
second = _trade(trade_id=second_trade_id, offset_ms=1)
repository.store_trade(
venue=VENUE,
trade=first,
observed_at=OBSERVED_AT,
)
repository.store_trade(
venue=VENUE,
trade=second,
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
state_store = TradeStreamStateStore()
state = TradeStreamStateHydrator(
checkpoint_storage=repository,
state_store=state_store,
venue=VENUE,
deduplication_window_size=2,
).hydrate(symbols=(SYMBOL,))[0]
assert tuple(state._trade_window) == (
first_trade_id,
second_trade_id,
)
assert state.last_trade == second
assert state.accept(
replace(first, source="dzengi_recovery"),
) is None
next_trade = _trade(trade_id=next_trade_id, offset_ms=2)
assert state.accept(next_trade) is next_trade
assert state.last_trade is next_trade
def test_real_multi_symbol_failure_does_not_publish_partial_state(
migrated_postgres_pool: PostgresConnectionPool,
) -> None:
repository = _repository(migrated_postgres_pool)
repository.store_trade(
venue=VENUE,
trade=_trade(trade_id=10),
observed_at=OBSERVED_AT,
)
repository.store_trade(
venue=VENUE,
trade=_trade(
symbol=SECOND_SYMBOL,
trade_id=20,
offset_ms=1,
),
observed_at=OBSERVED_AT + timedelta(seconds=1),
)
repository.store_trade(
venue=VENUE,
trade=_trade(
symbol=SECOND_SYMBOL,
trade_id=19,
offset_ms=2,
),
observed_at=OBSERVED_AT + timedelta(seconds=2),
)
state_store = TradeStreamStateStore()
with pytest.raises(MarketDataCheckpointIntegrityError):
TradeStreamStateHydrator(
checkpoint_storage=repository,
state_store=state_store,
venue=VENUE,
deduplication_window_size=2,
).hydrate(symbols=(SYMBOL, SECOND_SYMBOL))
assert state_store.is_empty() is True
assert state_store.contains(SYMBOL) is False
assert state_store.contains(SECOND_SYMBOL) is False
adopted_first_symbol = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
assert adopted_first_symbol is not None
assert adopted_first_symbol.trade.trade_id == 10
assert repository.load_checkpoint(
venue=VENUE,
symbol=SECOND_SYMBOL,
) is None

View File

@@ -1,6 +1,7 @@
from __future__ import annotations
import ipaddress
import math
import re
import time
from collections.abc import Mapping
@@ -177,28 +178,26 @@ def release_postgres_test_lock(
def count_other_test_connections(
connection: psycopg.Connection[Any],
) -> int:
"""Посчитать оставшиеся соединения стенда и пула с тестовой базой."""
"""Посчитать клиентские соединения с проверенной тестовой БД."""
database_name = _validated_postgres_test_control_database_name(connection)
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT COUNT(*)
FROM pg_stat_activity
WHERE datname = current_database()
AND application_name = %s
FROM pg_catalog.pg_stat_activity
WHERE datname = %s
AND backend_type = 'client backend'
AND application_name IS DISTINCT FROM %s
""",
(POSTGRES_TEST_APPLICATION_NAME,),
(database_name, POSTGRES_TEST_CONTROL_APPLICATION_NAME),
)
row = cursor.fetchone()
if (
not isinstance(row, tuple)
or len(row) != 1
or isinstance(row[0], bool)
or not isinstance(row[0], int)
):
raise RuntimeError("PostgreSQL returned an invalid connection count")
return row[0]
return _validated_postgres_count(
row,
error_message="PostgreSQL returned an invalid connection count",
)
def wait_for_postgres_advisory_lock_waiters(
@@ -260,6 +259,138 @@ def wait_for_postgres_advisory_lock_waiters(
time.sleep(0.01)
def wait_for_postgres_relation_lock_waiters(
connection: psycopg.Connection[Any],
*,
relation_name: str,
expected_count: int,
timeout_seconds: float = 5.0,
) -> None:
"""Дождаться точного числа ожидающих блокировку заданной таблицы."""
if not isinstance(relation_name, str) or not relation_name.strip():
raise ValueError("relation_name must be a non-empty string")
if (
isinstance(expected_count, bool)
or not isinstance(expected_count, int)
or expected_count <= 0
):
raise ValueError("expected_count must be a positive integer")
if (
isinstance(timeout_seconds, bool)
or not isinstance(timeout_seconds, (int, float))
or not math.isfinite(float(timeout_seconds))
or timeout_seconds <= 0
):
raise ValueError("timeout_seconds must be a positive finite number")
database_name = _validated_postgres_test_control_database_name(connection)
with connection.cursor() as cursor:
cursor.execute(
"SELECT pg_catalog.to_regclass(%s)::oid",
(relation_name,),
)
relation_row = cursor.fetchone()
if (
not isinstance(relation_row, tuple)
or len(relation_row) != 1
or isinstance(relation_row[0], bool)
or not isinstance(relation_row[0], int)
or relation_row[0] <= 0
):
raise RuntimeError(
"PostgreSQL did not resolve the requested test relation"
)
relation_oid = relation_row[0]
deadline = time.monotonic() + timeout_seconds
while True:
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT COUNT(*)
FROM pg_catalog.pg_locks
WHERE locktype = 'relation'
AND database = (
SELECT oid
FROM pg_catalog.pg_database
WHERE datname = %s
)
AND relation = %s
AND NOT granted
""",
(database_name, relation_oid),
)
row = cursor.fetchone()
observed_count = _validated_postgres_count(
row,
error_message=(
"PostgreSQL returned an invalid relation-lock waiter count"
),
)
if observed_count == expected_count:
return
if time.monotonic() >= deadline:
raise TimeoutError(
"PostgreSQL did not observe all relation-lock callers; "
f"expected {expected_count}, observed {observed_count}."
)
time.sleep(0.01)
def _validated_postgres_test_control_database_name(
connection: psycopg.Connection[Any],
) -> str:
"""Повторно подтвердить безопасную БД и управляющее соединение."""
with connection.cursor() as cursor:
cursor.execute(
"""
SELECT current_database(), current_setting('application_name')
"""
)
identity = cursor.fetchone()
if (
not isinstance(identity, tuple)
or len(identity) != 2
or not isinstance(identity[0], str)
or not _SAFE_DATABASE_NAME.fullmatch(identity[0])
or identity[1] != POSTGRES_TEST_CONTROL_APPLICATION_NAME
):
raise RuntimeError(
"PostgreSQL connection is not the validated test control "
"connection."
)
return identity[0]
def _validated_postgres_count(
row: object,
*,
error_message: str,
) -> int:
"""Проверить форму и тип результата PostgreSQL COUNT(*)."""
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 RuntimeError(error_message)
return row[0]
def _validate_local_endpoint(parameters: Mapping[str, object]) -> None:
service = str(parameters.get("service", "")).strip()
host = str(parameters.get("host", "")).strip()

View File

@@ -30,12 +30,17 @@ RUNTIME_CLEANUP_TIMEOUT_SECONDS = 5.0
OWNED_TASK_NAMES = frozenset(
{
"application-shutdown",
"market-data-storage-shutdown",
"market-data-storage-startup",
"telegram-polling",
"trade-stream-market-processing",
"trade-stream-receive",
"trade-stream-runtime",
"trade-stream-runtime-recovery",
"trade-stream-scheduler",
"trade-stream-startup-recovery",
"trade-stream-state-hydration",
"trade-stream-startup",
}
)
@@ -173,7 +178,10 @@ def active_owned_task_names() -> tuple[str, ...]:
for task in asyncio.all_tasks()
if task is not current_task
and not task.done()
and task.get_name() in OWNED_TASK_NAMES
and (
task.get_name() in OWNED_TASK_NAMES
or task.get_name().startswith("persistent-application-")
)
)
)

View File

@@ -37,7 +37,10 @@ class RecordingJournal:
del event, message, context
def make_settings() -> SimpleNamespace:
def make_settings(
*,
storage_enabled: bool = True,
) -> SimpleNamespace:
return SimpleNamespace(
bot_token="test-token",
bot_parse_mode="HTML",
@@ -46,7 +49,9 @@ def make_settings() -> SimpleNamespace:
exchange_name="dzengi",
default_symbol="BTC/USD_LEVERAGE",
trade_stream=SimpleNamespace(enabled=True),
market_data_storage=SimpleNamespace(enabled=True),
market_data_storage=SimpleNamespace(
enabled=storage_enabled,
),
)
@@ -59,9 +64,11 @@ def test_create_app_builds_one_application_composition(
runtime = object()
storage_lifecycle = object()
storage_sink = object()
storage_repository = object()
storage = SimpleNamespace(
lifecycle=storage_lifecycle,
trade_observation_sink=storage_sink,
trade_repository=storage_repository,
)
journal = RecordingJournal()
observed_runtime_settings: list[object] = []
@@ -103,9 +110,11 @@ def test_create_app_builds_one_application_composition(
received_settings: object,
*,
trade_observation_sink: object,
checkpoint_storage: object,
) -> object:
observed_runtime_settings.append(received_settings)
assert trade_observation_sink is storage_sink
assert checkpoint_storage is storage_repository
return runtime
monkeypatch.setattr(
@@ -162,7 +171,7 @@ def test_runtime_build_error_is_fatal(
monkeypatch.setattr(
app_factory,
"load_settings",
make_settings,
lambda: make_settings(storage_enabled=False),
)
monkeypatch.setattr(
app_factory,
@@ -189,8 +198,11 @@ def test_runtime_build_error_is_fatal(
settings: object,
*,
trade_observation_sink: object,
checkpoint_storage: object,
) -> None:
del settings, trade_observation_sink
del settings
assert trade_observation_sink is None
assert checkpoint_storage is None
raise expected
monkeypatch.setattr(

View File

@@ -557,6 +557,8 @@ def test_application_cancellation_performs_full_cleanup() -> None:
"telegram-polling",
"trade-stream-runtime",
"application-shutdown",
"market-data-storage-startup",
"market-data-storage-shutdown",
}
}
@@ -637,6 +639,8 @@ def test_simultaneous_root_failures_are_awaited_deterministically() -> None:
"telegram-polling",
"trade-stream-runtime",
"application-shutdown",
"market-data-storage-startup",
"market-data-storage-shutdown",
}
}
@@ -647,6 +651,7 @@ def test_repeated_cancellation_does_not_interrupt_cleanup() -> None:
async def scenario() -> None:
dispatcher = FakeDispatcher()
runtime = BlockingStopRuntime()
storage = FakeStorageLifecycle()
bot = FakeBot()
task = asyncio.create_task(
run_application(
@@ -654,6 +659,7 @@ def test_repeated_cancellation_does_not_interrupt_cleanup() -> None:
dispatcher=dispatcher,
runtime=runtime,
bot=bot,
storage_lifecycle=storage,
)
)
)
@@ -664,6 +670,8 @@ def test_repeated_cancellation_does_not_interrupt_cleanup() -> None:
await runtime.stop_entered.wait()
task.cancel()
await asyncio.sleep(0)
assert storage.stop_calls == 0
runtime.stop_release.set()
with pytest.raises(asyncio.CancelledError):
@@ -673,6 +681,7 @@ def test_repeated_cancellation_does_not_interrupt_cleanup() -> None:
assert dispatcher.cancelled.is_set()
assert runtime.stop_calls == 1
assert runtime.stopped.is_set()
assert storage.stop_calls == 1
assert bot.session.close_calls == 1
await asyncio.sleep(0)
assert not {
@@ -685,6 +694,8 @@ def test_repeated_cancellation_does_not_interrupt_cleanup() -> None:
"telegram-polling",
"trade-stream-runtime",
"application-shutdown",
"market-data-storage-startup",
"market-data-storage-shutdown",
}
}

View File

@@ -15,6 +15,10 @@ from src.core.config import (
Settings,
TradeStreamSettings,
)
from src.market_data.storage.contracts import (
TradeCheckpointStorageProtocol,
TradeStorageProtocol,
)
def make_settings(
@@ -159,6 +163,18 @@ def test_builds_shared_graph_without_opening_pool() -> None:
composition.trade_observation_sink._trade_storage
is composition.trade_repository
)
assert (
composition.trade_observation_sink._checkpoint_storage
is composition.trade_repository
)
assert isinstance(
composition.trade_repository,
TradeStorageProtocol,
)
assert isinstance(
composition.trade_repository,
TradeCheckpointStorageProtocol,
)
assert composition.trade_observation_sink._venue == "dzengi"
conninfo = conninfo_to_dict(composition.connection_pool._conninfo)

View File

@@ -2,6 +2,10 @@ from __future__ import annotations
from typing import Any
import pytest
import src.bootstrap.trade_stream_runtime as runtime_factory
from src.bootstrap.market_data_storage import build_market_data_storage
from src.bootstrap.trade_stream_runtime import (
build_trade_stream_production_runtime,
)
@@ -21,7 +25,16 @@ class RecordingTradeObservationSink:
def __init__(self) -> None:
self.observations: list[Trade] = []
def persist(self, trade: Trade) -> None:
def persist_accepted(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
del expected_trade
self.observations.append(trade)
def persist_duplicate(self, trade: Trade) -> None:
self.observations.append(trade)
@@ -29,6 +42,9 @@ def make_settings(
*,
enabled: bool = True,
api_key: str = "api-key",
storage_enabled: bool = False,
subscription_ack_timeout_seconds: float = 12.5,
startup_market_buffer_capacity: int = 1_234,
) -> Settings:
return Settings(
bot_token="test-token",
@@ -58,6 +74,12 @@ def make_settings(
heartbeat_timeout_seconds=31.0,
scheduler_interval_seconds=6.0,
recovery_window_ms=123_456,
subscription_ack_timeout_seconds=(
subscription_ack_timeout_seconds
),
startup_market_buffer_capacity=(
startup_market_buffer_capacity
),
),
db_host="localhost",
db_port=5432,
@@ -65,7 +87,7 @@ def make_settings(
db_user="test",
db_password="test",
market_data_storage=MarketDataStorageSettings(
enabled=False,
enabled=storage_enabled,
pool_min_size=1,
pool_max_size=4,
pool_timeout_seconds=10.0,
@@ -81,6 +103,16 @@ def test_disabled_feature_does_not_build_runtime() -> None:
assert build_trade_stream_production_runtime(settings) is None
def test_enabled_storage_requires_enabled_trade_stream() -> None:
settings = make_settings(
enabled=False,
storage_enabled=True,
)
with pytest.raises(RuntimeError, match="Trade Stream"):
build_trade_stream_production_runtime(settings)
def test_builds_runtime_without_starting_lifecycle() -> None:
runtime = build_trade_stream_production_runtime(
make_settings(),
@@ -120,6 +152,116 @@ def test_uses_one_shared_stateful_dependency_graph() -> None:
assert runtime_graph._runtime_scheduler.runtime_supervisor is (
runtime_graph._runtime_supervisor
)
assert runtime_graph._startup_recovery_coordinator is None
def test_persistent_runtime_reuses_one_storage_graph_without_io() -> None:
settings = make_settings(storage_enabled=True)
storage = build_market_data_storage(settings)
assert storage is not None
runtime = build_trade_stream_production_runtime(
settings,
trade_observation_sink=storage.trade_observation_sink,
checkpoint_storage=storage.trade_repository,
)
assert isinstance(runtime, TradeStreamProductionRuntime)
runtime_graph: Any = runtime
startup_recovery = runtime_graph._startup_recovery_coordinator
live_controller = (
runtime_graph._trade_stream_service._consistency_controller
)
assert startup_recovery is not None
assert (
startup_recovery._state_hydrator._checkpoint_storage
is storage.trade_repository
)
assert (
storage.trade_observation_sink._trade_storage
is storage.trade_repository
)
assert (
storage.trade_observation_sink._checkpoint_storage
is storage.trade_repository
)
assert live_controller._trade_observation_sink is (
storage.trade_observation_sink
)
assert startup_recovery._state_hydrator._state_store is (
live_controller._state_store
)
assert startup_recovery._state_hydrator._venue == "dzengi"
assert storage.trade_observation_sink._venue == "dzengi"
assert storage.connection_pool.is_open is False
assert storage.lifecycle.started is False
assert startup_recovery._hydration_task is None
assert startup_recovery._recovery_task is None
assert runtime_graph._startup_task is None
assert runtime_graph._receive_task is None
assert runtime_graph._scheduler_task is None
@pytest.mark.parametrize(
("include_sink", "include_checkpoint"),
(
(False, False),
(True, False),
(False, True),
),
)
def test_enabled_storage_rejects_incomplete_runtime_graph_before_transport(
monkeypatch: pytest.MonkeyPatch,
include_sink: bool,
include_checkpoint: bool,
) -> None:
settings = make_settings(storage_enabled=True)
storage = build_market_data_storage(settings)
assert storage is not None
def unexpected_transport(**kwargs: object) -> None:
del kwargs
raise AssertionError("Transport graph must not be created.")
monkeypatch.setattr(
runtime_factory,
"DzengiWebSocketTransport",
unexpected_transport,
)
with pytest.raises(RuntimeError, match="requires both"):
build_trade_stream_production_runtime(
settings,
trade_observation_sink=(
storage.trade_observation_sink
if include_sink
else None
),
checkpoint_storage=(
storage.trade_repository
if include_checkpoint
else None
),
)
assert storage.connection_pool.is_open is False
def test_disabled_storage_rejects_checkpoint_dependency() -> None:
settings = make_settings()
persistent_settings = make_settings(storage_enabled=True)
storage = build_market_data_storage(persistent_settings)
assert storage is not None
with pytest.raises(RuntimeError, match="requires enabled"):
build_trade_stream_production_runtime(
settings,
checkpoint_storage=storage.trade_repository,
)
def test_optional_storage_sink_is_shared_by_live_and_recovery() -> None:
@@ -172,6 +314,8 @@ def test_applies_explicit_transport_and_runtime_settings() -> None:
"ETH/USD_LEVERAGE",
)
assert runtime_graph._runtime_scheduler.interval_seconds == 6.0
assert runtime_graph._subscription_ack_timeout_seconds == 12.5
assert runtime_graph._startup_market_buffer_capacity == 1_234
assert (
runtime_graph._runtime_supervisor._heartbeat_monitor.timeout_seconds
== 31.0

View File

@@ -15,6 +15,8 @@ _TRADE_STREAM_VARIABLES = (
"TRADE_STREAM_HEARTBEAT_TIMEOUT_SECONDS",
"TRADE_STREAM_SCHEDULER_INTERVAL_SECONDS",
"TRADE_STREAM_RECOVERY_WINDOW_MS",
"TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS",
"TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY",
)
_MARKET_DATA_STORAGE_VARIABLES = (
@@ -69,6 +71,8 @@ def test_trade_stream_is_disabled_by_default(
assert settings.trade_stream.enabled is False
assert settings.trade_stream.websocket_url == ""
assert settings.trade_stream.symbols == ()
assert settings.trade_stream.subscription_ack_timeout_seconds == 10.0
assert settings.trade_stream.startup_market_buffer_capacity == 10_000
def test_disabled_trade_stream_ignores_dependent_values(
@@ -76,6 +80,14 @@ def test_disabled_trade_stream_ignores_dependent_values(
) -> None:
prepare_environment(monkeypatch)
monkeypatch.setenv("TRADE_STREAM_OPEN_TIMEOUT_SECONDS", "invalid")
monkeypatch.setenv(
"TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS",
"invalid",
)
monkeypatch.setenv(
"TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY",
"invalid",
)
monkeypatch.setenv("TRADE_STREAM_SYMBOLS", ",")
settings = load_settings()
@@ -144,7 +156,7 @@ def test_enabled_trade_stream_parses_independent_settings(
enable_trade_stream(monkeypatch)
monkeypatch.setenv(
"TRADE_STREAM_SYMBOLS",
" ETH/USD_LEVERAGE, BTC/USD_LEVERAGE,ETH/USD_LEVERAGE ",
" eth/usd_leverage, BTC/USD_LEVERAGE,ETH/USD_LEVERAGE ",
)
monkeypatch.setenv("TRADE_STREAM_OPEN_TIMEOUT_SECONDS", "11.5")
monkeypatch.setenv("TRADE_STREAM_PROBE_TIMEOUT_SECONDS", "21")
@@ -152,6 +164,14 @@ def test_enabled_trade_stream_parses_independent_settings(
monkeypatch.setenv("TRADE_STREAM_HEARTBEAT_TIMEOUT_SECONDS", "31")
monkeypatch.setenv("TRADE_STREAM_SCHEDULER_INTERVAL_SECONDS", "6")
monkeypatch.setenv("TRADE_STREAM_RECOVERY_WINDOW_MS", "123456")
monkeypatch.setenv(
"TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS",
"12.5",
)
monkeypatch.setenv(
"TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY",
"1234",
)
settings = load_settings()
trade_stream = settings.trade_stream
@@ -168,6 +188,20 @@ def test_enabled_trade_stream_parses_independent_settings(
assert trade_stream.heartbeat_timeout_seconds == 31.0
assert trade_stream.scheduler_interval_seconds == 6.0
assert trade_stream.recovery_window_ms == 123_456
assert trade_stream.subscription_ack_timeout_seconds == 12.5
assert trade_stream.startup_market_buffer_capacity == 1_234
def test_enabled_trade_stream_uses_startup_boundary_defaults(
monkeypatch: pytest.MonkeyPatch,
) -> None:
prepare_environment(monkeypatch)
enable_trade_stream(monkeypatch)
trade_stream = load_settings().trade_stream
assert trade_stream.subscription_ack_timeout_seconds == 10.0
assert trade_stream.startup_market_buffer_capacity == 10_000
@pytest.mark.parametrize(
@@ -180,6 +214,12 @@ def test_enabled_trade_stream_parses_independent_settings(
("TRADE_STREAM_SCHEDULER_INTERVAL_SECONDS", "invalid"),
("TRADE_STREAM_RECOVERY_WINDOW_MS", "1.5"),
("TRADE_STREAM_RECOVERY_WINDOW_MS", "0"),
("TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS", "0"),
("TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS", "nan"),
("TRADE_STREAM_SUBSCRIPTION_ACK_TIMEOUT_SECONDS", "invalid"),
("TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY", "0"),
("TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY", "-1"),
("TRADE_STREAM_STARTUP_MARKET_BUFFER_CAPACITY", "1.5"),
),
)
def test_enabled_trade_stream_rejects_invalid_numeric_settings(

View File

@@ -0,0 +1,557 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from decimal import Decimal
import pytest
from src.market_data.acquisition.checkpoint.trade_stream_state_hydrator import (
TradeStreamStateHydrator,
TradeStreamStateHydratorProtocol,
)
from src.market_data.acquisition.consistency.trade_stream_state_store import (
TradeStreamStateStore,
)
from src.market_data.acquisition.consistency.trade_stream_state_store_exceptions import (
TradeStreamStateStoreInitializationError,
)
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.storage.contracts import (
PersistentTradeCheckpoint,
TradeCheckpointStorageProtocol,
)
from src.market_data.storage.exceptions import (
MarketDataCheckpointIntegrityError,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
SECOND_SYMBOL = "ETH/USD_LEVERAGE"
BASE_TIME = datetime(2026, 8, 1, 10, 0, tzinfo=timezone.utc)
def make_trade(
*,
symbol: str = SYMBOL,
trade_id: int = 100,
second: int = 0,
source: str = "dzengi_websocket_trade",
) -> Trade:
return Trade(
symbol=symbol,
trade_id=trade_id,
price=Decimal("65000.25"),
quantity=Decimal("0.001"),
executed_at=BASE_TIME + timedelta(seconds=second),
aggressor_side=TradeAggressorSide.BUY,
source=source,
)
def make_checkpoint(
trade: Trade,
*,
venue: str = VENUE,
revision: int = 1,
) -> PersistentTradeCheckpoint:
return PersistentTradeCheckpoint(
venue=venue,
trade=trade,
revision=revision,
updated_at=BASE_TIME + timedelta(minutes=1),
)
CheckpointResult = PersistentTradeCheckpoint | None | BaseException
TailResult = tuple[Trade, ...] | object | BaseException
AdoptionResult = PersistentTradeCheckpoint | BaseException
class RecordingCheckpointStorage:
"""Настраиваемый fake полного checkpoint-контракта."""
def __init__(
self,
*,
checkpoint_results: dict[
str,
tuple[CheckpointResult, ...],
]
| None = None,
checkpoint_tails: dict[str, TailResult] | None = None,
latest_tails: dict[str, TailResult] | None = None,
adoption_results: dict[str, AdoptionResult] | None = None,
) -> None:
self.checkpoint_results = checkpoint_results or {}
self.checkpoint_tails = checkpoint_tails or {}
self.latest_tails = latest_tails or {}
self.adoption_results = adoption_results or {}
self.calls: list[tuple[object, ...]] = []
self._checkpoint_offsets: dict[str, int] = {}
@property
def operation_names(self) -> list[str]:
return [str(call[0]) for call in self.calls]
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
self.calls.append(("load_checkpoint", venue, symbol))
results = self.checkpoint_results.get(symbol, (None,))
offset = self._checkpoint_offsets.get(symbol, 0)
self._checkpoint_offsets[symbol] = offset + 1
result = results[min(offset, len(results) - 1)]
if isinstance(result, BaseException):
raise result
return result
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
self.calls.append(
(
"load_checkpoint_tail",
venue,
checkpoint,
limit,
)
)
result = self.checkpoint_tails.get(
checkpoint.trade.symbol,
(checkpoint.trade,),
)
if isinstance(result, BaseException):
raise result
return result # type: ignore[return-value]
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
self.calls.append(
("load_latest_trade_tail", venue, symbol, limit)
)
result = self.latest_tails.get(symbol, ())
if isinstance(result, BaseException):
raise result
return result # type: ignore[return-value]
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
self.calls.append(
("adopt_existing_trade_as_checkpoint", venue, trade)
)
result = self.adoption_results.get(
trade.symbol,
make_checkpoint(trade, venue=venue),
)
if isinstance(result, BaseException):
raise result
return result
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
raise AssertionError(
"Hydrator не должен использовать writer checkpoint."
)
def make_hydrator(
storage: RecordingCheckpointStorage,
*,
state_store: TradeStreamStateStore | None = None,
window_size: int = 3,
) -> tuple[TradeStreamStateHydrator, TradeStreamStateStore]:
store = state_store or TradeStreamStateStore()
return (
TradeStreamStateHydrator(
checkpoint_storage=storage,
state_store=store,
venue=VENUE,
deduplication_window_size=window_size,
),
store,
)
def test_implements_protocol_uses_slots_and_constructor_has_no_io() -> None:
storage = RecordingCheckpointStorage()
hydrator, _ = make_hydrator(storage)
assert isinstance(storage, TradeCheckpointStorageProtocol)
assert isinstance(hydrator, TradeStreamStateHydratorProtocol)
assert not hasattr(hydrator, "__dict__")
assert storage.calls == []
def test_existing_checkpoint_restores_bounded_tail_and_publishes_state() -> None:
trades = (
make_trade(trade_id=100, second=0),
make_trade(trade_id=101, second=1),
make_trade(trade_id=102, second=2),
)
checkpoint = make_checkpoint(trades[-1], revision=7)
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (checkpoint,)},
checkpoint_tails={SYMBOL: trades},
)
hydrator, store = make_hydrator(storage)
states = hydrator.hydrate(symbols=(SYMBOL,))
assert len(states) == 1
assert states[0].last_trade is trades[-1]
assert states[0].last_trade_id == 102
assert store.get(SYMBOL) is states[0]
assert storage.calls == [
("load_checkpoint", VENUE, SYMBOL),
("load_checkpoint_tail", VENUE, checkpoint, 3),
]
def test_first_adoption_reloads_tail_before_publishing_state() -> None:
candidate_tail = (
make_trade(trade_id=100, second=0),
make_trade(trade_id=101, second=1),
)
adopted = make_checkpoint(candidate_tail[-1])
reloaded_tail = (
make_trade(
trade_id=99,
second=-1,
source="postgres_trade_history",
),
make_trade(
trade_id=100,
second=0,
source="postgres_trade_history",
),
make_trade(
trade_id=101,
second=1,
source="postgres_trade_history",
),
)
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (None,)},
latest_tails={SYMBOL: candidate_tail},
adoption_results={SYMBOL: adopted},
checkpoint_tails={SYMBOL: reloaded_tail},
)
hydrator, store = make_hydrator(storage)
(state,) = hydrator.hydrate(symbols=(SYMBOL,))
assert state.last_trade is reloaded_tail[-1]
assert store.get(SYMBOL) is state
assert storage.calls == [
("load_checkpoint", VENUE, SYMBOL),
("load_latest_trade_tail", VENUE, SYMBOL, 3),
(
"adopt_existing_trade_as_checkpoint",
VENUE,
candidate_tail[-1],
),
("load_checkpoint_tail", VENUE, adopted, 3),
]
def test_empty_history_rechecks_checkpoint_and_uses_concurrent_value() -> None:
trade = make_trade(trade_id=100)
concurrent_checkpoint = make_checkpoint(trade, revision=2)
storage = RecordingCheckpointStorage(
checkpoint_results={
SYMBOL: (None, concurrent_checkpoint),
},
latest_tails={SYMBOL: ()},
checkpoint_tails={SYMBOL: (trade,)},
)
hydrator, store = make_hydrator(storage)
(state,) = hydrator.hydrate(symbols=(SYMBOL,))
assert state.last_trade is trade
assert store.get(SYMBOL) is state
assert storage.calls == [
("load_checkpoint", VENUE, SYMBOL),
("load_latest_trade_tail", VENUE, SYMBOL, 3),
("load_checkpoint", VENUE, SYMBOL),
(
"load_checkpoint_tail",
VENUE,
concurrent_checkpoint,
3,
),
]
def test_empty_history_without_checkpoint_publishes_empty_state() -> None:
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (None, None)},
latest_tails={SYMBOL: ()},
)
hydrator, store = make_hydrator(storage)
(state,) = hydrator.hydrate(symbols=(SYMBOL,))
assert state.last_trade is None
assert state.last_trade_id is None
assert store.get(SYMBOL) is state
assert storage.operation_names == [
"load_checkpoint",
"load_latest_trade_tail",
"load_checkpoint",
]
@pytest.mark.parametrize(
"malformed_tail",
(
None,
[],
),
)
def test_rejects_falsey_non_tuple_latest_tail(
malformed_tail: object,
) -> None:
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (None, None)},
latest_tails={SYMBOL: malformed_tail},
)
hydrator, store = make_hydrator(storage)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="non-tuple latest Trade tail",
):
hydrator.hydrate(symbols=(SYMBOL,))
assert store.is_empty() is True
assert storage.operation_names == [
"load_checkpoint",
"load_latest_trade_tail",
]
@pytest.mark.parametrize(
("previous_trade_id", "checkpoint_trade_id"),
(
(2**31 - 1, -(2**31)),
(-1, 0),
),
)
def test_restores_tail_across_signed_rollover(
previous_trade_id: int,
checkpoint_trade_id: int,
) -> None:
trades = (
make_trade(trade_id=previous_trade_id, second=0),
make_trade(trade_id=checkpoint_trade_id, second=1),
)
checkpoint = make_checkpoint(trades[-1])
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (checkpoint,)},
checkpoint_tails={SYMBOL: trades},
)
hydrator, _ = make_hydrator(storage)
(state,) = hydrator.hydrate(symbols=(SYMBOL,))
assert state.last_trade is trades[-1]
assert state.last_trade_id == checkpoint_trade_id
def test_rejects_tail_that_does_not_end_at_checkpoint() -> None:
checkpoint_trade = make_trade(trade_id=102, second=2)
checkpoint = make_checkpoint(checkpoint_trade)
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (checkpoint,)},
checkpoint_tails={
SYMBOL: (
make_trade(trade_id=100, second=0),
make_trade(trade_id=101, second=1),
)
},
)
hydrator, store = make_hydrator(storage)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="does not end",
):
hydrator.hydrate(symbols=(SYMBOL,))
assert store.contains(SYMBOL) is False
def test_rejects_tail_larger_than_deduplication_window() -> None:
trades = tuple(
make_trade(trade_id=trade_id, second=trade_id - 100)
for trade_id in range(100, 103)
)
checkpoint = make_checkpoint(trades[-1])
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (checkpoint,)},
checkpoint_tails={SYMBOL: trades},
)
hydrator, store = make_hydrator(storage, window_size=2)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="cannot hydrate",
):
hydrator.hydrate(symbols=(SYMBOL,))
assert store.contains(SYMBOL) is False
@pytest.mark.parametrize(
"malformed_tail",
(
[make_trade()],
(object(),),
),
)
def test_rejects_malformed_checkpoint_tail(
malformed_tail: object,
) -> None:
checkpoint = make_checkpoint(make_trade())
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (checkpoint,)},
checkpoint_tails={SYMBOL: malformed_tail},
)
hydrator, store = make_hydrator(storage)
with pytest.raises(MarketDataCheckpointIntegrityError):
hydrator.hydrate(symbols=(SYMBOL,))
assert store.contains(SYMBOL) is False
def test_storage_error_is_not_swallowed_and_state_is_not_published() -> None:
failure = RuntimeError("storage failed")
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (failure,)},
)
hydrator, store = make_hydrator(storage)
with pytest.raises(RuntimeError, match="storage failed"):
hydrator.hydrate(symbols=(SYMBOL,))
assert store.contains(SYMBOL) is False
assert storage.operation_names == ["load_checkpoint"]
def test_adoption_error_is_not_swallowed_and_state_is_not_published() -> None:
trade = make_trade()
failure = RuntimeError("adoption failed")
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (None,)},
latest_tails={SYMBOL: (trade,)},
adoption_results={SYMBOL: failure},
)
hydrator, store = make_hydrator(storage)
with pytest.raises(RuntimeError, match="adoption failed"):
hydrator.hydrate(symbols=(SYMBOL,))
assert store.contains(SYMBOL) is False
assert storage.operation_names == [
"load_checkpoint",
"load_latest_trade_tail",
"adopt_existing_trade_as_checkpoint",
]
def test_rejects_adopted_checkpoint_for_different_candidate() -> None:
candidate = make_trade(trade_id=100)
different_trade = make_trade(trade_id=101, second=1)
storage = RecordingCheckpointStorage(
checkpoint_results={SYMBOL: (None,)},
latest_tails={SYMBOL: (candidate,)},
adoption_results={
SYMBOL: make_checkpoint(different_trade),
},
)
hydrator, store = make_hydrator(storage)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="does not match candidate",
):
hydrator.hydrate(symbols=(SYMBOL,))
assert store.contains(SYMBOL) is False
assert "load_checkpoint_tail" not in storage.operation_names
def test_second_symbol_failure_does_not_publish_first_state() -> None:
first_trade = make_trade(symbol=SYMBOL)
first_checkpoint = make_checkpoint(first_trade)
second_failure = RuntimeError("second symbol failed")
storage = RecordingCheckpointStorage(
checkpoint_results={
SYMBOL: (first_checkpoint,),
SECOND_SYMBOL: (second_failure,),
},
checkpoint_tails={SYMBOL: (first_trade,)},
)
hydrator, store = make_hydrator(storage)
with pytest.raises(RuntimeError, match="second symbol failed"):
hydrator.hydrate(symbols=(SYMBOL, SECOND_SYMBOL))
assert store.contains(SYMBOL) is False
assert store.contains(SECOND_SYMBOL) is False
assert storage.operation_names == [
"load_checkpoint",
"load_checkpoint_tail",
"load_checkpoint",
]
def test_nonempty_store_is_rejected_before_storage_io() -> None:
storage = RecordingCheckpointStorage()
store = TradeStreamStateStore()
existing_state = store.get_or_create(SYMBOL)
hydrator, _ = make_hydrator(storage, state_store=store)
with pytest.raises(TradeStreamStateStoreInitializationError):
hydrator.hydrate(symbols=(SYMBOL, SECOND_SYMBOL))
assert store.get(SYMBOL) is existing_state
assert store.contains(SECOND_SYMBOL) is False
assert storage.calls == []

View File

@@ -78,12 +78,27 @@ class RecordingTradeObservationSink:
) -> None:
self.error = error
self.observations: list[Trade] = []
self.accepted: list[tuple[Trade, Trade | None]] = []
self.duplicates: list[Trade] = []
def persist(
def persist_accepted(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
self.observations.append(trade)
self.accepted.append((trade, expected_trade))
if self.error is not None:
raise self.error
def persist_duplicate(
self,
trade: Trade,
) -> None:
self.observations.append(trade)
self.duplicates.append(trade)
if self.error is not None:
raise self.error
@@ -221,6 +236,8 @@ def test_persists_trade_before_advancing_checkpoint(
assert result is trade
assert sink.observations == [trade]
assert sink.accepted == [(trade, None)]
assert sink.duplicates == []
assert state.last_trade is trade
@@ -247,6 +264,10 @@ def test_persistence_failure_leaves_checkpoint_unchanged(
state = state_store.get(first_trade.symbol)
assert error_info.value is storage_error
assert sink.accepted == [
(first_trade, None),
(failed_trade, first_trade),
]
assert state.last_trade is first_trade
assert state.last_trade_id == first_trade.trade_id
assert failed_trade.trade_id not in state._trades
@@ -280,10 +301,59 @@ def test_valid_duplicate_is_persisted_without_checkpoint_advance(
websocket_trade,
rest_duplicate,
]
assert sink.accepted == [(websocket_trade, None)]
assert sink.duplicates == [rest_duplicate]
assert state.last_trade is websocket_trade
assert state.last_trade_id == websocket_trade.trade_id
def test_duplicate_persistence_failure_keeps_checkpoint(
state_store: TradeStreamStateStore,
) -> None:
storage_error = RuntimeError("storage failed")
sink = RecordingTradeObservationSink()
controller = TradeStreamConsistencyController(
state_store=state_store,
trade_observation_sink=sink,
)
original = _trade(source="dzengi_websocket_trade")
duplicate = _trade(source="dzengi")
controller.accept(original)
sink.error = storage_error
with pytest.raises(RuntimeError, match="storage failed") as error_info:
controller.accept(duplicate)
state = state_store.get(original.symbol)
assert error_info.value is storage_error
assert sink.accepted == [(original, None)]
assert sink.duplicates == [duplicate]
assert state.last_trade is original
assert state.last_trade_id == original.trade_id
def test_rollover_advance_uses_previous_trade_as_expected_checkpoint(
state_store: TradeStreamStateStore,
) -> None:
sink = RecordingTradeObservationSink()
controller = TradeStreamConsistencyController(
state_store=state_store,
trade_observation_sink=sink,
)
previous = _trade(trade_id=2**31 - 1)
current = _trade(trade_id=-(2**31))
controller.accept(previous)
controller.accept(current)
assert sink.accepted == [
(previous, None),
(current, previous),
]
assert sink.duplicates == []
def test_invalid_trades_do_not_reach_persistence_sink(
state_store: TradeStreamStateStore,
) -> None:

View File

@@ -420,6 +420,88 @@ def test_checkpoint_trade_id_matches_last_trade_id() -> None:
assert state.last_trade.trade_id == state.last_trade_id
def test_checkpoint_callback_receives_expected_previous_trade() -> None:
state = TradeStreamState(symbol="BTCUSD")
first_trade = _trade(trade_id=100)
second_trade = _trade(trade_id=101)
calls: list[tuple[Trade, Trade | None]] = []
def before_checkpoint(
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
calls.append((trade, expected_trade))
assert state.last_trade is expected_trade
state.accept(
first_trade,
before_checkpoint=before_checkpoint,
)
state.accept(
second_trade,
before_checkpoint=before_checkpoint,
)
assert calls == [
(first_trade, None),
(second_trade, first_trade),
]
assert state.last_trade is second_trade
def test_duplicate_uses_only_duplicate_callback() -> None:
state = TradeStreamState(
symbol="BTCUSD",
deduplication_window_size=3,
)
first_trade = _trade(trade_id=100)
latest_trade = _trade(trade_id=101)
duplicate = _trade(
trade_id=100,
source="dzengi_websocket_trade",
)
checkpoint_calls: list[tuple[Trade, Trade | None]] = []
duplicate_calls: list[Trade] = []
state.accept(first_trade)
state.accept(latest_trade)
result = state.accept(
duplicate,
before_checkpoint=lambda trade, expected_trade: (
checkpoint_calls.append((trade, expected_trade))
),
on_duplicate=duplicate_calls.append,
)
assert result is None
assert checkpoint_calls == []
assert duplicate_calls == [duplicate]
assert state.last_trade is latest_trade
def test_duplicate_callback_failure_keeps_checkpoint() -> None:
state = TradeStreamState(symbol="BTCUSD")
original = _trade(source="dzengi_websocket_trade")
duplicate = _trade(source="dzengi")
storage_error = RuntimeError("storage failed")
state.accept(original)
def fail_duplicate(trade: Trade) -> None:
assert trade is duplicate
raise storage_error
with pytest.raises(RuntimeError, match="storage failed") as error_info:
state.accept(
duplicate,
on_duplicate=fail_duplicate,
)
assert error_info.value is storage_error
assert state.last_trade is original
assert state.last_trade_id == original.trade_id
def test_rejects_empty_symbol() -> None:
with pytest.raises(ValueError):
TradeStreamState(symbol="")
@@ -440,3 +522,184 @@ def test_rejects_non_positive_window_size(
symbol="BTCUSD",
deduplication_window_size=window_size,
)
def test_from_history_builds_empty_state() -> None:
state = TradeStreamState.from_history(
symbol="BTCUSD",
trades=(),
deduplication_window_size=3,
)
assert state.symbol == "BTCUSD"
assert state.deduplication_window_size == 3
assert state.last_trade_id is None
assert state.last_trade is None
def test_from_history_restores_valid_deduplication_window() -> None:
first_trade = _trade(trade_id=100)
second_trade = _trade(trade_id=101)
last_trade = _trade(trade_id=102)
state = TradeStreamState.from_history(
symbol="BTCUSD",
trades=(
first_trade,
second_trade,
last_trade,
),
deduplication_window_size=3,
)
assert state.last_trade is last_trade
assert state.last_trade_id == last_trade.trade_id
assert state.accept(_trade(trade_id=100)) is None
assert state.last_trade is last_trade
def test_from_history_keeps_exact_window_boundary() -> None:
first_trade = _trade(trade_id=100)
second_trade = _trade(trade_id=101)
third_trade = _trade(trade_id=102)
state = TradeStreamState.from_history(
symbol="BTCUSD",
trades=(
first_trade,
second_trade,
third_trade,
),
deduplication_window_size=3,
)
next_trade = _trade(trade_id=103)
assert state.accept(next_trade) is next_trade
assert state.accept(_trade(trade_id=101)) is None
with pytest.raises(TradeOrderingError):
state.accept(_trade(trade_id=100))
@pytest.mark.parametrize(
("first_trade_id", "next_trade_id"),
(
(2**31 - 1, -(2**31)),
(-1, 0),
),
)
def test_from_history_restores_signed_rollover_sequence(
first_trade_id: int,
next_trade_id: int,
) -> None:
first_trade = _trade(trade_id=first_trade_id)
next_trade = _trade(trade_id=next_trade_id)
state = TradeStreamState.from_history(
symbol="BTCUSD",
trades=(first_trade, next_trade),
deduplication_window_size=2,
)
assert state.last_trade is next_trade
assert state.last_trade_id == next_trade_id
assert state.accept(_trade(trade_id=first_trade_id)) is None
def test_from_history_accepts_next_trade_after_restoration() -> None:
last_history_trade = _trade(trade_id=101)
state = TradeStreamState.from_history(
symbol="BTCUSD",
trades=(
_trade(trade_id=100),
last_history_trade,
),
deduplication_window_size=3,
)
next_trade = _trade(trade_id=102)
result = state.accept(next_trade)
assert result is next_trade
assert state.last_trade is next_trade
assert state.last_trade_id == next_trade.trade_id
def test_from_history_rejects_identical_duplicate_strictly() -> None:
original = _trade(trade_id=100)
duplicate = _trade(
trade_id=100,
source="dzengi_websocket_trade",
)
with pytest.raises(
TradeConsistencyError,
match="duplicate Trade",
):
TradeStreamState.from_history(
symbol="BTCUSD",
trades=(original, duplicate),
deduplication_window_size=2,
)
def test_from_history_rejects_conflicting_trade_id() -> None:
original = _trade(
trade_id=100,
price=Decimal("50000.00"),
)
conflict = _trade(
trade_id=100,
price=Decimal("50001.00"),
)
with pytest.raises(TradeConsistencyError):
TradeStreamState.from_history(
symbol="BTCUSD",
trades=(original, conflict),
deduplication_window_size=2,
)
def test_from_history_rejects_reverse_sequence() -> None:
with pytest.raises(TradeOrderingError):
TradeStreamState.from_history(
symbol="BTCUSD",
trades=(
_trade(trade_id=101),
_trade(trade_id=100),
),
deduplication_window_size=2,
)
def test_from_history_rejects_half_cycle_sequence() -> None:
with pytest.raises(
ValueError,
match="exactly half",
):
TradeStreamState.from_history(
symbol="BTCUSD",
trades=(
_trade(trade_id=0),
_trade(trade_id=-(2**31)),
),
deduplication_window_size=2,
)
def test_from_history_rejects_history_larger_than_window() -> None:
with pytest.raises(
ValueError,
match="fit deduplication_window_size",
):
TradeStreamState.from_history(
symbol="BTCUSD",
trades=(
_trade(trade_id=100),
_trade(trade_id=101),
_trade(trade_id=102),
),
deduplication_window_size=2,
)

View File

@@ -2,6 +2,8 @@
from __future__ import annotations
from typing import cast
import pytest
from src.market_data.acquisition.consistency.trade_stream_state import (
@@ -12,6 +14,7 @@ from src.market_data.acquisition.consistency.trade_stream_state_store import (
)
from src.market_data.acquisition.consistency.trade_stream_state_store_exceptions import (
TradeStreamStateNotFoundError,
TradeStreamStateStoreInitializationError,
)
from src.market_data.acquisition.consistency.trade_stream_state_store_protocol import (
TradeStreamStateStoreProtocol,
@@ -140,4 +143,106 @@ def test_clear_is_idempotent_for_empty_store() -> None:
store.clear()
store.clear()
assert store.contains("BTCUSD") is False
assert store.contains("BTCUSD") is False
def test_initialize_publishes_all_states_preserving_identity() -> None:
store = TradeStreamStateStore()
btc_state = TradeStreamState(symbol="BTCUSD")
eth_state = TradeStreamState(symbol="ETHUSD")
store.initialize((btc_state, eth_state))
assert store.get("BTCUSD") is btc_state
assert store.get("ETHUSD") is eth_state
def test_initialize_rejects_duplicate_symbols_atomically() -> None:
store = TradeStreamStateStore()
with pytest.raises(
TradeStreamStateStoreInitializationError,
match="повторно",
):
store.initialize(
(
TradeStreamState(symbol="BTCUSD"),
TradeStreamState(symbol="BTCUSD"),
)
)
assert store.contains("BTCUSD") is False
def test_initialize_rejects_invalid_item_atomically() -> None:
store = TradeStreamStateStore()
with pytest.raises(
TypeError,
match="TradeStreamState",
):
store.initialize(
cast(
tuple[TradeStreamState, ...],
(
TradeStreamState(symbol="BTCUSD"),
object(),
),
),
)
assert store.contains("BTCUSD") is False
def test_initialize_rejects_non_empty_store() -> None:
store = TradeStreamStateStore()
existing_state = store.get_or_create("BTCUSD")
with pytest.raises(
TradeStreamStateStoreInitializationError,
match="уже содержит",
):
store.initialize(
(TradeStreamState(symbol="ETHUSD"),)
)
assert store.get("BTCUSD") is existing_state
assert store.contains("ETHUSD") is False
def test_initialize_can_be_called_only_once() -> None:
store = TradeStreamStateStore()
original_state = TradeStreamState(symbol="BTCUSD")
store.initialize((original_state,))
with pytest.raises(
TradeStreamStateStoreInitializationError,
match="уже содержит",
):
store.initialize(
(TradeStreamState(symbol="ETHUSD"),)
)
assert store.get("BTCUSD") is original_state
assert store.contains("ETHUSD") is False
def test_empty_initialize_is_still_one_time_initialization() -> None:
store = TradeStreamStateStore()
store.initialize(())
with pytest.raises(TradeStreamStateStoreInitializationError):
store.initialize(())
def test_clear_allows_store_to_be_initialized_again() -> None:
store = TradeStreamStateStore()
store.initialize((TradeStreamState(symbol="BTCUSD"),))
store.clear()
eth_state = TradeStreamState(symbol="ETHUSD")
store.initialize((eth_state,))
assert store.contains("BTCUSD") is False
assert store.get("ETHUSD") is eth_state

View File

@@ -266,8 +266,8 @@ def test_reconnect_restore_boundary_and_recovery_order() -> None:
order,
) = create_coordinator(
symbols=(
f" {ETH} ",
BTC,
f" {ETH.lower()} ",
BTC.lower(),
ETH,
),
)

View File

@@ -119,6 +119,18 @@ class FakeStateStore:
self.remove_calls: list[str] = []
self.clear_calls = 0
def initialize(
self,
states: tuple[TradeStreamState, ...],
) -> None:
self._states = {
state.symbol: state
for state in states
}
def is_empty(self) -> bool:
return not self._states
def get_or_create(
self,
symbol: str,

View File

@@ -0,0 +1,462 @@
from __future__ import annotations
import asyncio
import threading
import pytest
from src.market_data.acquisition.consistency.trade_stream_state import (
TradeStreamState,
)
from src.market_data.acquisition.recovery.trade_recovery_result import (
TradeRecoveryResult,
)
from src.market_data.acquisition.runtime.live_processing_gate import (
RuntimeLiveProcessingGate,
)
from src.market_data.acquisition.runtime.runtime_startup_recovery_coordinator import (
RuntimeStartupRecoveryCoordinator,
RuntimeStartupRecoveryProtocol,
)
BTC = "BTC/USD_LEVERAGE"
ETH = "ETH/USD_LEVERAGE"
RECOVERY_END_TIME_MS = 1_785_326_405_123
HYDRATION_TASK_NAME = "trade-stream-state-hydration"
RECOVERY_TASK_NAME = "trade-stream-startup-recovery"
class FakeStateHydrator:
def __init__(
self,
*,
states: tuple[TradeStreamState, ...],
error: Exception | None = None,
started: threading.Event | None = None,
release: threading.Event | None = None,
) -> None:
self._states = states
self._error = error
self._started = started
self._release = release
self.calls: list[tuple[str, ...]] = []
self.thread_ids: list[int] = []
def hydrate(
self,
*,
symbols: tuple[str, ...],
) -> tuple[TradeStreamState, ...]:
self.calls.append(symbols)
self.thread_ids.append(threading.get_ident())
if self._started is not None:
self._started.set()
if self._release is not None and not self._release.wait(
timeout=2.0,
):
raise AssertionError("hydration release was not signalled")
if self._error is not None:
raise self._error
return self._states
class FakeRecoveryCoordinator:
def __init__(
self,
*,
error: Exception | None = None,
started: threading.Event | None = None,
release: threading.Event | None = None,
) -> None:
self._error = error
self._started = started
self._release = release
self.calls: list[tuple[str, int]] = []
self.thread_ids: list[int] = []
def recover(
self,
*,
symbol: str,
recovery_end_time: int,
) -> TradeRecoveryResult:
self.calls.append(
(
symbol,
recovery_end_time,
)
)
self.thread_ids.append(threading.get_ident())
if self._started is not None:
self._started.set()
if self._release is not None and not self._release.wait(
timeout=2.0,
):
raise AssertionError("recovery release was not signalled")
if self._error is not None:
raise self._error
return TradeRecoveryResult(
symbol=symbol,
requested_start_time=recovery_end_time,
requested_end_time=recovery_end_time,
recovered_trades=(),
)
class RecordingClock:
def __init__(
self,
value: object = RECOVERY_END_TIME_MS,
) -> None:
self._value = value
self.calls = 0
def __call__(self) -> int:
self.calls += 1
return self._value # type: ignore[return-value]
def create_coordinator(
*,
symbols: tuple[str, ...] = (BTC,),
states: tuple[TradeStreamState, ...] | None = None,
hydration_error: Exception | None = None,
hydration_started: threading.Event | None = None,
hydration_release: threading.Event | None = None,
recovery_error: Exception | None = None,
recovery_started: threading.Event | None = None,
recovery_release: threading.Event | None = None,
clock_value: object = RECOVERY_END_TIME_MS,
) -> tuple[
RuntimeStartupRecoveryCoordinator,
FakeStateHydrator,
FakeRecoveryCoordinator,
RuntimeLiveProcessingGate,
RecordingClock,
]:
hydrated_states = states or (
TradeStreamState(symbol=BTC),
)
hydrator = FakeStateHydrator(
states=hydrated_states,
error=hydration_error,
started=hydration_started,
release=hydration_release,
)
recovery = FakeRecoveryCoordinator(
error=recovery_error,
started=recovery_started,
release=recovery_release,
)
gate = RuntimeLiveProcessingGate()
clock = RecordingClock(clock_value)
coordinator = RuntimeStartupRecoveryCoordinator(
state_hydrator=hydrator,
recovery_coordinator=recovery,
live_processing_gate=gate,
symbols=symbols,
clock=clock,
)
return coordinator, hydrator, recovery, gate, clock
async def wait_until(
predicate: object,
) -> None:
for _ in range(100):
if callable(predicate) and predicate():
return
await asyncio.sleep(0)
raise AssertionError("condition was not reached")
def test_implements_protocol_uses_slots_and_constructor_has_no_io() -> None:
coordinator, hydrator, recovery, gate, clock = create_coordinator(
symbols=(
f" {ETH.lower()} ",
BTC.lower(),
ETH,
),
)
assert isinstance(coordinator, RuntimeStartupRecoveryProtocol)
assert not hasattr(coordinator, "__dict__")
assert coordinator.live_processing_gate is gate
assert coordinator.symbols == (BTC, ETH)
assert hydrator.calls == []
assert recovery.calls == []
assert clock.calls == 0
assert gate.locked is False
@pytest.mark.parametrize(
("symbols", "error_type"),
[
([], TypeError),
((), ValueError),
(("", " "), ValueError),
((BTC, 1), TypeError),
],
)
def test_rejects_invalid_symbols(
symbols: object,
error_type: type[Exception],
) -> None:
with pytest.raises(error_type):
create_coordinator(
symbols=symbols, # type: ignore[arg-type]
)
def test_hydration_is_idempotent_and_runs_off_event_loop() -> None:
async def scenario() -> tuple[
tuple[TradeStreamState, ...],
tuple[TradeStreamState, ...],
int,
]:
states = (
TradeStreamState(symbol=BTC),
)
coordinator, hydrator, *_ = create_coordinator(
states=states,
)
event_loop_thread_id = threading.get_ident()
first = await coordinator.hydrate_once()
second = await coordinator.hydrate_once()
assert hydrator.calls == [(BTC,)]
assert len(hydrator.thread_ids) == 1
return first, second, event_loop_thread_id
first, second, event_loop_thread_id = asyncio.run(scenario())
assert first is second
assert first[0].symbol == BTC
coordinator, hydrator, *_ = create_coordinator()
asyncio.run(coordinator.hydrate_once())
assert hydrator.thread_ids[0] != event_loop_thread_id
def test_recovery_uses_one_clock_boundary_and_symbol_order() -> None:
async def scenario() -> tuple[
tuple[TradeRecoveryResult, ...],
FakeRecoveryCoordinator,
RecordingClock,
int,
]:
coordinator, _, recovery, gate, clock = create_coordinator(
symbols=(
f" {ETH.lower()} ",
BTC.lower(),
ETH,
),
)
event_loop_thread_id = threading.get_ident()
async with gate:
results = await coordinator.recover_after_ack()
return results, recovery, clock, event_loop_thread_id
results, recovery, clock, event_loop_thread_id = asyncio.run(
scenario()
)
assert tuple(result.symbol for result in results) == (BTC, ETH)
assert recovery.calls == [
(BTC, RECOVERY_END_TIME_MS),
(ETH, RECOVERY_END_TIME_MS),
]
assert clock.calls == 1
assert len(set(recovery.thread_ids)) == 1
assert recovery.thread_ids[0] != event_loop_thread_id
def test_cancelled_hydration_waits_and_caches_success() -> None:
async def scenario() -> tuple[
RuntimeStartupRecoveryCoordinator,
FakeStateHydrator,
tuple[TradeStreamState, ...],
]:
started = threading.Event()
release = threading.Event()
states = (
TradeStreamState(symbol=BTC),
)
coordinator, hydrator, *_ = create_coordinator(
states=states,
hydration_started=started,
hydration_release=release,
)
task = asyncio.create_task(
coordinator.hydrate_once(),
)
await wait_until(started.is_set)
joined_task = asyncio.create_task(
coordinator.hydrate_once(),
)
await asyncio.sleep(0)
task.cancel()
await asyncio.sleep(0)
assert task.done() is False
assert joined_task.done() is False
release.set()
with pytest.raises(asyncio.CancelledError):
await task
cached = await joined_task
assert await coordinator.hydrate_once() is cached
return coordinator, hydrator, cached
coordinator, hydrator, cached = asyncio.run(scenario())
assert cached[0].symbol == BTC
assert hydrator.calls == [(BTC,)]
assert coordinator._hydration_task is None
def test_cancelled_recovery_waits_for_worker_completion() -> None:
async def scenario() -> tuple[
RuntimeStartupRecoveryCoordinator,
FakeRecoveryCoordinator,
]:
started = threading.Event()
release = threading.Event()
coordinator, _, recovery, gate, _ = create_coordinator(
recovery_started=started,
recovery_release=release,
)
async def run_recovery() -> None:
async with gate:
await coordinator.recover_after_ack()
task = asyncio.create_task(run_recovery())
await wait_until(started.is_set)
task.cancel()
await asyncio.sleep(0)
assert task.done() is False
assert gate.locked is True
release.set()
with pytest.raises(asyncio.CancelledError):
await task
assert gate.locked is False
return coordinator, recovery
coordinator, recovery = asyncio.run(scenario())
assert recovery.calls == [(BTC, RECOVERY_END_TIME_MS)]
assert coordinator._recovery_task is None
@pytest.mark.parametrize(
"operation",
[
"hydration",
"recovery",
],
)
def test_worker_error_identity_is_preserved(
operation: str,
) -> None:
error = RuntimeError(f"{operation} failed")
async def scenario() -> RuntimeStartupRecoveryCoordinator:
coordinator, *_ = create_coordinator(
hydration_error=(
error if operation == "hydration" else None
),
recovery_error=(
error if operation == "recovery" else None
),
)
with pytest.raises(RuntimeError) as raised:
if operation == "hydration":
await coordinator.hydrate_once()
else:
await coordinator.recover_after_ack()
assert raised.value is error
return coordinator
coordinator = asyncio.run(scenario())
assert coordinator._hydration_task is None
assert coordinator._recovery_task is None
@pytest.mark.parametrize(
("clock_value", "error_type"),
[
(True, TypeError),
(1.5, TypeError),
("1", TypeError),
(-1, ValueError),
],
)
def test_invalid_clock_result_fails_before_recovery(
clock_value: object,
error_type: type[Exception],
) -> None:
async def scenario() -> tuple[
FakeRecoveryCoordinator,
RecordingClock,
RuntimeStartupRecoveryCoordinator,
]:
coordinator, _, recovery, _, clock = create_coordinator(
clock_value=clock_value,
)
with pytest.raises(error_type):
await coordinator.recover_after_ack()
return recovery, clock, coordinator
recovery, clock, coordinator = asyncio.run(scenario())
assert recovery.calls == []
assert clock.calls == 1
assert coordinator._recovery_task is None
def test_completed_operations_leave_no_owned_tasks() -> None:
async def scenario() -> None:
coordinator, *_ = create_coordinator()
await coordinator.hydrate_once()
await coordinator.recover_after_ack()
pending_owned_tasks = {
task.get_name()
for task in asyncio.all_tasks()
if task is not asyncio.current_task()
and not task.done()
and task.get_name()
in {
HYDRATION_TASK_NAME,
RECOVERY_TASK_NAME,
}
}
assert pending_owned_tasks == set()
assert coordinator._hydration_task is None
assert coordinator._recovery_task is None
asyncio.run(scenario())

View File

@@ -24,10 +24,15 @@ from src.market_data.acquisition.consistency.trade_stream_consistency_controller
from src.market_data.acquisition.consistency.trade_stream_state_store import (
TradeStreamStateStore,
)
from src.market_data.acquisition.consistency.trade_stream_state import (
TradeStreamState,
)
from src.market_data.acquisition.exceptions import (
WebSocketControlMessageError,
WebSocketMessageDecodeError,
WebSocketMessageRoutingError,
WebSocketStartupMarketBufferOverflowError,
WebSocketSubscriptionAckTimeoutError,
WebSocketTransportError,
)
from src.market_data.acquisition.models.trade import Trade
@@ -37,6 +42,9 @@ from src.market_data.acquisition.runtime.runtime_events import (
DisconnectedEvent,
MessageReceivedEvent,
)
from src.market_data.acquisition.recovery.trade_recovery_result import (
TradeRecoveryResult,
)
from src.market_data.acquisition.runtime.live_processing_gate import (
RuntimeLiveProcessingGate,
)
@@ -66,6 +74,7 @@ from src.market_data.acquisition.runtime.websocket_protocol import (
from src.market_data.acquisition.trade_stream_acquisition_service import (
TradeStreamAcquisitionService,
)
from src.market_data.acquisition.symbols import normalize_symbol
SYMBOL = "BTC/USD_LEVERAGE"
@@ -153,6 +162,8 @@ class FakeTransport:
self._incoming = deque(incoming)
self._message_available = asyncio.Event()
self.receive_calls = 0
self.active_receivers = 0
self.max_active_receivers = 0
self.probe_calls = 0
self.waiting = asyncio.Event()
@@ -171,24 +182,32 @@ class FakeTransport:
async def receive(self) -> str | bytes:
self.receive_calls += 1
self._calls.append("transport.receive")
self.active_receivers += 1
self.max_active_receivers = max(
self.max_active_receivers,
self.active_receivers,
)
try:
while not self._incoming:
self.waiting.set()
await self._message_available.wait()
self._message_available.clear()
except asyncio.CancelledError:
self._calls.append(
"transport.receive.cancelled",
)
raise
try:
while not self._incoming:
self.waiting.set()
await self._message_available.wait()
self._message_available.clear()
except asyncio.CancelledError:
self._calls.append(
"transport.receive.cancelled",
)
raise
result = self._incoming.popleft()
result = self._incoming.popleft()
if isinstance(result, BaseException):
raise result
if isinstance(result, BaseException):
raise result
return result
return result
finally:
self.active_receivers -= 1
def feed(
self,
@@ -302,9 +321,9 @@ class FakeReconnectRecoveryCoordinator:
self._symbols = tuple(
sorted(
{
symbol.strip()
normalize_symbol(symbol)
for symbol in valid_symbols
if symbol.strip()
if normalize_symbol(symbol)
}
)
)
@@ -369,6 +388,80 @@ class FakeReconnectRecoveryCoordinator:
raise self._error
class FakeStartupRecoveryCoordinator:
def __init__(
self,
*,
calls: list[str],
live_processing_gate: RuntimeLiveProcessingGate,
symbols: tuple[str, ...],
hydration_error: Exception | None = None,
recovery_error: Exception | None = None,
hydration_release: asyncio.Event | None = None,
recovery_release: asyncio.Event | None = None,
) -> None:
self._calls = calls
self._live_processing_gate = live_processing_gate
self._symbols = tuple(
sorted(
{
normalize_symbol(symbol)
for symbol in symbols
if normalize_symbol(symbol)
}
)
)
self._hydration_error = hydration_error
self._recovery_error = recovery_error
self._hydration_release = hydration_release
self._recovery_release = recovery_release
self.hydration_entered = asyncio.Event()
self.recovery_entered = asyncio.Event()
self.hydrate_calls = 0
self.recover_calls = 0
@property
def live_processing_gate(
self,
) -> RuntimeLiveProcessingGate:
return self._live_processing_gate
@property
def symbols(self) -> tuple[str, ...]:
return self._symbols
async def hydrate_once(self) -> tuple[TradeStreamState, ...]:
self.hydrate_calls += 1
self._calls.append("startup_recovery.hydrate")
self.hydration_entered.set()
if self._hydration_release is not None:
await self._hydration_release.wait()
if self._hydration_error is not None:
raise self._hydration_error
return ()
async def recover_after_ack(self) -> tuple[TradeRecoveryResult, ...]:
self.recover_calls += 1
self._calls.append("startup_recovery.recover")
self.recovery_entered.set()
if not self._live_processing_gate.locked:
raise AssertionError(
"Production Runtime должен удерживать startup gate."
)
if self._recovery_release is not None:
await self._recovery_release.wait()
if self._recovery_error is not None:
raise self._recovery_error
return ()
class FakeTradeStreamService:
def __init__(
self,
@@ -623,6 +716,15 @@ class RuntimeDependencies:
scheduler_uses_different_transport: bool = False,
symbols: tuple[str, ...] = (SYMBOL,),
recovery_symbols: tuple[str, ...] | None = None,
startup_recovery_enabled: bool = False,
startup_hydration_error: Exception | None = None,
startup_recovery_error: Exception | None = None,
startup_hydration_release: asyncio.Event | None = None,
startup_recovery_release: asyncio.Event | None = None,
startup_recovery_symbols: tuple[str, ...] | None = None,
startup_recovery_uses_different_gate: bool = False,
subscription_ack_timeout_seconds: float = 10.0,
startup_market_buffer_capacity: int = 10_000,
) -> None:
self.calls: list[str] = []
self.session = FakeSession(
@@ -663,6 +765,27 @@ class RuntimeDependencies:
else recovery_symbols
),
)
self.startup_recovery = (
FakeStartupRecoveryCoordinator(
calls=self.calls,
live_processing_gate=(
RuntimeLiveProcessingGate()
if startup_recovery_uses_different_gate
else self.live_processing_gate
),
symbols=(
symbols
if startup_recovery_symbols is None
else startup_recovery_symbols
),
hydration_error=startup_hydration_error,
recovery_error=startup_recovery_error,
hydration_release=startup_hydration_release,
recovery_release=startup_recovery_release,
)
if startup_recovery_enabled
else None
)
self.supervisor = FakeRuntimeSupervisor(
calls=self.calls,
start_error=supervisor_start_error,
@@ -711,6 +834,13 @@ class RuntimeDependencies:
runtime_supervisor=self.supervisor,
runtime_scheduler=self.scheduler,
symbols=symbols,
startup_recovery_coordinator=self.startup_recovery,
subscription_ack_timeout_seconds=(
subscription_ack_timeout_seconds
),
startup_market_buffer_capacity=(
startup_market_buffer_capacity
),
)
@@ -773,6 +903,21 @@ def test_rejects_invalid_symbols(
)
def test_canonicalizes_case_variants_across_runtime_graph() -> None:
dependencies = RuntimeDependencies(
symbols=(
f" {SYMBOL.lower()} ",
SYMBOL,
),
startup_recovery_enabled=True,
)
assert dependencies.runtime._symbols == (SYMBOL,)
assert dependencies.reconnect_recovery.symbols == (SYMBOL,)
assert dependencies.startup_recovery is not None
assert dependencies.startup_recovery.symbols == (SYMBOL,)
def test_rejects_symbols_different_from_recovery_coordinator() -> None:
with pytest.raises(
ValueError,
@@ -804,6 +949,577 @@ def test_rejects_scheduler_with_different_transport() -> None:
)
def test_rejects_startup_recovery_with_different_gate() -> None:
with pytest.raises(
ValueError,
match="startup recovery coordinator must share one",
):
RuntimeDependencies(
startup_recovery_enabled=True,
startup_recovery_uses_different_gate=True,
)
def test_rejects_startup_recovery_with_different_symbols() -> None:
with pytest.raises(
ValueError,
match="startup recovery coordinator must use the same symbols",
):
RuntimeDependencies(
startup_recovery_enabled=True,
startup_recovery_symbols=(ETH_SYMBOL,),
)
@pytest.mark.parametrize(
("timeout", "error_type"),
[
(True, TypeError),
("10", TypeError),
(0.0, ValueError),
(-1.0, ValueError),
(float("inf"), ValueError),
(float("nan"), ValueError),
],
)
def test_rejects_invalid_subscription_ack_timeout(
timeout: object,
error_type: type[Exception],
) -> None:
with pytest.raises(error_type):
RuntimeDependencies(
subscription_ack_timeout_seconds=timeout, # type: ignore[arg-type]
)
@pytest.mark.parametrize(
("capacity", "error_type"),
[
(True, TypeError),
(1.0, TypeError),
(0, ValueError),
(-1, ValueError),
],
)
def test_rejects_invalid_startup_market_buffer_capacity(
capacity: object,
error_type: type[Exception],
) -> None:
with pytest.raises(error_type):
RuntimeDependencies(
startup_market_buffer_capacity=capacity, # type: ignore[arg-type]
)
def test_startup_recovery_preserves_boundary_order_and_one_receiver() -> None:
first_document = {
"destination": "internal.trade",
"payload": {
"symbol": SYMBOL,
"sequence": 1,
},
}
second_document = {
"destination": "internal.trade",
"payload": {
"symbol": SYMBOL,
"sequence": 2,
},
}
post_ack_document = {
"destination": "internal.trade",
"payload": {
"symbol": SYMBOL,
"sequence": 3,
},
}
async def scenario() -> RuntimeDependencies:
recovery_release = asyncio.Event()
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
startup_recovery_release=recovery_release,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
await dependencies.service.subscribe_entered.wait()
correlation_id = (
dependencies.service.subscribe_correlation_ids[0]
)
assert isinstance(correlation_id, str)
dependencies.transport.feed(json.dumps(first_document))
dependencies.transport.feed(json.dumps(second_document))
dependencies.transport.feed(
make_control_message(
correlation_id=correlation_id,
)
)
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
await startup_recovery.recovery_entered.wait()
assert dependencies.live_processing_gate.locked is True
assert dependencies.service.documents == []
assert dependencies.transport.receive_calls == 3
assert dependencies.supervisor.start_calls == 0
assert dependencies.scheduler.start_calls == 0
dependencies.transport.feed(
json.dumps(post_ack_document),
)
await asyncio.sleep(0)
assert dependencies.transport.receive_calls == 3
recovery_release.set()
await dependencies.scheduler.started.wait()
await wait_until(
lambda: dependencies.service.documents
== [
first_document,
second_document,
post_ack_document,
],
)
await wait_until(
lambda: dependencies.transport.receive_calls == 5,
)
assert dependencies.service.documents == [
first_document,
second_document,
post_ack_document,
]
assert dependencies.transport.max_active_receivers == 1
await dependencies.runtime.stop()
await runtime_task
return dependencies
dependencies = asyncio.run(scenario())
assert dependencies.calls.index(
"startup_recovery.hydrate"
) < dependencies.calls.index("session.start")
assert dependencies.calls.index(
"service.subscribe"
) < dependencies.calls.index("startup_recovery.recover")
assert dependencies.calls.index(
"startup_recovery.recover"
) < dependencies.calls.index("service.handle_message")
assert dependencies.calls.index(
"service.handle_message"
) < dependencies.calls.index("supervisor.start")
def test_startup_ack_timeout_is_terminal() -> None:
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
subscription_ack_timeout_seconds=0.01,
)
with pytest.raises(
WebSocketSubscriptionAckTimeoutError,
match="Истекло время",
):
await dependencies.runtime.run()
return dependencies
dependencies = asyncio.run(scenario())
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
assert startup_recovery.recover_calls == 0
assert dependencies.live_processing_gate.failed is True
assert dependencies.supervisor.start_calls == 0
assert dependencies.scheduler.start_calls == 0
assert dependencies.transport.max_active_receivers == 1
assert dependencies.runtime.state is (
TradeStreamProductionRuntimeState.FAILED
)
def test_startup_hydration_failure_prevents_network_io() -> None:
hydration_error = RuntimeError("startup hydration failed")
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
startup_hydration_error=hydration_error,
)
with pytest.raises(
RuntimeError,
match="startup hydration failed",
) as error_info:
await dependencies.runtime.run()
assert error_info.value is hydration_error
return dependencies
dependencies = asyncio.run(scenario())
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
assert startup_recovery.hydrate_calls == 1
assert startup_recovery.recover_calls == 0
assert dependencies.session.start_calls == 0
assert dependencies.service.subscribe_calls == []
assert dependencies.live_processing_gate.failed is True
assert dependencies.runtime.state is (
TradeStreamProductionRuntimeState.FAILED
)
def test_persistent_startup_connect_failure_marks_gate_failed() -> None:
connect_error = RuntimeError("persistent connect failed")
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
start_error=connect_error,
)
with pytest.raises(
RuntimeError,
match="persistent connect failed",
) as error_info:
await dependencies.runtime.run()
assert error_info.value is connect_error
return dependencies
dependencies = asyncio.run(scenario())
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
assert startup_recovery.hydrate_calls == 1
assert startup_recovery.recover_calls == 0
assert dependencies.service.subscribe_calls == []
assert dependencies.live_processing_gate.failed is True
assert dependencies.runtime.state is (
TradeStreamProductionRuntimeState.FAILED
)
def test_negative_startup_ack_is_terminal_before_recovery() -> None:
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
await dependencies.service.subscribe_entered.wait()
correlation_id = (
dependencies.service.subscribe_correlation_ids[0]
)
assert isinstance(correlation_id, str)
dependencies.transport.feed(
make_control_message(
correlation_id=correlation_id,
status="ERROR",
)
)
with pytest.raises(
WebSocketControlMessageError,
match="отклонил",
):
await runtime_task
return dependencies
dependencies = asyncio.run(scenario())
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
assert startup_recovery.recover_calls == 0
assert dependencies.live_processing_gate.failed is True
assert dependencies.supervisor.start_calls == 0
@pytest.mark.parametrize(
("control_document_factory", "error_match"),
[
(
lambda correlation_id: {
"correlationId": "unknown-request",
"destination": "trades.subscribe",
"status": "OK",
},
"неизвестным correlationId",
),
(
lambda correlation_id: {
"correlationId": correlation_id,
"destination": "trades.subscribe",
},
"непустой строковый status",
),
(
lambda correlation_id: {
"correlationId": correlation_id,
"destination": "unknown.destination",
"status": "OK",
},
"неизвестным destination",
),
],
)
def test_invalid_startup_control_message_is_terminal(
control_document_factory: Callable[[str], object],
error_match: str,
) -> None:
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
await dependencies.service.subscribe_entered.wait()
correlation_id = (
dependencies.service.subscribe_correlation_ids[0]
)
assert isinstance(correlation_id, str)
dependencies.transport.feed(
json.dumps(
control_document_factory(correlation_id),
)
)
with pytest.raises(
WebSocketMessageRoutingError,
match=error_match,
):
await runtime_task
return dependencies
dependencies = asyncio.run(scenario())
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
assert startup_recovery.recover_calls == 0
assert dependencies.service.documents == []
assert dependencies.live_processing_gate.failed is True
assert dependencies.reconnect_recovery.calls == []
def test_startup_transport_error_is_terminal_without_reconnect() -> None:
transport_error = WebSocketTransportError(
"startup transport failed",
)
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
incoming=(transport_error,),
)
with pytest.raises(
WebSocketTransportError,
match="startup transport failed",
) as error_info:
await dependencies.runtime.run()
assert error_info.value is transport_error
return dependencies
dependencies = asyncio.run(scenario())
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
assert startup_recovery.recover_calls == 0
assert dependencies.reconnect_recovery.calls == []
assert dependencies.live_processing_gate.failed is True
assert dependencies.supervisor.start_calls == 0
def test_startup_market_buffer_overflow_is_terminal() -> None:
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
startup_market_buffer_capacity=1,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
await dependencies.service.subscribe_entered.wait()
dependencies.transport.feed(MARKET_MESSAGE)
dependencies.transport.feed(MARKET_MESSAGE)
with pytest.raises(
WebSocketStartupMarketBufferOverflowError,
match="Переполнен буфер",
):
await runtime_task
return dependencies
dependencies = asyncio.run(scenario())
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
assert startup_recovery.recover_calls == 0
assert dependencies.service.documents == []
assert dependencies.live_processing_gate.failed is True
def test_startup_recovery_failure_prevents_fifo_drain() -> None:
recovery_error = RuntimeError("startup recovery failed")
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
startup_recovery_error=recovery_error,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
await dependencies.service.subscribe_entered.wait()
correlation_id = (
dependencies.service.subscribe_correlation_ids[0]
)
assert isinstance(correlation_id, str)
dependencies.transport.feed(MARKET_MESSAGE)
dependencies.transport.feed(
make_control_message(
correlation_id=correlation_id,
)
)
with pytest.raises(
RuntimeError,
match="startup recovery failed",
) as error_info:
await runtime_task
assert error_info.value is recovery_error
return dependencies
dependencies = asyncio.run(scenario())
assert dependencies.service.documents == []
assert dependencies.live_processing_gate.failed is True
assert dependencies.supervisor.start_calls == 0
assert dependencies.scheduler.start_calls == 0
def test_stop_during_startup_ack_wait_releases_boundary() -> None:
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
await dependencies.service.subscribe_entered.wait()
await dependencies.transport.waiting.wait()
await asyncio.wait_for(
dependencies.runtime.stop(),
timeout=1.0,
)
await runtime_task
return dependencies
dependencies = asyncio.run(scenario())
assert dependencies.live_processing_gate.locked is False
assert dependencies.live_processing_gate.failed is False
assert dependencies.supervisor.start_calls == 0
assert dependencies.transport.active_receivers == 0
assert dependencies.runtime.state is (
TradeStreamProductionRuntimeState.STOPPED
)
def test_stop_during_startup_hydration_prevents_connection() -> None:
async def scenario() -> RuntimeDependencies:
hydration_release = asyncio.Event()
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
startup_hydration_release=hydration_release,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
await startup_recovery.hydration_entered.wait()
await asyncio.wait_for(
dependencies.runtime.stop(),
timeout=1.0,
)
await runtime_task
return dependencies
dependencies = asyncio.run(scenario())
assert dependencies.session.start_calls == 0
assert dependencies.live_processing_gate.failed is False
assert dependencies.runtime.state is (
TradeStreamProductionRuntimeState.STOPPED
)
def test_stop_during_startup_recovery_prevents_fifo_drain() -> None:
async def scenario() -> RuntimeDependencies:
recovery_release = asyncio.Event()
dependencies = RuntimeDependencies(
startup_recovery_enabled=True,
startup_recovery_release=recovery_release,
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)
await dependencies.service.subscribe_entered.wait()
correlation_id = (
dependencies.service.subscribe_correlation_ids[0]
)
assert isinstance(correlation_id, str)
dependencies.transport.feed(MARKET_MESSAGE)
dependencies.transport.feed(
make_control_message(
correlation_id=correlation_id,
)
)
startup_recovery = dependencies.startup_recovery
assert startup_recovery is not None
await startup_recovery.recovery_entered.wait()
await asyncio.wait_for(
dependencies.runtime.stop(),
timeout=1.0,
)
await runtime_task
return dependencies
dependencies = asyncio.run(scenario())
assert dependencies.service.documents == []
assert dependencies.live_processing_gate.locked is False
assert dependencies.live_processing_gate.failed is False
assert dependencies.supervisor.start_calls == 0
assert dependencies.runtime.state is (
TradeStreamProductionRuntimeState.STOPPED
)
def test_locked_gate_does_not_leave_runtime_partially_started() -> None:
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies()
@@ -893,7 +1609,12 @@ def test_scheduler_claim_blocks_external_start_during_runtime_startup() -> None:
def test_startup_receive_and_shutdown_order() -> None:
async def scenario() -> RuntimeDependencies:
dependencies = RuntimeDependencies()
dependencies = RuntimeDependencies(
symbols=(
f" {SYMBOL.lower()} ",
SYMBOL,
),
)
runtime_task = asyncio.create_task(
dependencies.runtime.run(),
)

View File

@@ -1,6 +1,7 @@
from __future__ import annotations
import asyncio
from collections.abc import Awaitable
from typing import Any
import pytest
@@ -73,6 +74,12 @@ class FakeConnection:
async def recv(self) -> str | bytes:
return ""
async def ping(self) -> Awaitable[float]:
async def wait_for_pong() -> float:
return 0.001
return wait_for_pong()
class RecordingConnector:
def __init__(

View File

@@ -13,9 +13,15 @@ import pytest
from src.market_data.acquisition.adapters.dzengi.rest import (
DzengiTradesDocumentSource,
)
from src.market_data.acquisition.checkpoint.trade_stream_state_hydrator import (
TradeStreamStateHydratorProtocol,
)
from src.market_data.acquisition.consistency.trade_observation_sink_protocol import (
TradeObservationSinkProtocol,
)
from src.market_data.acquisition.consistency.trade_stream_state import (
TradeStreamState,
)
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
@@ -44,6 +50,12 @@ from src.market_data.acquisition.runtime.runtime_events import (
from src.market_data.acquisition.runtime.runtime_recovery_protocol import (
RuntimeRecoveryProtocol,
)
from src.market_data.acquisition.recovery.trade_recovery_result import (
TradeRecoveryResult,
)
from src.market_data.acquisition.runtime.runtime_startup_recovery_coordinator import (
RuntimeStartupRecoveryProtocol,
)
from src.market_data.acquisition.runtime.scheduler import (
RuntimeSchedulerProtocol,
)
@@ -66,6 +78,10 @@ from src.market_data.acquisition.trade_stream_runtime_composition import (
TradeStreamRuntimeComposition,
build_trade_stream_runtime_composition,
)
from src.market_data.storage.contracts import (
PersistentTradeCheckpoint,
TradeCheckpointStorageProtocol,
)
SYMBOL = "BTC/USD_LEVERAGE"
@@ -326,17 +342,152 @@ class RecordingTradeObservationSink:
) -> None:
self._fail_on_trade_id = fail_on_trade_id
self.observations: list[Trade] = []
self.accepted: list[tuple[Trade, Trade | None]] = []
self.duplicates: list[Trade] = []
def persist(
def persist_accepted(
self,
trade: Trade,
*,
expected_trade: Trade | None,
) -> None:
self.observations.append(trade)
self.accepted.append((trade, expected_trade))
if trade.trade_id == self._fail_on_trade_id:
raise RuntimeError("storage failed")
def persist_duplicate(
self,
trade: Trade,
) -> None:
self.observations.append(trade)
self.duplicates.append(trade)
if trade.trade_id == self._fail_on_trade_id:
raise RuntimeError("storage failed")
class RecordingCheckpointStorage:
"""Checkpoint storage, фиксирующий нежелательный I/O."""
def __init__(self) -> None:
self.calls: list[str] = []
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
self.calls.append("load_checkpoint")
raise AssertionError("Composition не должна читать checkpoint.")
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
self.calls.append("load_checkpoint_tail")
raise AssertionError("Composition не должна читать Trade tail.")
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
self.calls.append("load_latest_trade_tail")
raise AssertionError("Composition не должна читать Trade tail.")
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
self.calls.append("adopt_existing_trade_as_checkpoint")
raise AssertionError("Composition не должна создавать checkpoint.")
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
self.calls.append("store_trade_and_advance_checkpoint")
raise AssertionError("Composition не должна записывать checkpoint.")
class CheckpointBackedStorage:
"""Checkpoint storage для проверки общего канонического ключа."""
def __init__(
self,
checkpoint: PersistentTradeCheckpoint,
) -> None:
self._checkpoint = checkpoint
self.calls: list[tuple[object, ...]] = []
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
self.calls.append(("load_checkpoint", venue, symbol))
return self._checkpoint
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
self.calls.append(
(
"load_checkpoint_tail",
venue,
checkpoint,
limit,
)
)
return (checkpoint.trade,)
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
raise AssertionError("Latest Trade tail не должен запрашиваться.")
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
raise AssertionError("Checkpoint не должен создаваться.")
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
raise AssertionError("Checkpoint не должен записываться.")
@dataclass(slots=True)
class CompositionDependencies:
session: FakeSession
@@ -352,6 +503,7 @@ class CompositionDependencies:
def create_composition(
*,
symbols: tuple[str, ...] = (SYMBOL,),
trade: Trade | None = None,
recovery_document: object = (),
heartbeat_timeout_seconds: float = 10.0,
@@ -359,6 +511,8 @@ def create_composition(
max_recovery_window_ms: int = 3_599_999,
probe_results: tuple[bool, ...] = (True,),
trade_observation_sink: TradeObservationSinkProtocol | None = None,
checkpoint_storage: TradeCheckpointStorageProtocol | None = None,
checkpoint_venue: str | None = None,
) -> tuple[
TradeStreamRuntimeComposition,
CompositionDependencies,
@@ -390,7 +544,7 @@ def create_composition(
recovery_document_source=(
dependencies.recovery_document_source
),
symbols=(SYMBOL,),
symbols=symbols,
heartbeat_timeout_seconds=heartbeat_timeout_seconds,
scheduler_interval_seconds=scheduler_interval_seconds,
trade_observation_sink=trade_observation_sink,
@@ -400,6 +554,8 @@ def create_composition(
dependencies.recovery_end_time_clock
),
scheduler_sleep=dependencies.scheduler_sleep,
checkpoint_storage=checkpoint_storage,
checkpoint_venue=checkpoint_venue,
)
return (
@@ -467,6 +623,196 @@ def test_components_implement_public_protocols() -> None:
)
def test_checkpoint_components_are_absent_without_configuration() -> None:
composition, *_ = create_composition()
assert composition.state_hydrator is None
assert composition.runtime_startup_recovery_coordinator is None
def test_checkpoint_components_implement_public_protocols() -> None:
storage = RecordingCheckpointStorage()
composition, *_ = create_composition(
checkpoint_storage=storage,
checkpoint_venue="Dzengi",
)
assert isinstance(storage, TradeCheckpointStorageProtocol)
assert isinstance(
composition.state_hydrator,
TradeStreamStateHydratorProtocol,
)
assert isinstance(
composition.runtime_startup_recovery_coordinator,
RuntimeStartupRecoveryProtocol,
)
def test_checkpoint_runtime_reuses_shared_dependency_graph() -> None:
storage = RecordingCheckpointStorage()
composition, dependencies = create_composition(
checkpoint_storage=storage,
checkpoint_venue="Dzengi",
)
state_hydrator = composition.state_hydrator
startup_recovery = (
composition.runtime_startup_recovery_coordinator
)
assert state_hydrator is not None
assert startup_recovery is not None
assert state_hydrator._state_store is composition.state_store
assert startup_recovery._state_hydrator is state_hydrator
assert (
startup_recovery._recovery_coordinator
is composition.runtime_recovery_coordinator
)
assert (
startup_recovery.live_processing_gate
is composition.live_processing_gate
)
assert (
startup_recovery.symbols
== composition.runtime_reconnect_recovery_coordinator.symbols
== (SYMBOL,)
)
assert (
startup_recovery._clock
is dependencies.recovery_end_time_clock
)
assert (
composition.runtime_reconnect_recovery_coordinator._clock
is dependencies.recovery_end_time_clock
)
def test_lowercase_symbol_uses_one_key_for_hydration_and_recovery() -> None:
checkpoint = PersistentTradeCheckpoint(
venue="dzengi",
trade=make_trade(),
revision=1,
updated_at=CHECKPOINT_TIME,
)
storage = CheckpointBackedStorage(checkpoint)
composition, dependencies = create_composition(
symbols=(f" {SYMBOL.lower()} ",),
recovery_document=[],
checkpoint_storage=storage,
checkpoint_venue="Dzengi",
)
startup_recovery = (
composition.runtime_startup_recovery_coordinator
)
assert startup_recovery is not None
async def scenario() -> tuple[
tuple[TradeStreamState, ...],
tuple[TradeRecoveryResult, ...],
]:
states = await startup_recovery.hydrate_once()
async with composition.live_processing_gate:
results = await startup_recovery.recover_after_ack()
return states, results
states, results = asyncio.run(scenario())
assert startup_recovery.symbols == (SYMBOL,)
assert (
composition.runtime_reconnect_recovery_coordinator.symbols
== (SYMBOL,)
)
assert states[0].symbol == SYMBOL
assert composition.state_store.get(SYMBOL) is states[0]
assert composition.state_store.contains(SYMBOL.lower()) is False
assert results[0].symbol == SYMBOL
assert storage.calls[0] == (
"load_checkpoint",
"dzengi",
SYMBOL,
)
assert dependencies.recovery_document_source.calls
assert dependencies.recovery_document_source.calls[0][0] == SYMBOL
def test_case_variant_duplicates_collapse_before_hydration_worker() -> None:
checkpoint = PersistentTradeCheckpoint(
venue="dzengi",
trade=make_trade(),
revision=1,
updated_at=CHECKPOINT_TIME,
)
storage = CheckpointBackedStorage(checkpoint)
composition, _ = create_composition(
symbols=(
SYMBOL.lower(),
SYMBOL,
f" {SYMBOL.lower()} ",
),
checkpoint_storage=storage,
checkpoint_venue="dzengi",
)
startup_recovery = (
composition.runtime_startup_recovery_coordinator
)
assert startup_recovery is not None
states = asyncio.run(startup_recovery.hydrate_once())
assert startup_recovery.symbols == (SYMBOL,)
assert len(states) == 1
assert states[0].symbol == SYMBOL
assert [call[0] for call in storage.calls] == [
"load_checkpoint",
"load_checkpoint_tail",
]
@pytest.mark.parametrize(
("checkpoint_storage", "checkpoint_venue"),
[
(RecordingCheckpointStorage(), None),
(None, "dzengi"),
],
)
def test_partial_checkpoint_configuration_is_rejected_without_io(
checkpoint_storage: TradeCheckpointStorageProtocol | None,
checkpoint_venue: str | None,
) -> None:
with pytest.raises(
ValueError,
match="должны быть переданы вместе",
):
create_composition(
checkpoint_storage=checkpoint_storage,
checkpoint_venue=checkpoint_venue,
)
if isinstance(checkpoint_storage, RecordingCheckpointStorage):
assert checkpoint_storage.calls == []
def test_checkpoint_composition_has_no_io_or_background_tasks() -> None:
storage = RecordingCheckpointStorage()
composition, dependencies = create_composition(
checkpoint_storage=storage,
checkpoint_venue="dzengi",
)
startup_recovery = (
composition.runtime_startup_recovery_coordinator
)
assert startup_recovery is not None
assert storage.calls == []
assert dependencies.recovery_document_source.calls == []
assert dependencies.recovery_end_time_clock.calls == 0
assert startup_recovery._hydration_task is None
assert startup_recovery._recovery_task is None
def test_external_dependencies_are_reused() -> None:
composition, dependencies = create_composition()
@@ -548,6 +894,11 @@ def test_live_and_recovery_share_optional_persistence_sink() -> None:
trade.trade_id
for trade in sink.observations
] == [100, recovered_trade_id]
assert sink.accepted == [
(live_trade, None),
(recovery_result.last_trade, live_trade),
]
assert sink.duplicates == []
def test_recovery_duplicate_updates_persistence_without_checkpoint_change(
@@ -589,6 +940,8 @@ def test_recovery_duplicate_updates_persistence_without_checkpoint_change(
assert len(sink.observations) == 2
assert sink.observations[0] is live_trade
assert sink.observations[1].source == "dzengi"
assert sink.accepted == [(live_trade, None)]
assert sink.duplicates == [sink.observations[1]]
assert state.last_trade is live_trade

View File

@@ -0,0 +1,260 @@
from __future__ import annotations
from dataclasses import replace
from datetime import datetime, timezone
from decimal import Decimal
from typing import Any
import pytest
from src.market_data.acquisition.models.trade import (
Trade,
TradeAggressorSide,
)
from src.market_data.storage import (
MarketDataCheckpointConflictError,
MarketDataCheckpointIntegrityError,
MarketDataStorageConflictError,
MarketDataStorageError,
PersistentTradeCheckpoint,
TradeCheckpointStorageProtocol,
)
VENUE = "DZENGI"
SYMBOL = "BTC/USD_LEVERAGE"
EXECUTED_AT = datetime(2026, 8, 1, 10, 0, tzinfo=timezone.utc)
UPDATED_AT = datetime(2026, 8, 1, 10, 1, tzinfo=timezone.utc)
def make_trade(
*,
symbol: str = SYMBOL,
trade_id: int = 123,
executed_at: datetime = EXECUTED_AT,
) -> 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="dzengi_websocket_trade",
)
def make_checkpoint() -> PersistentTradeCheckpoint:
return PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(),
revision=7,
updated_at=UPDATED_AT,
)
class RecordingCheckpointStorage:
def __init__(self) -> None:
self.checkpoint = make_checkpoint()
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
return self.checkpoint
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
return (checkpoint.trade,)
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
return (self.checkpoint.trade,)
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
return replace(
self.checkpoint,
venue=venue,
trade=trade,
revision=1,
)
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
return replace(
self.checkpoint,
trade=trade,
revision=self.checkpoint.revision + 1,
updated_at=observed_at,
)
def test_checkpoint_preserves_full_trade_and_durable_identity() -> None:
checkpoint = make_checkpoint()
assert checkpoint.venue == VENUE
assert checkpoint.trade is not None
assert checkpoint.revision == 7
assert checkpoint.updated_at == UPDATED_AT
assert checkpoint.checkpoint_schema_version == 1
assert checkpoint.identity == (
VENUE,
SYMBOL,
123,
EXECUTED_AT,
)
def test_checkpoint_storage_protocol_is_runtime_checkable() -> None:
assert isinstance(
RecordingCheckpointStorage(),
TradeCheckpointStorageProtocol,
)
@pytest.mark.parametrize("venue", ("", " ", "\t"))
def test_checkpoint_rejects_empty_venue(venue: str) -> None:
with pytest.raises(ValueError, match="venue"):
PersistentTradeCheckpoint(
venue=venue,
trade=make_trade(),
revision=1,
updated_at=UPDATED_AT,
)
def test_checkpoint_rejects_non_string_venue() -> None:
with pytest.raises(TypeError, match="venue"):
PersistentTradeCheckpoint(
venue=123, # type: ignore[arg-type]
trade=make_trade(),
revision=1,
updated_at=UPDATED_AT,
)
def test_checkpoint_rejects_non_trade_payload() -> None:
with pytest.raises(TypeError, match="Canonical Trade"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=object(), # type: ignore[arg-type]
revision=1,
updated_at=UPDATED_AT,
)
def test_checkpoint_rejects_empty_trade_symbol() -> None:
with pytest.raises(ValueError, match="trade.symbol"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(symbol=" "),
revision=1,
updated_at=UPDATED_AT,
)
def test_checkpoint_rejects_naive_trade_time() -> None:
with pytest.raises(ValueError, match="trade.executed_at"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(
executed_at=datetime(2026, 8, 1, 10, 0),
),
revision=1,
updated_at=UPDATED_AT,
)
@pytest.mark.parametrize("trade_id", (-2_147_483_649, 2_147_483_648))
def test_checkpoint_rejects_trade_id_outside_signed_range(
trade_id: int,
) -> None:
with pytest.raises(ValueError, match="signed 32-bit"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(trade_id=trade_id),
revision=1,
updated_at=UPDATED_AT,
)
@pytest.mark.parametrize("revision", (0, -1))
def test_checkpoint_rejects_non_positive_revision(revision: int) -> None:
with pytest.raises(ValueError, match="revision"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(),
revision=revision,
updated_at=UPDATED_AT,
)
@pytest.mark.parametrize("revision", (True, 1.5, "1", None))
def test_checkpoint_rejects_non_integer_revision(
revision: Any,
) -> None:
with pytest.raises(TypeError, match="revision"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(),
revision=revision,
updated_at=UPDATED_AT,
)
def test_checkpoint_rejects_naive_updated_at() -> None:
with pytest.raises(ValueError, match="updated_at"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(),
revision=1,
updated_at=datetime(2026, 8, 1, 10, 1),
)
@pytest.mark.parametrize("version", (0, -1))
def test_checkpoint_rejects_non_positive_schema_version(
version: int,
) -> None:
with pytest.raises(ValueError, match="checkpoint_schema_version"):
PersistentTradeCheckpoint(
venue=VENUE,
trade=make_trade(),
revision=1,
updated_at=UPDATED_AT,
checkpoint_schema_version=version,
)
def test_checkpoint_errors_preserve_storage_hierarchy() -> None:
assert issubclass(
MarketDataCheckpointConflictError,
MarketDataStorageConflictError,
)
assert issubclass(
MarketDataCheckpointIntegrityError,
MarketDataStorageError,
)

View File

@@ -129,6 +129,9 @@ class PartitionCursor:
if normalized.startswith("ALTER TABLE") and "ADD CONSTRAINT" in normalized:
return
if normalized.startswith("ALTER TABLE") and "DROP CONSTRAINT" in normalized:
return
if normalized.startswith("WITH moved_rows AS"):
self.rowcount = self._connection.moved_row_count
return
@@ -481,6 +484,70 @@ def test_manager_creates_partition_moves_default_rows_and_registers_it() -> None
assert all("::timestamptz" in statement for statement, _ in ddl_calls)
def test_trade_partition_rebuilds_checkpoint_foreign_key_around_move() -> None:
_, connection, provider = _dependencies()
manager = PostgresMarketDataPartitionManager(
connection_provider=provider
)
manager.ensure_month_partition(
data_type=MarketDataPartitionType.TRADES,
month=JULY,
)
statements = tuple(statement for statement, _ in connection.calls)
parent_lock_index = statements.index(
'LOCK TABLE "market_data"."trades" '
"IN SHARE ROW EXCLUSIVE MODE"
)
default_lock_index = statements.index(
'LOCK TABLE "market_data"."trades_default" '
"IN ACCESS EXCLUSIVE MODE"
)
drop_index = statements.index(
'ALTER TABLE "market_data"."trade_stream_checkpoints" '
'DROP CONSTRAINT "trade_stream_checkpoints_trade_fk"'
)
move_index = next(
index
for index, statement in enumerate(statements)
if statement.startswith("WITH moved_rows AS")
)
restore_index = next(
index
for index, statement in enumerate(statements)
if "ADD CONSTRAINT \"trade_stream_checkpoints_trade_fk\"" in statement
)
assert (
parent_lock_index
< default_lock_index
< drop_index
< move_index
< restore_index
)
restored_sql = statements[restore_index]
assert "ON UPDATE NO ACTION ON DELETE NO ACTION" in restored_sql
assert "DEFERRABLE INITIALLY DEFERRED" in restored_sql
def test_non_trade_partition_does_not_touch_checkpoint_foreign_key() -> None:
_, connection, provider = _dependencies()
manager = PostgresMarketDataPartitionManager(
connection_provider=provider
)
manager.ensure_month_partition(
data_type=MarketDataPartitionType.QUOTES,
month=JULY,
)
assert not any(
"trade_stream_checkpoints_trade_fk" in statement
for statement, _ in connection.calls
)
def test_manager_is_idempotent_after_registered_partition_exists() -> None:
_, connection, provider = _dependencies()
manager = PostgresMarketDataPartitionManager(
@@ -718,6 +785,22 @@ def test_retention_drops_complete_month_and_deletes_partial_history() -> None:
for statement, _ in connection.calls
)
statements = tuple(statement for statement, _ in connection.calls)
drop_index = statements.index(
'ALTER TABLE "market_data"."trade_stream_checkpoints" '
'DROP CONSTRAINT "trade_stream_checkpoints_trade_fk"'
)
delete_index = statements.index(
'DELETE FROM "market_data"."trades" '
'WHERE "executed_at" < %s'
)
restore_index = next(
index
for index, statement in enumerate(statements)
if "ADD CONSTRAINT \"trade_stream_checkpoints_trade_fk\"" in statement
)
assert drop_index < delete_index < restore_index
def test_unconfigured_data_type_is_unlimited_and_not_touched() -> None:
database, connection, provider = _dependencies()

View File

@@ -0,0 +1,817 @@
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.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 (
MarketDataCheckpointConflictError,
MarketDataCheckpointIntegrityError,
MarketDataStorageOperationError,
MarketDataStorageValidationError,
PersistentTradeCheckpoint,
PostgresTradeRepository,
TradeCheckpointStorageProtocol,
)
VENUE = "dzengi"
SYMBOL = "BTC/USD_LEVERAGE"
EXECUTED_AT = datetime(2026, 8, 1, 10, 0, tzinfo=timezone.utc)
OBSERVED_AT = EXECUTED_AT + timedelta(seconds=1)
def _trade(
*,
trade_id: int = 100,
executed_at: datetime = EXECUTED_AT,
price: Decimal = Decimal("65000.25"),
source: str = "dzengi_websocket_trade",
) -> Trade:
return Trade(
symbol=SYMBOL,
trade_id=trade_id,
price=price,
quantity=Decimal("0.001"),
executed_at=executed_at,
aggressor_side=TradeAggressorSide.BUY,
source=source,
)
def _trade_row(trade: Trade) -> tuple[object, ...]:
return (
trade.symbol,
trade.trade_id,
trade.executed_at,
trade.price,
trade.quantity,
trade.aggressor_side.value,
trade.source,
1,
)
def _checkpoint_row(
trade: Trade,
*,
revision: int,
updated_at: datetime = OBSERVED_AT,
) -> tuple[object, ...]:
return (
VENUE,
trade.symbol,
trade.trade_id,
trade.executed_at,
revision,
updated_at,
1,
VENUE,
trade.symbol,
trade.trade_id,
trade.executed_at,
trade.price,
trade.quantity,
trade.aggressor_side.value,
trade.source,
1,
)
def _existing_trade_row(trade: Trade) -> tuple[object, ...]:
return (
trade.price,
trade.quantity,
trade.aggressor_side.value,
OBSERVED_AT,
OBSERVED_AT,
[trade.source],
1,
)
@dataclass(frozen=True)
class SqlStep:
starts_with: str
fetchone: object = None
fetchall: tuple[tuple[object, ...], ...] = ()
error: BaseException | None = None
class ScriptedCursor:
def __init__(self, connection: ScriptedConnection) -> None:
self._connection = connection
self._fetchone: object = None
self._fetchall: tuple[tuple[object, ...], ...] = ()
def __enter__(self) -> ScriptedCursor:
return self
def __exit__(self, *args: object) -> None:
return None
def execute(
self,
statement: str,
parameters: tuple[Any, ...],
) -> None:
normalized = " ".join(statement.split())
self._connection.calls.append((normalized, parameters))
if not self._connection.steps:
raise AssertionError(f"Unexpected SQL: {normalized}")
step = self._connection.steps.pop(0)
if not normalized.startswith(step.starts_with):
raise AssertionError(
f"Expected SQL starting with {step.starts_with!r}, "
f"received {normalized!r}."
)
if step.error is not None:
raise step.error
self._fetchone = step.fetchone
self._fetchall = step.fetchall
def fetchone(self) -> object:
return self._fetchone
def fetchall(self) -> tuple[tuple[object, ...], ...]:
return self._fetchall
class ScriptedConnection:
def __init__(self, steps: tuple[SqlStep, ...]) -> None:
self.steps = list(steps)
self.calls: list[tuple[str, tuple[Any, ...]]] = []
self.exit_exception_types: list[type[BaseException] | None] = []
def __enter__(self) -> ScriptedConnection:
return self
def __exit__(
self,
exception_type: type[BaseException] | None,
exception: BaseException | None,
traceback: object,
) -> None:
self.exit_exception_types.append(exception_type)
return None
def cursor(self) -> ScriptedCursor:
return ScriptedCursor(self)
class RecordingProvider:
def __init__(self, connection: ScriptedConnection) -> None:
self.connection = connection
self.calls = 0
def __call__(self) -> ScriptedConnection:
self.calls += 1
return self.connection
def _repository(
*steps: SqlStep,
) -> tuple[PostgresTradeRepository, ScriptedConnection, RecordingProvider]:
connection = ScriptedConnection(steps)
provider = RecordingProvider(connection)
repository = PostgresTradeRepository(
connection_provider=provider,
)
return repository, connection, provider
def test_repository_matches_checkpoint_storage_protocol() -> None:
repository, _, _ = _repository()
assert isinstance(repository, TradeCheckpointStorageProtocol)
def test_load_checkpoint_returns_none_without_persistent_state() -> None:
repository, connection, provider = _repository(
SqlStep(starts_with="SELECT checkpoint.venue"),
)
result = repository.load_checkpoint(
venue=" dzengi ",
symbol=" btc/usd_leverage ",
)
assert result is None
assert provider.calls == 1
assert connection.calls[0][1] == (VENUE, SYMBOL)
assert connection.steps == []
def test_load_checkpoint_restores_exact_canonical_trade() -> None:
trade = _trade()
repository, connection, _ = _repository(
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(trade, revision=7),
),
)
result = repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
assert result == PersistentTradeCheckpoint(
venue=VENUE,
trade=trade,
revision=7,
updated_at=OBSERVED_AT,
)
assert "FOR UPDATE" not in connection.calls[0][0]
def test_load_checkpoint_rejects_missing_durable_trade() -> None:
trade = _trade()
orphan_row = list(_checkpoint_row(trade, revision=1))
orphan_row[7:] = [None] * 9
repository, connection, _ = _repository(
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=tuple(orphan_row),
),
)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="exact Canonical Trade",
):
repository.load_checkpoint(
venue=VENUE,
symbol=SYMBOL,
)
assert connection.exit_exception_types == [
MarketDataCheckpointIntegrityError
]
def test_checkpoint_tail_is_returned_oldest_to_checkpoint_across_rollover(
) -> None:
previous = _trade(
trade_id=SIGNED_TRADE_ID_MAX,
executed_at=EXECUTED_AT,
)
current = _trade(
trade_id=SIGNED_TRADE_ID_MIN,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
checkpoint = PersistentTradeCheckpoint(
venue=VENUE,
trade=current,
revision=2,
updated_at=OBSERVED_AT,
)
repository, connection, _ = _repository(
SqlStep(
starts_with="WITH reference_trade AS",
fetchall=(
_trade_row(current),
_trade_row(previous),
),
),
)
result = repository.load_checkpoint_tail(
venue=VENUE,
checkpoint=checkpoint,
limit=2,
)
assert result == (previous, current)
parameters = connection.calls[0][1]
assert parameters[:4] == (
VENUE,
SYMBOL,
SIGNED_TRADE_ID_MIN,
current.executed_at,
)
assert parameters[-1] == 2
def test_checkpoint_tail_rejects_non_strict_trade_id_order() -> None:
previous = _trade(trade_id=100)
conflicting = _trade(
trade_id=100,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
checkpoint = PersistentTradeCheckpoint(
venue=VENUE,
trade=conflicting,
revision=2,
updated_at=OBSERVED_AT,
)
repository, connection, _ = _repository(
SqlStep(
starts_with="WITH reference_trade AS",
fetchall=(
_trade_row(conflicting),
_trade_row(previous),
),
),
)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="strictly rollover-ordered",
):
repository.load_checkpoint_tail(
venue=VENUE,
checkpoint=checkpoint,
limit=2,
)
assert connection.exit_exception_types == [None]
@pytest.mark.parametrize("limit", (0, -1, True, 1.5, "1"))
def test_tail_rejects_invalid_limit_without_io(limit: object) -> None:
repository, _, provider = _repository()
with pytest.raises(MarketDataStorageValidationError, match="limit"):
repository.load_latest_trade_tail(
venue=VENUE,
symbol=SYMBOL,
limit=limit, # type: ignore[arg-type]
)
assert provider.calls == 0
def test_latest_tail_returns_empty_without_trade_history() -> None:
repository, connection, _ = _repository(
SqlStep(starts_with="SELECT symbol, trade_id"),
)
result = repository.load_latest_trade_tail(
venue=VENUE,
symbol=SYMBOL,
limit=10,
)
assert result == ()
assert len(connection.calls) == 1
def test_adopts_existing_trade_without_mutating_trade_history() -> None:
trade = _trade()
repository, connection, provider = _repository(
SqlStep(
starts_with="SELECT symbol, trade_id",
fetchone=_trade_row(trade),
),
SqlStep(
starts_with="INSERT INTO market_data.trade_stream_checkpoints",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(trade, revision=1),
),
)
result = repository.adopt_existing_trade_as_checkpoint(
venue=" dzengi ",
trade=trade,
)
assert result.trade == trade
assert result.revision == 1
assert provider.calls == 1
assert connection.calls[0][1] == (
VENUE,
SYMBOL,
trade.trade_id,
trade.executed_at,
)
assert "FOR SHARE" in connection.calls[0][0]
assert not any(
statement.startswith("INSERT INTO market_data.trades")
or statement.startswith("UPDATE market_data.trades")
for statement, _ in connection.calls
)
assert connection.exit_exception_types == [None]
assert connection.steps == []
def test_repeated_adoption_of_same_trade_is_idempotent() -> None:
trade = _trade()
repository, connection, _ = _repository(
SqlStep(
starts_with="SELECT symbol, trade_id",
fetchone=_trade_row(trade),
),
SqlStep(
starts_with="INSERT INTO market_data.trade_stream_checkpoints",
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(trade, revision=1),
),
)
result = repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=trade,
)
assert result.trade == trade
assert result.revision == 1
assert not any(
statement.startswith(
"UPDATE market_data.trade_stream_checkpoints"
)
for statement, _ in connection.calls
)
assert connection.exit_exception_types == [None]
def test_adoption_rejects_missing_durable_trade() -> None:
repository, connection, _ = _repository(
SqlStep(starts_with="SELECT symbol, trade_id"),
)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="absent from durable",
):
repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=_trade(),
)
assert len(connection.calls) == 1
assert connection.exit_exception_types == [
MarketDataCheckpointIntegrityError
]
def test_adoption_rejects_conflicting_durable_payload() -> None:
candidate = _trade()
durable = _trade(price=Decimal("65000.26"))
repository, connection, _ = _repository(
SqlStep(
starts_with="SELECT symbol, trade_id",
fetchone=_trade_row(durable),
),
)
with pytest.raises(
MarketDataCheckpointConflictError,
match="conflicts with durable",
):
repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=candidate,
)
assert len(connection.calls) == 1
assert connection.exit_exception_types == [
MarketDataCheckpointConflictError
]
def test_adoption_rejects_checkpoint_of_another_trade() -> None:
candidate = _trade(trade_id=100)
existing = _trade(
trade_id=101,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
repository, connection, _ = _repository(
SqlStep(
starts_with="SELECT symbol, trade_id",
fetchone=_trade_row(candidate),
),
SqlStep(
starts_with="INSERT INTO market_data.trade_stream_checkpoints",
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(existing, revision=1),
),
)
with pytest.raises(
MarketDataCheckpointConflictError,
match="already points",
):
repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=candidate,
)
assert connection.exit_exception_types == [
MarketDataCheckpointConflictError
]
def test_adoption_database_error_rolls_back_transaction() -> None:
trade = _trade()
repository, connection, _ = _repository(
SqlStep(
starts_with="SELECT symbol, trade_id",
fetchone=_trade_row(trade),
),
SqlStep(
starts_with="INSERT INTO market_data.trade_stream_checkpoints",
error=RuntimeError("database failed"),
),
)
with pytest.raises(
MarketDataStorageOperationError,
match="adopt existing Trade",
) as error_info:
repository.adopt_existing_trade_as_checkpoint(
venue=VENUE,
trade=trade,
)
assert isinstance(error_info.value.__cause__, RuntimeError)
assert connection.exit_exception_types == [RuntimeError]
def test_first_checkpoint_is_inserted_with_trade_in_one_transaction() -> None:
trade = _trade()
repository, connection, provider = _repository(
SqlStep(
starts_with="INSERT INTO market_data.trades",
fetchone=(1,),
),
SqlStep(
starts_with="INSERT INTO market_data.trade_stream_checkpoints",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(trade, revision=1),
),
)
result = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=trade,
observed_at=OBSERVED_AT,
)
assert result.trade == trade
assert result.revision == 1
assert provider.calls == 1
assert connection.exit_exception_types == [None]
assert "FOR UPDATE OF checkpoint" in connection.calls[-1][0]
assert connection.steps == []
def test_first_checkpoint_mismatch_rolls_back_transaction() -> None:
candidate = _trade(trade_id=100)
unexpected = _trade(
trade_id=101,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
repository, connection, _ = _repository(
SqlStep(
starts_with="INSERT INTO market_data.trades",
fetchone=(1,),
),
SqlStep(
starts_with="INSERT INTO market_data.trade_stream_checkpoints",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(unexpected, revision=1),
),
)
with pytest.raises(
MarketDataCheckpointIntegrityError,
match="First checkpoint",
):
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=None,
trade=candidate,
observed_at=OBSERVED_AT,
)
assert connection.exit_exception_types == [
MarketDataCheckpointIntegrityError
]
def test_existing_checkpoint_advances_with_revision_cas() -> None:
previous = _trade(trade_id=100)
current = _trade(
trade_id=101,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
repository, connection, _ = _repository(
SqlStep(
starts_with="INSERT INTO market_data.trades",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(previous, revision=7),
),
SqlStep(
starts_with="UPDATE market_data.trade_stream_checkpoints",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(current, revision=8),
),
)
result = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=previous,
trade=current,
observed_at=OBSERVED_AT,
)
assert result.trade == current
assert result.revision == 8
update_parameters = connection.calls[2][1]
assert update_parameters[-1] == 7
assert connection.exit_exception_types == [None]
def test_retry_after_committed_candidate_keeps_original_revision() -> None:
previous = _trade(trade_id=100)
current = _trade(
trade_id=101,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
repository, connection, _ = _repository(
SqlStep(starts_with="INSERT INTO market_data.trades"),
SqlStep(
starts_with="SELECT price, quantity",
fetchone=_existing_trade_row(current),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(current, revision=8),
),
)
result = repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=previous,
trade=current,
observed_at=OBSERVED_AT,
)
assert result.revision == 8
assert not any(
statement.startswith(
"UPDATE market_data.trade_stream_checkpoints"
)
for statement, _ in connection.calls
)
assert connection.exit_exception_types == [None]
def test_stale_expected_checkpoint_rolls_back_candidate_trade() -> None:
stale = _trade(trade_id=100)
database_current = _trade(
trade_id=101,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
candidate = _trade(
trade_id=102,
executed_at=EXECUTED_AT + timedelta(milliseconds=2),
)
repository, connection, _ = _repository(
SqlStep(
starts_with="INSERT INTO market_data.trades",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(database_current, revision=2),
),
)
with pytest.raises(
MarketDataCheckpointConflictError,
match="differs from expected",
):
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=stale,
trade=candidate,
observed_at=OBSERVED_AT,
)
assert connection.exit_exception_types == [
MarketDataCheckpointConflictError
]
def test_checkpoint_database_error_rolls_back_whole_transaction() -> None:
previous = _trade(trade_id=100)
current = _trade(
trade_id=101,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
repository, connection, _ = _repository(
SqlStep(
starts_with="INSERT INTO market_data.trades",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(previous, revision=1),
),
SqlStep(
starts_with="UPDATE market_data.trade_stream_checkpoints",
error=RuntimeError("database failed"),
),
)
with pytest.raises(
MarketDataStorageOperationError,
match="atomically",
) as error_info:
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=previous,
trade=current,
observed_at=OBSERVED_AT,
)
assert isinstance(error_info.value.__cause__, RuntimeError)
assert connection.exit_exception_types == [RuntimeError]
def test_half_cycle_candidate_is_rejected_as_ambiguous() -> None:
previous = _trade(trade_id=0)
ambiguous = _trade(
trade_id=SIGNED_TRADE_ID_MIN,
executed_at=EXECUTED_AT + timedelta(milliseconds=1),
)
repository, connection, _ = _repository(
SqlStep(
starts_with="INSERT INTO market_data.trades",
fetchone=(1,),
),
SqlStep(
starts_with="SELECT checkpoint.venue",
fetchone=_checkpoint_row(previous, revision=1),
),
)
with pytest.raises(
MarketDataCheckpointConflictError,
match="ambiguous",
):
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=previous,
trade=ambiguous,
observed_at=OBSERVED_AT,
)
assert connection.exit_exception_types == [
MarketDataCheckpointConflictError
]
def test_invalid_expected_trade_is_rejected_without_io() -> None:
repository, _, provider = _repository()
with pytest.raises(
MarketDataStorageValidationError,
match="expected_trade",
):
repository.store_trade_and_advance_checkpoint(
venue=VENUE,
expected_trade=object(), # type: ignore[arg-type]
trade=_trade(),
observed_at=OBSERVED_AT,
)
assert provider.calls == 0

View File

@@ -17,6 +17,7 @@ from src.market_data.storage import (
MarketDataStorageValidationError,
MarketDataWriteResult,
MarketDataWriteStatus,
PersistentTradeCheckpoint,
TradeStorageObservationSink,
)
@@ -42,13 +43,18 @@ class RecordingTradeStorage:
self,
*,
result: object | None = None,
checkpoint_result: object | None = None,
error: Exception | None = None,
) -> None:
self.result = result or MarketDataWriteResult(
status=MarketDataWriteStatus.INSERTED,
)
self.checkpoint_result = checkpoint_result
self.error = error
self.calls: list[tuple[str, Trade, datetime]] = []
self.checkpoint_calls: list[
tuple[str, Trade | None, Trade, datetime]
] = []
def store_trade(
self,
@@ -73,12 +79,91 @@ class RecordingTradeStorage:
) -> MarketDataBatchWriteResult:
raise AssertionError("Runtime persists observations one by one")
def load_checkpoint(
self,
*,
venue: str,
symbol: str,
) -> PersistentTradeCheckpoint | None:
raise AssertionError((venue, symbol))
def load_checkpoint_tail(
self,
*,
venue: str,
checkpoint: PersistentTradeCheckpoint,
limit: int,
) -> tuple[Trade, ...]:
raise AssertionError((venue, checkpoint, limit))
def load_latest_trade_tail(
self,
*,
venue: str,
symbol: str,
limit: int,
) -> tuple[Trade, ...]:
raise AssertionError((venue, symbol, limit))
def adopt_existing_trade_as_checkpoint(
self,
*,
venue: str,
trade: Trade,
) -> PersistentTradeCheckpoint:
raise AssertionError((venue, trade))
def store_trade_and_advance_checkpoint(
self,
*,
venue: str,
expected_trade: Trade | None,
trade: Trade,
observed_at: datetime,
) -> PersistentTradeCheckpoint:
self.checkpoint_calls.append(
(venue, expected_trade, trade, observed_at)
)
if self.error is not None:
raise self.error
if self.checkpoint_result is not None:
return self.checkpoint_result # type: ignore[return-value]
return PersistentTradeCheckpoint(
venue=venue,
trade=trade,
revision=1,
updated_at=observed_at,
)
class RecordingWriteOnlyTradeStorage:
def store_trade(
self,
*,
venue: str,
trade: Trade,
observed_at: datetime,
) -> MarketDataWriteResult:
raise AssertionError((venue, trade, observed_at))
def store_trades(
self,
*,
venue: str,
trades: tuple[Trade, ...],
observed_at: datetime,
) -> MarketDataBatchWriteResult:
raise AssertionError((venue, trades, observed_at))
@pytest.mark.parametrize(
"status",
tuple(MarketDataWriteStatus),
)
def test_forwards_observation_and_accepts_all_success_statuses(
def test_duplicate_forwards_observation_and_accepts_all_success_statuses(
status: MarketDataWriteStatus,
) -> None:
storage = RecordingTradeStorage(
@@ -91,7 +176,7 @@ def test_forwards_observation_and_accepts_all_success_statuses(
)
trade = make_trade()
result = sink.persist(trade)
result = sink.persist_duplicate(trade)
assert result is None
assert storage.calls == [
@@ -101,6 +186,42 @@ def test_forwards_observation_and_accepts_all_success_statuses(
OBSERVED_AT,
)
]
assert storage.checkpoint_calls == []
def test_accepted_trade_atomically_advances_checkpoint() -> None:
storage = RecordingTradeStorage()
sink = TradeStorageObservationSink(
trade_storage=storage,
venue=" dzengi ",
clock=lambda: OBSERVED_AT,
)
previous_trade = make_trade()
trade = Trade(
symbol=previous_trade.symbol,
trade_id=previous_trade.trade_id + 1,
price=previous_trade.price,
quantity=previous_trade.quantity,
executed_at=previous_trade.executed_at,
aggressor_side=previous_trade.aggressor_side,
source=previous_trade.source,
)
result = sink.persist_accepted(
trade,
expected_trade=previous_trade,
)
assert result is None
assert storage.checkpoint_calls == [
(
VENUE,
previous_trade,
trade,
OBSERVED_AT,
)
]
assert storage.calls == []
def test_implements_acquisition_side_sink_protocol() -> None:
@@ -113,7 +234,11 @@ def test_implements_acquisition_side_sink_protocol() -> None:
assert not hasattr(sink, "__dict__")
def test_storage_error_is_not_wrapped() -> None:
@pytest.mark.parametrize(
"operation",
("accepted", "duplicate"),
)
def test_storage_error_is_not_wrapped(operation: str) -> None:
storage_error = RuntimeError("storage failed")
sink = TradeStorageObservationSink(
trade_storage=RecordingTradeStorage(error=storage_error),
@@ -125,7 +250,13 @@ def test_storage_error_is_not_wrapped() -> None:
RuntimeError,
match="storage failed",
) as error_info:
sink.persist(make_trade())
if operation == "accepted":
sink.persist_accepted(
make_trade(),
expected_trade=None,
)
else:
sink.persist_duplicate(make_trade())
assert error_info.value is storage_error
@@ -141,7 +272,26 @@ def test_rejects_invalid_storage_result() -> None:
TypeError,
match="MarketDataWriteResult",
):
sink.persist(make_trade())
sink.persist_duplicate(make_trade())
def test_rejects_invalid_checkpoint_result() -> None:
sink = TradeStorageObservationSink(
trade_storage=RecordingTradeStorage(
checkpoint_result=object(),
),
venue=VENUE,
clock=lambda: OBSERVED_AT,
)
with pytest.raises(
TypeError,
match="PersistentTradeCheckpoint",
):
sink.persist_accepted(
make_trade(),
expected_trade=None,
)
def test_rejects_invalid_dependencies() -> None:
@@ -151,6 +301,15 @@ def test_rejects_invalid_dependencies() -> None:
venue=VENUE,
)
with pytest.raises(
TypeError,
match="TradeCheckpointStorageProtocol",
):
TradeStorageObservationSink(
trade_storage=RecordingWriteOnlyTradeStorage(),
venue=VENUE,
)
with pytest.raises(TypeError, match="clock must be callable"):
TradeStorageObservationSink(
trade_storage=RecordingTradeStorage(),

View File

@@ -105,6 +105,7 @@ def test_default_migrations_have_stable_order_and_names() -> None:
(5, "add_trade_observation_sources"),
(6, "add_quote_and_candle_observation_sources"),
(7, "create_market_data_partition_registry"),
(8, "create_trade_stream_checkpoints"),
)
@@ -119,8 +120,8 @@ def test_default_schema_defines_partitions_identities_and_constraints() -> None:
assert "CREATE TABLE market_data.trades" in sql
assert "PRIMARY KEY (venue, symbol, trade_id, executed_at)" in sql
assert "trade_id BETWEEN -2147483648 AND 2147483647" in sql
assert sql.count("CHECK (BTRIM(venue) <> '')") == 3
assert sql.count("CHECK (BTRIM(symbol) <> '')") == 3
assert sql.count("CHECK (BTRIM(venue) <> '')") == 4
assert sql.count("CHECK (BTRIM(symbol) <> '')") == 4
assert sql.count("CHECK (BTRIM(source) <> '')") == 3
assert "PARTITION BY RANGE (executed_at)" in sql
assert "CREATE TABLE market_data.quotes" in sql
@@ -143,6 +144,18 @@ def test_default_schema_defines_partitions_identities_and_constraints() -> None:
assert "partition_bound TEXT NOT NULL" in sql
assert "BTRIM(partition_bound) <> ''" in sql
assert "range_end > range_start" in sql
assert "CREATE TABLE market_data.trade_stream_checkpoints" in sql
assert "PRIMARY KEY (venue, symbol)" in sql
assert "revision BIGINT NOT NULL" in sql
assert "checkpoint_schema_version INTEGER NOT NULL DEFAULT 1" in sql
assert "CONSTRAINT trade_stream_checkpoints_trade_fk" in sql
assert "FOREIGN KEY (" in sql
assert ") REFERENCES market_data.trades (" in sql
assert "ON UPDATE NO ACTION" in sql
assert "ON DELETE NO ACTION" in sql
assert "DEFERRABLE INITIALLY DEFERRED" in sql
assert "CHECK (revision > 0)" in sql
assert "CHECK (checkpoint_schema_version > 0)" in sql
def test_run_locks_and_applies_every_pending_migration_in_order() -> None:
@@ -150,7 +163,7 @@ def test_run_locks_and_applies_every_pending_migration_in_order() -> None:
result = runner.run()
assert result == (1, 2, 3, 4, 5, 6, 7)
assert result == (1, 2, 3, 4, 5, 6, 7, 8)
assert provider.calls == 1
assert connection.entered == 1
assert connection.exited == 1
@@ -167,7 +180,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)
assert inserted_versions == (1, 2, 3, 4, 5, 6, 7, 8)
def test_run_skips_already_applied_migrations() -> None:
@@ -196,7 +209,7 @@ def test_run_applies_only_migrations_after_existing_prefix() -> None:
result = runner.run()
assert result == (3, 4, 5, 6, 7)
assert result == (3, 4, 5, 6, 7, 8)
inserted_versions = tuple(
parameters[0]
for statement, parameters in cursor.calls
@@ -205,7 +218,7 @@ def test_run_applies_only_migrations_after_existing_prefix() -> None:
)
and isinstance(parameters, tuple)
)
assert inserted_versions == (3, 4, 5, 6, 7)
assert inserted_versions == (3, 4, 5, 6, 7, 8)
def test_run_rejects_unknown_applied_version() -> None:

View File

@@ -5,11 +5,14 @@ from typing import Any
import pytest
from psycopg.conninfo import conninfo_to_dict
from tests.support import postgres_market_data
from tests.support.postgres_market_data import (
POSTGRES_TEST_APPLICATION_NAME,
POSTGRES_TEST_CONTROL_APPLICATION_NAME,
count_other_test_connections,
load_postgres_test_settings,
reset_postgres_test_database,
wait_for_postgres_relation_lock_waiters,
)
@@ -59,6 +62,52 @@ class RecordingControlConnection:
)
class QueuedCursor:
def __init__(self, connection: QueuedControlConnection) -> None:
self._connection = connection
def __enter__(self) -> QueuedCursor:
return self
def __exit__(self, *_: object) -> None:
return None
def execute(
self,
statement: object,
parameters: object = None,
) -> None:
self._connection.statements.append(
(" ".join(str(statement).split()), parameters)
)
def fetchone(self) -> object:
if not self._connection.rows:
raise AssertionError("Не подготовлен ответ PostgreSQL")
return self._connection.rows.pop(0)
class QueuedControlConnection:
def __init__(self, *, rows: list[object]) -> None:
self.rows = list(rows)
self.statements: list[tuple[str, object]] = []
def cursor(self) -> QueuedCursor:
return QueuedCursor(self)
class FakeClock:
def __init__(self) -> None:
self.current = 0.0
def monotonic(self) -> float:
return self.current
def sleep(self, seconds: float) -> None:
self.current += seconds
def test_postgres_harness_is_disabled_without_explicit_flag() -> None:
assert load_postgres_test_settings({}) is None
@@ -188,3 +237,230 @@ def test_postgres_reset_refuses_unvalidated_connection_before_drop(
assert connection.statements == [
"SELECT current_database(), current_setting('application_name')",
]
def test_connection_count_includes_every_non_control_database_session() -> None:
database_name = "dzentra_test_all_connections"
connection: Any = QueuedControlConnection(
rows=[
(database_name, POSTGRES_TEST_CONTROL_APPLICATION_NAME),
(3,),
]
)
result = count_other_test_connections(connection)
assert result == 3
assert connection.statements == [
(
"SELECT current_database(), "
"current_setting('application_name')",
None,
),
(
"SELECT COUNT(*) FROM pg_catalog.pg_stat_activity "
"WHERE datname = %s "
"AND backend_type = 'client backend' "
"AND application_name IS DISTINCT FROM %s",
(database_name, POSTGRES_TEST_CONTROL_APPLICATION_NAME),
),
]
@pytest.mark.parametrize(
"identity",
(
("production", POSTGRES_TEST_CONTROL_APPLICATION_NAME),
("dzentra_test_all_connections", "another-application"),
None,
),
)
def test_connection_count_revalidates_control_connection(
identity: object,
) -> None:
connection: Any = QueuedControlConnection(rows=[identity, (0,)])
with pytest.raises(RuntimeError, match="validated test control"):
count_other_test_connections(connection)
assert len(connection.statements) == 1
@pytest.mark.parametrize(
"count_row",
(
(True,),
(-1,),
("1",),
[1],
None,
),
)
def test_connection_count_rejects_invalid_postgres_result(
count_row: object,
) -> None:
connection: Any = QueuedControlConnection(
rows=[
(
"dzentra_test_all_connections",
POSTGRES_TEST_CONTROL_APPLICATION_NAME,
),
count_row,
]
)
with pytest.raises(RuntimeError, match="invalid connection count"):
count_other_test_connections(connection)
def test_relation_waiter_uses_validated_database_and_exact_count(
monkeypatch: pytest.MonkeyPatch,
) -> None:
database_name = "dzentra_test_relation_waiter"
connection: Any = QueuedControlConnection(
rows=[
(database_name, POSTGRES_TEST_CONTROL_APPLICATION_NAME),
(16_384,),
(0,),
(2,),
]
)
clock = FakeClock()
monkeypatch.setattr(postgres_market_data.time, "monotonic", clock.monotonic)
monkeypatch.setattr(postgres_market_data.time, "sleep", clock.sleep)
wait_for_postgres_relation_lock_waiters(
connection,
relation_name="market_data.trade_stream_checkpoints",
expected_count=2,
timeout_seconds=1.0,
)
assert connection.rows == []
assert connection.statements[1] == (
"SELECT pg_catalog.to_regclass(%s)::oid",
("market_data.trade_stream_checkpoints",),
)
lock_statement, lock_parameters = connection.statements[2]
assert "locktype = 'relation'" in lock_statement
assert "AND NOT granted" in lock_statement
assert lock_parameters == (database_name, 16_384)
assert connection.statements[3] == (
lock_statement,
lock_parameters,
)
def test_relation_waiter_times_out_with_observed_count(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection: Any = QueuedControlConnection(
rows=[
(
"dzentra_test_relation_timeout",
POSTGRES_TEST_CONTROL_APPLICATION_NAME,
),
(42,),
(0,),
(0,),
]
)
clock = FakeClock()
monkeypatch.setattr(postgres_market_data.time, "monotonic", clock.monotonic)
monkeypatch.setattr(postgres_market_data.time, "sleep", clock.sleep)
with pytest.raises(
TimeoutError,
match=r"expected 1, observed 0\.",
):
wait_for_postgres_relation_lock_waiters(
connection,
relation_name="market_data.trade_stream_checkpoints",
expected_count=1,
timeout_seconds=0.01,
)
@pytest.mark.parametrize(
"relation_row",
(
None,
(),
(None,),
(True,),
("42",),
(0,),
),
)
def test_relation_waiter_rejects_unresolved_or_invalid_relation(
relation_row: object,
) -> None:
connection: Any = QueuedControlConnection(
rows=[
(
"dzentra_test_relation_shape",
POSTGRES_TEST_CONTROL_APPLICATION_NAME,
),
relation_row,
]
)
with pytest.raises(RuntimeError, match="requested test relation"):
wait_for_postgres_relation_lock_waiters(
connection,
relation_name="market_data.trade_stream_checkpoints",
expected_count=1,
)
@pytest.mark.parametrize(
"waiter_row",
(
None,
(),
(True,),
(-1,),
("1",),
[1],
),
)
def test_relation_waiter_rejects_invalid_count_result(
waiter_row: object,
) -> None:
connection: Any = QueuedControlConnection(
rows=[
(
"dzentra_test_waiter_shape",
POSTGRES_TEST_CONTROL_APPLICATION_NAME,
),
(42,),
waiter_row,
]
)
with pytest.raises(RuntimeError, match="invalid relation-lock waiter"):
wait_for_postgres_relation_lock_waiters(
connection,
relation_name="market_data.trade_stream_checkpoints",
expected_count=1,
)
@pytest.mark.parametrize(
"timeout_seconds",
(0.0, -1.0, float("inf"), float("nan"), True),
)
def test_relation_waiter_requires_positive_finite_timeout(
timeout_seconds: float,
) -> None:
connection: Any = QueuedControlConnection(rows=[])
with pytest.raises(ValueError, match="positive finite"):
wait_for_postgres_relation_lock_waiters(
connection,
relation_name="market_data.trade_stream_checkpoints",
expected_count=1,
timeout_seconds=timeout_seconds,
)
assert connection.statements == []

View File

@@ -5,6 +5,7 @@ import asyncio
import pytest
from tests.support.trade_stream_runtime import (
active_owned_task_names,
run_scenario,
wait_until_or_runtime_exit,
)
@@ -90,3 +91,30 @@ def test_wait_until_or_runtime_exit_rejects_clean_early_runtime_exit() -> None:
)
run_scenario(scenario())
def test_owned_task_detection_covers_application_and_processing_tasks() -> None:
async def scenario() -> None:
release = asyncio.Event()
names = (
"application-shutdown",
"telegram-polling",
"trade-stream-market-processing",
"persistent-application-verification",
)
tasks = tuple(
asyncio.create_task(release.wait(), name=name)
for name in names
)
try:
await asyncio.sleep(0)
assert active_owned_task_names() == tuple(sorted(names))
finally:
release.set()
await asyncio.gather(*tasks)
assert active_owned_task_names() == ()
run_scenario(scenario())