diff --git a/apps/paper_trader/backtest_sim.py b/apps/paper_trader/backtest_sim.py index 34b0fc6..1fc30a7 100644 --- a/apps/paper_trader/backtest_sim.py +++ b/apps/paper_trader/backtest_sim.py @@ -235,6 +235,20 @@ def _convert_from_runner( except Exception: pass + # Snapshot coverage info from the canonical snapshot manifest + snapshot_coverage_end_date: str | None = None + snapshot_last_refresh_utc: str | None = None + try: + import json as _json + from pathlib import Path as _Path + _snap_manifest = _Path("data/parquet") / (config.canonical_snapshot_id or config.dataset_snapshot_id) / "manifest.json" + if _snap_manifest.exists(): + _snap_meta = _json.loads(_snap_manifest.read_text()) + snapshot_coverage_end_date = _snap_meta.get("coverage_end_date") + snapshot_last_refresh_utc = _snap_meta.get("last_refresh_utc") + except Exception: + pass + return { "session_name": session_name, "config_path": config_path, @@ -252,6 +266,8 @@ def _convert_from_runner( "sharpe": sharpe, }, "metrics_bundle": metrics_bundle, + "snapshot_coverage_end_date": snapshot_coverage_end_date, + "snapshot_last_refresh_utc": snapshot_last_refresh_utc, } @@ -414,13 +430,32 @@ async def _refresh_snapshot( if console: console.print(f" [yellow]Label generator skipped: {exc}[/]") - # Step 2: Rebuild canonical snapshot + # Step 2: Update canonical snapshot (incremental first, full rebuild as fallback) if console: console.print(" [dim]4/4 Exporting snapshot...[/]") try: - from libs.export.canonical_snapshots import build_canonical_snapshot + from libs.export.canonical_snapshots import ( + build_canonical_snapshot, + incremental_update_canonical_snapshot, + ) + + snapshot_path = _resolve_snapshot_path(resolution.canonical_snapshot_id) + use_incremental = snapshot_path is not None and snapshot_path.exists() + + if use_incremental: + if console: + console.print(" [dim]Incremental update (new events only)...[/]") + try: + await incremental_update_canonical_snapshot(snapshot_id) + except Exception as inc_exc: + if console: + console.print(f" [yellow]Incremental failed ({inc_exc}), falling back to full rebuild...[/]") + await build_canonical_snapshot(snapshot_id, manual=manual) + else: + if console: + console.print(" [dim]Full rebuild (no existing snapshot)...[/]") + await build_canonical_snapshot(snapshot_id, manual=manual) - await build_canonical_snapshot(snapshot_id, manual=manual) if console: console.print(" [green]Snapshot refreshed.[/]") # Write marker to avoid re-refreshing today diff --git a/libs/export/canonical_snapshots.py b/libs/export/canonical_snapshots.py index ffea0e6..d3a9358 100644 --- a/libs/export/canonical_snapshots.py +++ b/libs/export/canonical_snapshots.py @@ -107,7 +107,8 @@ def _run_enrichment_step(step_name: str, input_dir: Path, output_dir: Path) -> N str(output_dir), ] logger.info("canonical_snapshot_enrichment_start", step=step_name, cmd=cmd) - subprocess.run(cmd, check=True) + env = {**__import__("os").environ, "PYTHONUNBUFFERED": "1"} + subprocess.run(cmd, check=True, env=env) async def _materialize_runtime_backfills(snapshot_dir: Path) -> list[str]: diff --git a/scripts/enrich_tier2_features.py b/scripts/enrich_tier2_features.py index 24d026e..4beaf94 100644 --- a/scripts/enrich_tier2_features.py +++ b/scripts/enrich_tier2_features.py @@ -9,9 +9,11 @@ from __future__ import annotations import argparse import asyncio +import concurrent.futures import datetime as dt import json import sys +from collections import defaultdict from pathlib import Path import asyncpg @@ -231,7 +233,10 @@ 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) + start_str = start_dt.strftime("%Y-%m-%d") + # Use full-history cache to avoid per-row HTTP calls + all_bars_raw = _fetch_full_bars_raw(ticker) + bars_raw = [b for b in all_bars_raw if start_str <= b.get("date", "") <= event_date] bars = [] for b in bars_raw: try: @@ -243,6 +248,44 @@ def fetch_bars(ticker: str, event_date: str) -> list[PriceBar]: return bars +def _prefetch_ticker_bars(tickers: list, max_workers: int = 16) -> None: + """Pre-warm full-history bar cache for all unique tickers in parallel.""" + unique = sorted({str(t) for t in tickers if t}) + print(f" Prefetching price bars for {len(unique)} tickers...") + with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as pool: + for loaded, _ in enumerate(pool.map(_fetch_full_bars_raw, unique), start=1): + if loaded % 100 == 0 or loaded == len(unique): + print(f" bars prefetch {loaded}/{len(unique)}") + + +async def _prefetch_short_ratio_db_batch(tickers: list[str]) -> None: + """Bulk-fetch short ratio history for all tickers in a single DB query.""" + dsn = get_settings().postgres_dsn.replace("+asyncpg", "") + conn = await asyncpg.connect(dsn=dsn) + try: + rows = await conn.fetch( + """ + SELECT ticker_raw, + trade_date::text AS date, + short_volume::double precision / NULLIF(total_volume::double precision, 0.0) AS short_ratio + FROM short_sale_daily + WHERE ticker_raw = ANY($1) + AND total_volume IS NOT NULL + AND total_volume > 0 + ORDER BY trade_date DESC + """, + tickers, + ) + finally: + await conn.close() + grouped: dict[str, list[dict]] = defaultdict(list) + for row in rows: + if row["short_ratio"] is not None: + grouped[row["ticker_raw"]].append({"date": row["date"], "short_ratio": float(row["short_ratio"])}) + for ticker in tickers: + _short_ratio_db_cache[ticker] = grouped.get(ticker, []) + + def fetch_short_ratio(ticker: str, event_date: str) -> float | None: """Fetch average short ratio over 5 days before event.""" points = _fetch_short_ratio_history_from_db(ticker) @@ -377,6 +420,15 @@ def enrich_split(input_path: Path, output_path: Path): results = {f: [None] * n for f in FEATURES} success = 0 + # Pre-fetch price bars and short ratio history for all unique tickers in batch + _prefetch_ticker_bars(tickers) + unique_tickers = sorted({str(t) for t in tickers if t}) + try: + asyncio.run(_prefetch_short_ratio_db_batch(unique_tickers)) + print(f" Short ratio DB prefetch done for {len(unique_tickers)} tickers") + except Exception as e: + print(f" Short ratio DB prefetch failed (will fall back per-row): {e}") + for i in range(n): ticker = tickers[i] event_date = str(event_dates[i])