"""
evz_full_backtest_v1.py
=========================
Full re-simulation of the EURUSD macro model with EVZ-adaptive TP levels.
This is the proper test — not MFE approximations, but actual trade-by-trade
execution with wider TPs, tracking real outcomes when trades pass 0.20% TP
but do not reach the wider target.

Steps:
  1. Full backtest with each TP config (2008-2026, EVZ coverage period)
  2. Walk-forward validation (IS: 2008-2019, OOS: 2020+)
  3. Monte Carlo FTMO pass rate analysis
  4. Head-to-head comparison: baseline vs best EVZ-adaptive config

Run from project root (~45 mins):
  python src/research/evz_full_backtest_v1.py
"""

import pandas as pd
import numpy as np
from pathlib import Path
import sys
import warnings
warnings.filterwarnings('ignore')

BASE_PATH  = Path(__file__).resolve().parents[2]
SRC_PATH   = BASE_PATH / "src"
PROC_PATH  = BASE_PATH / "data" / "processed"
RATES_PATH = BASE_PATH / "data" / "raw" / "rates"
if str(SRC_PATH) not in sys.path:
    sys.path.insert(0, str(SRC_PATH))

# ── Validated model constants ─────────────────────────────────────────────────
THRESHOLD   = 2.75
FIB         = 0.786
STOP_PCT    = 0.0025
ZSCORE_EXIT = 1.5
SPREAD_COST = 0.0001
HOLD_HOURS  = 52
NOTIONAL    = 300_000
ZSCORE_BANDS = [(2.75,3.50,1.0),(3.50,4.50,1.5),(4.50,99.,2.0)]

# ── TP configs to test ────────────────────────────────────────────────────────
TP_CONFIGS = [
    (0.0020,0.0020,0.0020,0.0020, "Baseline    (0.20% flat)"),
    (0.0020,0.0020,0.0025,0.0030, "Conservative(0.20/0.20/0.25/0.30)"),
    (0.0020,0.0023,0.0027,0.0035, "Dynamic     (0.20/0.23/0.27/0.35)"),
    (0.0020,0.0025,0.0035,0.0050, "Bold        (0.20/0.25/0.35/0.50)"),
    (0.0020,0.0030,0.0040,0.0060, "Wider       (0.20/0.30/0.40/0.60)"),
]

WF_SPLIT    = pd.Timestamp('2020-01-01')  # Walk-forward OOS start
EVZ_START   = pd.Timestamp('2008-01-01')  # EVZ data available from

def get_mult(z):
    for lo,hi,m in ZSCORE_BANDS:
        if lo<=z<hi: return m
    return 2.0

def get_tp(evz_val, tp_calm, tp_norm, tp_elev, tp_stress):
    if evz_val is None: return tp_calm
    if evz_val >= 15:   return tp_stress
    if evz_val >= 11:   return tp_elev
    if evz_val >= 8:    return tp_norm
    return tp_calm

def evz_regime(v):
    if v is None:  return "Unknown"
    if v < 8:      return "Calm"
    if v < 11:     return "Normal"
    if v < 15:     return "Elevated"
    return                "Stress"


