r"""
validated_model_on_live_window.py
==================================
Runs the EXACT validated backtest model (15M frequency, spot_lag_v3 signal
generation, combined_candidate_matrix entry logic) on the LIVE DEPLOYMENT 
PERIOD (April 11 - May 13 2026, or whatever the user specifies).

Compares against what the live monitor actually fired.

This tells us: how many trades did we MISS by running the wrong (hourly) 
model during this paper-trading period?

Run on home machine:
  cd C:\Users\paul_\OneDrive\fx_macro_intraday
  python src\research\validated_model_on_live_window.py
"""
import sys
import warnings
from pathlib import Path

import pandas as pd
import numpy as np

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 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

# ── EU validated parameters ──────────────────────────────────────────────
EU_THRESHOLD     = 2.75
EU_FIB           = 0.786
EU_HOLD_HOURS    = 52
EU_STOP          = 0.0025
EU_TP            = 0.0020
EU_ALLOWED_HOURS = set(range(7, 17))

# ── UJ validated parameters ──────────────────────────────────────────────
UJ_THRESHOLD     = 2.0
UJ_FIB           = 0.786
UJ_HOLD_HOURS    = 24
UJ_TP            = 0.0070
UJ_SL            = 0.0040
UJ_Z_WINDOW      = 30

# ── Period defaults ──────────────────────────────────────────────────────
DEFAULT_START = "2026-04-11"
DEFAULT_END   = "2026-05-13"

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


# ──────────────────────────────────────────────────────────────────────────
# Replicate the full backtest signal+entry+exit pipeline for EU
# ──────────────────────────────────────────────────────────────────────────
def run_eurusd_validated(start_date, end_date):
    print(f"\n  EU validated model from {start_date.date()} to {end_date.date()}")
    print(f"  Using signal frequency: every 15M bar (matches validated model)")
    print(f"  Session filter: UTC 7-16 (range(7,17))")
    
    signals = build_frozen_signals(
        threshold=EU_THRESHOLD,
        allowed_hours=EU_ALLOWED_HOURS,
    )
    signals["datetime"] = pd.to_datetime(signals["datetime"])
    window = signals[
        (signals["datetime"] >= start_date) &
        (signals["datetime"] <= end_date)
    ].copy().reset_index(drop=True)
    
    print(f"  Signal-bars in window: {len(window)}")
    
    if len(window) == 0:
        print(f"  No signals in window — note that spot_lag_v3 data may end before {end_date.date()}")
        return [], window
    
    m15 = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)
    
    trades = []
    last_exit_time = None
    
    for _, sig in window.iterrows():
        if last_exit_time is not None and sig["datetime"] < last_exit_time:
            continue
        entry = find_entry_pullback(
            m15=m15, signal_row=sig, fib=EU_FIB, wait_hours=6
        )
        if entry is None:
            continue
        
        entry_idx = get_first_m15_idx_at_or_after(m15, entry["entry_time"])
        if entry_idx is None:
            continue
        hold_bars = EU_HOLD_HOURS * 4
        exit_idx = min(entry_idx + hold_bars, len(m15) - 1)
        path = m15.iloc[entry_idx:exit_idx + 1].copy()
        if path.empty:
            continue
        
        entry_price = entry["entry_price"]
        signal = int(sig["signal"])
        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 >= EU_TP
                stop_hit = (float(bar["low"]) - entry_price) / entry_price <= -EU_STOP
            else:
                tp_hit = (entry_price - float(bar["low"])) / entry_price >= EU_TP
                stop_hit = -((float(bar["high"]) - entry_price) / entry_price) <= -EU_STOP
            if tp_hit:
                exit_price = entry_price*(1+EU_TP) if signal==1 else entry_price*(1-EU_TP)
                exit_time = bar["datetime"]
                exit_reason = "tp"
                break
            if stop_hit:
                exit_price = entry_price*(1-EU_STOP) if signal==1 else entry_price*(1+EU_STOP)
                exit_time = bar["datetime"]
                exit_reason = "stop"
                break
        
        last_exit_time = exit_time
        ret = ((exit_price - entry_price) / entry_price) * signal
        trades.append({
            "pair": "EURUSD",
            "signal_time": sig["datetime"],
            "entry_time": entry["entry_time"],
            "exit_time": exit_time,
            "signal": signal,
            "zscore": float(sig["lag_zscore_24h_v3"]),
            "entry_price": entry_price,
            "exit_price": exit_price,
            "exit_reason": exit_reason,
            "return": ret,
        })
    
    return trades, window


