"""
simulate_regimes.py

[SIMULATION-ONLY 2026-09-04] Standalone runner -- NOT part of the live bot's
startup path (monitor.py never imports this). Classifies every symbol in a
day's snapshot (see build_snapshot.py / data/snapshots/<date>/) using
regime_detector.classify_regime(), writes the results back into that same
snapshot folder, and cross-references the day's actual trades against their
detected regime so a strategy designer can see which regimes the live bot
already handled well or badly before building regime-specific entry rules.

Usage:
    python3 simulate_regimes.py [YYYY-MM-DD]   # defaults to today's snapshot
"""

import csv
import json
import sys
from collections import defaultdict
from pathlib import Path

from regime_detector import classify_regime, bars_from_price_rows, opening_window, ALL_REGIMES

BASE_DIR = Path(__file__).resolve().parent

# How much of the session actually matters for an entry-strategy question:
# this bot's live trades on 2026-09-04 all landed within the first 61
# minutes of the regular session (12 of 13 within the first 27). Classifying
# on the full 6.5-hour session instead washes out exactly the early
# spike/fail patterns that matter -- see regime_detector.opening_window()'s
# docstring for the concrete CHPT example. Widen this only if a later
# day's entries spread further into the session.
CLASSIFICATION_WINDOW_MINUTES = 90


def load_snapshot(date: str):
    snap_dir = BASE_DIR / "data" / "snapshots" / date
    if not snap_dir.exists():
        raise FileNotFoundError(f"No snapshot at {snap_dir} -- run build_snapshot.py for this date first.")
    meta = json.loads((snap_dir / "meta.json").read_text())
    prices_dir = snap_dir / "prices"
    symbols = sorted(p.stem for p in prices_dir.glob("*.csv"))
    return snap_dir, meta, symbols


def load_trades(snap_dir: Path):
    trades_path = snap_dir / "trades.jsonl"
    trades = {}
    if trades_path.exists():
        with open(trades_path) as f:
            for line in f:
                line = line.strip()
                if not line:
                    continue
                t = json.loads(line)
                trades[t["symbol"]] = t
    return trades


def run(date: str):
    snap_dir, meta, symbols = load_snapshot(date)
    trades = load_trades(snap_dir)

    results = {}
    for sym in symbols:
        prices_path = snap_dir / "prices" / f"{sym}.csv"
        with open(prices_path) as f:
            rows = list(csv.DictReader(f))
        bars = bars_from_price_rows(rows)
        window_bars = opening_window(bars, window_minutes=CLASSIFICATION_WINDOW_MINUTES)
        prev_close = (meta.get(sym) or {}).get("previous_close")
        reading = classify_regime(sym, window_bars, prev_close)
        results[sym] = {
            "regime": reading.regime,
            "confidence": reading.confidence,
            "gap_pct": reading.gap_pct,
            "range_pct": reading.range_pct,
            "net_change_pct": reading.net_change_pct,
            "trend_fit_pct": reading.trend_fit_pct,
            "atr_pct": reading.atr_pct,
            "breakout": reading.breakout,
            "reclaim": reading.reclaim,
            "whipsaw": reading.whipsaw,
            "insufficient_data": reading.insufficient_data,
            "traded": sym in trades,
            "trade_pl": trades[sym]["current_pl"] if sym in trades else None,
            "trade_exit_reason": trades[sym]["exit_reason"] if sym in trades else None,
        }

    out_json = snap_dir / "regimes.json"
    payload = {
        "classification_window_minutes": CLASSIFICATION_WINDOW_MINUTES,
        "note": "regime is classified over the first CLASSIFICATION_WINDOW_MINUTES of the regular "
                "session (bot entry decisions all happen in this window), not the full day.",
        "symbols": results,
    }
    out_json.write_text(json.dumps(payload, indent=2))

    out_csv = snap_dir / "regimes.csv"
    with open(out_csv, "w", newline="") as f:
        w = csv.writer(f)
        w.writerow(["symbol", "regime", "confidence", "gap_pct", "range_pct", "net_change_pct",
                    "trend_fit_pct", "atr_pct", "broke_out", "breakout_failed", "reclaim_detected",
                    "traded", "trade_pl", "trade_exit_reason"])
        for sym, r in sorted(results.items()):
            w.writerow([
                sym, r["regime"], r["confidence"], r["gap_pct"], r["range_pct"], r["net_change_pct"],
                r["trend_fit_pct"], r["atr_pct"],
                r["breakout"].get("broke_out"), r["breakout"].get("failed"),
                r["reclaim"].get("detected"),
                r["traded"], r["trade_pl"], r["trade_exit_reason"],
            ])

    # ---- console summary ----
    by_regime = defaultdict(list)
    for sym, r in results.items():
        by_regime[r["regime"]].append(sym)

    print(f"Snapshot: {date}  ({len(symbols)} symbols, classified on first {CLASSIFICATION_WINDOW_MINUTES} min of regular session)\n")
    print("Regime distribution (full candidate universe):")
    for regime in ALL_REGIMES:
        syms = by_regime.get(regime, [])
        print(f"  {regime:38} {len(syms):3} symbols")

    print("\nActual trades by detected regime:")
    traded_by_regime = defaultdict(list)
    for sym, r in results.items():
        if r["traded"]:
            traded_by_regime[r["regime"]].append((sym, r["trade_pl"], r["trade_exit_reason"]))

    for regime in ALL_REGIMES:
        entries = traded_by_regime.get(regime, [])
        if not entries:
            continue
        total_pl = sum(pl for _, pl, _ in entries)
        print(f"  {regime}:")
        for sym, pl, reason in sorted(entries, key=lambda x: x[1]):
            print(f"      {sym:6} P/L={pl:>8.2f}  exit={reason}")
        print(f"      -> regime subtotal: ${total_pl:.2f} over {len(entries)} trade(s)")

    print(f"\nWrote {out_json.relative_to(BASE_DIR)} and {out_csv.relative_to(BASE_DIR)}")


if __name__ == "__main__":
    date = sys.argv[1] if len(sys.argv) > 1 else "2026-09-04"
    run(date)
