"""
threshold_regime_filter_v1.py
==============================
Solves the challenge pacing problem (15 month median) by testing two levers:

  Lever 1 — Threshold reduction
    Current: 2.75 → ~64 trades/year
    Test:    2.00, 2.25, 2.50, 2.75
    Lower threshold = more signals = more trades = faster to target
    Risk: lower quality trades enter. Regime filter compensates.

  Lever 2 — Regime filter (built entirely from existing pipeline data)
    Uses data already computed in your pipeline — no new collection needed.

    Filter A: Beta alignment
      beta_2y < BETA_THRESHOLD (e.g. -0.015)
      Meaning: EUR is actively responding to spread changes right now.
      When beta is near zero, the macro transmission is broken — skip.

    Filter B: 5-day spread momentum confirmation
      spread_2y_change_5d same sign as signal direction
      Meaning: the spread move has 5-day persistence, not one-day noise.

    Filter C: 10Y tenor agreement
      spread_10y_change_1d same sign as spread_2y_change_1d
      Meaning: both the short and long end of the curve agree.
      When only 2Y moves but 10Y is silent, signal is weaker.

    Filters tested individually and in combinations.

Validation
----------
  For each threshold × regime_filter combination:
    - Trade count and annualised frequency
    - Win rate, avg return, final equity, max DD
    - FTMO Monte Carlo (2000 sims) with signal scaling
    - Estimated challenge duration

  Accept if: 95th pct DD < 10%
  Rank by:   avg trades to pass (lower = faster challenge)

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

Run from project root:
  python src/research/threshold_regime_filter_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 (
    get_first_m15_idx_at_or_after,
    find_entry_pullback,
    simulate_trade,
)
from features.spot_lag_v3 import get_model_ready_spot_lag_v3
from models.rolling_beta_model import build_rolling_beta_model
from features.yield_spreads import build_spread_features
from ingestion.price_loader_15m import load_eurusd_15m

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

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

# ── FTMO ──────────────────────────────────────────────────────────────────────
ACCOUNT_START    = 100_000.0
PROFIT_TARGET    = ACCOUNT_START * 0.10
MAX_OVERALL_LOSS = ACCOUNT_START * 0.10
MAX_DAILY_LOSS   = ACCOUNT_START * 0.05

BASE_RISK_PCT    = 0.003
BASE_NOTIONAL    = (ACCOUNT_START * BASE_RISK_PCT) / STOP   # $120,000

# Signal scaling bands (validated)
ZSCORE_BANDS = [
    (2.00, 2.75, 0.75),   # sub-threshold trades get 0.75x if threshold lowered
    (2.75, 3.50, 1.00),
    (3.50, 4.50, 1.50),
    (4.50, 99.0, 2.00),
]

NOTIONAL_HARD_CAP = BASE_NOTIONAL * 3.0

N_SIMS      = 2000
RANDOM_SEED = 42

# ── Thresholds to test ────────────────────────────────────────────────────────
THRESHOLDS = [2.00, 2.25, 2.50, 2.75]

# ── Regime filter parameter ───────────────────────────────────────────────────
BETA_THRESHOLD = -0.015   # beta_2y must be below this (active macro transmission)

YEARS_IN_DATA  = 22.0     # 2003-2026


# ── Build enriched signal dataframe ──────────────────────────────────────────
def build_enriched_signals() -> pd.DataFrame:
    """
    Builds the full signal dataframe with regime filter columns attached.
    All data comes from existing pipeline — no new sources needed.
    """
    # Core lag model (has lag_zscore_24h_v3, datetime, close, open, high, low)
    spot = get_model_ready_spot_lag_v3().copy()
    spot["datetime"] = pd.to_datetime(spot["datetime"])
    spot["date"]     = spot["datetime"].dt.normalize()

    # Rolling beta model (has beta_2y, beta_10y, spread_2y_change_1d,
    # spread_10y_change_1d)
    beta = build_rolling_beta_model(window=120, smooth_span=20).copy()
    beta["date"] = pd.to_datetime(beta["date"])

    # Spread features (has spread_2y_change_5d, spread_2y_zscore_20d,
    # spread_10y_change_5d)
    spreads = build_spread_features().copy()
    spreads["date"] = pd.to_datetime(spreads["date"])

    # Merge beta onto hourly bars (daily → hourly via date)
    beta_cols = ["date", "beta_2y", "beta_10y",
                 "spread_2y_change_1d", "spread_10y_change_1d"]
    spot = spot.merge(beta[beta_cols], on="date", how="left", suffixes=("", "_beta"))

    # Merge spread features
    spread_cols = ["date", "spread_2y_change_5d", "spread_2y_zscore_20d",
                   "spread_10y_change_5d", "spread_10y_zscore_20d"]
    spot = spot.merge(spreads[spread_cols], on="date", how="left")

    # Forward fill daily fields across intraday bars
    daily_fill_cols = [
        "beta_2y", "beta_10y",
        "spread_2y_change_1d", "spread_10y_change_1d",
        "spread_2y_change_5d", "spread_2y_zscore_20d",
        "spread_10y_change_5d", "spread_10y_zscore_20d",
    ]
    spot[daily_fill_cols] = spot[daily_fill_cols].ffill()

    # Session filter
    spot["hour"] = spot["datetime"].dt.hour
    spot = spot[spot["hour"].isin(ALLOWED_HOURS)].copy()

    return spot.sort_values("datetime").reset_index(drop=True)


# ── Apply regime filters ──────────────────────────────────────────────────────
def apply_regime_filters(df: pd.DataFrame, filters: list) -> pd.DataFrame:
    """
    Applies a list of named regime filters to the dataframe.
    Returns rows that pass ALL specified filters.
    """
    mask = pd.Series(True, index=df.index)

    for f in filters:
        if f == "beta_active":
            # beta_2y must be sufficiently negative (active macro transmission)
            mask &= df["beta_2y"] < BETA_THRESHOLD

        elif f == "momentum_confirm":
            # 5-day spread momentum must confirm signal direction
            signal_dir = np.sign(df["lag_zscore_24h_v3"])
            spread_dir = np.sign(df["spread_2y_change_5d"])
            mask &= (signal_dir == spread_dir)

        elif f == "tenor_agree":
            # Both 2Y and 10Y daily changes must agree in direction
            mask &= (
                np.sign(df["spread_2y_change_1d"]) ==
                np.sign(df["spread_10y_change_1d"])
            )

        elif f == "all_three":
            signal_dir = np.sign(df["lag_zscore_24h_v3"])
            spread_dir = np.sign(df["spread_2y_change_5d"])
            mask &= (df["beta_2y"] < BETA_THRESHOLD)
            mask &= (signal_dir == spread_dir)
            mask &= (
                np.sign(df["spread_2y_change_1d"]) ==
                np.sign(df["spread_10y_change_1d"])
            )

    return df[mask].copy()


# ── Run one combination ───────────────────────────────────────────────────────
def run_combination(
    enriched_df : pd.DataFrame,
    m15         : pd.DataFrame,
    threshold   : float,
    filter_name : str,
    filters     : list,
) -> dict | None:
    """
    Runs the full trade simulation for one threshold × regime_filter combination.
    Returns summary dict.
    """
    # Apply threshold
    df = enriched_df.copy()
    df["signal"] = 0
    df.loc[df["lag_zscore_24h_v3"] >=  threshold, "signal"] =  1
    df.loc[df["lag_zscore_24h_v3"] <= -threshold, "signal"] = -1
    df = df[df["signal"] != 0].copy()

    # Apply regime filters
    df = apply_regime_filters(df, filters)
    df = df.reset_index(drop=True)

    if len(df) < 50:
        return None   # too few trades to be meaningful

    # Trade simulation
    trades         = []
    last_exit_time = None

    for _, sig in df.iterrows():
        signal_time = sig["datetime"]
        if last_exit_time is not None and signal_time < 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,
            entry_time=entry["entry_time"],
            entry_price=entry["entry_price"],
            signal=int(sig["signal"]),
            hold_hours=HOLD_HOURS,
            stop=STOP,
            spread_cost=SPREAD_COST,
        )
        if trade is None:
            continue

        # Attach z-score for signal scaling
        trade["zscore_abs"] = abs(float(sig["lag_zscore_24h_v3"]))
        trades.append(trade)
        last_exit_time = trade["exit_time"]

    if len(trades) < 30:
        return None

    t = pd.DataFrame(trades)
    t["equity"]   = (1 + t["return"]).cumprod()
    t["peak"]     = t["equity"].cummax()
    t["drawdown"] = t["equity"] / t["peak"] - 1

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

    win_rate   = (t["return"] > 0).mean()
    avg_ret    = t["return"].mean()
    final_eq   = t["equity"].iloc[-1]
    max_dd     = t["drawdown"].min()
    n_trades   = len(t)
    per_year   = n_trades / YEARS_IN_DATA

    return {
        "threshold"  : threshold,
        "filter"     : filter_name,
        "n_trades"   : n_trades,
        "per_year"   : per_year,
        "win_rate"   : win_rate,
        "avg_return" : avg_ret,
        "final_eq"   : final_eq,
        "max_dd"     : max_dd,
        "max_streak" : max_streak,
        "trade_df"   : t,
    }


# ── Signal scaling multiplier ─────────────────────────────────────────────────
def get_multiplier(zscore_abs: float) -> float:
    for min_z, max_z, mult in ZSCORE_BANDS:
        if min_z <= zscore_abs < max_z:
            return mult
    return ZSCORE_BANDS[-1][2]


# ── FTMO Monte Carlo ──────────────────────────────────────────────────────────
def run_ftmo_mc(trade_df: pd.DataFrame, n_sims: int = N_SIMS) -> dict:
    returns     = trade_df["return"].values
    zscores     = trade_df["zscore_abs"].values
    exit_dates  = trade_df["exit_time"].values
    multipliers = np.array([get_multiplier(z) for z in zscores])

    rng     = np.random.default_rng(RANDOM_SEED)
    results = []

    for _ in range(n_sims):
        idx      = rng.permutation(len(returns))
        sh_ret   = returns[idx]
        sh_mult  = multipliers[idx]

        balance          = ACCOUNT_START
        peak             = ACCOUNT_START
        max_dd           = 0.0
        daily_pnl        = {}
        outcome          = "INCOMPLETE"
        trading_days     = set()
        trades_taken     = 0

        for ret, mult, ex_dt in zip(sh_ret, sh_mult, exit_dates):
            date_key = str(pd.Timestamp(ex_dt).date())
            if date_key not in daily_pnl:
                daily_pnl[date_key] = 0.0

            notional    = min(BASE_NOTIONAL * mult, NOTIONAL_HARD_CAP)
            dollar_pnl  = ret * notional

            balance     += dollar_pnl
            trades_taken += 1
            trading_days.add(date_key)
            daily_pnl[date_key] += dollar_pnl

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

            if daily_pnl[date_key] < -MAX_DAILY_LOSS:
                outcome = "BREACH_DAILY_DD"
                break
            if balance <= (ACCOUNT_START - MAX_OVERALL_LOSS):
                outcome = "BREACH_OVERALL_DD"
                break
            if balance >= (ACCOUNT_START + PROFIT_TARGET):
                outcome = "PASS"
                break

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

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

    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_50"          : mc["max_dd_pct"].quantile(0.50),
        "dd_75"          : mc["max_dd_pct"].quantile(0.75),
        "dd_95"          : mc["max_dd_pct"].quantile(0.95),
        "danger_pct"     : (mc["max_dd_pct"] > 8.0).mean() * 100,
        "avg_trades_pass": passing["trades"].mean() if len(passing) > 0 else np.nan,
        "med_trades_pass": passing["trades"].median() if len(passing) > 0 else np.nan,
    }


# ── Main ──────────────────────────────────────────────────────────────────────
def main():
    print("=" * 75)
    print("THRESHOLD + REGIME FILTER TEST V1")
    print(f"  Hold: {HOLD_HOURS}h  |  Fib: {FIB}  |  Stop: {STOP:.2%}")
    print(f"  Base risk: {BASE_RISK_PCT:.2%}  |  Signal scaling: ON")
    print(f"  MC sims: {N_SIMS:,}  |  DD limit: 10% (95th pct)")
    print("=" * 75)

    # Build enriched signal dataframe once
    print("\nBuilding enriched signal dataframe...")
    enriched = build_enriched_signals()
    print(f"  Hourly bars with regime data: {len(enriched):,}")
    print(f"  Beta_2y coverage: "
          f"{enriched['beta_2y'].notna().mean():.1%}")
    print(f"  Beta_2y active (< {BETA_THRESHOLD}): "
          f"{(enriched['beta_2y'] < BETA_THRESHOLD).mean():.1%} of bars")

    # Load 15M data once
    print("\nLoading 15M data...")
    m15 = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)

    # ── Define filter combinations ────────────────────────────────────────────
    filter_combos = {
        "none"            : [],
        "beta_active"     : ["beta_active"],
        "momentum"        : ["momentum_confirm"],
        "tenor_agree"     : ["tenor_agree"],
        "beta+momentum"   : ["beta_active", "momentum_confirm"],
        "beta+tenor"      : ["beta_active", "tenor_agree"],
        "momentum+tenor"  : ["momentum_confirm", "tenor_agree"],
        "all_three"       : ["all_three"],
    }

    # ── Run all combinations ──────────────────────────────────────────────────
    print("\nRunning threshold × regime filter combinations...")
    print(f"  {'Threshold':<12}{'Filter':<20}{'Trades':>7}{'Per/yr':>7}"
          f"{'WR':>7}{'AvgRet':>9}{'FinalEq':>10}{'MaxDD':>8}")
    print(f"  {'─'*82}")

    all_combos   = []
    viable_combos = []

    for threshold in THRESHOLDS:
        for filter_name, filters in filter_combos.items():
            result = run_combination(
                enriched_df=enriched,
                m15=m15,
                threshold=threshold,
                filter_name=filter_name,
                filters=filters,
            )

            if result is None:
                print(f"  {threshold:<12}{filter_name:<20}  — too few trades")
                continue

            print(f"  {threshold:<12}{filter_name:<20}"
                  f"{result['n_trades']:>7}"
                  f"{result['per_year']:>7.1f}"
                  f"{result['win_rate']:>6.2%}"
                  f"{result['avg_return']:>9.6f}"
                  f"{result['final_eq']:>10.4f}"
                  f"{result['max_dd']:>7.2%}")

            all_combos.append(result)

            # Pre-screen: only run MC on combos with positive avg return
            # and at least 40 trades/year to be viable for pacing
            if result["avg_return"] > 0 and result["per_year"] >= 40:
                viable_combos.append(result)

    # ── Monte Carlo on viable combinations ────────────────────────────────────
    print(f"\n{'='*75}")
    print(f"MONTE CARLO — VIABLE COMBINATIONS ({len(viable_combos)} combos)")
    print(f"(positive avg return + ≥40 trades/year)")
    print(f"{'='*75}")

    print(f"\n  {'Threshold':<12}{'Filter':<20}{'Trades/yr':>10}"
          f"{'Pass%':>7}{'95pctDD':>8}{'AvgTrades':>11}"
          f"{'MedTrades':>11}{'AvgMths':>9}{'MedMths':>9}{'Status':>8}")
    print(f"  {'─'*100}")

    mc_results = []

    for combo in viable_combos:
        mc   = run_ftmo_mc(combo["trade_df"])
        valid = mc["dd_95"] < 10.0

        per_yr   = combo["per_year"]
        avg_mths = (mc["avg_trades_pass"] / per_yr * 12) if per_yr > 0 else np.nan
        med_mths = (mc["med_trades_pass"] / per_yr * 12) if per_yr > 0 else np.nan
        status   = "✓" if valid else "✗"

        print(f"  {combo['threshold']:<12}{combo['filter']:<20}"
              f"{per_yr:>10.1f}"
              f"{mc['pass_rate']:>7.2f}"
              f"{mc['dd_95']:>8.3f}"
              f"{mc['avg_trades_pass']:>11.0f}"
              f"{mc['med_trades_pass']:>11.0f}"
              f"{avg_mths:>9.1f}"
              f"{med_mths:>9.1f}"
              f"{status:>8}")

        mc_results.append({
            **{k: v for k, v in combo.items() if k != "trade_df"},
            **mc,
            "avg_months": avg_mths,
            "med_months": med_mths,
            "valid"     : valid,
        })

    # ── Find best valid combination ────────────────────────────────────────────
    print(f"\n{'='*75}")
    print("BEST VALID COMBINATIONS (95th pct DD < 10%, ranked by avg months)")
    print(f"{'='*75}")

    valid_results = [r for r in mc_results if r["valid"]]

    if not valid_results:
        print("\n  No combinations passed the DD limit.")
        print("  Showing best available (lowest 95th pct DD):")
        valid_results = sorted(mc_results, key=lambda r: r["dd_95"])[:5]

    valid_sorted = sorted(valid_results, key=lambda r: r["avg_months"])

    # Reference — current validated baseline
    baseline = next(
        (r for r in mc_results
         if r["threshold"] == 2.75 and r["filter"] == "none"),
        None,
    )

    print(f"\n  Rank  {'Threshold':<12}{'Filter':<20}{'Trades/yr':>10}"
          f"{'WR':>7}{'Pass%':>7}{'95DD':>7}"
          f"{'AvgMths':>9}{'MedMths':>9}")
    print(f"  {'─'*90}")

    for rank, r in enumerate(valid_sorted[:10], 1):
        flag = " ◄ BASELINE" if (r["threshold"] == 2.75 and r["filter"] == "none") else ""
        print(f"  {rank:<6}{r['threshold']:<12}{r['filter']:<20}"
              f"{r['per_year']:>10.1f}"
              f"{r['win_rate']:>6.2%}"
              f"{r['pass_rate']:>7.2f}"
              f"{r['dd_95']:>7.3f}"
              f"{r['avg_months']:>9.1f}"
              f"{r['med_months']:>9.1f}{flag}")

    # ── Detailed breakdown of top 3 ───────────────────────────────────────────
    print(f"\n{'='*75}")
    print("TOP 3 DETAILED BREAKDOWN")
    print(f"{'='*75}")

    for i, r in enumerate(valid_sorted[:3], 1):
        print(f"\n  #{i}: threshold {r['threshold']} | filter: {r['filter']}")
        print(f"    Trades/year     : {r['per_year']:.1f}")
        print(f"    Win rate        : {r['win_rate']:.4%}")
        print(f"    Avg return      : {r['avg_return']:.6f}")
        print(f"    Final equity    : {r['final_eq']:.4f}")
        print(f"    Backtest max DD : {r['max_dd']:.4%}")
        print(f"    MC pass rate    : {r['pass_rate']:.2f}%")
        print(f"    MC 95th pct DD  : {r['dd_95']:.3f}%")
        print(f"    Avg months/pass : {r['avg_months']:.1f}")
        print(f"    Med months/pass : {r['med_months']:.1f}")

        if baseline:
            mths_saved = baseline.get("avg_months", np.nan) - r["avg_months"]
            trades_saved = baseline.get("avg_trades_pass", np.nan) - r["avg_trades_pass"]
            print(f"\n    vs baseline (2.75, no filter):")
            print(f"      Trades/yr gain   : +{r['per_year']-baseline['per_year']:.1f}/yr")
            print(f"      Win rate change  : {(r['win_rate']-baseline['win_rate'])*100:+.2f}pp")
            print(f"      Avg months saved : {mths_saved:+.1f} months")

    # ── Save summary ──────────────────────────────────────────────────────────
    summary_df = pd.DataFrame([
        {k: v for k, v in r.items() if k != "trade_df"}
        for r in mc_results
    ])
    out_path = OUTPUT_DIR / "threshold_regime_filter_results.csv"
    summary_df.to_csv(out_path, index=False)
    print(f"\n  Results saved: {out_path}")

    print(f"\n{'='*75}")
    print("NEXT STEP")
    print(f"{'='*75}")
    print("""
  Take the best valid combination (lowest avg months, 95th pct DD < 10%)
  and freeze it as the updated model. Then re-run the full improvement
  testing script (improvement_testing_v1.py) with the new threshold and
  filter applied to confirm signal scaling still holds on the new universe.
""")


if __name__ == "__main__":
    main()