# ──────────────────────────────────────────────────────────────────────────
# UJ validated model (matches usdjpy_real_costs_v1.py)
# ──────────────────────────────────────────────────────────────────────────
def run_usdjpy_validated(start_date, end_date):
    print(f"\n  UJ validated model from {start_date.date()} to {end_date.date()}")
    
    # Load rates
    us2y = pd.read_csv(RATES_DIR / "us2y.csv")
    us2y.columns = [c.lower() for c in us2y.columns]
    us2y["date"] = pd.to_datetime(us2y["date"])
    us2y["us2y"] = pd.to_numeric(us2y["us2y"], errors="coerce")
    us2y = us2y[["date", "us2y"]].dropna()
    
    jp2y = pd.read_csv(RATES_DIR / "jp2y.csv")
    jp2y.columns = [c.lower() for c in jp2y.columns]
    jp2y["date"] = pd.to_datetime(jp2y["date"])
    jp2y["jp2y"] = pd.to_numeric(jp2y["jp2y"], errors="coerce")
    jp2y = jp2y[["date", "jp2y"]].dropna()
    
    rates = pd.merge(us2y, jp2y, on="date", how="outer").sort_values("date")
    rates = rates.set_index("date")
    rates = rates.reindex(
        pd.date_range(rates.index.min(), rates.index.max(), freq="D")
    ).ffill().dropna()
    rates["spread"] = rates["us2y"] - rates["jp2y"]
    z = (rates["spread"] - rates["spread"].rolling(UJ_Z_WINDOW).mean()) \
        / rates["spread"].rolling(UJ_Z_WINDOW).std()
    z_hourly = z.reindex(
        pd.date_range(z.index.min(), z.index.max(), freq="h")
    ).ffill()
    
    fpath = next(BASE.rglob("USDJPY_15M*.csv"), None)
    if fpath is None:
        print("  No USDJPY 15M data found")
        return [], pd.DataFrame()
    
    df = pd.read_csv(fpath)
    df.columns = df.columns.str.strip().str.lower()
    dt_col = next(c for c in df.columns if "date" in c or "time" in c)
    df[dt_col] = pd.to_datetime(df[dt_col])
    m15 = df.rename(columns={dt_col: "datetime"}).set_index("datetime").sort_index()
    
    # Count signal-hours in window
    z_in_window = z_hourly[(z_hourly.index >= start_date) & (z_hourly.index <= end_date)]
    breach_hours = ((z_in_window.abs() >= UJ_THRESHOLD) &
                    z_in_window.index.map(lambda d: 7 <= (d.hour + 2) % 24 <= 17)).sum()
    print(f"  Signal-hours in window (|z|>=2.0 + session): {breach_hours}")
    
    trades = []
    last_exit = None
    
    for dt in pd.date_range(start_date, end_date, freq="h"):
        if dt not in z_hourly.index:
            continue
        z_val = float(z_hourly.at[dt])
        if abs(z_val) < UJ_THRESHOLD:
            continue
        if not (7 <= (dt.hour + 2) % 24 <= 17):
            continue
        if last_exit is not None and dt <= last_exit:
            continue
        
        direction = 1 if z_val >= UJ_THRESHOLD else -1
        try:
            slice_15m = m15.loc[dt: dt + pd.Timedelta(hours=1) - pd.Timedelta(minutes=1)]
        except Exception:
            continue
        if slice_15m.empty:
            continue
        bh = float(slice_15m["high"].max())
        bl = float(slice_15m["low"].min())
        bc = float(slice_15m["close"].iloc[-1])
        
        if direction == 1:
            pull = bc - bl
            if pull <= 0.05:
                continue
            target = bc - UJ_FIB * pull
        else:
            pull = bh - bc
            if pull <= 0.05:
                continue
            target = bc + UJ_FIB * pull
        
        entry_price = None
        entry_time = None
        try:
            entry_window = m15.loc[
                dt + pd.Timedelta(minutes=1):
                dt + pd.Timedelta(hours=6)
            ]
        except Exception:
            continue
        
        for edt, ebar in entry_window.iterrows():
            if direction == 1 and float(ebar["low"]) <= target:
                entry_price = target; entry_time = edt; break
            elif direction == -1 and float(ebar["high"]) >= target:
                entry_price = target; entry_time = edt; break
        if entry_price is None:
            continue
        
        tp_px = entry_price * (1 + UJ_TP) if direction == 1 else entry_price * (1 - UJ_TP)
        sl_px = entry_price * (1 - UJ_SL) if direction == 1 else entry_price * (1 + UJ_SL)
        hold_bars = m15.loc[
            entry_time + pd.Timedelta(minutes=1):
            entry_time + pd.Timedelta(hours=UJ_HOLD_HOURS)
        ]
        
        exit_price = None
        exit_reason = "hold_expiry"
        exit_time = None
        for hdt, hbar in hold_bars.iterrows():
            if direction == 1:
                if float(hbar["high"]) >= tp_px:
                    exit_price = tp_px; exit_reason = "tp"; exit_time = hdt; break
                if float(hbar["low"]) <= sl_px:
                    exit_price = sl_px; exit_reason = "stop"; exit_time = hdt; break
            else:
                if float(hbar["low"]) <= tp_px:
                    exit_price = tp_px; exit_reason = "tp"; exit_time = hdt; break
                if float(hbar["high"]) >= sl_px:
                    exit_price = sl_px; exit_reason = "stop"; exit_time = hdt; break
        if exit_price is None and not hold_bars.empty:
            exit_price = float(hold_bars["close"].iloc[-1])
            exit_time = hold_bars.index[-1]
        if exit_time is None:
            continue
        
        last_exit = exit_time
        ret = ((exit_price - entry_price) / entry_price) * direction
        trades.append({
            "pair": "USDJPY",
            "signal_time": dt,
            "entry_time": entry_time,
            "exit_time": exit_time,
            "signal": direction,
            "zscore": z_val,
            "entry_price": entry_price,
            "exit_price": exit_price,
            "exit_reason": exit_reason,
            "return": ret,
        })
    
    return trades, z_in_window


