"""Enrich snapshot with Tier 3 features: OU theta, Gravitational Pull, Market Temperature. All computed from price bars (no extra API needed beyond what Tier 2 already fetches). Usage: PYTHONUNBUFFERED=1 uv run python3 scripts/enrich_tier3_features.py \ --input data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_tier2 \ --output data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_tier3 """ from __future__ import annotations import argparse, datetime as dt, json, sys from pathlib import Path import pyarrow as pa, pyarrow.parquet as pq import requests sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from libs.features.market_features import pre_event_ou_theta, pre_event_gravitational_pull, pre_event_market_temperature from libs.oracle_client.models import PriceBar ORACLE_URL = "http://localhost:18001" FEATURES = ["pre_event_ou_theta_60d", "pre_event_gravitational_pull", "pre_event_market_temperature"] 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) resp = requests.get(f"{ORACLE_URL}/api/v1/price/data/{ticker}", params={"start_date": start_dt.strftime("%Y-%m-%d"), "end_date": event_date}, timeout=30) if resp.status_code != 200: return [] data = resp.json() bars = [] for b in (data.get("bars") or data.get("data") or []): 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: continue return bars def enrich_split(input_path: Path, output_path: Path): table = pq.read_table(input_path) n = len(table) if n == 0 or "ticker" not in table.column_names or "event_date" not in table.column_names: print(" Empty or schema-less split, copying") pq.write_table(table, output_path) return existing = set(table.column_names) if all(f in existing for f in FEATURES): nulls = {f: table.column(f).null_count for f in FEATURES} if all(v < n * 0.3 for v in nulls.values()): print(f" Already enriched, copying"); pq.write_table(table, output_path); return tickers = table.column("ticker").to_pylist() dates = table.column("event_date").to_pylist() results = {f: [None]*n for f in FEATURES} ok = 0 for i in range(n): if i % 200 == 0: print(f" {i}/{n} ({ok} ok)...") try: bars = fetch_bars(tickers[i], str(dates[i])) if bars: results["pre_event_ou_theta_60d"][i] = pre_event_ou_theta(bars, str(dates[i]), 60) results["pre_event_gravitational_pull"][i] = pre_event_gravitational_pull(bars, str(dates[i])) results["pre_event_market_temperature"][i] = pre_event_market_temperature(bars, str(dates[i])) if results["pre_event_ou_theta_60d"][i] is not None: ok += 1 except: pass print(f" Done: {ok}/{n}") for f in FEATURES: arr = pa.array(results[f], type=pa.float64()) table = table.set_column(table.column_names.index(f), f, arr) if f in existing else table.append_column(f, arr) pq.write_table(table, output_path) def main(): parser = argparse.ArgumentParser() parser.add_argument("--input", required=True); parser.add_argument("--output", required=True) args = parser.parse_args() inp, out = Path(args.input), Path(args.output) out.mkdir(parents=True, exist_ok=True) m = inp / "manifest.json" if m.exists(): manifest = json.loads(m.read_text()) source_snapshot_id = manifest.get("snapshot_id") manifest["snapshot_id"] = out.name manifest["output_dir"] = str(out.resolve()) manifest["created_at_utc"] = dt.datetime.now(dt.UTC).isoformat() manifest["enrichment_source_snapshot_id"] = source_snapshot_id manifest["export_enrichments"] = sorted(set((manifest.get("export_enrichments") or []) + FEATURES)) (out / "manifest.json").write_text(json.dumps(manifest, indent=2)) for split in ["train","valid","test"]: p = inp / f"{split}.parquet" if not p.exists(): continue print(f"Enriching {split}...") enrich_split(p, out / f"{split}.parquet") print("Done!") if __name__ == "__main__": main()