def run_backtest(signal_df, m15_idx, z_lookup, evz_ff,
                 tp_calm, tp_norm, tp_elev, tp_stress,
                 label, evz_start=None, evz_end=None):
    """
    Full trade simulation with EVZ-adaptive TP.
    Identical to validated backtest except TP varies by EVZ level at entry.
    """
    trades    = []
    last_exit = None

    for _, row in signal_df.iterrows():
        dt = row["datetime"]
        z  = float(row["lag_zscore_24h_v3"])

        if abs(z) < THRESHOLD:
            continue
        if last_exit is not None and dt <= last_exit:
            continue

        # Restrict to EVZ coverage period if specified
        if evz_start and dt < evz_start:
            continue
        if evz_end and dt > evz_end:
            continue

        direction = 1 if z >= THRESHOLD else -1

        # Get EVZ at signal date
        signal_date = dt.normalize()
        evz_val = None
        for lag in range(5):
            d = signal_date - pd.Timedelta(days=lag)
            if d in evz_ff.index:
                evz_val = float(evz_ff[d])
                break

        # Select TP based on EVZ
        tp_pct = get_tp(evz_val, tp_calm, tp_norm, tp_elev, tp_stress)

        # Signal bar OHLC
        bar_end   = dt + pd.Timedelta(hours=1) - pd.Timedelta(minutes=1)
        bar_slice = m15_idx.loc[dt:bar_end]
        if bar_slice.empty:
            continue

        bh = float(bar_slice["high"].max())
        bl = float(bar_slice["low"].min())
        bc = float(bar_slice["close"].iloc[-1])

        if direction == 1:
            pull = bc - bl
            if pull <= 0.00005: continue
            target = bc - FIB * pull
        else:
            pull = bh - bc
            if pull <= 0.00005: continue
            target = bc + FIB * pull

        # 6-hour entry window
        entry_slice = m15_idx.loc[
            dt + pd.Timedelta(minutes=1):dt + pd.Timedelta(hours=6)]
        entry_price = entry_time = None
        for edt, ebar in entry_slice.iterrows():
            if direction == 1 and float(ebar["low"]) <= target:
                entry_price = target; entry_time = edt; break
            elif direction == -1 and float(ebar["high"]) >= target:
                entry_price = target; entry_time = edt; break

        if entry_price is None:
            continue

        tp = entry_price*(1+tp_pct) if direction==1 else entry_price*(1-tp_pct)
        sl = entry_price*(1-STOP_PCT) if direction==1 else entry_price*(1+STOP_PCT)

        hold = m15_idx.loc[
            entry_time + pd.Timedelta(minutes=1):
            entry_time + pd.Timedelta(hours=HOLD_HOURS)]

        exit_price  = None
        exit_reason = "hold_expiry"

        for hdt, hbar in hold.iterrows():
            hr = hdt.replace(minute=0, second=0, microsecond=0)
            if hr in z_lookup.index:
                hz = float(z_lookup[hr])
                if ((direction==1 and hz<=-ZSCORE_EXIT) or
                        (direction==-1 and hz>=ZSCORE_EXIT)):
                    exit_price  = float(hbar["close"])
                    exit_reason = "z_exit"; break
            if direction == 1:
                if float(hbar["high"]) >= tp:
                    exit_price = tp; exit_reason = "tp"; break
                if float(hbar["low"])  <= sl:
                    exit_price = sl; exit_reason = "stop"; break
            else:
                if float(hbar["low"])  <= tp:
                    exit_price = tp; exit_reason = "tp"; break
                if float(hbar["high"]) >= sl:
                    exit_price = sl; exit_reason = "stop"; break

        if exit_price is None:
            exit_price  = float(hold["close"].iloc[-1]) if not hold.empty else entry_price
            exit_reason = "hold_expiry"

        mult = get_mult(abs(z))
        pnl  = round(((exit_price-entry_price)/entry_price
                       * direction - SPREAD_COST) * mult * NOTIONAL, 2)

        trades.append({
            "entry_time"  : entry_time,
            "exit_time"   : entry_time + pd.Timedelta(hours=HOLD_HOURS),
            "direction"   : direction,
            "zscore"      : round(z, 4),
            "mult"        : mult,
            "evz"         : round(evz_val,1) if evz_val else None,
            "evz_regime"  : evz_regime(evz_val),
            "tp_pct"      : tp_pct,
            "exit_reason" : exit_reason,
            "pnl"         : pnl,
        })
        last_exit = entry_time + pd.Timedelta(hours=HOLD_HOURS)

    return pd.DataFrame(trades)


def score(df, label, years=None):
    if df is None or len(df) == 0:
        return {"label":label,"n":0,"wr":0,"sharpe":0,"pf":0,"dd":0,
                "total":0,"per_yr":0}
    pnl    = df["pnl"].values
    n      = len(pnl)
    wr     = (pnl > 0).mean() * 100
    total  = pnl.sum()
    eq     = 100_000 + np.cumsum(pnl)
    peak   = np.maximum.accumulate(eq)
    max_dd = ((eq-peak)/peak*100).min()
    yrs    = years if years else n/182
    per_yr = n / yrs
    exc    = pnl/100_000 - 0.04/per_yr
    sharpe = exc.mean()/exc.std()*np.sqrt(per_yr) if exc.std()>0 else 0
    gp = pnl[pnl>0].sum(); gl = abs(pnl[pnl<0].sum())
    pf = round(gp/gl,2) if gl>0 else 999
    tp_pct = (df["exit_reason"]=="tp").mean()*100 if "exit_reason" in df else 0
    sl_pct = (df["exit_reason"]=="stop").mean()*100 if "exit_reason" in df else 0
    return {"label":label,"n":n,"wr":round(wr,1),"sharpe":round(sharpe,2),
            "pf":pf,"dd":round(max_dd,2),"total":int(total),
            "per_yr":round(per_yr,1),"tp_pct":round(tp_pct,1),
            "sl_pct":round(sl_pct,1)}


