r"""
formula_comparison_eu_backtest.py
==================================
Side-by-side EURUSD backtest comparing two z-score formulas:

  Formula A — "15M" (what the backtest code currently does):
    - lag_gap computed on 15M bars
    - actual_return = close.pct_change(24) on 15M = 6-hour rolling return  
    - z-score = rolling(60, 60).mean/std on 15M = 15-hour normalisation window
    - Signal evaluated every 15-minute bar

  Formula B — "1H" (what the v5.1 doc describes / what the live monitor runs):
    - lag_gap computed on 1H bars
    - actual_return = close.pct_change(24) on 1H = 24-hour rolling return
    - z-score = rolling(60, 60).mean/std on 1H = 60-hour normalisation window
    - Signal evaluated every 1-hour bar

Common to both:
  - Predicted return: daily, from 120-day OLS rolling-beta model, 20-day EWM smooth
    (same daily series ffilled onto each frequency)
  - Threshold: |z| >= 2.75
  - Session: UTC 6-15 (matches v5.1 doc's 08-17 EET in winter flat-tz)
  - Entry: 0.786 Fibonacci pullback within 6h on 15M bars
  - Exit: TP 0.20% / SL 0.25% / z-exit ±1.5 / max hold 52h
  - Dynamic sizing tiers: 1.0× / 1.5× / 2.0× by |z| band
  - Real-world cost model: variable spread + slippage by exit type
  - Full period: 2003-2026

What we measure (per formula):
  Total trades, Win rate, Profit factor, Avg return, Sharpe (daily),
  Worst calendar year DD, Max DD, Exit reason breakdown, Trades/year,
  2024-2026 trade rate (recent regime).

Run on home machine (long run — ~10-20 min total):
  cd C:\Users\paul_\OneDrive\fx_macro_intraday
  python src\research\formula_comparison_eu_backtest.py
"""
from __future__ import annotations
import sys
import warnings
from pathlib import Path
from datetime import datetime

import numpy as np
import pandas as pd
import statsmodels.api as sm

warnings.filterwarnings("ignore")

BASE = Path(__file__).resolve().parents[2]
SRC  = BASE / "src"
if str(SRC) not in sys.path:
    sys.path.insert(0, str(SRC))

from ingestion.price_loader_15m import load_eurusd_15m
# from ingestion.rates_loader import load_rate_csv  # not used

# ── Validated parameters (match real_world_costs_v1.py) ────────────────────
THRESHOLD     = 2.75
FIB           = 0.786
HOLD_HOURS    = 52
STOP          = 0.0025
TP            = 0.0020
ZSCORE_EXIT   = 1.5
SESSION_LO_UTC = 6   # 08:00 EET = 06:00 UTC (winter, flat-tz)
SESSION_HI_UTC = 15  # 17:00 EET = 15:00 UTC inclusive

BETA_WINDOW = 120
SMOOTH_SPAN = 20
BETA_2Y_CLIP_LOW, BETA_2Y_CLIP_HIGH = -0.10, 0.03
BETA_10Y_CLIP_LOW, BETA_10Y_CLIP_HIGH = -0.08, 0.05

# Sizing tiers (z-band -> multiplier)
ZSCORE_BANDS = [
    (2.75, 3.50, 1.0),
    (3.50, 4.50, 1.5),
    (4.50, 99.0, 2.0),
]
BASE_NOTIONAL = 300_000   # 0.75% risk on $100k at 0.25% stop

# Cost model (matches real_world_costs_v1.py)
PIP                 = 0.0001
SPREAD_OPTIONS      = [0.6, 0.8, 1.0, 1.2, 1.5, 2.0, 3.0]
SPREAD_WEIGHTS      = [0.20, 0.30, 0.25, 0.12, 0.07, 0.04, 0.02]
ENTRY_SLIPPAGE_PIPS = 0.5
EXIT_TP_SLIPPAGE    = 0.0
EXIT_STOP_SLIPPAGE  = 0.8
EXIT_Z_SLIPPAGE     = 0.5

RNG_SEED = 42

RATES_DIR = BASE / "data" / "raw" / "rates"


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


def sample_spread(rng: np.random.Generator) -> float:
    return rng.choice(SPREAD_OPTIONS, p=SPREAD_WEIGHTS) * PIP


