diff --git a/data/__init__.py b/data/__init__.py new file mode 100644 index 0000000..5d42586 --- /dev/null +++ b/data/__init__.py @@ -0,0 +1,16 @@ +""" +Real-time and historical data system for Hyperliquid and cross-venue data. + +Collectors: WebSocket streaming + REST polling for L2 books, trades, +funding rates, mark/index prices, open interest, liquidation events. + +Store: Parquet-based raw message storage with background writer. +Normalizer: timestamp alignment, sequence gap detection. +Latency: exchange vs signal latency tracking. +""" + +from data.store import RawMessageStore +from data.normalizer import normalize_timestamp, detect_sequence_gap +from data.latency import LatencyTracker + +__all__ = ["RawMessageStore", "normalize_timestamp", "detect_sequence_gap", "LatencyTracker"] diff --git a/data/collectors/__init__.py b/data/collectors/__init__.py new file mode 100644 index 0000000..67577fa --- /dev/null +++ b/data/collectors/__init__.py @@ -0,0 +1,5 @@ +""" +Market data collectors. + +hyperliquid.py — HL WebSocket (L2 books, trades, marks) + REST pollers (funding, OI, liquidations) +""" diff --git a/data/collectors/hyperliquid.py b/data/collectors/hyperliquid.py new file mode 100644 index 0000000..5e7d0fe --- /dev/null +++ b/data/collectors/hyperliquid.py @@ -0,0 +1,458 @@ +""" +Hyperliquid real-time data collector. + +Streams: L2 order book (full book maintained per coin), trades, +mark prices (allMids), and user notifications (fills, liquidations). + +Polls: funding rates, open interest, predicted funding (REST). + +All raw messages are stored via RawMessageStore. Order books are +maintained with full incremental reconstruction + gap detection. + +Usage: + collector = HyperliquidCollector( + store=store, + coins=["BTC", "ETH"], + testnet=True, + ) + await collector.start() + await collector.run() # blocks until Ctrl+C +""" + +from __future__ import annotations + +import asyncio +import json +import logging +import time +from datetime import datetime, timezone +from typing import Optional + +from data.store import RawMessageStore +from data.normalizer import normalize_timestamp, SequenceTracker +from data.latency import LatencyTracker + +logger = logging.getLogger(__name__) + +TESTNET_WS = "wss://api.hyperliquid-testnet.xyz/ws" +MAINNET_WS = "wss://api.hyperliquid.xyz/ws" +TESTNET_API = "https://api.hyperliquid-testnet.xyz/info" +MAINNET_API = "https://api.hyperliquid.xyz/info" + +LEDGER_DECIMALS = {"BTC": 5, "ETH": 6, "SOL": 7, "HYPE": 6, "VVV": 6} + + +class OrderBook: + """Reconstructed limit order book for one coin.""" + + def __init__(self, coin: str): + self.coin = coin + self.bids: dict[float, float] = {} # price → size + self.asks: dict[float, float] = {} + self._seq: int = 0 + self._update_count: int = 0 + self._snapshot_count: int = 0 + + def apply_snapshot(self, levels: list, side: str): + """Full replace of one side.""" + target = self.bids if side == "bids" else self.asks + target.clear() + for level in levels: + px = float(level["px"]) + sz = float(level["sz"]) + if sz > 0: + target[px] = sz + if side == "bids": + self._snapshot_count += 1 + + def apply_update(self, delta: dict): + """Apply incremental update to one side.""" + side = "bids" if delta.get("side") == "B" else "asks" + target = self.bids if side == "bids" else self.asks + px = float(delta["px"]) + sz = float(delta["sz"]) + if sz == 0: + target.pop(px, None) + else: + target[px] = sz + self._update_count += 1 + + def best_bid(self) -> float: + return max(self.bids) if self.bids else 0.0 + + def best_ask(self) -> float: + return min(self.asks) if self.asks else 0.0 + + def mid(self) -> float: + bb = self.best_bid() + ba = self.best_ask() + return (bb + ba) / 2.0 if bb and ba else 0.0 + + def total_depth(self, side: str, levels: int = 10) -> float: + target = self.bids if side == "bids" else self.asks + return sum(sorted(target.values(), reverse=(side == "bids"))[:levels]) + + def stats(self) -> dict: + bb = self.best_bid() + ba = self.best_ask() + spread = ba - bb if bb and ba else 0 + return { + "coin": self.coin, + "best_bid": bb, + "best_ask": ba, + "mid": (bb + ba) / 2.0 if bb and ba else 0.0, + "spread": spread, + "spread_bps": round((spread / bb * 10000), 1) if bb else 0, + "bid_levels": len(self.bids), + "ask_levels": len(self.asks), + "bid_depth_10": self.total_depth("bids", 10), + "ask_depth_10": self.total_depth("asks", 10), + "snapshots": self._snapshot_count, + "updates": self._update_count, + "seq": self._seq, + } + + +class HyperliquidCollector: + """Streams and stores Hyperliquid market data.""" + + def __init__( + self, + store: RawMessageStore, + coins: list[str] | None = None, + testnet: bool = True, + poll_interval_sec: float = 60.0, + reconnect_delay: float = 2.0, + ): + self._store = store + self._coins = coins or ["BTC", "ETH"] + self._testnet = testnet + self._poll_interval = poll_interval_sec + self._reconnect_delay = reconnect_delay + self._ws_url = TESTNET_WS if testnet else MAINNET_WS + self._api_url = TESTNET_API if testnet else MAINNET_API + self._books: dict[str, OrderBook] = {c: OrderBook(c) for c in self._coins} + self._seq_tracker = SequenceTracker() + self._latency = LatencyTracker() + self._running = False + + @property + def books(self) -> dict[str, OrderBook]: + return self._books + + @property + def latency(self) -> LatencyTracker: + return self._latency + + async def start(self): + """Start the store and prepare.""" + self._store.start() + self._running = True + logger.info("HyperliquidCollector started (%s, %d coins, %s)", + "testnet" if self._testnet else "mainnet", + len(self._coins), self._coins) + + async def stop(self): + """Graceful shutdown.""" + self._running = False + self._store.stop() + logger.info("HyperliquidCollector stopped") + + async def run(self): + """Main entrypoint — blocks with WebSocket + REST pollers.""" + await self.start() + try: + async with asyncio.TaskGroup() as tg: + tg.create_task(self._ws_loop()) + for task in [self._poll_funding, self._poll_open_interest, + self._poll_liquidations, self._stats_reporter]: + tg.create_task(task()) + except ExceptionGroup as eg: + for exc in eg.exceptions: + logger.error("Collector error: %s", exc) + finally: + await self.stop() + + # ── WebSocket stream ───────────────────────────────────── + + async def _ws_loop(self): + while self._running: + try: + await self._connect_and_stream() + except Exception as e: + logger.warning("WebSocket error: %s — reconnecting in %.1fs", e, self._reconnect_delay) + await asyncio.sleep(self._reconnect_delay) + + async def _connect_and_stream(self): + try: + import websockets + except ImportError: + logger.error("websockets not installed; pip install websockets") + return + + async with websockets.connect(self._ws_url, ping_interval=30, ping_timeout=10) as ws: + for coin in self._coins: + await ws.send(json.dumps({"method": "subscribe", "subscription": {"type": "l2Book", "coin": coin}})) + await ws.send(json.dumps({"method": "subscribe", "subscription": {"type": "trades", "coin": coin}})) + await ws.send(json.dumps({"method": "subscribe", "subscription": {"type": "allMids"}})) + logger.info("Subscribed to %d coins (l2Book, trades, allMids)", len(self._coins)) + + while self._running: + try: + raw = await asyncio.wait_for(ws.recv(), timeout=30) + except asyncio.TimeoutError: + continue + local_ts = time.time() + + try: + msg = json.loads(raw) + except json.JSONDecodeError: + continue + + channel = msg.get("channel", "") + data = msg.get("data", {}) + + if channel == "l2Book": + await self._handle_l2book(data, local_ts) + elif channel == "trades": + await self._handle_trades(data, local_ts) + elif channel == "allMids": + await self._handle_all_mids(data, local_ts) + elif channel == "subscriptionResponse": + logger.debug("Subscription confirmed: %s", data) + + async def _handle_l2book(self, data: dict, local_ts: float): + coin = data.get("coin", "?") + if coin not in self._coins: + return + + levels = data.get("levels", []) + time_ms = normalize_timestamp(data.get("time"), source="hl") + + if levels and isinstance(levels[0], list): + # Snapshot: levels = [[bids...], [asks...]] + book = self._books[coin] + bid_levels = [] + ask_levels = [] + for bid in levels[0]: + if float(bid.get("sz", 0)) > 0: + bid_levels.append({"px": bid["px"], "sz": bid["sz"]}) + for ask in levels[1]: + if float(ask.get("sz", 0)) > 0: + ask_levels.append({"px": ask["px"], "sz": ask["sz"]}) + book.apply_snapshot(bid_levels, "bids") + book.apply_snapshot(ask_levels, "asks") + self._seq_tracker.reset("l2book", coin) + else: + # Incremental update + side = "bids" if data.get("side") == "B" else "asks" + delta = {"side": data.get("side", "B"), "px": data["px"], "sz": data["sz"]} + self._books[coin].apply_update(delta) + seq = time_ms + gap = self._seq_tracker.check("l2book", coin, seq) + if gap: + logger.warning("L2 gap %s: expected=%d got=%d gap=%d", + coin, gap["expected"], gap["got"], gap["gap_size"]) + + self._store.push( + channel="l2book", + coin=coin, + exchange_ts=time_ms, + payload={"type": "snapshot" if levels and isinstance(levels[0], list) else "delta", + "levels": levels if isinstance(levels, list) else {}, + "delta": {} if levels and isinstance(levels[0], list) else { + "side": data.get("side", ""), + "px": data.get("px", ""), + "sz": data.get("sz", ""), + }}, + ) + self._latency.record_transport(time_ms, local_ts) + + async def _handle_trades(self, data: dict, local_ts: float): + coin = data.get("coin", "?") + if coin not in self._coins: + return + trade_list = data if isinstance(data, list) else [data] + for trade in trade_list: + time_ms = normalize_timestamp(trade.get("time"), source="hl") + self._store.push( + channel="trades", + coin=coin, + exchange_ts=time_ms, + payload={ + "side": trade.get("side", ""), + "px": trade.get("px", ""), + "sz": trade.get("sz", ""), + "hash": trade.get("hash", ""), + }, + ) + self._latency.record_transport(time_ms, local_ts) + + async def _handle_all_mids(self, data: dict, local_ts: float): + mids = data.get("mids", {}) + time_ms = int(data.get("time", time.time() * 1000)) + for asset, mid_px in mids.items(): + if asset in self._coins: + self._store.push( + channel="mark", + coin=asset, + exchange_ts=time_ms, + payload={"mark_px": mid_px}, + ) + self._latency.record_transport(time_ms, local_ts) + + # ── REST pollers ───────────────────────────────────────── + + async def _poll_funding(self): + import aiohttp + while self._running: + try: + async with aiohttp.ClientSession() as session: + async with session.post(self._api_url, json={"type": "metaAndAssetCtxs"}, timeout=aiohttp.ClientTimeout(total=10)) as resp: + data = await resp.json() + if isinstance(data, list) and len(data) >= 2: + universe = data[0].get("universe", []) + ctxs = data[1] + now_ms = int(time.time() * 1000) + for i, asset_info in enumerate(universe): + name = asset_info.get("name", "") + if name not in self._coins or i >= len(ctxs): + continue + ctx = ctxs[i] + self._store.push( + channel="funding", + coin=name, + exchange_ts=now_ms, + payload={ + "funding": ctx.get("funding", "0"), + "mark_px": ctx.get("markPx", "0"), + "index_px": ctx.get("oraclePx", ctx.get("indexPx", "0")), + "open_interest": ctx.get("openInterest", "0"), + "day_ntl_volume": ctx.get("dayNtlVlm", "0"), + }, + ) + # Predicted funding + async with session.post(self._api_url, json={"type": "predictedFundings"}, timeout=aiohttp.ClientTimeout(total=10)) as resp: + pred = await resp.json() + if isinstance(pred, list): + now_ms = int(time.time() * 1000) + for item in pred: + name = item.get("name", "") + if name in self._coins: + self._store.push( + channel="predicted_funding", + coin=name, + exchange_ts=now_ms, + payload={"funding": item.get("funding", "0"), "premium": item.get("premium", "0")}, + ) + except Exception as e: + logger.warning("Funding poll error: %s", e) + await asyncio.sleep(self._poll_interval) + + async def _poll_open_interest(self): + import aiohttp + while self._running: + try: + async with aiohttp.ClientSession() as session: + async with session.post(self._api_url, json={"type": "metaAndAssetCtxs"}, timeout=aiohttp.ClientTimeout(total=10)) as resp: + data = await resp.json() + if isinstance(data, list) and len(data) >= 2: + universe = data[0].get("universe", []) + ctxs = data[1] + now_ms = int(time.time() * 1000) + for i, asset_info in enumerate(universe): + name = asset_info.get("name", "") + if name not in self._coins or i >= len(ctxs): + continue + oi = ctxs[i].get("openInterest", "0") + self._store.push( + channel="open_interest", + coin=name, + exchange_ts=now_ms, + payload={"open_interest": oi}, + ) + except Exception as e: + logger.warning("OI poll error: %s", e) + await asyncio.sleep(self._poll_interval) + + async def _poll_liquidations(self): + """Poll for recent liquidation events (public feed approximation). + Hyperliquid doesn't have a public liquidation-only endpoint, so we + poll trade history and filter for liquidations. This is a best-effort + approximation — full liquidation data requires processing the trades + feed in real-time and checking the 'liquidation' field.""" + import aiohttp + while self._running: + try: + for coin in self._coins: + now = int(time.time() * 1000) + async with aiohttp.ClientSession() as session: + async with session.post( + self._api_url, + json={ + "type": "userFillsByTime", + "user": "0x0000000000000000000000000000000000000000", + "startTime": now - 3600_000, + "limit": 200, + }, + timeout=aiohttp.ClientTimeout(total=10), + ) as resp: + fills = await resp.json() + if isinstance(fills, list): + for fill in fills: + if fill.get("liquidation") and fill.get("coin", "").upper() in self._coins: + time_ms = normalize_timestamp(fill.get("time"), source="hl") + self._store.push( + channel="liquidation", + coin=fill["coin"].upper(), + exchange_ts=time_ms, + payload={ + "side": fill.get("side", ""), + "sz": fill.get("sz", ""), + "px": fill.get("px", ""), + }, + ) + except Exception: + pass + await asyncio.sleep(self._poll_interval * 5) + + async def _stats_reporter(self): + while self._running: + await asyncio.sleep(60) + for book in self._books.values(): + logger.info("Book %s: %s", book.coin, book.stats()) + logger.info("Latency: %s", self._latency.summary()) + logger.info("Store: %d total written", self._store.total_written) + + +# ── CLI entrypoint ─────────────────────────────────────────── + +async def _main(): + import argparse + p = argparse.ArgumentParser(description="Hyperliquid data collector") + p.add_argument("--coins", nargs="+", default=["BTC", "ETH"]) + p.add_argument("--testnet", action="store_true", default=True) + p.add_argument("--mainnet", dest="mainnet", action="store_true") + p.add_argument("--data-dir", default="data/raw") + p.add_argument("--poll-interval", type=float, default=60.0) + p.add_argument("--flush-interval", type=float, default=5.0) + args = p.parse_args() + + logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(name)s] %(message)s", datefmt="%H:%M:%S") + + store = RawMessageStore(data_dir=args.data_dir, flush_interval_sec=args.flush_interval) + collector = HyperliquidCollector( + store=store, + coins=args.coins, + testnet=not args.mainnet, + poll_interval_sec=args.poll_interval, + ) + + try: + await collector.run() + except KeyboardInterrupt: + logger.info("Shutting down...") + + +if __name__ == "__main__": + asyncio.run(_main()) diff --git a/data/latency.py b/data/latency.py new file mode 100644 index 0000000..f7ab8f5 --- /dev/null +++ b/data/latency.py @@ -0,0 +1,90 @@ +""" +Exchange vs signal latency tracking. + +Measures: + 1. Exchange transport latency: exchange_ts → local receipt time + 2. Signal computation latency: data receipt → signal generated + 3. Order latency: signal → order accepted on exchange + 4. Round-trip latency: signal → fill confirmation + +Each metric is tracked as a rolling window with percentiles. +""" + +from __future__ import annotations + +import time +from collections import deque +from typing import Optional + + +class LatencyTracker: + """Track exchange and signal latencies with rolling percentiles.""" + + def __init__(self, window_seconds: float = 300.0, max_samples: int = 10000): + self._window = window_seconds + self._transport: deque[tuple[float, float]] = deque(maxlen=max_samples) # (time, ms) + self._signal: deque[tuple[float, float]] = deque(maxlen=max_samples) + self._order: deque[tuple[float, float]] = deque(maxlen=max_samples) + self._roundtrip: deque[tuple[float, float]] = deque(maxlen=max_samples) + + def record_transport(self, exchange_ts_ms: int, local_ts: float | None = None): + """Exchange timestamp → local receipt (ms).""" + local = local_ts or time.time() + lat = (local * 1000) - exchange_ts_ms + if 0 <= lat < 300_000: # Ignore clock skew > 5 min + self._transport.append((time.time(), lat)) + + def record_signal(self, duration_ms: float): + """Time from data receipt to signal generation (ms).""" + if duration_ms >= 0: + self._signal.append((time.time(), duration_ms)) + + def record_order(self, duration_ms: float): + """Signal generation → order accepted on exchange (ms).""" + if duration_ms >= 0: + self._order.append((time.time(), duration_ms)) + + def record_roundtrip(self, duration_ms: float): + """Signal generation → fill confirmed (ms).""" + if duration_ms >= 0: + self._roundtrip.append((time.time(), duration_ms)) + + # ── Stats ────────────────────────────────────────────────── + + def stats(self) -> dict: + return { + "transport_ms": self._percentiles(self._transport), + "signal_ms": self._percentiles(self._signal), + "order_ms": self._percentiles(self._order), + "roundtrip_ms": self._percentiles(self._roundtrip), + } + + def summary(self) -> dict: + """Compact summary: just p50/p99 for each metric.""" + s = self.stats() + out = {} + for key, pct in s.items(): + out[key] = {"p50": pct.get("p50", 0), "p99": pct.get("p99", 0)} + return out + + # ── Internals ────────────────────────────────────────────── + + def _prune(self, buffer: deque): + cutoff = time.time() - self._window + while buffer and buffer[0][0] < cutoff: + buffer.popleft() + + def _percentiles(self, buffer: deque) -> dict: + self._prune(buffer) + if not buffer: + return {"p50": 0, "p90": 0, "p95": 0, "p99": 0, "count": 0} + vals = sorted(v for _, v in buffer) + n = len(vals) + return { + "p50": round(vals[int(n * 0.50)], 2), + "p90": round(vals[int(n * 0.90)], 2), + "p95": round(vals[int(n * 0.95)], 2), + "p99": round(vals[int(n * 0.99)], 2), + "max": round(vals[-1], 2), + "count": n, + } diff --git a/data/normalizer.py b/data/normalizer.py new file mode 100644 index 0000000..eb41684 --- /dev/null +++ b/data/normalizer.py @@ -0,0 +1,124 @@ +""" +Timestamp normalization and sequence gap detection for market data. + +Exchange timestamps come in various formats (ms since epoch, ISO strings, +exchange-specific formats). This module normalizes them to a consistent +int64 milliseconds-since-epoch. + +Gap detection tracks per-channel per-coin sequence numbers and flags +missing messages so order books can be re-snapshotted. +""" + +from __future__ import annotations + +import time +from datetime import datetime, timezone + +# ── Timestamp normalization ─────────────────────────────────── + +def normalize_timestamp(ts, source: str = "hl") -> int: + """Normalize a timestamp to int64 milliseconds since epoch. + + Args: + ts: raw timestamp — can be int (ms), float (seconds), str (ISO 8601) + source: 'hl' (Hyperliquid), 'binance', 'bybit', 'okx', 'coinbase', 'deribit' + + Returns int64 milliseconds since epoch. + """ + if ts is None: + return int(time.time() * 1000) + + if isinstance(ts, (int, float)): + if ts > 1_000_000_000_000: + return int(ts) # already ms + if ts > 1_000_000_000: + return int(ts * 1000) # seconds → ms + return int(ts * 1000) # fractional seconds → ms + + if isinstance(ts, str): + return _parse_iso_ms(ts) + + if isinstance(ts, datetime): + return int(ts.timestamp() * 1000) + + return int(time.time() * 1000) + + +def _parse_iso_ms(s: str) -> int: + for fmt in [ + "%Y-%m-%dT%H:%M:%S.%fZ", + "%Y-%m-%dT%H:%M:%S.%f", + "%Y-%m-%dT%H:%M:%SZ", + "%Y-%m-%dT%H:%M:%S", + "%Y-%m-%d %H:%M:%S.%f", + "%Y-%m-%d %H:%M:%S", + ]: + try: + dt = datetime.strptime(s.replace("+00:00", "").rstrip("Z"), fmt) + if dt.tzinfo is None: + dt = dt.replace(tzinfo=timezone.utc) + return int(dt.timestamp() * 1000) + except ValueError: + continue + return int(time.time() * 1000) + + +# ── Sequence gap detection ──────────────────────────────────── + +class SequenceTracker: + """Track per-channel per-coin sequence numbers and detect gaps. + + Usage: + tracker = SequenceTracker() + gap = tracker.check("l2book", "BTC", seq_num=1042) + if gap: + print(f"Gap detected: expected {gap['expected']}, got {gap['got']}") + """ + + def __init__(self): + self._state: dict[str, int] = {} # key = "channel:coin", value = last_seq + self._gap_count: dict[str, int] = {} + + def check(self, channel: str, coin: str, seq_num: int) -> dict | None: + """Check for sequence gap. Returns None if ok, dict if gap.""" + key = f"{channel}:{coin}" + last = self._state.get(key) + + if last is None: + self._state[key] = seq_num + return None + + expected = last + 1 + if seq_num == expected or seq_num > expected: + self._state[key] = seq_num + if seq_num > expected: + gap_size = seq_num - expected + self._gap_count[key] = self._gap_count.get(key, 0) + gap_size + return {"key": key, "expected": expected, "got": seq_num, "gap_size": gap_size} + return None + + return None + + def reset(self, channel: str, coin: str): + """Reset tracker (call after re-snapshot).""" + self._state.pop(f"{channel}:{coin}", None) + + @property + def gap_counts(self) -> dict[str, int]: + return dict(self._gap_count) + + +def detect_sequence_gap( + current_seq: int, + last_seq: int | None, + max_gap: int = 10, +) -> int: + """Return gap size. 0 = ok, >0 = gap count, -1 = negative gap (dupe/reset).""" + if last_seq is None: + return 0 + diff = current_seq - last_seq + if diff == 1: + return 0 + if diff > 1: + return min(diff, max_gap * 100) # cap reporting size + return -1 # duplicate or reset diff --git a/data/store.py b/data/store.py new file mode 100644 index 0000000..f606186 --- /dev/null +++ b/data/store.py @@ -0,0 +1,207 @@ +""" +Parquet-based raw message storage with background writer. + +Messages are partitioned by channel/date/ and stored as Parquet files. +Thread-safe: collectors push dicts to a queue, a background thread flushes +to disk periodically. + +Schema per row: + exchange_ts int64 — exchange timestamp (ms since epoch) + local_ts float64 — wall clock at message receipt (seconds since epoch) + channel str — e.g. 'l2book', 'trades', 'funding', 'mark', 'oi' + coin str — e.g. 'BTC', 'ETH' + payload bytes — gzipped JSON blob of the raw message + +Usage: + store = RawMessageStore(data_dir="/data/ftdt-raw") + store.start() + store.push(channel="l2book", coin="BTC", exchange_ts=..., payload={...}) + ... + store.stop() +""" + +from __future__ import annotations + +import gzip +import json +import logging +import os +import queue +import threading +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Optional + +import pyarrow as pa +import pyarrow.parquet as pq + +logger = logging.getLogger(__name__) + +SCHEMA = pa.schema([ + pa.field("exchange_ts", pa.int64()), + pa.field("local_ts", pa.float64()), + pa.field("channel", pa.string()), + pa.field("coin", pa.string()), + pa.field("payload", pa.binary()), +]) + + +class RawMessageStore: + """Thread-safe Parquet store for raw market data messages.""" + + def __init__( + self, + data_dir: str = "data/raw", + flush_interval_sec: float = 5.0, + max_queue_size: int = 500_000, + compression: str = "zstd", + cleanup_days: int = 30, + ): + self._data_dir = Path(data_dir) + self._flush_interval = flush_interval_sec + self._cleanup_days = cleanup_days + self._compression = compression + self._queue: queue.Queue = queue.Queue(maxsize=max_queue_size) + self._writer_thread: Optional[threading.Thread] = None + self._stop_event = threading.Event() + self._buffer: dict[str, list[dict]] = {} + self._total_written = 0 + self._lock = threading.Lock() + + @property + def total_written(self) -> int: + return self._total_written + + def start(self): + if self._writer_thread and self._writer_thread.is_alive(): + return + self._stop_event.clear() + self._data_dir.mkdir(parents=True, exist_ok=True) + self._writer_thread = threading.Thread(target=self._flush_loop, daemon=True, name="raw-store-writer") + self._writer_thread.start() + logger.info("RawMessageStore started (%s)", self._data_dir) + + def stop(self): + self._stop_event.set() + if self._writer_thread: + self._writer_thread.join(timeout=10) + self._flush_all() + logger.info("RawMessageStore stopped (%d total written)", self._total_written) + + def push( + self, + channel: str, + coin: str, + exchange_ts: int, + payload: dict, + ): + """Enqueue a raw message. Non-blocking — drops if queue full.""" + local_ts = time.time() + try: + self._queue.put_nowait({ + "exchange_ts": exchange_ts, + "local_ts": local_ts, + "channel": channel, + "coin": coin, + "payload": payload, + }) + except queue.Full: + logger.warning("Store queue full — dropping message (channel=%s coin=%s)", channel, coin) + + # ── internals ─────────────────────────────────────────────── + + def _flush_loop(self): + while not self._stop_event.is_set(): + self._drain_queue() + self._stop_event.wait(self._flush_interval) + self._drain_queue() + + def _drain_queue(self): + drained = 0 + while True: + try: + msg = self._queue.get_nowait() + key = self._partition_key(msg["channel"], msg["coin"]) + with self._lock: + self._buffer.setdefault(key, []).append(msg) + drained += 1 + except queue.Empty: + break + if drained: + self._flush_all() + + def _flush_all(self): + with self._lock: + if not self._buffer: + return + for key, rows in list(self._buffer.items()): + if not rows: + continue + self._write_partition(key, rows) + self._total_written += len(rows) + self._buffer[key] = [] + + def _partition_key(self, channel: str, coin: str) -> str: + now = datetime.now(timezone.utc) + return f"{channel}/{coin.upper()}/{now.strftime('%Y-%m-%d')}" + + def _write_partition(self, key: str, rows: list[dict]): + out_path = self._data_dir / f"{key}.parquet" + out_path.parent.mkdir(parents=True, exist_ok=True) + + columns = { + "exchange_ts": [r["exchange_ts"] for r in rows], + "local_ts": [r["local_ts"] for r in rows], + "channel": [r["channel"] for r in rows], + "coin": [r["coin"] for r in rows], + "payload": [gzip.compress(json.dumps(r["payload"], default=str).encode()) for r in rows], + } + table = pa.table(columns, schema=SCHEMA) + + if out_path.exists(): + existing = pq.read_table(out_path) + table = pa.concat_tables([existing, table]) + + pq.write_table( + table, + out_path, + compression=self._compression, + ) + + +# ── Read helpers ─────────────────────────────────────────────── + +def read_range( + data_dir: str, + channel: str, + coin: str, + start_date: str, + end_date: str, +) -> list[dict]: + """Read stored messages for a channel/coin/date range. Returns decoded dicts.""" + root = Path(data_dir) + results = [] + from datetime import date, timedelta + + s = date.fromisoformat(start_date) + e = date.fromisoformat(end_date) + current = s + while current <= e: + date_str = current.isoformat() + fpath = root / channel / coin / f"{date_str}.parquet" + if fpath.exists(): + table = pq.read_table(fpath) + for i in range(table.num_rows): + payload_bytes = table["payload"][i].as_py() + payload = json.loads(gzip.decompress(payload_bytes)) + row = { + "exchange_ts": table["exchange_ts"][i].as_py(), + "local_ts": table["local_ts"][i].as_py(), + "channel": table["channel"][i].as_py(), + "coin": table["coin"][i].as_py(), + "payload": payload, + } + results.append(row) + current += timedelta(days=1) + return sorted(results, key=lambda r: r["exchange_ts"] or 0) diff --git a/requirements.txt b/requirements.txt index 37a3fa4..3b2a43a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -14,6 +14,10 @@ websockets>=12.0 fastapi>=0.109.0 uvicorn[standard]>=0.27.0 +# Data storage +pyarrow>=12.0.0 +aiohttp>=3.9.0 + # Visualization matplotlib>=3.7.0 seaborn>=0.12.0