def print_table(rows, title):
    print(f"\n{'='*78}")
    print(title)
    print(f"{'='*78}")
    print(f"  {'Config':<40}{'N':>5}{'WR%':>6}{'Sharpe':>8}"
          f"{'PF':>6}{'MaxDD%':>8}{'Total$':>11}{'Tr/yr':>7}")
    print(f"  {'─'*78}")
    base_sh = next((r['sharpe'] for r in rows if 'Baseline' in r['label']), 0)
    for r in rows:
        if r['n'] == 0: continue
        gain = r['sharpe'] - base_sh
        mark = f" ◄+{gain:.2f}" if gain>=0.3 and 'Baseline' not in r['label'] else ""
        print(f"  {r['label']:<40}{r['n']:>5}{r['wr']:>6.1f}"
              f"{r['sharpe']:>8.2f}{r['pf']:>6.2f}{r['dd']:>8.2f}"
              f"{r['total']:>11,.0f}{r['per_yr']:>7.1f}{mark}")


def mc_ftmo(pnl_series, n_sims=2000, account=100_000,
            daily_limit=0.03, total_limit=0.10,
            profit_target=0.10, seed=42):
    """Monte Carlo simulation of FTMO 1-Step challenge."""
    np.random.seed(seed)
    pnl = np.array(pnl_series)
    passes = 0; daily_b = 0; total_b = 0; no_profit = 0

    for _ in range(n_sims):
        sim   = np.random.choice(pnl, size=len(pnl), replace=True)
        eq    = np.concatenate([[account], account + np.cumsum(sim)])
        peak  = np.maximum.accumulate(eq)
        dd    = (eq - peak) / account

        # Daily DD: approximate each trade as one day
        daily_chg = np.diff(eq) / account
        daily_hit = (daily_chg < -daily_limit).any()
        total_hit = (dd[1:] < -total_limit).any()
        profit_ok = (eq[-1]-account)/account >= profit_target

        if daily_hit:                    daily_b += 1
        if total_hit:                    total_b += 1
        if profit_ok and not daily_hit and not total_hit: passes += 1
        if not profit_ok and not daily_hit and not total_hit: no_profit += 1

    return {
        "pass_rate"    : round(passes/n_sims*100,1),
        "daily_breach" : round(daily_b/n_sims*100,1),
        "total_breach" : round(total_b/n_sims*100,1),
        "no_profit"    : round(no_profit/n_sims*100,1),
    }


