"""
regime_segmentation_v1.py  (v2 — volatility overlay)
======================================================
Exact 4 regimes as specified + volatility overlay per regime.

  R1: Pre-QE / Early cycle    2003–2008
  R2: QE / Low vol            2009–2019
  R3: Post-COVID / high vol   2020–2022
  R4: Hiking cycle            2022–2026

Volatility overlay:
  Realised vol of daily P&L per regime (proxy for market vol)
  Trade frequency per regime
  Separates: is it the MACRO that drives performance, or just the VOLATILITY?
"""
import sys, warnings
from pathlib import Path
import numpy as np
import pandas as pd
warnings.filterwarnings('ignore')

BASE    = Path(__file__).resolve().parents[2]
EU_FILE = BASE / 'data/processed/trades/trades_real_costs.csv'
UJ_FILE = BASE / 'data/processed/usdjpy_trades_real_costs.csv'
RATES   = BASE / 'data/raw/rates'
ACCOUNT = 100_000

REGIMES = [
    ('R1: Pre-QE 2003-08',    2003, 2008, 'Rising rates, no QE, carry trades'),
    ('R2: QE/Low vol 2009-19',2009, 2019, 'ZIRP, QE, ECB NIRP, BoJ YCC'),
    ('R3: COVID/HiVol 2020-22',2020, 2022,'Pandemic, emergency QE, inflation shock'),
    ('R4: Hiking 2022-26',     2022, 2026, 'Fed +525bp, fastest hikes in 40yr'),
]

def load_trades(path, pair):
    df = pd.read_csv(path)
    df.columns = [c.lower().strip() for c in df.columns]
    for col in ['dollar_pnl_real','pnl_real','pnl','net_pnl']:
        if col in df.columns:
            df['pnl'] = pd.to_numeric(df[col], errors='coerce'); break
    for col in ['entry_time','date','entry_date']:
        if col in df.columns:
            df['date'] = pd.to_datetime(df[col], errors='coerce'); break
    df = df.dropna(subset=['pnl','date']).sort_values('date').reset_index(drop=True)
    if df['pnl'].abs().median() < 1.0: df['pnl'] *= ACCOUNT
    df['year']  = df['date'].dt.year
    df['pair']  = pair
    return df

def load_rate(fname, col):
    p = RATES / fname
    if not p.exists(): return None
    df = pd.read_csv(p)
    df.columns = [c.lower() for c in df.columns]
    df['date'] = pd.to_datetime(df['date'], errors='coerce')
    df[col]    = pd.to_numeric(df[col], errors='coerce')
    return df[['date',col]].dropna().set_index('date')

def regime_metrics(df, label, y0, y1):
    n = len(df)
    if n < 5:
        return dict(label=label, n=n, wr=np.nan, avg=np.nan,
                    sharpe=np.nan, maxdd=np.nan, total=np.nan,
                    n_years=0, pos_years=0, trades_per_yr=0,
                    realised_vol=np.nan, vol_adj_sharpe=np.nan)

    wr    = (df['pnl'] > 0).mean()
    total = df['pnl'].sum()
    avg   = df['pnl'].mean()
    n_yrs = max(y1 - y0, 1)

    # Daily Sharpe
    daily = df.groupby(df['date'].dt.date)['pnl'].sum()
    idx   = pd.date_range(daily.index.min(), daily.index.max(), freq='B')
    daily = daily.reindex(idx, fill_value=0)
    ret   = daily / ACCOUNT
    sharpe= (ret.mean() / ret.std() * np.sqrt(252)) if ret.std() > 0 else 0.0

    # Max DD
    eq  = df['pnl'].cumsum() + ACCOUNT
    pk  = eq.expanding().max()
    mdd = ((eq - pk) / pk * 100).min()

    # Yearly consistency
    yearly    = df.groupby('year')['pnl'].sum()
    pos_years = int((yearly > 0).sum())

    # ── Volatility overlay ────────────────────────────────────────────────────
    # Realised vol of daily active returns (annualised)
    active_ret = ret[ret != 0]
    realised_vol = active_ret.std() * np.sqrt(252) * 100  # as %
    trades_per_yr = n / max(n_yrs, 1)

    # Vol-adjusted Sharpe: if vol is high, some of the Sharpe is just from
    # higher absolute moves, not better signal quality
    vol_adj = sharpe / (realised_vol + 1e-9) if realised_vol > 0 else np.nan

    return dict(label=label, n=n, wr=wr, avg=avg, sharpe=sharpe,
                maxdd=mdd, total=total,
                n_years=len(yearly), pos_years=pos_years,
                trades_per_yr=trades_per_yr, realised_vol=realised_vol,
                vol_adj_sharpe=vol_adj)

eu = load_trades(EU_FILE, 'EURUSD')
uj = load_trades(UJ_FILE, 'USDJPY')

print("=" * 85)
print("  REGIME SEGMENTATION v2  —  Metrics + Volatility Overlay")
print("=" * 85)
print(f"\n  EURUSD: {len(eu):,} trades  |  USDJPY: {len(uj):,} trades")