def apply_real_costs(raw_return, exit_reason, entry_price, rng):
    spread     = sample_spread(rng)
    entry_slip = ENTRY_SLIPPAGE_PIPS * PIP
    if exit_reason == "tp":
        exit_slip = EXIT_TP_SLIPPAGE * PIP
    elif exit_reason == "stop":
        exit_slip = EXIT_STOP_SLIPPAGE * PIP
    else:
        exit_slip = EXIT_Z_SLIPPAGE * PIP
    total_cost = spread + entry_slip + exit_slip
    cost_as_return = total_cost / entry_price
    return raw_return - cost_as_return, total_cost / PIP


def load_rate(fname: str, col: str) -> pd.DataFrame:
    p = RATES_DIR / fname
    df = pd.read_csv(p)
    df.columns = [c.lower() for c in df.columns]
    df["date"] = pd.to_datetime(df.iloc[:, 0], errors="coerce")
    df[col]    = pd.to_numeric(df.iloc[:, 1], errors="coerce")
    return df[["date", col]].dropna().sort_values("date").reset_index(drop=True)


# ──────────────────────────────────────────────────────────────────────────
# Build daily predicted_return_1d (shared by both formulas)
# ──────────────────────────────────────────────────────────────────────────
def build_daily_predicted_return(prices_15m: pd.DataFrame) -> pd.DataFrame:
    """Build daily predicted_return_1d series using rolling-beta model.
    
    Matches rolling_beta_model.py exactly (120-day OLS, 20-day EWM, +/-clip, shift 1).
    Returns DataFrame with columns: date, predicted_return_1d, spread_2y_change_1d, spread_10y_change_1d
    """
    print("  Building daily predicted_return_1d series (shared by both formulas)...")
    
    daily_px = (prices_15m.assign(date=prices_15m["datetime"].dt.normalize())
                .groupby("date", as_index=False)
                .agg(close=("close", "last"))
                .sort_values("date").reset_index(drop=True))
    daily_px["eurusd_return_1d"] = daily_px["close"].pct_change()
    daily_px["date"] = pd.to_datetime(daily_px["date"])
    
    us2y  = load_rate("us2y.csv",  "us2y")
    us10y = load_rate("us10y.csv", "us10y")
    de2y  = load_rate("de2y.csv",  "de2y")
    de10y = load_rate("de10y.csv", "de10y")
    
    date_range = pd.DataFrame({"date": pd.date_range(
        daily_px["date"].min(), daily_px["date"].max(), freq="D")})
    
    def ffill(rate_df, col):
        left  = date_range.copy().astype({"date": "datetime64[us]"})
        right = rate_df.copy().astype({"date": "datetime64[us]"})
        return pd.merge_asof(left, right, on="date", direction="backward")
    
    spreads = (ffill(us2y, "us2y")
               .merge(ffill(us10y, "us10y"), on="date")
               .merge(ffill(de2y, "de2y"),   on="date")
               .merge(ffill(de10y, "de10y"), on="date"))
    spreads["spread_2y"]            = spreads["us2y"]  - spreads["de2y"]
    spreads["spread_10y"]           = spreads["us10y"] - spreads["de10y"]
    spreads["spread_2y_change_1d"]  = spreads["spread_2y"].diff(1)
    spreads["spread_10y_change_1d"] = spreads["spread_10y"].diff(1)
    spreads = spreads[["date", "spread_2y_change_1d", "spread_10y_change_1d"]].dropna()
    
    model_df = daily_px.merge(spreads, on="date", how="left").dropna(
        subset=["eurusd_return_1d", "spread_2y_change_1d"]).reset_index(drop=True)
    
    print(f"    OLS across {len(model_df)} daily bars...")
    n = len(model_df)
    b2y_raw  = np.full(n, np.nan)
    b10y_raw = np.full(n, np.nan)
    for i in range(BETA_WINDOW, n):
        sample = model_df.iloc[i - BETA_WINDOW : i].copy()
        sample = sample[(sample["spread_2y_change_1d"].abs()  > 0) |
                        (sample["spread_10y_change_1d"].abs() > 0)]
        if len(sample) < 30:
            continue
        X = sm.add_constant(sample[["spread_2y_change_1d", "spread_10y_change_1d"]])
        try:
            res = sm.OLS(sample["eurusd_return_1d"], X).fit()
            b2y_raw[i]  = res.params.get("spread_2y_change_1d",  np.nan)
            b10y_raw[i] = res.params.get("spread_10y_change_1d", np.nan)
        except Exception:
            pass
    
    model_df["beta_2y"]  = (pd.Series(b2y_raw).ewm(span=SMOOTH_SPAN, adjust=False).mean()
                            .shift(1).clip(BETA_2Y_CLIP_LOW, BETA_2Y_CLIP_HIGH).values)
    model_df["beta_10y"] = (pd.Series(b10y_raw).ewm(span=SMOOTH_SPAN, adjust=False).mean()
                            .shift(1).clip(BETA_10Y_CLIP_LOW, BETA_10Y_CLIP_HIGH).values)
    model_df["predicted_return_1d"] = (
        model_df["beta_2y"]  * model_df["spread_2y_change_1d"] +
        model_df["beta_10y"] * model_df["spread_10y_change_1d"]
    )
    
    out = model_df[["date", "predicted_return_1d", "spread_2y_change_1d", "spread_10y_change_1d"]]
    print(f"    Daily series built: {len(out)} rows")
    return out