def main():
    print("="*78)
    print("EVZ FULL BACKTEST v1 — Proper re-simulation with adaptive TP")
    print("="*78)

    # ── Load data ─────────────────────────────────────────────────────────
    print("\nLoading EVZ data...")
    evz_path = next((p for p in [
        RATES_PATH/"EVZCLS.csv", RATES_PATH/"evz.csv"
    ] if p.exists()), None)
    evz_raw  = pd.read_csv(evz_path)
    dc = [c for c in evz_raw.columns if 'date' in c.lower()][0]
    vc = [c for c in evz_raw.columns if c!=dc][0]
    evz_raw[dc] = pd.to_datetime(evz_raw[dc])
    evz_s    = evz_raw.set_index(dc)[vc].rename("evz").sort_index()
    full_idx = pd.date_range(evz_s.index.min(), evz_s.index.max(), freq='D')
    evz_ff   = evz_s.reindex(full_idx).ffill()
    print(f"  EVZ: {evz_s.index[0].date()} → {evz_s.index[-1].date()}")

    print("\nLoading signal series...")
    try:
        from features.spot_lag_v3 import get_model_ready_spot_lag_v3
        sig = get_model_ready_spot_lag_v3().copy()
        sig["datetime"] = pd.to_datetime(sig["datetime"])
        sig = sig.sort_values("datetime").reset_index(drop=True)
        z_lookup = sig.set_index("datetime")["lag_zscore_24h_v3"]
        print(f"  Signals: {len(sig):,} rows")
    except Exception as e:
        print(f"  ERROR: {e}"); sys.exit(1)

    print("\nLoading 15M price data...")
    try:
        from ingestion.price_loader_15m import load_eurusd_15m
        m15     = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)
        m15["datetime"] = pd.to_datetime(m15["datetime"])
        m15_idx = m15.set_index("datetime").sort_index()
        print(f"  15M bars: {len(m15):,}")
    except Exception as e:
        print(f"  ERROR: {e}"); sys.exit(1)

    # ── Full period backtest (EVZ coverage: 2008+) ────────────────────────
    print(f"\n{'='*78}")
    print("STEP 1: FULL PERIOD BACKTEST (2008-2026, with EVZ)")
    print(f"{'='*78}")

    all_results = []
    all_trades  = {}

    for cfg in TP_CONFIGS:
        tp_c,tp_n,tp_e,tp_s,label = cfg
        print(f"\n  Running: {label}...")
        df_trades = run_backtest(
            sig, m15_idx, z_lookup, evz_ff,
            tp_c, tp_n, tp_e, tp_s, label,
            evz_start=EVZ_START
        )
        years = (df_trades["entry_time"].max() -
                 df_trades["entry_time"].min()).days / 365.25 if len(df_trades) else 17
        r = score(df_trades, label, years)
        all_results.append(r)
        all_trades[label] = df_trades
        print(f"    Trades={r['n']:,}  WR={r['wr']:.1f}%  "
              f"Sharpe={r['sharpe']:.2f}  Total=${r['total']:,.0f}")

    print_table(all_results, "FULL PERIOD — EVZ-Adaptive TP vs Baseline")

    # ── Walk-forward: IS and OOS ──────────────────────────────────────────
    print(f"\n{'='*78}")
    print(f"STEP 2: WALK-FORWARD VALIDATION (IS: 2008-2019, OOS: 2020+)")
    print(f"{'='*78}")

    wf_results = []
    for cfg in TP_CONFIGS:
        tp_c,tp_n,tp_e,tp_s,label = cfg
        df = all_trades[label]
        if len(df) == 0: continue

        df_is  = df[df["entry_time"] < WF_SPLIT]
        df_oos = df[df["entry_time"] >= WF_SPLIT]

        yrs_is  = (WF_SPLIT - EVZ_START).days / 365.25
        yrs_oos = (pd.Timestamp('2026-01-01') - WF_SPLIT).days / 365.25

        r_is  = score(df_is,  f"IS  | {label}", yrs_is)
        r_oos = score(df_oos, f"OOS | {label}", yrs_oos)
        wf_results.extend([r_is, r_oos])

    print_table(wf_results, "WALK-FORWARD — In-Sample (2008-2019) vs OOS (2020+)")

    # ── OOS only comparison ───────────────────────────────────────────────
    print(f"\n{'─'*78}")
    print("OOS ONLY (2020+) — Most honest comparison")
    print(f"{'─'*78}")
    oos_results = []
    for cfg in TP_CONFIGS:
        tp_c,tp_n,tp_e,tp_s,label = cfg
        df     = all_trades[label]
        df_oos = df[df["entry_time"] >= WF_SPLIT]
        yrs    = (pd.Timestamp('2026-03-01') - WF_SPLIT).days / 365.25
        oos_results.append(score(df_oos, label, yrs))

    print_table(oos_results, "OOS RESULTS (2020+)")

    # ── Monte Carlo FTMO ─────────────────────────────────────────────────
    print(f"\n{'='*78}")
    print("STEP 3: MONTE CARLO — FTMO 1-STEP CHALLENGE")
    print(f"{'='*78}")
    print(f"\n  Account: $100k | Profit target: 10% | Daily DD: 3% | Total DD: 10%")
    print(f"  Using OOS trades (2020+) as sampling distribution")
    print(f"  2,000 simulations per config\n")

    print(f"  {'Config':<40}{'Pass%':>8}{'Daily%':>9}{'Total%':>9}{'NoProf%':>10}")
    print(f"  {'─'*72}")

    best_pass = 0
    best_label = ""
    for cfg in TP_CONFIGS:
        tp_c,tp_n,tp_e,tp_s,label = cfg
        df     = all_trades[label]
        df_oos = df[df["entry_time"] >= WF_SPLIT]
        if len(df_oos) < 20:
            print(f"  {label:<40}  insufficient OOS trades")
            continue
        mc = mc_ftmo(df_oos["pnl"].values)
        mark = " ◄ BEST" if mc['pass_rate'] > best_pass else ""
        if mc['pass_rate'] > best_pass:
            best_pass  = mc['pass_rate']
            best_label = label
        print(f"  {label:<40}{mc['pass_rate']:>8.1f}%"
              f"{mc['daily_breach']:>9.1f}%"
              f"{mc['total_breach']:>9.1f}%"
              f"{mc['no_profit']:>10.1f}%{mark}")

    # ── By EVZ regime ─────────────────────────────────────────────────────
    print(f"\n{'='*78}")
    print("STEP 4: BY EVZ REGIME — Where does the edge come from?")
    print(f"{'='*78}")

    base_label = TP_CONFIGS[0][4]
    best_non_base = max(TP_CONFIGS[1:], key=lambda c: 
        score(all_trades.get(c[4], pd.DataFrame()), c[4])['sharpe'])[4]

    for lbl in [base_label, best_non_base]:
        df = all_trades.get(lbl)
        if df is None or len(df)==0: continue
        print(f"\n  {lbl}:")
        print(f"  {'Regime':<16}{'N':>5}{'WR%':>7}{'Sharpe':>9}"
              f"{'Total$':>12}{'TP%':>7}{'SL%':>7}")
        print(f"  {'─'*60}")
        for regime in ['Calm','Normal','Elevated','Stress']:
            sub = df[df["evz_regime"]==regime]
            if len(sub)==0: continue
            r   = score(sub, regime)
            tp  = (sub["exit_reason"]=="tp").mean()*100
            sl  = (sub["exit_reason"]=="stop").mean()*100
            print(f"  {regime:<16}{len(sub):>5}{r['wr']:>7.1f}"
                  f"{r['sharpe']:>9.2f}{r['total']:>12,.0f}"
                  f"{tp:>7.1f}{sl:>7.1f}")

    # ── Final verdict ─────────────────────────────────────────────────────
    print(f"\n{'='*78}")
    print("FINAL VERDICT")
    print(f"{'='*78}")

    base_full = all_results[0]
    best_full = max(all_results[1:], key=lambda r: r['sharpe'])
    best_oos  = max(oos_results[1:], key=lambda r: r['sharpe'])

    sharpe_gain_full = best_full['sharpe'] - base_full['sharpe']
    sharpe_gain_oos  = best_oos['sharpe']  - oos_results[0]['sharpe']

    print(f"\n  Full period  baseline : Sharpe {base_full['sharpe']:.2f}")
    print(f"  Full period  best EVZ : Sharpe {best_full['sharpe']:.2f}"
          f"  (+{sharpe_gain_full:.2f})  [{best_full['label'].strip()}]")
    print(f"\n  OOS baseline          : Sharpe {oos_results[0]['sharpe']:.2f}")
    print(f"  OOS best EVZ          : Sharpe {best_oos['sharpe']:.2f}"
          f"  (+{sharpe_gain_oos:.2f})  [{best_oos['label'].strip()}]")
    print(f"  Best FTMO pass rate   : {best_pass:.1f}%  [{best_label.strip()}]")

    if sharpe_gain_oos >= 0.3 and best_pass >= best_pass:
        verdict = "IMPLEMENT — OOS confirms the edge, FTMO risk acceptable"
        detail  = (f"Start with {best_oos['label'].strip()} in live trading.\n"
                   f"  Monitor first 20 live trades before expanding to all EVZ regimes.")
    elif sharpe_gain_oos >= 0.1:
        verdict = "CAUTIOUS — marginal OOS improvement, monitor closely"
        detail  = "Consider Stress regime only (EVZ 15+) as first live test."
    else:
        verdict = "DO NOT IMPLEMENT — OOS does not confirm in-sample edge"
        detail  = "The wider TP is curve-fitted to historical data."

    print(f"\n  VERDICT : {verdict}")
    print(f"  DETAIL  : {detail}")

    # Save all trades
    for label, df in all_trades.items():
        safe = label.replace(' ','_').replace('/','_').replace('(','').replace(')','')[:30]
        out  = PROC_PATH / f"evz_trades_{safe}.csv"
        df.to_csv(out, index=False)
    print(f"\n  Trade files saved to data/processed/")


if __name__ == "__main__":
    main()
