Files
robinhood/backend/app/streaming/live_store.py

174 lines
6.2 KiB
Python

from __future__ import annotations
import asyncio
import os
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime, timedelta
from typing import Dict, List, Tuple, Any
import glob
import shutil
import pyarrow as pa
import pyarrow.parquet as pq
@dataclass
class _Series:
bars: List[Dict[str, Any]]
last_flushed_ts: int
class LiveStore:
def __init__(self, max_bars: int = 5000, root: str = "data/parquet/live") -> None:
self._series: Dict[Tuple[str, str], _Series] = {}
self._max_bars = max_bars
self._root = root
self._lock = asyncio.Lock()
@property
def root(self) -> str:
return self._root
def _get_series(self, symbol: str, timeframe: str) -> _Series:
key = (symbol, timeframe)
s = self._series.get(key)
if not s:
s = _Series(bars=[], last_flushed_ts=0)
self._series[key] = s
return s
def get_history(self, symbol: str, timeframe: str) -> List[Dict[str, Any]]:
s = self._get_series(symbol, timeframe)
return list(s.bars)
def ingest_bar(self, symbol: str, timeframe: str, bar: Dict[str, Any]) -> None:
s = self._get_series(symbol, timeframe)
if s.bars and s.bars[-1]["time"] == bar["time"]:
# update last
last = s.bars[-1]
last["high"] = max(last["high"], bar["high"])
last["low"] = min(last["low"], bar["low"])
last["close"] = bar["close"]
last["volume"] = last.get("volume", 0.0) + bar.get("volume", 0.0)
else:
s.bars.append(bar)
if len(s.bars) > self._max_bars:
s.bars.pop(0)
async def flush_parquet(self) -> None:
# Write new bars since last flush, partitioned by date
async with self._lock:
for (symbol, timeframe), s in self._series.items():
new_rows = [b for b in s.bars if b["time"] > s.last_flushed_ts]
if not new_rows:
continue
# Partition by date
rows_by_date: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
for r in new_rows:
dt = datetime.utcfromtimestamp(int(r["time"]))
rows_by_date[dt.strftime("%Y-%m-%d")].append(r)
for date_str, rows in rows_by_date.items():
table = pa.Table.from_pylist(rows)
base = os.path.join(self._root, symbol, timeframe)
out = os.path.join(base, f"date={date_str}")
os.makedirs(out, exist_ok=True)
# write one file per flush to this partition
pq.write_table(table, os.path.join(out, f"part-{int(datetime.utcnow().timestamp())}.parquet"))
s.last_flushed_ts = max(b["time"] for b in new_rows)
# Singleton store
live_store = LiveStore()
async def periodic_flush(interval_sec: int = 60):
while True:
try:
await live_store.flush_parquet()
except Exception:
pass
await asyncio.sleep(interval_sec)
def _iter_partitions(root: str):
"""Yield (symbol, timeframe, partition_path, date_str) for existing partitions."""
# root/symbol/timeframe/date=YYYY-MM-DD
for sym_dir in glob.glob(f"{root}/*"):
if not os.path.isdir(sym_dir):
continue
symbol = os.path.basename(sym_dir)
for tf_dir in glob.glob(f"{sym_dir}/*"):
if not os.path.isdir(tf_dir):
continue
timeframe = os.path.basename(tf_dir)
for part_dir in glob.glob(f"{tf_dir}/date=*" ):
if not os.path.isdir(part_dir):
continue
date_str = os.path.basename(part_dir).split("=", 1)[-1]
yield (symbol, timeframe, part_dir, date_str)
def prune_old_partitions(root: str, retention_days: int = 7) -> int:
"""Delete partition directories older than retention_days. Returns count deleted."""
now = datetime.utcnow()
deleted = 0
for symbol, timeframe, part_dir, date_str in list(_iter_partitions(root)):
try:
y, m, d = map(int, date_str.split("-"))
dt = datetime(y, m, d)
if now - dt > timedelta(days=retention_days):
shutil.rmtree(part_dir, ignore_errors=True)
deleted += 1
except Exception:
# Skip unparsable date partitions
continue
return deleted
def compact_partition(part_dir: str, max_files_threshold: int = 20) -> bool:
"""If too many small part files exist, compact them into a single file.
Returns True if compaction performed.
"""
part_files = sorted(glob.glob(os.path.join(part_dir, "part-*.parquet")))
if len(part_files) < max_files_threshold:
return False
try:
tables: List[pa.Table] = []
for p in part_files:
tables.append(pq.read_table(p))
if not tables:
return False
combined = pa.concat_tables(tables, promote=True)
out_file = os.path.join(part_dir, f"compact-{int(datetime.utcnow().timestamp())}.parquet")
pq.write_table(combined, out_file)
# remove old parts
for p in part_files:
try:
os.remove(p)
except Exception:
pass
return True
except Exception:
return False
def compact_all(root: str, max_files_threshold: int = 20) -> int:
"""Run compaction across all partitions. Returns number of partitions compacted."""
compacted = 0
for _, _, part_dir, _ in list(_iter_partitions(root)):
if compact_partition(part_dir, max_files_threshold=max_files_threshold):
compacted += 1
return compacted
async def periodic_maintenance(retention_days: int = 7, compact_threshold_files: int = 20, interval_sec: int = 900):
"""Periodically prune old partitions and compact small files."""
while True:
try:
prune_old_partitions(live_store.root, retention_days=retention_days)
compact_all(live_store.root, max_files_threshold=compact_threshold_files)
except Exception:
pass
await asyncio.sleep(interval_sec)