"""
entry_method_comparison_v1.py
==============================
Tests four entry methods side by side to find the true optimal:

  Method A: Signal bar close  (matches original validated backtest)
  Method B: Next 15M bar open (immediate market execution)
  Method C: 0.618 Fibonacci   (halfway pullback)
  Method D: 0.786 Fibonacci   (current live system)

The original backtest produced 4,019 trades at 75% WR using Method A.
The live system uses Method D producing 1,607 trades at 72.8% WR.
This script finds which method best replicates the backtest and
whether the Fibonacci filter adds genuine value vs Method A.

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

Run from project root:
  python src/research/entry_method_comparison_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"
THRESHOLD     = 2.75
STOP          = 0.0025
TP            = 0.0020
ZSCORE_EXIT   = 1.5
SPREAD_COST   = 0.0001
ALLOWED_HOURS = set(range(7, 17))
ACCOUNT_START = 100_000.0
BASE_NOTIONAL = 300_000.0
YEARS         = 22.0
HOLD_HOURS    = 52

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

def get_multiplier(z):
    for lo, hi, mult in ZSCORE_BANDS:
        if lo <= z < hi:
            return mult
    return ZSCORE_BANDS[-1][2]


def simulate_trade(m15, signal_df, entry_time, entry_price, signal_dir):
    """Simulate a single trade. Returns (ret, exit_reason)."""
    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 calc_metrics(trades_list, label):
    if not trades_list:
        return None
    df = pd.DataFrame(trades_list)
    r  = df["ret"].values
    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

    exits = df["exit_reason"].value_counts(normalize=True)*100

    return {
        "label"     : label,
        "n"         : n,
        "per_week"  : n/YEARS/52,
        "win_rate"  : df["win"].mean()*100,
        "avg_ret"   : r.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"    : exits.get("tp",0),
        "stop_pct"  : exits.get("stop",0),
        "zexit_pct" : exits.get("zexit",0),
        "monthly_pnl": d.sum()/YEARS/12,
        "p1_months" : 10000/(d.sum()/YEARS/12) if d.sum()>0 else 99,
    }


def run_method(label, signals, m15, signal_df, entry_fn):
    """Run backtest for a given entry method function."""
    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

        result = entry_fn(sig, sig_time, signal_dir, m15, signal_df)
        if result is None:
            continue

        entry_time, entry_price = result
        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,
        })
        last_exit = entry_time + pd.Timedelta(hours=HOLD_HOURS)

    return calc_metrics(trades, label)


def main():
    print("="*70)
    print("ENTRY METHOD COMPARISON  —  Find True Optimal Entry")
    print("="*70)

    print("\nLoading data...")
    m15     = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)
    signals = build_frozen_signals(threshold=THRESHOLD, allowed_hours=ALLOWED_HOURS)

    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)
    except Exception as e:
        print(f"  Warning: {e}")
        signal_df = pd.DataFrame(columns=["datetime","lag_zscore_24h_v3"])

    print(f"  {len(signals):,} signals | {len(m15):,} 15M bars\n")

    # ── Define entry methods ──────────────────────────────────────────────────

    def entry_signal_bar_close(sig, sig_time, sig_dir, m15, sig_df):
        """Method A: Enter at the hourly signal bar's close price on the 1H bar.
        This is what the original backtest did — enter at close of bar where
        z-score first crossed threshold."""
        # Find the 15M bar that corresponds to signal close
        idx = get_first_m15_idx_at_or_after(m15, sig_time)
        if idx is None:
            return None
        # Use the next completed 15M bar after signal (matches hourly close)
        # Look for the bar that closes at or just after the signal
        for i in range(idx, min(idx+4, len(m15))):
            bar = m15.iloc[i]
            if bar["datetime"] >= sig_time:
                return bar["datetime"], float(bar["close"])
        return None

    def entry_next_bar_open(sig, sig_time, sig_dir, m15, sig_df):
        """Method B: Enter at open of very next 15M bar after signal."""
        idx = get_first_m15_idx_at_or_after(m15, sig_time)
        if idx is None or idx+1 >= len(m15):
            return None
        bar = m15.iloc[idx+1]
        return bar["datetime"], float(bar["open"])

    def entry_fib_618(sig, sig_time, sig_dir, m15, sig_df):
        """Method C: Fibonacci 0.618 pullback within 6h."""
        entry = find_entry_pullback(m15=m15, signal_row=sig,
                                     fib=0.618, wait_hours=6)
        if entry is None:
            return None
        return entry["entry_time"], entry["entry_price"]

    def entry_fib_786(sig, sig_time, sig_dir, m15, sig_df):
        """Method D: Fibonacci 0.786 pullback within 6h (current live system)."""
        entry = find_entry_pullback(m15=m15, signal_row=sig,
                                     fib=0.786, wait_hours=6)
        if entry is None:
            return None
        return entry["entry_time"], entry["entry_price"]

    def entry_fib_500(sig, sig_time, sig_dir, m15, sig_df):
        """Method E: Fibonacci 0.500 pullback within 6h (50% retracement)."""
        entry = find_entry_pullback(m15=m15, signal_row=sig,
                                     fib=0.500, wait_hours=6)
        if entry is None:
            return None
        return entry["entry_time"], entry["entry_price"]

    # ── Run all methods ───────────────────────────────────────────────────────
    methods = [
        ("A: Signal bar close  (ORIGINAL BACKTEST)", entry_signal_bar_close),
        ("B: Next bar open     (immediate market)", entry_next_bar_open),
        ("C: Fib 0.500         (50% pullback)",    entry_fib_500),
        ("D: Fib 0.618         (61.8% pullback)",  entry_fib_618),
        ("E: Fib 0.786         (CURRENT LIVE)",    entry_fib_786),
    ]

    results = []
    for label, fn in methods:
        print(f"  Testing {label}...", end="", flush=True)
        m = run_method(label, signals, m15, signal_df, fn)
        if m:
            results.append(m)
            print(f"  n={m['n']:,}  WR={m['win_rate']:.1f}%  "
                  f"Sharpe={m['sharpe']:.2f}  {m['per_week']:.1f}/wk")
        else:
            print("  no results")

    # ── Print comparison table ────────────────────────────────────────────────
    print(f"\n{'='*95}")
    print("ENTRY METHOD COMPARISON — FULL RESULTS")
    print(f"{'='*95}")
    print(f"\n  {'Method':<38}{'Trades':>7}{'WR%':>7}{'Sharpe':>8}{'PF':>7}"
          f"{'CAGR%':>7}{'MaxDD%':>8}{'P1 Mo':>7}{'TP%':>6}{'Stop%':>7}"
          f"{'Per Wk':>8}")
    print(f"  {'─'*95}")

    best_sharpe = max(r["sharpe"] for r in results)
    backtest_n  = 4019

    for r in results:
        flag = " <- BEST" if r["sharpe"] == best_sharpe else ""
        n_diff = r["n"] - backtest_n
        match  = f" (Δ{n_diff:+,} vs backtest)"
        print(f"  {r['label']:<38}{r['n']:>7}"
              f"{r['win_rate']:>7.1f}{r['sharpe']:>8.2f}"
              f"{r['pf']:>7.2f}{r['cagr']:>7.2f}"
              f"{r['max_dd']:>8.2f}{r['p1_months']:>7.1f}"
              f"{r['tp_pct']:>6.1f}{r['stop_pct']:>7.1f}"
              f"{r['per_week']:>8.1f}{flag}")

    # ── Analysis ─────────────────────────────────────────────────────────────
    best    = max(results, key=lambda r: r["sharpe"])
    orig    = next((r for r in results if "ORIGINAL" in r["label"]), None)
    current = next((r for r in results if "CURRENT"  in r["label"]), None)

    print(f"\n{'='*70}")
    print("ANALYSIS AND RECOMMENDATION")
    print(f"{'='*70}")

    if orig:
        print(f"""
  ORIGINAL BACKTEST METHOD (Signal bar close):
    Trades/week : {orig['per_week']:.1f}
    Win rate    : {orig['win_rate']:.1f}%
    Sharpe      : {orig['sharpe']:.2f}
    P1 months   : {orig['p1_months']:.1f}
    Trade count : {orig['n']:,}  (backtest had 4,019 — gap is position overlap logic)