# ── Main metrics table ────────────────────────────────────────────────────────
for pair, df in [('EURUSD', eu), ('USDJPY', uj)]:
    all_m = regime_metrics(df, 'FULL 2003-26', 2003, 2026)
    print(f"\n{'─'*85}")
    print(f"  {pair}  (all-time total ${all_m['total']:+,.0f})")
    print(f"{'─'*85}")
    print(f"  {'Regime':>26}  {'n':>5}  {'WR':>6}  {'Avg $':>7}  "
          f"{'Sharpe':>7}  {'MaxDD':>7}  {'Total':>10}  "
          f"{'Pos yrs':>8}  {'% total':>9}")
    print(f"  {'─'*95}")

    for label, y0, y1, _ in REGIMES:
        sub = df[(df['year'] >= y0) & (df['year'] <= y1)]
        m   = regime_metrics(sub, label, y0, y1)
        if m['n'] < 5:
            print(f"  {label:>26}  {'< 5 trades':>70}")
            continue
        pct   = m['total'] / all_m['total'] * 100 if all_m['total'] else 0
        flag  = ' ⚠️ DOMINANT' if pct > 50 else ''
        pyr   = f"{m['pos_years']}/{m['n_years']}"
        print(f"  {label:>26}  {m['n']:>5}  {m['wr']:5.1%}  {m['avg']:>+7.0f}  "
              f"{m['sharpe']:>7.2f}  {m['maxdd']:>6.1f}%  {m['total']:>+10,.0f}  "
              f"{pyr:>8}  {pct:>7.1f}%{flag}")

    # Full row
    pyr = f"{all_m['pos_years']}/{all_m['n_years']}"
    print(f"  {'─'*95}")
    print(f"  {'FULL 2003-26':>26}  {all_m['n']:>5}  {all_m['wr']:5.1%}  "
          f"{all_m['avg']:>+7.0f}  {all_m['sharpe']:>7.2f}  {all_m['maxdd']:>6.1f}%  "
          f"{all_m['total']:>+10,.0f}  {pyr:>8}  {'100.0%':>9}")

# ── Volatility overlay ────────────────────────────────────────────────────────
print(f"\n{'─'*85}")
print("  VOLATILITY OVERLAY")
print("  Key question: is performance driven by MACRO edge or just by market vol?")
print("  If Sharpe is high only in high-vol regimes → it's vol, not macro ❌")
print("  If Sharpe is consistent across vol levels → it's the macro edge ✅")
print(f"{'─'*85}")

for pair, df in [('EURUSD', eu), ('USDJPY', uj)]:
    print(f"\n  {pair}:")
    print(f"  {'Regime':>26}  {'Trades/yr':>10}  {'Realised vol':>13}  "
          f"{'Sharpe':>8}  {'Vol-adj Sharpe':>15}  {'Interpretation':>22}")
    print(f"  {'─'*100}")
    vol_sharpes = []
    for label, y0, y1, _ in REGIMES:
        sub = df[(df['year'] >= y0) & (df['year'] <= y1)]
        m   = regime_metrics(sub, label, y0, y1)
        if m['n'] < 5: continue
        vol_sharpes.append((m['realised_vol'], m['sharpe']))
        # Interpretation: high vol + high Sharpe = vol-driven not necessarily good
        if m['realised_vol'] > 0 and not np.isnan(m['realised_vol']):
            if m['sharpe'] > 0 and m['vol_adj_sharpe'] > 0.5:
                interp = 'Edge genuine ✅'
            elif m['sharpe'] > 0 and m['vol_adj_sharpe'] <= 0.5:
                interp = 'Possibly vol-driven ⚠️'
            elif m['sharpe'] <= 0:
                interp = 'Regime failed ❌'
            else:
                interp = '—'
        else:
            interp = '—'
        print(f"  {label:>26}  {m['trades_per_yr']:>10.1f}  "
              f"{m['realised_vol']:>12.2f}%  {m['sharpe']:>8.2f}  "
              f"{m['vol_adj_sharpe']:>15.4f}  {interp:>22}")

    # Correlation between regime vol and Sharpe
    if len(vol_sharpes) >= 3:
        vols   = np.array([x[0] for x in vol_sharpes])
        sharps = np.array([x[1] for x in vol_sharpes])
        if np.std(vols) > 0:
            corr = np.corrcoef(vols, sharps)[0,1]
            print(f"\n  Correlation(regime vol, Sharpe): {corr:+.3f}")
            if corr > 0.7:
                print("  ⚠️  HIGH positive correlation → Sharpe depends on market vol")
            elif corr > 0.3:
                print("  ⚠️  Moderate correlation → some vol dependency")
            else:
                print("  ✅  Low correlation → Sharpe is not driven by vol levels")

# ── Rate spread vol by regime ─────────────────────────────────────────────────
print(f"\n{'─'*85}")
print("  RATE SPREAD VOLATILITY BY REGIME")
print("  Shows whether macro signal quality tracks market vol or is independent")
print(f"{'─'*85}")

