import pandas as pd
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 features.spot_lag_v3 import get_model_ready_spot_lag_v3
from ingestion.price_loader_15m import load_eurusd_15m


def build_frozen_signals(
    threshold: float = 2.0,
    allowed_hours: set | None = None,
):
    """
    Same frozen benchmark signal layer:
    - lag_zscore_24h_v3 from the validated model
    - London + NY session filter
    - signal rows only
    """
    df = get_model_ready_spot_lag_v3().copy()
    df = df.sort_values("datetime").reset_index(drop=True)

    df["hour"] = pd.to_datetime(df["datetime"]).dt.hour
    if allowed_hours is not None:
        df = df[df["hour"].isin(allowed_hours)].copy()

    df["signal"] = 0
    df.loc[df["lag_zscore_24h_v3"] >= threshold, "signal"] = 1
    df.loc[df["lag_zscore_24h_v3"] <= -threshold, "signal"] = -1

    signals = df[df["signal"] != 0].copy().reset_index(drop=True)

    keep_cols = [
        "datetime",
        "open",
        "high",
        "low",
        "close",
        "lag_zscore_24h_v3",
        "signal",
    ]
    return signals[keep_cols].copy()


def get_first_m15_idx_at_or_after(m15: pd.DataFrame, ts: pd.Timestamp):
    idx = m15["datetime"].searchsorted(ts, side="left")
    if idx >= len(m15):
        return None
    return int(idx)


def find_entry(
    m15: pd.DataFrame,
    signal_row: pd.Series,
    entry_type,
    wait_hours: int = 6,
):
    """
    Entry refinement only.

    entry_type:
      - "immediate"
      - 0.236 / 0.382 / 0.5 / 0.618 / 0.786

    Fib pullbacks are anchored to the SIGNAL HOUR CANDLE RANGE,
    which is known at signal time and avoids lookahead.

    Long:
      target = signal_close - fib * (signal_close - signal_low)

    Short:
      target = signal_close + fib * (signal_high - signal_close)
    """
    signal_time = signal_row["datetime"]
    signal = int(signal_row["signal"])
    signal_close = float(signal_row["close"])
    signal_high = float(signal_row["high"])
    signal_low = float(signal_row["low"])

    start_idx = get_first_m15_idx_at_or_after(m15, signal_time)
    if start_idx is None:
        return None

    wait_bars = wait_hours * 4
    window = m15.iloc[start_idx : start_idx + wait_bars].copy()
    if window.empty:
        return None

    if entry_type == "immediate":
        entry_row = window.iloc[0]
        return {
            "entry_time": entry_row["datetime"],
            "entry_price": float(entry_row["close"]),
        }

    fib = float(entry_type)

    if signal == 1:
        pullback_range = signal_close - signal_low
        if pullback_range <= 0:
            return None

        target_price = signal_close - fib * pullback_range

        hit = window[window["low"] <= target_price]
        if hit.empty:
            return None

        hit_row = hit.iloc[0]
        return {
            "entry_time": hit_row["datetime"],
            "entry_price": target_price,  # limit-style fill at target
        }

    else:
        pullback_range = signal_high - signal_close
        if pullback_range <= 0:
            return None

        target_price = signal_close + fib * pullback_range

        hit = window[window["high"] >= target_price]
        if hit.empty:
            return None

        hit_row = hit.iloc[0]
        return {
            "entry_time": hit_row["datetime"],
            "entry_price": target_price,  # limit-style fill at target
        }


def simulate_trade(
    m15: pd.DataFrame,
    entry_time: pd.Timestamp,
    entry_price: float,
    signal: int,
    hold_hours: int = 48,
    stop: float = 0.0025,
    spread_cost: float = 0.0001,
):
    """
    Contiguous 15m truth path from the chosen entry time onward.
    """
    entry_idx = get_first_m15_idx_at_or_after(m15, entry_time)
    if entry_idx is None:
        return None

    hold_bars = 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:
        return None

    exit_price = float(path.iloc[-1]["close"])

    if signal == 1:
        adverse = (path["low"].min() - entry_price) / entry_price
        stop_hit = adverse < -stop
        ret = exit_price / entry_price - 1
        if stop_hit:
            ret = -stop
    else:
        adverse = (path["high"].max() - entry_price) / entry_price
        stop_hit = adverse > stop
        ret = -(exit_price / entry_price - 1)
        if stop_hit:
            ret = -stop

    ret -= spread_cost

    return {
        "entry_time": entry_time,
        "exit_time": path.iloc[-1]["datetime"],
        "signal": signal,
        "entry_price": entry_price,
        "exit_price": exit_price,
        "return": ret,
    }


def run_backtest_for_entry_type(
    entry_type,
    threshold: float = 2.0,
    hold_hours: int = 48,
    stop: float = 0.0025,
    spread_cost: float = 0.0001,
    wait_hours: int = 6,
):
    # frozen benchmark session
    allowed_hours = set(range(7, 17))

    signals = build_frozen_signals(
        threshold=threshold,
        allowed_hours=allowed_hours,
    )

    m15 = load_eurusd_15m().copy()
    m15 = m15.sort_values("datetime").reset_index(drop=True)

    trades = []
    last_exit_time = None

    for _, sig in signals.iterrows():
        signal_time = sig["datetime"]

        # one active trade at a time
        if last_exit_time is not None and signal_time < last_exit_time:
            continue

        entry = find_entry(
            m15=m15,
            signal_row=sig,
            entry_type=entry_type,
            wait_hours=wait_hours,
        )
        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

        trades.append(trade)
        last_exit_time = trade["exit_time"]

    trades = pd.DataFrame(trades)

    if trades.empty:
        return {
            "entry_type": entry_type,
            "trades": 0,
            "win_rate": None,
            "avg_return": None,
            "final_equity": None,
            "max_drawdown": None,
        }

    trades["equity_curve"] = (1 + trades["return"]).cumprod()
    trades["running_peak"] = trades["equity_curve"].cummax()
    trades["drawdown"] = trades["equity_curve"] / trades["running_peak"] - 1

    return {
        "entry_type": entry_type,
        "trades": len(trades),
        "win_rate": (trades["return"] > 0).mean(),
        "avg_return": trades["return"].mean(),
        "final_equity": trades["equity_curve"].iloc[-1],
        "max_drawdown": trades["drawdown"].min(),
    }


def run_test():
    entry_types = ["immediate", 0.236, 0.382, 0.5, 0.618, 0.786]
    results = []

    for entry_type in entry_types:
        res = run_backtest_for_entry_type(entry_type)
        results.append(res)

    result_df = pd.DataFrame(results)

    print("\n=== ENTRY REFINEMENT TEST V1 ===\n")
    print(result_df.to_string(index=False))

    print("\nBest by final equity:")
    print(result_df.sort_values("final_equity", ascending=False).head(1).to_string(index=False))

    print("\nBest by lowest drawdown:")
    print(result_df.sort_values("max_drawdown", ascending=False).head(1).to_string(index=False))


if __name__ == "__main__":
    run_test()