"""
zscore_exit_tuning_v1.py
=========================
Tunes the z-score reversal exit threshold to find the optimal sensitivity.

Finding from dynamic_exit_v1.py
---------------------------------
Method C (0.25% stop + z-score reversal at 1.5) is the only method
within the 10% DD limit but fires on 55% of trades — too aggressive.
The threshold of 1.5 cuts trades where the z-score temporarily dips
rather than genuinely reverses.

Goal
----
Find the threshold where z-score exit fires only on genuine thesis
failures — not normal oscillation. Target:
  - Z-exit fires on ~15-30% of trades
  - Stop still handles the fast catastrophic moves
  - Time exit captures the clean winners (40-50%)
  - Win rate 40-48% (improvement over 31% baseline)
  - Avg return preserved closer to baseline 0.000740
  - 95th pct DD stays below 10%

Thresholds tested: 1.0, 1.5, 2.0, 2.5, 3.0, 3.5

Place this file in:
  C:\\Users\\paul_\\OneDrive\\fx_macro_intraday\\src\\research\\zscore_exit_tuning_v1.py

Run from project root:
  python src/research/zscore_exit_tuning_v1.py
"""

import pandas as pd
import numpy as np
from pathlib import Path
import sys

BASE_PATH = Path(__file__).resolve().parents[2]
SRC_PATH  = BASE_PATH / "src"
if str(SRC_PATH) not in sys.path:
    sys.path.append(str(SRC_PATH))

from research.combined_candidate_matrix_v1 import (
    build_frozen_signals,
    find_entry_pullback,
    get_first_m15_idx_at_or_after,
)
from ingestion.price_loader_15m import load_eurusd_15m
from features.spot_lag_v3 import get_model_ready_spot_lag_v3

TRADES_DIR = BASE_PATH / "data" / "processed" / "trades"
TRADES_DIR.mkdir(parents=True, exist_ok=True)

# ── Frozen parameters ─────────────────────────────────────────────────────────
THRESHOLD     = 2.75
FIB           = 0.786
HOLD_HOURS    = 52
STOP          = 0.0025
SPREAD_COST   = 0.0001
ALLOWED_HOURS = set(range(7, 17))

ACCOUNT_START = 100_000.0
PROFIT_TARGET = ACCOUNT_START * 0.10
MAX_OVERALL   = ACCOUNT_START * 0.10
MAX_DAILY     = ACCOUNT_START * 0.05
BASE_RISK_PCT = 0.003
BASE_NOTIONAL = (ACCOUNT_START * BASE_RISK_PCT) / STOP

ZSCORE_BANDS = [
    (2.75, 3.50, 1.0),
    (3.50, 4.50, 1.5),
    (4.50, 99.0, 2.0),
]

N_SIMS      = 2000
RANDOM_SEED = 42
YEARS       = 22.0

# ── Thresholds to test ────────────────────────────────────────────────────────
THRESHOLDS = [1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 99.0]
# 99.0 = never fires = pure time exit with 0.25% stop (true baseline)


def get_multiplier(z: float) -> float:
    for lo, hi, mult in ZSCORE_BANDS:
        if lo <= z < hi:
            return mult
    return ZSCORE_BANDS[-1][2]