# ──────────────────────────────────────────────────────────────────────────
# Compute z-score for each formula
# ──────────────────────────────────────────────────────────────────────────
def compute_zscore_15m(prices_15m: pd.DataFrame, daily_pred: pd.DataFrame) -> pd.DataFrame:
    """Formula A: lag_gap on 15M bars, pct_change(24) = 6h, rolling(60) = 15h window."""
    df = prices_15m.copy()
    df["date"] = df["datetime"].dt.normalize()
    df = df.merge(daily_pred, on="date", how="left")
    df["predicted_return_1d"] = df["predicted_return_1d"].ffill()
    df["eurusd_return_24h"] = df["close"].pct_change(24)
    df["lag_gap"] = df["predicted_return_1d"] - df["eurusd_return_24h"]
    m = df["lag_gap"].rolling(60, min_periods=60).mean()
    s = df["lag_gap"].rolling(60, min_periods=60).std()
    df["zscore"] = (df["lag_gap"] - m) / s
    df["formula"] = "A_15M"
    return df.dropna(subset=["zscore"]).reset_index(drop=True)


def compute_zscore_1h(prices_15m: pd.DataFrame, daily_pred: pd.DataFrame) -> pd.DataFrame:
    """Formula B: resample 15M to 1H, lag_gap on 1H, pct_change(24) = 24h, rolling(60) = 60h window."""
    df = (prices_15m.set_index("datetime")
          .resample("1h")
          .agg(open=("open", "first"), high=("high", "max"),
               low=("low", "min"), close=("close", "last"),
               volume=("volume", "sum"))
          .dropna()
          .reset_index())
    df["date"] = df["datetime"].dt.normalize()
    df = df.merge(daily_pred, on="date", how="left")
    df["predicted_return_1d"] = df["predicted_return_1d"].ffill()
    df["eurusd_return_24h"] = df["close"].pct_change(24)
    df["lag_gap"] = df["predicted_return_1d"] - df["eurusd_return_24h"]
    m = df["lag_gap"].rolling(60, min_periods=60).mean()
    s = df["lag_gap"].rolling(60, min_periods=60).std()
    df["zscore"] = (df["lag_gap"] - m) / s
    df["formula"] = "B_1H"
    return df.dropna(subset=["zscore"]).reset_index(drop=True)


