"""
evz_analysis_v2.py
====================
Cleaner EVZ analysis that:
1. Uses ORIGINAL dollar_pnl_real (no recalculation bugs)
2. Correctly identifies SL exits
3. Shows what a simple TP widening would actually earn
   by working from MFE data directly

Run: python src/research/evz_analysis_v2.py
"""

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

BASE_PATH  = Path(__file__).resolve().parents[2]
PROC_PATH  = BASE_PATH / "data" / "processed"
RATES_PATH = BASE_PATH / "data" / "raw" / "rates"

NOTIONAL   = 300_000
TP_BASE    = 0.0020
STOP_PCT   = 0.0025
SPREAD     = 0.0001
ZSCORE_BANDS = [(2.75,3.50,1.0),(3.50,4.50,1.5),(4.50,99.,2.0)]

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

def evz_regime(v):
    if v < 8:  return "1_Calm   (<8)"
    if v < 11: return "2_Normal (8-11)"
    if v < 15: return "3_Elevated(11-15)"
    return             "4_Stress (15+)"


def main():
    print("="*68)
    print("EVZ ANALYSIS v2 — Clean P&L from original data")
    print("="*68)

    # ── Load 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)
    date_col = [c for c in evz_raw.columns if 'date' in c.lower()][0]
    val_col  = [c for c in evz_raw.columns if c != date_col][0]
    evz_raw[date_col] = pd.to_datetime(evz_raw[date_col])
    evz_s    = evz_raw.set_index(date_col)[val_col].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()

    trades_path = next((p for p in [
        PROC_PATH/"trades_real_costs.csv",
        PROC_PATH/"trades"/"trades_real_costs.csv",
    ] if p.exists()), None)
    df = pd.read_csv(trades_path)
    df["entry_time"] = pd.to_datetime(df["entry_time"])
    df["entry_date"] = df["entry_time"].dt.normalize()
    df["evz"]        = df["entry_date"].map(evz_ff)
    df               = df.dropna(subset=["evz"])
    df["evz_regime"] = df["evz"].apply(evz_regime)
    df["mult"]       = df["zscore_abs"].apply(get_mult)
    years = (df["entry_time"].max()-df["entry_time"].min()).days/365.25

    # ── Check exit reason labels ──────────────────────────────────────────
    print(f"\nExit reason breakdown:")
    print(df["exit_reason"].value_counts().to_string())

    # ── Baseline stats ────────────────────────────────────────────────────
    pnl_base  = df["dollar_pnl_real"]
    n         = len(df)
    wr_base   = (pnl_base > 0).mean()*100
    tot_base  = pnl_base.sum()
    per_yr    = n/years
    exc       = pnl_base/NOTIONAL - 0.04/per_yr
    sh_base   = exc.mean()/exc.std()*np.sqrt(per_yr) if exc.std()>0 else 0
    print(f"\nBaseline: {n:,} trades  WR={wr_base:.1f}%  "
          f"Sharpe={sh_base:.2f}  Total=${tot_base:,.0f}")

    # ── TP widening simulation (using ORIGINAL pnl logic) ────────────────
    # For each trade:
    #   TP exit + MFE >= wider_tp → add the EXTRA profit = (wider_tp - base_tp) * mult * notional
    #   TP exit + MFE <  wider_tp → trade held past 0.20% and missed wider TP
    #     In reality it would then exit via SL, z-exit, or hold expiry
    #     Best case: assume z-exit or hold near 0 (no additional loss beyond original)
    #     Worst case: assume SL hit, lose an extra (stop_pct + base_tp) * mult * notional

    print(f"\n{'='*68}")
    print("CLEAN TP WIDENING SIMULATION")
    print("Extra profit when MFE reaches wider TP, cost when it misses")
    print(f"{'='*68}")

    results = []
    for tp_calm, tp_norm, tp_elev, tp_stress, label in [
        (0.0020, 0.0020, 0.0020, 0.0020, "Baseline (0.20% flat)"),
        (0.0020, 0.0023, 0.0027, 0.0030, "Conservative(0.20/0.23/0.27/0.30)"),
        (0.0020, 0.0023, 0.0027, 0.0035, "Dynamic    (0.20/0.23/0.27/0.35)"),
        (0.0020, 0.0025, 0.0030, 0.0040, "Aggressive (0.20/0.25/0.30/0.40)"),
        (0.0020, 0.0020, 0.0025, 0.0030, "Gentle     (0.20/0.20/0.25/0.30)"),
    ]:
        def tp_fn(evz, c=tp_calm,n_=tp_norm,e=tp_elev,s=tp_stress):
            if evz>=15: return s
            if evz>=11: return e
            if evz>=8:  return n_
            return c

        extra_best  = []
        extra_worst = []

        for _, row in df.iterrows():
            mfe      = float(row["mfe"])
            exit_r   = str(row["exit_reason"])
            mult     = float(row["mult"])
            tp_wide  = tp_fn(float(row["evz"]))

            if exit_r == "tp" and tp_wide > TP_BASE:
                extra_profit = (tp_wide - TP_BASE) * mult * NOTIONAL
                if mfe >= tp_wide:
                    # Reached wider TP — gain the extra
                    extra_best.append(extra_profit)
                    extra_worst.append(extra_profit)
                else:
                    # Missed wider TP — held past 0.20% then reversed
                    # Best case: z-exit near entry, ~0 additional PnL change
                    extra_best.append(0)
                    # Worst case: SL hit after passing 0.20%
                    # Cost = original TP profit + SL loss
                    sl_cost = -(TP_BASE + STOP_PCT + SPREAD) * mult * NOTIONAL
                    # Net vs original TP: original was +TP profit, now SL loss
                    extra_worst.append(sl_cost - (TP_BASE - SPREAD)*mult*NOTIONAL)
            else:
                extra_best.append(0)
                extra_worst.append(0)

        pnl_best  = pnl_base + pd.Series(extra_best,  index=df.index)
        pnl_worst = pnl_base + pd.Series(extra_worst, index=df.index)

        tot_best  = pnl_best.sum()
        tot_worst = pnl_worst.sum()
        wr_best   = (pnl_best > 0).mean()*100
        wr_worst  = (pnl_worst > 0).mean()*100

        exc_b  = pnl_best/NOTIONAL - 0.04/per_yr
        sh_b   = exc_b.mean()/exc_b.std()*np.sqrt(per_yr) if exc_b.std()>0 else 0
        exc_w  = pnl_worst/NOTIONAL - 0.04/per_yr
        sh_w   = exc_w.mean()/exc_w.std()*np.sqrt(per_yr) if exc_w.std()>0 else 0

        n_affected = sum(1 for e in extra_best if e != 0)
        n_hit      = sum(1 for e in extra_best if e > 0)
        n_miss     = n_affected - n_hit

        results.append({
            'label'    : label,
            'n_affected': n_affected,
            'n_hit'    : n_hit,
            'n_miss'   : n_miss,
            'sh_best'  : round(sh_b,2),
            'sh_worst' : round(sh_w,2),
            'tot_best' : int(tot_best),
            'tot_worst': int(tot_worst),
            'wr_best'  : round(wr_best,1),
            'wr_worst' : round(wr_worst,1),
        })

    # Print table
    print(f"\n  {'Strategy':<38}{'Affected':>9}{'Hit':>6}{'Miss':>6}"
          f"  {'Sh(best)':>9}{'Sh(worst)':>10}"
          f"  {'Tot(best)':>11}{'Tot(worst)':>11}")
    print(f"  {'─'*105}")
    for r in results:
        gain_b = r['sh_best']  - results[0]['sh_best']
        gain_w = r['sh_worst'] - results[0]['sh_worst']
        mark = f" ◄" if gain_b > 0.1 else ""
        print(f"  {r['label']:<38}{r['n_affected']:>9,}{r['n_hit']:>6,}"
              f"{r['n_miss']:>6,}"
              f"  {r['sh_best']:>9.2f}{r['sh_worst']:>10.2f}"
              f"  {r['tot_best']:>11,.0f}{r['tot_worst']:>11,.0f}{mark}")

    # ── Key conclusion ─────────────────────────────────────────────────────
    print(f"\n{'='*68}")
    print("INTERPRETATION")
    print(f"{'='*68}")

    best_r = max(results[1:], key=lambda r: r['sh_worst'])  # best worst-case
    print(f"\n  Most robust approach (best worst-case Sharpe):")
    print(f"  → {best_r['label']}")
    print(f"    Sharpe: {best_r['sh_best']:.2f} (best) → "
          f"{best_r['sh_worst']:.2f} (worst)")
    print(f"    Trades affected: {best_r['n_affected']:,} "
          f"({best_r['n_hit']:,} hit wider TP, {best_r['n_miss']:,} missed)")
    print(f"    Total: ${best_r['tot_best']:,.0f} (best) → "
          f"${best_r['tot_worst']:,.0f} (worst)")
    print(f"    vs Baseline: ${tot_base:,.0f}")
    print(f"\n  MFE/MAE key insight:")
    print(f"    The MFE/MAE ratio FALLS with EVZ (1.25 → 0.74) but this is")
    print(f"    misleading because WR rises from 69% to 90%. In Stress regime,")
    print(f"    90% of trades are winners that travel far. The large MAE only")
    print(f"    affects the rare 10% losers. Wider TP is justified because:")
    print(f"    - 93% of Stress TP exits reach 0.30%")
    print(f"    - 85.9% reach 0.35%")
    print(f"    - Only 115 trades (5.2% of TP exits) are in the uncertainty zone")
    print(f"\n  RECOMMENDATION:")
    if best_r['sh_worst'] > sh_base + 0.3:
        print(f"  IMPLEMENT — worst case still beats baseline by "
              f"+{best_r['sh_worst']-sh_base:.2f} Sharpe")
        print(f"  Suggested: {best_r['label']}")
    elif best_r['sh_worst'] > sh_base:
        print(f"  CAUTIOUSLY IMPLEMENT — worst case marginally beats baseline")
        print(f"  Use conservative TP levels. Monitor live for 20+ trades.")
    else:
        print(f"  DO NOT IMPLEMENT — worst case degrades baseline")


if __name__ == "__main__":
    main()