# ── Trade simulator ───────────────────────────────────────────────────────────
def simulate_trade(
    m15         : pd.DataFrame,
    signal_df   : pd.DataFrame,
    entry_time  : pd.Timestamp,
    entry_price : float,
    signal      : int,
    z_threshold : float,
    hold_hours  : int = HOLD_HOURS,
    stop        : float = STOP,
    spread_cost : float = SPREAD_COST,
) -> dict | None:
    entry_idx = get_first_m15_idx_at_or_after(m15, entry_time)
    if entry_idx is None:
        return None

    hold_bars = hold_hours * 4
    exit_idx  = min(entry_idx + hold_bars, len(m15) - 1)
    path      = m15.iloc[entry_idx : exit_idx + 1].copy()
    if path.empty:
        return None

    # Get z-score bars during the hold window
    use_z_exit = z_threshold < 90.0
    if use_z_exit:
        signal_window = signal_df[
            (signal_df["datetime"] >= entry_time) &
            (signal_df["datetime"] <= path.iloc[-1]["datetime"])
        ].copy()
    else:
        signal_window = pd.DataFrame()

    exit_price  = float(path.iloc[-1]["close"])
    exit_reason = "time"
    stop_hit    = False
    final_bar_i = len(path) - 1

    for bar_i in range(len(path)):
        bar = path.iloc[bar_i]

        # Stop check
        if signal == 1:
            adverse = (float(bar["low"]) - entry_price) / entry_price
        else:
            adverse = -((float(bar["high"]) - entry_price) / entry_price)

        if adverse <= -stop:
            exit_price  = entry_price * (1 - stop) if signal == 1 \
                          else entry_price * (1 + stop)
            exit_reason = "stop"
            stop_hit    = True
            final_bar_i = bar_i
            break

        # Z-score reversal exit
        if use_z_exit and not signal_window.empty:
            z_bars = signal_window[signal_window["datetime"] <= bar["datetime"]]
            if not z_bars.empty:
                current_z = float(z_bars.iloc[-1]["lag_zscore_24h_v3"])
                if signal == 1 and current_z <= -z_threshold:
                    exit_price  = float(bar["close"])
                    exit_reason = "zscore_reversal"
                    final_bar_i = bar_i
                    break
                elif signal == -1 and current_z >= z_threshold:
                    exit_price  = float(bar["close"])
                    exit_reason = "zscore_reversal"
                    final_bar_i = bar_i
                    break

    # MAE and MFE over full path (not just to exit — shows what was available)
    if signal == 1:
        mae = (path["low"].min()  - entry_price) / entry_price
        mfe = (path["high"].max() - entry_price) / entry_price
        if stop_hit:
            raw_ret = -stop
        else:
            raw_ret = (exit_price - entry_price) / entry_price
    else:
        mae = -((path["high"].max() - entry_price) / entry_price)
        mfe = (entry_price - path["low"].min()) / entry_price
        if stop_hit:
            raw_ret = -stop
        else:
            raw_ret = -(exit_price - entry_price) / entry_price

    exit_time = path.iloc[final_bar_i]["datetime"]

    return {
        "entry_time" : entry_time,
        "exit_time"  : exit_time,
        "signal"     : signal,
        "entry_price": entry_price,
        "exit_price" : exit_price,
        "stop_hit"   : stop_hit,
        "exit_reason": exit_reason,
        "mae"        : mae,
        "mfe"        : mfe,
        "return"     : raw_ret - spread_cost,
    }


# ── Run one threshold ─────────────────────────────────────────────────────────
def run_threshold(
    z_threshold : float,
    signals     : pd.DataFrame,
    m15         : pd.DataFrame,
    signal_df   : pd.DataFrame,
) -> pd.DataFrame:
    trades         = []
    last_exit_time = None

    for _, sig in signals.iterrows():
        if last_exit_time is not None and sig["datetime"] < last_exit_time:
            continue

        entry = find_entry_pullback(
            m15=m15, signal_row=sig, fib=FIB, wait_hours=6
        )
        if entry is None:
            continue

        trade = simulate_trade(
            m15         =m15,
            signal_df   =signal_df,
            entry_time  =entry["entry_time"],
            entry_price =entry["entry_price"],
            signal      =int(sig["signal"]),
            z_threshold =z_threshold,
        )
        if trade is None:
            continue

        trade["zscore_abs"] = abs(float(sig["lag_zscore_24h_v3"]))
        trades.append(trade)
        last_exit_time = trade["exit_time"]

    if not trades:
        return pd.DataFrame()

    df             = pd.DataFrame(trades)
    df["equity"]   = (1 + df["return"]).cumprod()
    df["peak"]     = df["equity"].cummax()
    df["drawdown"] = df["equity"] / df["peak"] - 1
    df["win"]      = (df["return"] > 0).astype(int)
    return df