def main():
    import argparse
    parser = argparse.ArgumentParser()
    parser.add_argument("--start", default=DEFAULT_START)
    parser.add_argument("--end",   default=DEFAULT_END)
    args = parser.parse_args()
    
    start_date = pd.Timestamp(args.start)
    end_date   = pd.Timestamp(args.end)
    days = (end_date - start_date).days
    
    print("=" * 80)
    print("  VALIDATED MODEL ON LIVE DEPLOYMENT PERIOD")
    print("=" * 80)
    print(f"  Period: {start_date.date()} → {end_date.date()} ({days} days)")
    print(f"  This shows what the VALIDATED model would have produced.")
    
    eu_trades, eu_signals = run_eurusd_validated(start_date, end_date)
    uj_trades, uj_signals = run_usdjpy_validated(start_date, end_date)
    
    # Display
    print()
    print("=" * 80)
    print("  EURUSD: validated model would have taken")
    print("=" * 80)
    if eu_trades:
        for t in eu_trades:
            print(f"  {t['signal_time']} | entry {t['entry_time']} | "
                  f"z={t['zscore']:+.2f} | "
                  f"{'L' if t['signal']==1 else 'S'} | "
                  f"{t['exit_reason']} | ret={t['return']*100:+.3f}%")
    else:
        print("  (none — see note above re: spot_lag_v3 data freshness)")
    
    print()
    print("=" * 80)
    print("  USDJPY: validated model would have taken")
    print("=" * 80)
    if uj_trades:
        for t in uj_trades:
            print(f"  {t['signal_time']} | entry {t['entry_time']} | "
                  f"z={t['zscore']:+.2f} | "
                  f"{'L' if t['signal']==1 else 'S'} | "
                  f"{t['exit_reason']} | ret={t['return']*100:+.3f}%")
    else:
        print("  (none)")
    
    # Aggregate
    print()
    print("=" * 80)
    print("  SUMMARY")
    print("=" * 80)
    print(f"  Period: {days} days")
    print(f"  EU validated model: {len(eu_trades)} trade(s)")
    print(f"  UJ validated model: {len(uj_trades)} trade(s)")
    print(f"  Combined:           {len(eu_trades) + len(uj_trades)} trade(s)")
    print()
    
    all_trades = eu_trades + uj_trades
    if all_trades:
        df = pd.DataFrame(all_trades)
        wins = (df['return'] > 0).sum()
        print(f"  Total trades:     {len(df)}")
        print(f"  Win rate:         {wins/len(df)*100:.1f}% ({wins} wins, {len(df)-wins} losses)")
        print(f"  Avg return:       {df['return'].mean()*100:+.3f}%")
        print(f"  Total return:     {df['return'].sum()*100:+.3f}%")
        print()
        print(f"  On $100k account at 0.75% risk ($300k notional):")
        gross_pnl = df['return'].sum() * 300_000
        print(f"  Estimated gross P&L: ${gross_pnl:+,.0f}")
    
    print()
    print("  Compare against ACTUAL live monitor in same period:")
    print("  - 1 EU trade (April 17, won)")
    print("  - 4 UJ trades (April 30 + May 5, all lost)")
    print()
    print("  Gap = (validated model trades) - (live monitor trades)")
    print("  This is the empirical cost of running the wrong model.")


if __name__ == "__main__":
    main()