# ──────────────────────────────────────────────────────────────────────────
# Run backtest for a given signal series
# ──────────────────────────────────────────────────────────────────────────
def run_backtest(signals: pd.DataFrame, m15: pd.DataFrame, label: str) -> pd.DataFrame:
    """Iterate signals (each bar where |z|>=threshold in session), attempt entries.
    
    signals must have columns: datetime, close, high, low, zscore
    """
    print(f"\n  Running backtest for {label}...")
    
    # Filter to in-session signal bars only
    signals = signals.copy()
    signals["hour_utc"] = signals["datetime"].dt.hour
    signals = signals[(signals["hour_utc"] >= SESSION_LO_UTC) &
                      (signals["hour_utc"] <= SESSION_HI_UTC)]
    signals = signals[signals["zscore"].abs() >= THRESHOLD]
    signals["signal"] = np.where(signals["zscore"] >= THRESHOLD, 1, -1)
    signals = signals.sort_values("datetime").reset_index(drop=True)
    print(f"    Signal bars (|z|>={THRESHOLD} in session): {len(signals):,}")
    
    m15_sorted = m15.sort_values("datetime").reset_index(drop=True)
    m15_times = m15_sorted["datetime"].values
    
    rng = np.random.default_rng(RNG_SEED)
    trades = []
    last_exit_time = None
    
    for _, sig in signals.iterrows():
        if last_exit_time is not None and sig["datetime"] < last_exit_time:
            continue
        
        # Fib target from SIGNAL bar's high/low/close
        sig_close = float(sig["close"])
        sig_high  = float(sig["high"])
        sig_low   = float(sig["low"])
        signal    = int(sig["signal"])
        z_abs     = abs(float(sig["zscore"]))
        
        if signal == 1:
            pull = sig_close - sig_low
            if pull <= 0:
                continue
            target = sig_close - FIB * pull
        else:
            pull = sig_high - sig_close
            if pull <= 0:
                continue
            target = sig_close + FIB * pull
        
        # Search next 6h of 15M for entry fill
        start_idx = int(np.searchsorted(m15_times, sig["datetime"], side="left"))
        if start_idx >= len(m15_sorted):
            continue
        wait_bars = 6 * 4
        window = m15_sorted.iloc[start_idx:start_idx + wait_bars]
        if window.empty:
            continue
        
        if signal == 1:
            hit = window[window["low"] <= target]
        else:
            hit = window[window["high"] >= target]
        if hit.empty:
            continue
        
        entry_time  = hit.iloc[0]["datetime"]
        entry_price = target
        
        # Simulate hold up to HOLD_HOURS
        entry_idx = int(np.searchsorted(m15_times, entry_time, side="left"))
        hold_bars = HOLD_HOURS * 4
        exit_idx_max = min(entry_idx + hold_bars, len(m15_sorted) - 1)
        path = m15_sorted.iloc[entry_idx:exit_idx_max + 1]
        
        # We also need z-score series at hourly granularity for z-exit;
        # for simplicity here we use the signal-series z at the closest signal time
        # (approximation, both formulas use this same approximation)
        
        exit_price  = float(path.iloc[-1]["close"])
        exit_time   = path.iloc[-1]["datetime"]
        exit_reason = "time"
        
        for _, bar in path.iterrows():
            if signal == 1:
                tp_hit   = (float(bar["high"]) - entry_price) / entry_price >= TP
                stop_hit = (float(bar["low"])  - entry_price) / entry_price <= -STOP
            else:
                tp_hit   = (entry_price - float(bar["low"])) / entry_price >= TP
                stop_hit = -((float(bar["high"]) - entry_price) / entry_price) <= -STOP
            if tp_hit:
                exit_price = entry_price*(1+TP) if signal==1 else entry_price*(1-TP)
                exit_time  = bar["datetime"]
                exit_reason = "tp"
                break
            if stop_hit:
                exit_price = entry_price*(1-STOP) if signal==1 else entry_price*(1+STOP)
                exit_time  = bar["datetime"]
                exit_reason = "stop"
                break
        
        last_exit_time = exit_time
        
        raw_ret = ((exit_price - entry_price) / entry_price) * signal
        if exit_reason == "stop":
            raw_ret = -STOP
        
        real_ret, cost_pips = apply_real_costs(raw_ret, exit_reason, entry_price, rng)
        
        trades.append({
            "entry_time":  entry_time,
            "exit_time":   exit_time,
            "signal":      signal,
            "entry_price": entry_price,
            "exit_price":  exit_price,
            "exit_reason": exit_reason,
            "zscore_abs":  z_abs,
            "raw_return":  raw_ret,
            "real_return": real_ret,
            "cost_pips":   cost_pips,
            "mult":        get_multiplier(z_abs),
        })
    
    df = pd.DataFrame(trades)
    print(f"    Trades taken: {len(df):,}")
    return df


