"""Enrich snapshot with Tier 2 features: FINRA short ratio, Hurst, Entropy, Sector momentum. Usage: PYTHONUNBUFFERED=1 uv run python3 scripts/enrich_tier2_features.py \ --input data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit \ --output data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_tier2 """ from __future__ import annotations import argparse import json import sys import time from pathlib import Path import pyarrow as pa import pyarrow.parquet as pq import requests sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from libs.features.market_features import pre_event_entropy, pre_event_hurst from libs.oracle_client.models import PriceBar ORACLE_URL = "http://localhost:18001" FEATURES = ["pre_event_hurst_60d", "pre_event_entropy_60d", "pre_event_short_ratio", "pre_event_sector_momentum_20d"] SECTOR_ETF = { "Technology": "XLK", "Health Care": "XLV", "Financials": "XLF", "Consumer Discretionary": "XLY", "Industrials": "XLI", "Energy": "XLE", "Utilities": "XLU", "Real Estate": "XLRE", "Materials": "XLB", "Communication Services": "XLC", "Consumer Staples": "XLP", "Healthcare": "XLV", } # Cache sector ETF bars and SPY bars _sector_etf_cache: dict[str, list[dict]] = {} _spy_cache: list[dict] = [] def _fetch_bars_raw(ticker: str, start: str, end: str) -> list[dict]: resp = requests.get(f"{ORACLE_URL}/api/v1/price/data/{ticker}", params={"start_date": start, "end_date": end}, timeout=30) if resp.status_code != 200: return [] data = resp.json() return data.get("bars") or data.get("data") or [] def fetch_bars(ticker: str, event_date: str) -> list[PriceBar]: from datetime import datetime, timedelta end_dt = datetime.strptime(event_date, "%Y-%m-%d") start_dt = end_dt - timedelta(days=120) bars_raw = _fetch_bars_raw(ticker, start_dt.strftime("%Y-%m-%d"), event_date) bars = [] for b in bars_raw: try: bars.append(PriceBar(date=b["date"], open=float(b.get("open", 0)), high=float(b.get("high", 0)), low=float(b.get("low", 0)), close=float(b.get("close", 0)), volume=int(b.get("volume", 0)))) except (KeyError, ValueError, TypeError): continue return bars def fetch_short_ratio(ticker: str, event_date: str) -> float | None: """Fetch average short ratio over 5 days before event.""" resp = requests.get(f"{ORACLE_URL}/api/v1/finra/short-ratio/{ticker}", params={"days": 30}, timeout=30) if resp.status_code != 200: return None data = resp.json() points = data.get("data") or [] if not points: return None # Filter to dates before event before = [p for p in points if p.get("date", "") < event_date and p.get("short_percent") is not None] before.sort(key=lambda p: p["date"], reverse=True) recent = before[:5] if not recent: return None return sum(p["short_percent"] for p in recent) / len(recent) def fetch_sector_momentum(sector: str, event_date: str) -> float | None: """Compute sector ETF 20d return minus SPY 20d return.""" from datetime import datetime, timedelta etf = SECTOR_ETF.get(sector) if not etf: return None end_dt = datetime.strptime(event_date, "%Y-%m-%d") start_dt = end_dt - timedelta(days=45) start_str = start_dt.strftime("%Y-%m-%d") # Fetch sector ETF (cached) cache_key = f"{etf}_{start_str}_{event_date}" if cache_key not in _sector_etf_cache: _sector_etf_cache[cache_key] = _fetch_bars_raw(etf, start_str, event_date) etf_bars = _sector_etf_cache[cache_key] etf_dated = {b["date"]: float(b.get("close", 0)) for b in etf_bars if b.get("close")} etf_dates = sorted(etf_dated.keys()) if event_date not in etf_dated or len(etf_dates) < 21: return None idx = etf_dates.index(event_date) if idx < 20: return None etf_ret = (etf_dated[etf_dates[idx]] - etf_dated[etf_dates[idx - 20]]) / etf_dated[etf_dates[idx - 20]] # Fetch SPY (cached) spy_key = f"SPY_{start_str}_{event_date}" if spy_key not in _sector_etf_cache: _sector_etf_cache[spy_key] = _fetch_bars_raw("SPY", start_str, event_date) spy_bars = _sector_etf_cache[spy_key] spy_dated = {b["date"]: float(b.get("close", 0)) for b in spy_bars if b.get("close")} spy_dates = sorted(spy_dated.keys()) if event_date not in spy_dated: return None spy_idx = spy_dates.index(event_date) if spy_idx < 20: return None spy_ret = (spy_dated[spy_dates[spy_idx]] - spy_dated[spy_dates[spy_idx - 20]]) / spy_dated[spy_dates[spy_idx - 20]] return etf_ret - spy_ret def compute_features_for_event(ticker: str, event_date: str, sector: str) -> dict[str, float | None]: bars = fetch_bars(ticker, event_date) result: dict[str, float | None] = {} # Hurst & Entropy from price bars result["pre_event_hurst_60d"] = pre_event_hurst(bars, event_date, 60) if bars else None result["pre_event_entropy_60d"] = pre_event_entropy(bars, event_date, 60) if bars else None # FINRA short ratio try: result["pre_event_short_ratio"] = fetch_short_ratio(ticker, event_date) except Exception: result["pre_event_short_ratio"] = None # Sector momentum try: result["pre_event_sector_momentum_20d"] = fetch_sector_momentum(sector, event_date) except Exception: result["pre_event_sector_momentum_20d"] = None return result def enrich_split(input_path: Path, output_path: Path): table = pq.read_table(input_path) n = len(table) existing_cols = set(table.column_names) if all(f in existing_cols for f in FEATURES): null_counts = {f: table.column(f).null_count for f in FEATURES} if all(v < n * 0.3 for v in null_counts.values()): print(f" Already enriched ({null_counts}), copying") pq.write_table(table, output_path) return tickers = table.column("ticker").to_pylist() event_dates = table.column("event_date").to_pylist() # Get sector info if available sectors = table.column("asset_type_proxy").to_pylist() if "asset_type_proxy" in table.column_names else ["UNKNOWN"] * n # Actually sector might be elsewhere — check for common sector columns for col_name in ["sector", "asset_type_proxy"]: if col_name in table.column_names: sectors = table.column(col_name).to_pylist() break results = {f: [None] * n for f in FEATURES} success = 0 for i in range(n): ticker = tickers[i] event_date = str(event_dates[i]) sector = str(sectors[i]) if sectors[i] else "UNKNOWN" if i % 100 == 0: print(f" {i}/{n} ({success} ok)...") try: feats = compute_features_for_event(ticker, event_date, sector) for f in FEATURES: results[f][i] = feats.get(f) if feats.get("pre_event_hurst_60d") is not None: success += 1 except Exception as e: if i < 5: print(f" Error {ticker} {event_date}: {e}") print(f" Done: {success}/{n}") for f in FEATURES: arr = pa.array(results[f], type=pa.float64()) if f in existing_cols: idx = table.column_names.index(f) table = table.set_column(idx, f, arr) else: table = table.append_column(f, arr) pq.write_table(table, output_path) print(f" Written to {output_path}") def main(): parser = argparse.ArgumentParser() parser.add_argument("--input", required=True) parser.add_argument("--output", required=True) args = parser.parse_args() input_dir = Path(args.input) output_dir = Path(args.output) output_dir.mkdir(parents=True, exist_ok=True) manifest_src = input_dir / "manifest.json" if manifest_src.exists(): manifest = json.loads(manifest_src.read_text()) manifest["snapshot_id"] = output_dir.name (output_dir / "manifest.json").write_text(json.dumps(manifest, indent=2)) for split in ["train", "valid", "test"]: input_path = input_dir / f"{split}.parquet" if not input_path.exists(): continue output_path = output_dir / f"{split}.parquet" print(f"Enriching {split}...") enrich_split(input_path, output_path) print("All done!") if __name__ == "__main__": main()