# ── FTMO Monte Carlo ──────────────────────────────────────────────────────────
def run_mc(df: pd.DataFrame) -> dict:
    rng     = np.random.default_rng(RANDOM_SEED)
    returns = df["return"].values
    zscores = df["zscore_abs"].values
    exits   = df["exit_time"].values
    per_yr  = len(df) / YEARS

    results = []
    for _ in range(N_SIMS):
        idx    = rng.permutation(len(returns))
        sh_ret = returns[idx]
        sh_mult= np.array([get_multiplier(z) for z in zscores[idx]])

        balance  = ACCOUNT_START
        peak     = ACCOUNT_START
        max_dd   = 0.0
        dpnl     = {}
        outcome  = "INCOMPLETE"
        ttrades  = 0

        for ret, mult, ex_dt in zip(sh_ret, sh_mult, exits):
            dk = str(pd.Timestamp(ex_dt).date())
            if dk not in dpnl:
                dpnl[dk] = 0.0
            notional  = min(BASE_NOTIONAL * mult, BASE_NOTIONAL * 3)
            pnl       = ret * notional
            balance  += pnl
            ttrades  += 1
            dpnl[dk] += pnl

            if balance > peak:
                peak = balance
            dd = (peak - balance) / ACCOUNT_START
            if dd > max_dd:
                max_dd = dd

            if dpnl[dk] < -MAX_DAILY:
                outcome = "BREACH_DAILY"; break
            if balance <= ACCOUNT_START - MAX_OVERALL:
                outcome = "BREACH_OVERALL"; break
            if balance >= ACCOUNT_START + PROFIT_TARGET:
                outcome = "PASS"; break

        if outcome == "INCOMPLETE":
            outcome = ("PASS" if balance >= ACCOUNT_START + PROFIT_TARGET
                       else "BREACH_OVERALL"
                       if (peak - balance) / ACCOUNT_START >= 0.10
                       else "INCOMPLETE")

        results.append({"outcome": outcome, "max_dd": max_dd * 100, "trades": ttrades})

    mc = pd.DataFrame(results)
    passing = mc[mc["outcome"] == "PASS"]
    return {
        "pass_rate"  : (mc["outcome"] == "PASS").mean() * 100,
        "breach_rate": mc["outcome"].str.startswith("BREACH").mean() * 100,
        "dd_95"      : mc["max_dd"].quantile(0.95),
        "dd_50"      : mc["max_dd"].quantile(0.50),
        "avg_months" : (passing["trades"].mean()   / per_yr * 12
                        if len(passing) > 0 else np.nan),
        "med_months" : (passing["trades"].median() / per_yr * 12
                        if len(passing) > 0 else np.nan),
    }