""")

    if current:
        print(f"""  CURRENT LIVE METHOD (Fib 0.786):
    Trades/week : {current['per_week']:.1f}
    Win rate    : {current['win_rate']:.1f}%
    Sharpe      : {current['sharpe']:.2f}
    P1 months   : {current['p1_months']:.1f}
    Trade count : {current['n']:,}
""")

    print(f"""  BEST METHOD ({best['label'].strip()}):
    Trades/week : {best['per_week']:.1f}
    Win rate    : {best['win_rate']:.1f}%
    Sharpe      : {best['sharpe']:.2f}
    P1 months   : {best['p1_months']:.1f}
""")

    # Decision
    if orig and orig["sharpe"] > current["sharpe"] and orig["per_week"] > current["per_week"] * 1.5:
        print("""  RECOMMENDATION: REVERT TO SIGNAL BAR CLOSE ENTRY
    The original backtest used signal bar close entry.
    This method produces more trades AND better Sharpe.
    The Fibonacci pullback reduces frequency without improving quality.
    Update live_signal_monitor.py to enter at signal bar close.
""")
    elif current and current["sharpe"] >= orig["sharpe"] * 0.9:
        print("""  RECOMMENDATION: KEEP FIBONACCI 0.786 ENTRY
    Fibonacci entry produces similar or better risk-adjusted returns.
    Lower trade frequency is offset by higher per-trade quality.
    The difference in P1 timeline is acceptable.
""")
    else:
        fib_best = next((r for r in results if best["label"] in r["label"]), None)
        if fib_best:
            print(f"""  RECOMMENDATION: SWITCH TO {best['label'].strip()}
    This method produces the best risk-adjusted returns.
    Consider updating the live script entry logic.
""")

    # Save
    pd.DataFrame(results).to_csv(
        TRADES_DIR / "entry_method_comparison.csv", index=False)
    print(f"  Results saved: entry_method_comparison.csv")


if __name__ == "__main__":
    main()
