"""
backtest_v2.py — Revisi 2 runner.

Reads data/real/<TAG>/<SYM>_4h.csv (pump gate) + <SYM>_15m.csv (entry eval).
Runs detector_v2 + simulator_v2, computes metrics, writes CSV.

Usage:
  python backtest_v2.py --tag IS
  python backtest_v2.py --tag OOS
"""
import os, sys, glob, argparse
import numpy as np
import pandas as pd

sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
from config.settings_v2 import default_params, V2Params
from strategy.detector_v2 import detect_v2
from strategy.simulator_v2 import simulate_v2
from metrics_v2 import compute_metrics, print_metrics


def load_csv(path):
    df = pd.read_csv(path)
    df = df.sort_values("time").reset_index(drop=True)
    return [{"time": int(r.time), "open": r.open, "high": r.high,
             "low": r.low, "close": r.close, "volume": r.volume,
             "quote_volume": r.quote_volume, "taker_buy_quote": r.taker_buy_quote}
            for r in df.itertuples()]


def run(tag, params: V2Params = None):
    params = params or default_params()
    base = os.path.join(os.path.dirname(__file__), "data", "real", tag)
    if not os.path.isdir(base):
        print(f"[!] no data dir: {base}")
        return {}
    # recursive: support subfolders per month (OOS/2026-03/ etc)
    f15 = sorted(glob.glob(os.path.join(base, "**", "*_15m.csv"), recursive=True))
    trades = []
    total_bars = 0
    all_times = []
    for f15p in f15:
        d = os.path.dirname(f15p)
        sym = os.path.basename(f15p).replace("_15m.csv", "")
        f4p = os.path.join(d, f"{sym}_4h.csv")
        if not os.path.exists(f4p):
            continue
        k4 = load_csv(f4p)
        k15 = load_csv(f15p)
        total_bars += len(k15)
        if k15:
            all_times.append((k15[0]["time"], k15[-1]["time"]))
        signals = detect_v2(k4, k15, params)
        for sig in signals:
            t = simulate_v2(sym, k15, sig, params, equity=100.0)
            if t and t.r_multiple is not None:
                trades.append({
                    "symbol": sym,
                    "entry_bar": t.entry_bar,
                    "exit_bar": t.exit_bar,
                    "entry_price": round(t.entry_price, 6),
                    "direction": t.direction,
                    "score": sig.get("squeeze_score", 0),
                    "r_multiple": round(t.r_multiple, 4),
                    "exit_reason": t.exit_reason,
                    "leverage": getattr(t, "leverage", 1.0),
                    "divergence": sig.get("divergence", "none"),
                })
    if not trades:
        print(f"[{tag}] 0 trades")
        return {"tag": tag, "trades": 0}
    df = pd.DataFrame(trades)
    wins = df[df.r_multiple > 0]
    losses = df[df.r_multiple <= 0]
    wr = len(wins) / len(df)
    pf = wins.r_multiple.sum() / abs(losses.r_multiple.sum()) if len(losses) else float("inf")

    # ── Simulasi equity dari modal $100 (compounding per trade) ──
    equity = 100.0
    curve = [equity]
    equity_rows = []
    for _, row in df.iterrows():
        risk_pct = params.risk_score4 if row["score"] >= 4 else params.risk_score3
        risk_usd = equity * risk_pct
        pnl_r = row["r_multiple"]
        L = row.get("leverage", 1.0)
        pnl_usd = pnl_r * risk_usd * L
        cost = risk_usd * L * (params.fee_pct * 2 + params.slippage_pct * 2)
        net = pnl_usd - cost
        equity += net
        curve.append(equity)
        equity_rows.append(round(equity, 2))
    df["equity_after"] = equity_rows

    # ── period_days dari span data ──
    if all_times:
        t0 = min(a[0] for a in all_times)
        t1 = max(a[1] for a in all_times)
        period_days = (t1 - t0) / (1000 * 3600 * 24)
    else:
        period_days = 30.0

    # ── Full metric suite ──
    m = compute_metrics(trades, curve, period_days, total_bars=total_bars)
    m["tag"] = tag

    out = os.path.join(os.path.dirname(__file__), "cache", f"v2_{tag}_results.csv")
    os.makedirs(os.path.dirname(out), exist_ok=True)
    df.to_csv(out, index=False)
    print(f"\n=== V2 BACKTEST [{tag}] ===")
    print(f"Trades         : {len(df)}")
    print(f"Win rate       : {wr*100:.1f}%")
    print(f"Profit factor  : {pf:.2f}")
    print(f"Avg R          : {df.r_multiple.mean():.3f}")
    print(f"Max win R      : {df.r_multiple.max():.2f}")
    print(f"Max loss R     : {df.r_multiple.min():.2f}")
    print(f"--- Equity sim (start $100) ---")
    print(f"Final equity   : ${m['final_equity']:.2f}")
    print(f"Return         : {m['total_return_pct']:+.1f}%")
    print(f"Max drawdown   : {m['max_drawdown_pct']:.1f}%")
    print(f"--- LONG vs SHORT ---")
    print(f"LONG  : {m['long_trades']:>3} trades | win WR {m['long_wr']:.1f}% | ratio {m['long_ratio']}")
    print(f"SHORT : {m['short_trades']:>3} trades | win WR {m['short_wr']:.1f}% | ratio {m['short_ratio']}")
    print(f"--- Full metrics ---")
    print_metrics(m)
    print(f"CSV            : {out}")
    print("\nExit breakdown:")
    print(df.exit_reason.value_counts().to_string())
    return m


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--tag", default="IS")
    ap.add_argument("--donchian_n", type=int, default=20)
    ap.add_argument("--trailing_k", type=float, default=2.5)
    ap.add_argument("--range_pctile_thr", type=float, default=20.0)
    ap.add_argument("--atr_ratio_thr", type=float, default=0.7)
    ap.add_argument("--volume_ratio_thr", type=float, default=0.8)
    ap.add_argument("--va_pctile_thr", type=float, default=25.0)
    ap.add_argument("--min_squeeze_duration", type=int, default=3)
    ap.add_argument("--partial_tp_r", type=float, default=2.0)
    ap.add_argument("--sl_mode", default="swing_high", help="vah | pump_high | swing_high")
    ap.add_argument("--entry_mode", default="market", help="market | limit_vah")
    a = ap.parse_args()
    pr = default_params()
    pr.donchian_n = a.donchian_n
    pr.trailing_k = a.trailing_k
    pr.range_pctile_thr = a.range_pctile_thr
    pr.atr_ratio_thr = a.atr_ratio_thr
    pr.volume_ratio_thr = a.volume_ratio_thr
    pr.va_pctile_thr = a.va_pctile_thr
    pr.min_squeeze_duration = a.min_squeeze_duration
    pr.partial_tp_r = a.partial_tp_r
    pr.sl_mode = a.sl_mode
    pr.entry_mode = a.entry_mode
    run(a.tag, pr)
