"""
news_trigger_backtest_v1.py
============================
Tests whether adding news-triggered z-score checks improves results
vs the current hourly-only schedule.

The key question: are there signals that crossed 2.75 during a high-impact
news event window but were MISSED by the hourly :01 check?

Logic:
  - Hourly schedule: checks at :01 past each hour (current live system)
  - News-triggered: also checks at the exact minute of each high-impact
    USD or EUR news event

Compares:
  - Trade count, win rate, Sharpe, profit factor
  - Which news event types generate the most additional signals
  - Average time between signal and nearest hourly check (the "miss window")

Place in:
  C:\\Users\\paul_\\OneDrive\\fx_macro_intraday\\src\\research\\news_trigger_backtest_v1.py

Run from project root:
  python src/research/news_trigger_backtest_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

TRADES_DIR    = BASE_PATH / "data" / "processed" / "trades"
NEWS_PATH     = BASE_PATH / "notebooks" / "forex_factory_cache.csv"
THRESHOLD     = 2.75
FIB           = 0.786
STOP          = 0.0025
TP            = 0.0020
ZSCORE_EXIT   = 1.5
SPREAD_COST   = 0.0001
ALLOWED_HOURS = set(range(7, 17))
BASE_NOTIONAL = 300_000.0
YEARS         = 22.0
ACCOUNT_START = 100_000.0
HOLD_HOURS    = 52

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

# High impact news events for USD and EUR that move yield spreads
HIGH_IMPACT_EVENTS = [
    "Non-Farm Payrolls", "NFP", "CPI", "Core CPI", "PCE",
    "Fed Interest Rate Decision", "FOMC", "Fed Rate",
    "GDP", "Unemployment Rate", "ISM Manufacturing",
    "Retail Sales", "PPI", "Producer Price",
    "ECB Interest Rate Decision", "ECB Rate",
    "German CPI", "German GDP", "German Ifo",
    "Flash Manufacturing PMI", "Flash Services PMI",
    "ADP", "Initial Jobless Claims", "Consumer Confidence",
    "Durable Goods", "Housing Starts", "Treasury",
    "Jackson Hole", "Fed Chair", "ECB President",
    "Inflation", "Employment", "Trade Balance",
]

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


def simulate_trade(m15, signal_df, entry_time, entry_price, signal_dir):
    entry_idx = get_first_m15_idx_at_or_after(m15, entry_time)
    if entry_idx is None: return None, None

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

    sig_window = signal_df[
        (signal_df["datetime"] >= entry_time) &
        (signal_df["datetime"] <= path.iloc[-1]["datetime"])
    ]

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

    for _, bar in path.iterrows():
        bt = bar["datetime"]
        if signal_dir == 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 = (entry_price-float(bar["high"]))/entry_price <= -STOP

        if tp_hit:
            exit_price  = entry_price*(1+TP) if signal_dir==1 else entry_price*(1-TP)
            exit_reason = "tp"; break
        if stop_hit:
            exit_price  = entry_price*(1-STOP) if signal_dir==1 else entry_price*(1+STOP)
            exit_reason = "stop"; break

        z_bars = sig_window[sig_window["datetime"] <= bt]
        if not z_bars.empty:
            cz = float(z_bars.iloc[-1]["lag_zscore_24h_v3"])
            if (signal_dir==1 and cz<=-ZSCORE_EXIT) or (signal_dir==-1 and cz>=ZSCORE_EXIT):
                exit_price  = float(bar["close"])
                exit_reason = "zexit"; break

    raw_ret = (exit_price-entry_price)/entry_price * signal_dir
    return raw_ret-SPREAD_COST, exit_reason


def run_backtest(signals, m15, signal_df, label):
    trades    = []
    last_exit = None

    for _, sig in signals.iterrows():
        sig_time   = sig["datetime"]
        signal_dir = int(sig["signal"])
        zscore_abs = abs(float(sig["lag_zscore_24h_v3"]))
        mult       = get_multiplier(zscore_abs)
        notional   = min(BASE_NOTIONAL*mult, BASE_NOTIONAL*3)

        if last_exit is not None and sig_time < last_exit:
            continue

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

        entry_time  = entry["entry_time"]
        entry_price = entry["entry_price"]

        ret, reason = simulate_trade(m15, signal_df, entry_time,
                                      entry_price, signal_dir)
        if ret is None:
            continue

        trades.append({
            "entry_time"  : entry_time,
            "ret"         : ret,
            "exit_reason" : reason,
            "win"         : int(ret>0),
            "dollar_pnl"  : ret*notional,
            "zscore_abs"  : zscore_abs,
            "sig_time"    : sig_time,
        })
        last_exit = entry_time + pd.Timedelta(hours=HOLD_HOURS)

    if not trades:
        return None

    df = pd.DataFrame(trades)
    d  = df["dollar_pnl"].values
    n  = len(df)
    equity = ACCOUNT_START + d.cumsum()
    peak   = np.maximum.accumulate(equity)
    dd     = (equity-peak)/peak*100

    per_yr = n/YEARS
    rf_pt  = (1+0.04)**(1/per_yr)-1
    excess = d/ACCOUNT_START-rf_pt
    sharpe = (excess.mean()/excess.std()*np.sqrt(per_yr) if excess.std()>0 else 0)
    gp = d[d>0].sum(); gl = abs(d[d<0].sum())
    pf = gp/gl if gl>0 else 999

    return {
        "label"    : label,
        "n"        : n,
        "per_week" : n/YEARS/52,
        "win_rate" : df["win"].mean()*100,
        "sharpe"   : sharpe,
        "pf"       : pf,
        "max_dd"   : dd.min(),
        "net_pnl"  : d.sum(),
        "cagr"     : (equity[-1]/ACCOUNT_START)**(1/YEARS)*100-100,
        "tp_pct"   : df["exit_reason"].value_counts(normalize=True).get("tp",0)*100,
        "p1_months": 10000/(d.sum()/YEARS/12) if d.sum()>0 else 99,
    }


def main():
    print("="*70)
    print("NEWS-TRIGGER BACKTEST")
    print("="*70)

    # ── Load data ─────────────────────────────────────────────────────────────
    print("\nLoading data...")
    m15     = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)
    print(f"  15M bars: {len(m15):,}")

    try:
        from features.spot_lag_v3 import get_model_ready_spot_lag_v3
        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)
        z_col = "lag_zscore_24h_v3"
    except Exception as e:
        print(f"  ERROR loading signal_df: {e}"); return

    # ── Load news data ────────────────────────────────────────────────────────
    print(f"  Loading news: {NEWS_PATH}")
    news = pd.read_csv(NEWS_PATH)
    news["DateTime"] = pd.to_datetime(news["DateTime"], utc=True).dt.tz_localize(None)
    news = news[
        (news["Currency"].isin(["USD","EUR"])) &
        (news["Impact"] == "High Impact Expected")
    ].copy()
    news = news.sort_values("DateTime").reset_index(drop=True)
    print(f"  High-impact USD/EUR events: {len(news):,}")

    # ── Build hourly signal times (current system) ────────────────────────────
    # Signal fires at :01 past each hour within session
    hourly_times = pd.date_range(
        start=signal_df["datetime"].min(),
        end=signal_df["datetime"].max(),
        freq="1h"
    )
    hourly_times = pd.DatetimeIndex([
        t.replace(minute=1) for t in hourly_times
        if t.hour in ALLOWED_HOURS
    ])

    # ── Build news-triggered signal times ────────────────────────────────────
    # News event times within session hours
    news_times = news[
        news["DateTime"].dt.hour.isin(ALLOWED_HOURS)
    ]["DateTime"].drop_duplicates().sort_values()

    # Combined: hourly + news times
    combined_times = pd.DatetimeIndex(
        sorted(set(hourly_times.tolist() + news_times.tolist()))
    )

    print(f"\n  Hourly check times      : {len(hourly_times):,}")
    print(f"  News event times        : {len(news_times):,}")
    print(f"  Combined check times    : {len(combined_times):,}")
    print(f"  Extra checks from news  : {len(combined_times)-len(hourly_times):,}")

    # ── Find signals at each check schedule ──────────────────────────────────
    # For each check time, look up the most recent z-score
    print("\n  Building signal sets...")

    def get_signals_at_times(check_times, label):
        """Find all z-score threshold crossings at the given check times."""
        signals_found = []
        prev_z = 0.0

        for ct in check_times:
            # Get z-score at this check time (most recent available)
            avail = signal_df[signal_df["datetime"] <= ct]
            if avail.empty:
                continue

            row    = avail.iloc[-1]
            curr_z = float(row[z_col])

            # Signal fires when z crosses threshold (was below, now above)
            crossed_long  = prev_z <  THRESHOLD and curr_z >=  THRESHOLD
            crossed_short = prev_z > -THRESHOLD and curr_z <= -THRESHOLD

            if (crossed_long or crossed_short) and ct.hour in ALLOWED_HOURS:
                signals_found.append({
                    "datetime"         : ct,
                    "signal"           : 1 if crossed_long else -1,
                    "lag_zscore_24h_v3": curr_z,
                    "close"            : float(row.get("close", 0)),
                    "high"             : float(row.get("high", 0)),
                    "low"              : float(row.get("low", 0)),
                })

            prev_z = curr_z

        return pd.DataFrame(signals_found) if signals_found else pd.DataFrame()

    hourly_sigs  = get_signals_at_times(hourly_times,  "hourly")
    combined_sigs = get_signals_at_times(combined_times, "combined")

    print(f"  Signals (hourly only)   : {len(hourly_sigs):,}")
    print(f"  Signals (hourly+news)   : {len(combined_sigs):,}")
    extra = len(combined_sigs) - len(hourly_sigs)
    print(f"  Additional signals      : {extra:,} ({extra/max(len(hourly_sigs),1)*100:.1f}%)")

    # ── Run both backtests ────────────────────────────────────────────────────
    print("\n  Running backtests...")
    print("  Hourly only...", end="", flush=True)
    r_hourly = run_backtest(hourly_sigs, m15, signal_df, "Hourly only (current)")
    if r_hourly:
        print(f"  {r_hourly['n']:,} trades  WR={r_hourly['win_rate']:.1f}%  Sharpe={r_hourly['sharpe']:.2f}")

    print("  Hourly + news...", end="", flush=True)
    r_combined = run_backtest(combined_sigs, m15, signal_df, "Hourly + news trigger")
    if r_combined:
        print(f"  {r_combined['n']:,} trades  WR={r_combined['win_rate']:.1f}%  Sharpe={r_combined['sharpe']:.2f}")

    # ── Results table ─────────────────────────────────────────────────────────
    print(f"\n{'='*75}")
    print("RESULTS COMPARISON")
    print(f"{'='*75}")
    print(f"\n  {'Method':<28}{'Trades':>7}{'WR%':>7}{'Sharpe':>8}{'PF':>7}"
          f"{'CAGR%':>7}{'MaxDD%':>8}{'P1 Mo':>7}{'Per Wk':>8}")
    print(f"  {'─'*75}")

    for r in [r_hourly, r_combined]:
        if r:
            print(f"  {r['label']:<28}{r['n']:>7}{r['win_rate']:>7.1f}"
                  f"{r['sharpe']:>8.2f}{r['pf']:>7.2f}{r['cagr']:>7.2f}"
                  f"{r['max_dd']:>8.2f}{r['p1_months']:>7.1f}{r['per_week']:>8.1f}")

    if r_hourly and r_combined:
        sharpe_diff = r_combined["sharpe"] - r_hourly["sharpe"]
        trade_diff  = r_combined["n"] - r_hourly["n"]
        wr_diff     = r_combined["win_rate"] - r_hourly["win_rate"]

        print(f"\n  Delta (news - hourly):")
        print(f"    Extra trades    : {trade_diff:+,}")
        print(f"    Sharpe change   : {sharpe_diff:+.2f}")
        print(f"    Win rate change : {wr_diff:+.2f}%")

    # ── Which news events generate most extra signals ─────────────────────────
    print(f"\n{'='*75}")
    print("WHICH NEWS EVENTS GENERATE EXTRA SIGNALS?")
    print(f"{'='*75}")

    # Find the extra signals in combined that are not in hourly
    if not hourly_sigs.empty and not combined_sigs.empty:
        hourly_set   = set(hourly_sigs["datetime"].dt.floor("h"))
        combined_set = set(combined_sigs["datetime"].dt.floor("h"))
        extra_hours  = combined_set - hourly_set

        extra_sig_times = combined_sigs[
            combined_sigs["datetime"].dt.floor("h").isin(extra_hours)
        ]["datetime"]

        if len(extra_sig_times) > 0:
            event_counts = {}
            for st in extra_sig_times:
                # Find news events within 30 minutes before signal
                window = news[
                    (news["DateTime"] >= st - pd.Timedelta(minutes=30)) &
                    (news["DateTime"] <= st)
                ]
                for _, ev in window.iterrows():
                    name = ev["Event"]
                    event_counts[name] = event_counts.get(name, 0) + 1

            if event_counts:
                sorted_events = sorted(event_counts.items(),
                                       key=lambda x: x[1], reverse=True)
                print(f"\n  {'Event':<45}{'Extra signals':>15}")
                print(f"  {'─'*62}")
                for ev, cnt in sorted_events[:15]:
                    print(f"  {ev:<45}{cnt:>15}")
        else:
            print("\n  No extra signals generated by news triggers")

    # ── Conclusion ────────────────────────────────────────────────────────────
    print(f"\n{'='*70}")
    print("CONCLUSION")
    print(f"{'='*70}")

    if r_hourly and r_combined:
        if sharpe_diff > 0.1 and trade_diff > 50:
            print(f"""
  NEWS TRIGGERS ADD MEANINGFUL VALUE
  Extra trades    : {trade_diff:+,}
  Sharpe gain     : {sharpe_diff:+.2f}

  RECOMMENDATION: Add news-event checks to live_signal_monitor.py
  The live system should trigger an out-of-schedule z-score check
  immediately after each high-impact USD/EUR news release.
""")
        elif sharpe_diff > 0 and trade_diff > 0:
            print(f"""
  NEWS TRIGGERS ADD MARGINAL VALUE
  Extra trades    : {trade_diff:+,}
  Sharpe gain     : {sharpe_diff:+.2f}

  RECOMMENDATION: Optional improvement. The hourly schedule captures
  most signals. News triggers would help at the margin but are not
  critical to model performance.
""")
        else:
            print(f"""
  NEWS TRIGGERS ADD NO VALUE
  Extra trades    : {trade_diff:+,}
  Sharpe change   : {sharpe_diff:+.2f}

  CONCLUSION: The hourly schedule already captures all meaningful
  signals. The z-score builds over hours not minutes — news events
  do not create the kind of instant threshold crossing that would
  be missed by a 59-minute window.
  Keep the current hourly schedule.
""")


if __name__ == "__main__":
    main()