# ──────────────────────────────────────────────────────────────────────────
# Metrics
# ──────────────────────────────────────────────────────────────────────────
def compute_metrics(trades: pd.DataFrame, label: str) -> dict:
    if trades.empty:
        return {"label": label, "trades": 0}
    
    df = trades.copy()
    df["entry_time"] = pd.to_datetime(df["entry_time"])
    df["dollar_pnl_real"] = df["real_return"] * BASE_NOTIONAL * df["mult"]
    df["dollar_pnl_flat"] = df["real_return"] * BASE_NOTIONAL
    
    period_days = (df["entry_time"].max() - df["entry_time"].min()).days
    years = period_days / 365.25
    
    # Daily P&L for Sharpe
    df["date"] = df["entry_time"].dt.normalize()
    daily_pnl = df.groupby("date")["dollar_pnl_real"].sum()
    # Reindex to fill non-trading days with 0 P&L for accurate daily Sharpe
    full_dates = pd.date_range(df["date"].min(), df["date"].max(), freq="D")
    daily_pnl = daily_pnl.reindex(full_dates, fill_value=0)
    sharpe_daily = daily_pnl.mean() / daily_pnl.std() * np.sqrt(252) if daily_pnl.std() > 0 else np.nan
    
    # Annual returns -> worst calendar year DD on a fresh $100k account
    df["year"] = df["entry_time"].dt.year
    yearly_pnl = df.groupby("year")["dollar_pnl_real"].sum()
    worst_year_dd = -yearly_pnl.min() / 100_000 * 100  # as % of $100k
    
    # Running-peak DD on cumulative equity
    df_sorted = df.sort_values("entry_time")
    equity = 100_000 + df_sorted["dollar_pnl_real"].cumsum()
    peak = equity.cummax()
    dd = (peak - equity) / peak * 100
    max_dd_running = dd.max()
    
    wins = (df["real_return"] > 0).sum()
    losses = (df["real_return"] <= 0).sum()
    wr = wins / len(df) * 100 if len(df) > 0 else np.nan
    gross_wins = df.loc[df["real_return"] > 0, "real_return"].sum()
    gross_losses = -df.loc[df["real_return"] <= 0, "real_return"].sum()
    pf = gross_wins / gross_losses if gross_losses > 0 else np.inf
    
    # Recent regime rate
    recent = df[df["entry_time"] >= "2024-01-01"]
    recent_months = max((df["entry_time"].max() - pd.Timestamp("2024-01-01")).days / 30.4, 1)
    recent_rate = len(recent) / recent_months
    
    # Exit breakdown
    exit_breakdown = df["exit_reason"].value_counts(normalize=True) * 100
    
    return {
        "label":             label,
        "trades":            len(df),
        "period_years":      round(years, 1),
        "trades_per_year":   round(len(df) / years, 1),
        "win_rate":          round(wr, 1),
        "profit_factor":     round(pf, 2),
        "avg_return_pct":    round(df["real_return"].mean() * 100, 4),
        "total_pnl":         round(df["dollar_pnl_real"].sum(), 0),
        "sharpe_daily":      round(sharpe_daily, 2),
        "worst_year_dd_pct": round(worst_year_dd, 2),
        "max_dd_running_pct": round(max_dd_running, 2),
        "recent_rate_per_mo": round(recent_rate, 1),
        "exit_tp_pct":       round(exit_breakdown.get("tp", 0), 1),
        "exit_stop_pct":     round(exit_breakdown.get("stop", 0), 1),
        "exit_time_pct":     round(exit_breakdown.get("time", 0), 1),
    }