us2y = load_rate('us2y.csv', 'us2y')
de2y = load_rate('de2y.csv', 'de2y')
jp2y = load_rate('jp2y.csv', 'jp2y')

if us2y is not None and de2y is not None:
    eu_sp = us2y.join(de2y, how='inner').dropna()
    eu_sp['spread'] = eu_sp['us2y'] - eu_sp['de2y']

    print(f"\n  EURUSD spread (US-DE 2Y):")
    print(f"  {'Regime':>28}  {'Spread mean':>12}  {'Spread std':>12}  "
          f"{'Spread range':>14}")
    print(f"  {'─'*75}")
    for label, y0, y1, _ in REGIMES:
        sub = eu_sp.loc[f'{y0}':f'{y1}']['spread']
        if len(sub) < 50: continue
        print(f"  {label:>28}  {sub.mean():>+11.2f}%  "
              f"{sub.std():>11.2f}%  "
              f"{sub.min():+.2f}% to {sub.max():+.2f}%")

if us2y is not None and jp2y is not None:
    uj_sp = us2y.join(jp2y, how='inner').dropna()
    uj_sp['spread'] = uj_sp['us2y'] - uj_sp['jp2y']

    print(f"\n  USDJPY spread (US-JP 2Y):")
    print(f"  {'Regime':>28}  {'Spread mean':>12}  {'Spread std':>12}  "
          f"{'Spread range':>14}")
    print(f"  {'─'*75}")
    for label, y0, y1, _ in REGIMES:
        sub = uj_sp.loc[f'{y0}':f'{y1}']['spread']
        if len(sub) < 50: continue
        print(f"  {label:>28}  {sub.mean():>+11.2f}%  "
              f"{sub.std():>11.2f}%  "
              f"{sub.min():+.2f}% to {sub.max():+.2f}%")

# ── Year-by-year ──────────────────────────────────────────────────────────────
print(f"\n{'─'*85}")
print("  YEAR-BY-YEAR  (no single year should carry everything)")
print(f"{'─'*85}")

for pair, df in [('EURUSD', eu), ('USDJPY', uj)]:
    total = df['pnl'].sum()
    print(f"\n  {pair}:")
    print(f"  {'Yr':>5}  {'R':>3}  {'n':>5}  {'WR':>6}  "
          f"{'Avg':>7}  {'Total':>10}  {'%total':>7}  {'Bar':>28}")
    print(f"  {'─'*78}")
    for yr in range(2003, 2027):
        sub = df[df['year']==yr]
        if len(sub) < 2: continue
        r   = ('R1' if yr<=2008 else 'R2' if yr<=2019 else
               'R3' if yr<=2022 else 'R4')
        wr  = (sub['pnl']>0).mean()
        avg = sub['pnl'].mean()
        tot = sub['pnl'].sum()
        pct = tot / total * 100
        bar = ('█'*int(min(abs(pct),28))) if tot >= 0 else ('░'*int(min(abs(pct),28)))
        sgn = '+' if tot >= 0 else '-'
        print(f"  {yr:>5}  {r:>3}  {len(sub):>5}  {wr:5.1%}  "
              f"{avg:>+7.0f}  {tot:>+10,.0f}  {pct:>6.1f}%  {sgn}{bar}")

# ── Verdict ───────────────────────────────────────────────────────────────────
eu_pos = sum(1 for _,y0,y1,_ in REGIMES
             if regime_metrics(eu[(eu['year']>=y0)&(eu['year']<=y1)],'',y0,y1)['total'] > 0
             and regime_metrics(eu[(eu['year']>=y0)&(eu['year']<=y1)],'',y0,y1)['n'] >= 5)
uj_pos = sum(1 for _,y0,y1,_ in REGIMES
             if regime_metrics(uj[(uj['year']>=y0)&(uj['year']<=y1)],'',y0,y1)['total'] > 0
             and regime_metrics(uj[(uj['year']>=y0)&(uj['year']<=y1)],'',y0,y1)['n'] >= 5)

print(f"\n{'='*85}")
print("  VERDICT")
print(f"{'='*85}")
print(f"""
  EURUSD: {eu_pos}/4 regimes profitable
  USDJPY: {uj_pos}/4 regimes profitable

  Volatility overlay interpretation:
    If vol-adjusted Sharpe is CONSISTENT across regimes:
      → Edge is macro-driven, not vol-harvesting ✅
    If vol-adjusted Sharpe spikes only in high-vol regimes:
      → Model may be a vol harvester dressed as a macro model ⚠️

  Reviewer challenge: "one regime carrying everything"
    Check the % total column — if any regime > 50%, investigate.
    Check the bar chart — should show consistent positive bars, not one spike.

  Rate spread analysis:
    Different regimes have different spread levels and volatility.
    If model works across ALL spread environments → universal ✅
    If model works only when spread is very wide (2022+) → conditional ⚠️
""")