# ── Main ──────────────────────────────────────────────────────────────────────
def main():
    print("=" * 70)
    print("Z-SCORE EXIT TUNING V1")
    print(f"  Testing reversal thresholds: {THRESHOLDS}")
    print(f"  Stop: {STOP:.2%}  |  Hold: {HOLD_HOURS}h  |  Threshold: {THRESHOLD}")
    print("=" * 70)

    print("\nLoading data...")
    m15       = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)
    signal_df = get_model_ready_spot_lag_v3().copy()
    signal_df["datetime"] = pd.to_datetime(signal_df["datetime"])
    signal_df = signal_df.sort_values("datetime").reset_index(drop=True)
    signals   = build_frozen_signals(
        threshold=THRESHOLD, allowed_hours=ALLOWED_HOURS
    )
    print(f"  15M bars  : {len(m15):,}")
    print(f"  Signal df : {len(signal_df):,}")
    print(f"  Signals   : {len(signals)}")

    rows = []
    print(f"\nRunning {len(THRESHOLDS)} threshold levels...")
    print(f"  (z=99.0 = no z-exit = pure baseline)\n")

    for z_thresh in THRESHOLDS:
        label = f"z={z_thresh:.1f}" if z_thresh < 90 else "NO_Z_EXIT"
        print(f"  Running threshold {label}...")
        df = run_threshold(z_thresh, signals, m15, signal_df)

        if df.empty:
            print(f"    [WARNING] No trades")
            continue

        mc = run_mc(df)

        n     = len(df)
        wr    = df["win"].mean() * 100
        ar    = df["return"].mean()
        fe    = df["equity"].iloc[-1]
        dd    = df["drawdown"].min() * 100

        streak = 0
        max_s  = 0
        for r in df["return"]:
            if r <= 0:
                streak += 1
                max_s   = max(max_s, streak)
            else:
                streak = 0

        # Exit reason breakdown
        er = df["exit_reason"].value_counts(normalize=True) * 100
        stop_pct  = er.get("stop",           0)
        z_pct     = er.get("zscore_reversal", 0)
        time_pct  = er.get("time",            0)

        rows.append({
            "threshold"  : z_thresh,
            "label"      : label,
            "trades"     : n,
            "per_yr"     : n / YEARS,
            "win_rate"   : wr,
            "avg_return" : ar,
            "final_eq"   : fe,
            "max_dd"     : dd,
            "max_streak" : max_s,
            "stop_pct"   : stop_pct,
            "z_pct"      : z_pct,
            "time_pct"   : time_pct,
            "pass_rate"  : mc["pass_rate"],
            "dd_95"      : mc["dd_95"],
            "dd_50"      : mc["dd_50"],
            "avg_months" : mc["avg_months"],
            "med_months" : mc["med_months"],
        })

        print(f"    Trades={n:,}  WR={wr:.1f}%  AvgRet={ar:.6f}  "
              f"FinalEq={fe:.4f}  MaxDD={dd:.2f}%")
        print(f"    Exits: stop={stop_pct:.1f}%  z-exit={z_pct:.1f}%  "
              f"time={time_pct:.1f}%")
        print(f"    MC: pass={mc['pass_rate']:.2f}%  "
              f"95pct_DD={mc['dd_95']:.3f}%  "
              f"med_months={mc['med_months']:.1f}")

    # ── Full results table ────────────────────────────────────────────────────
    print(f"\n{'='*70}")
    print("FULL RESULTS TABLE")
    print(f"{'='*70}")

    print(f"\n  {'Label':<14}{'Trades':>7}{'WR%':>7}{'AvgRet':>10}"
          f"{'FinalEq':>10}{'MaxDD%':>8}"
          f"{'Stop%':>7}{'Z%':>7}{'Time%':>7}"
          f"{'Pass%':>8}{'DD95':>8}{'MedMths':>9}")
    print(f"  {'─'*98}")

    for r in rows:
        valid = "✓" if r["dd_95"] < 10.0 else "✗"
        bl    = " ◄ BASELINE" if r["label"] == "NO_Z_EXIT" else ""
        print(f"  {r['label']:<14}{r['trades']:>7,}{r['win_rate']:>7.1f}"
              f"{r['avg_return']:>10.6f}{r['final_eq']:>10.4f}"
              f"{r['max_dd']:>8.2f}"
              f"{r['stop_pct']:>7.1f}{r['z_pct']:>7.1f}{r['time_pct']:>7.1f}"
              f"{r['pass_rate']:>8.2f}{r['dd_95']:>8.3f}"
              f"{r['med_months']:>9.1f}  {valid}{bl}")

    # ── Sweet spot analysis ───────────────────────────────────────────────────
    print(f"\n{'='*70}")
    print("SWEET SPOT ANALYSIS")
    print(f"{'='*70}")

    valid_rows = [r for r in rows if r["dd_95"] < 10.0]

    if valid_rows:
        # Best by avg_return among valid
        best_ret  = max(valid_rows, key=lambda r: r["avg_return"])
        best_pass = max(valid_rows, key=lambda r: r["pass_rate"])
        best_eq   = max(valid_rows, key=lambda r: r["final_eq"])
        best_mths = min(valid_rows, key=lambda r: r["med_months"])

        print(f"\n  Among thresholds with 95th pct DD < 10%:")
        print(f"    Best avg return : threshold {best_ret['label']}"
              f"  ({best_ret['avg_return']:.6f})")
        print(f"    Best pass rate  : threshold {best_pass['label']}"
              f"  ({best_pass['pass_rate']:.2f}%)")
        print(f"    Best final equity: threshold {best_eq['label']}"
              f"  ({best_eq['final_eq']:.4f})")
        print(f"    Fastest challenge: threshold {best_mths['label']}"
              f"  ({best_mths['med_months']:.1f} months)")

        print(f"\n  Target zone: z-exit fires on 15-30% of trades")
        for r in rows:
            in_zone = 15 <= r["z_pct"] <= 30
            valid   = r["dd_95"] < 10.0
            if in_zone:
                print(f"    {r['label']:<12}: z_exit={r['z_pct']:.1f}%  "
                      f"WR={r['win_rate']:.1f}%  "
                      f"AvgRet={r['avg_return']:.6f}  "
                      f"DD95={r['dd_95']:.3f}%  "
                      f"{'VALID' if valid else 'OVER DD'}")

        print(f"""
  Key insight:
    The optimal threshold is the HIGHEST z-score value where:
    1. 95th pct DD stays below 10%
    2. z-exit fires on 15-30% of trades (cuts genuine reversals,
       not normal oscillation)
    3. Time exit still captures the majority of winners cleanly
    4. Avg return stays above the baseline where possible

  Next step:
    Take the optimal threshold into the full FTMO simulation
    and compare directly against the frozen growth_52h baseline.
""")
    else:
        print("\n  No threshold achieves 95th pct DD < 10%")
        print("  The z-score exit alone cannot solve the DD problem")
        print("  Consider combining with reduced risk % (0.25% per trade)")

    # Save results
    result_df = pd.DataFrame(rows)
    out_path  = TRADES_DIR / "zscore_threshold_results.csv"
    result_df.to_csv(out_path, index=False)
    print(f"  Results saved: {out_path}")


if __name__ == "__main__":
    main()