# ──────────────────────────────────────────────────────────────────────────
# Main
# ──────────────────────────────────────────────────────────────────────────
def main():
    print("=" * 78)
    print("  EURUSD FORMULA COMPARISON BACKTEST")
    print("=" * 78)
    print(f"  Run time: {datetime.now()}")
    print(f"  Threshold |z| >= {THRESHOLD}")
    print(f"  Session: UTC {SESSION_LO_UTC}-{SESSION_HI_UTC} (matches v5.1 doc 08-17 EET)")
    print(f"  Entry: 0.786 Fib pullback within 6h on 15M")
    print(f"  Exit: TP {TP*100:.2f}% / SL {STOP*100:.2f}% / max hold {HOLD_HOURS}h")
    print()
    
    print("Loading EURUSD 15M data...")
    m15 = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)
    print(f"  {len(m15):,} rows | {m15['datetime'].min()} → {m15['datetime'].max()}")
    print()
    
    # Build shared daily predicted_return series
    daily_pred = build_daily_predicted_return(m15)
    
    # Formula A — 15M
    print("\n" + "─" * 78)
    print("  FORMULA A — 15M (signal every 15m, pct_change(24) = 6h, rolling(60) = 15h window)")
    print("─" * 78)
    sig_A = compute_zscore_15m(m15, daily_pred)
    print(f"    Z-score series: {len(sig_A):,} rows")
    trades_A = run_backtest(sig_A, m15, "Formula A")
    
    # Formula B — 1H
    print("\n" + "─" * 78)
    print("  FORMULA B — 1H (signal every 1h, pct_change(24) = 24h, rolling(60) = 60h window)")
    print("─" * 78)
    sig_B = compute_zscore_1h(m15, daily_pred)
    print(f"    Z-score series: {len(sig_B):,} rows")
    trades_B = run_backtest(sig_B, m15, "Formula B")
    
    # Metrics
    metrics_A = compute_metrics(trades_A, "Formula A (15M)")
    metrics_B = compute_metrics(trades_B, "Formula B (1H)")
    
    # Display comparison
    print()
    print("=" * 78)
    print("  RESULTS")
    print("=" * 78)
    
    metric_rows = [
        ("Total trades",            "trades",            "{:>10,}"),
        ("Period (years)",          "period_years",      "{:>10.1f}"),
        ("Trades per year",         "trades_per_year",   "{:>10.1f}"),
        ("Recent rate per month",   "recent_rate_per_mo","{:>10.1f}"),
        ("Win rate (%)",            "win_rate",          "{:>10.1f}"),
        ("Profit factor",           "profit_factor",     "{:>10.2f}"),
        ("Avg return per trade (%)","avg_return_pct",    "{:>10.4f}"),
        ("Total $ P&L (real)",      "total_pnl",         "${:>+9,.0f}"),
        ("Sharpe (daily)",          "sharpe_daily",      "{:>10.2f}"),
        ("Worst year DD (%)",       "worst_year_dd_pct", "{:>10.2f}"),
        ("Max running DD (%)",      "max_dd_running_pct","{:>10.2f}"),
        ("Exit TP (%)",             "exit_tp_pct",       "{:>10.1f}"),
        ("Exit Stop (%)",           "exit_stop_pct",     "{:>10.1f}"),
        ("Exit Time (%)",           "exit_time_pct",     "{:>10.1f}"),
    ]
    
    print(f"\n  {'Metric':<28} {'Formula A (15M)':>16} {'Formula B (1H)':>16}")
    print("  " + "─" * 60)
    for label, key, fmt in metric_rows:
        try:
            v_A = fmt.format(metrics_A.get(key, "—"))
        except (TypeError, ValueError):
            v_A = str(metrics_A.get(key, "—"))
        try:
            v_B = fmt.format(metrics_B.get(key, "—"))
        except (TypeError, ValueError):
            v_B = str(metrics_B.get(key, "—"))
        print(f"  {label:<28} {v_A:>16} {v_B:>16}")
    
    # Save trade logs
    out_dir = BASE / "data" / "research"
    out_dir.mkdir(parents=True, exist_ok=True)
    if not trades_A.empty:
        trades_A.to_csv(out_dir / "formula_A_15M_trades.csv", index=False)
    if not trades_B.empty:
        trades_B.to_csv(out_dir / "formula_B_1H_trades.csv", index=False)
    
    print()
    print("  Trade logs saved to:")
    print(f"    {out_dir / 'formula_A_15M_trades.csv'}")
    print(f"    {out_dir / 'formula_B_1H_trades.csv'}")
    print()
    print("  This is Phase 1 — formula comparison only.")
    print("  If one formula is clearly stronger, Phase 2 will run full validation")
    print("  (FTMO Monte Carlo, permutation tests, walk-forward OOS) on the winner.")


if __name__ == "__main__":
    main()
