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_hourly_signals(
    threshold: float = 2.0,
    allowed_hours: set | None = None,
):
    """
    Signal layer stays on the validated hourly model.
    """
    df = get_model_ready_spot_lag_v3().copy()
    df["hour"] = pd.to_datetime(df["datetime"]).dt.hour

    if allowed_hours is not None:
        df = df[df["hour"].isin(allowed_hours)].copy()

    df = df.sort_values("datetime").reset_index(drop=True)

    signals = []
    for _, row in df.iterrows():
        signal = 0
        if row["lag_zscore_24h_v3"] >= threshold:
            signal = 1
        elif row["lag_zscore_24h_v3"] <= -threshold:
            signal = -1

        if signal != 0:
            signals.append(
                {
                    "signal_time": row["datetime"],
                    "signal": signal,
                    "signal_price": row["close"],
                    "lag_zscore": row["lag_zscore_24h_v3"],
                }
            )

    return pd.DataFrame(signals)


def get_entry_idx_immediate(m15_window: pd.DataFrame):
    if m15_window.empty:
        return None
    return 0


def get_entry_idx_pullback(
    m15_window: pd.DataFrame,
    signal: int,
    signal_price: float,
    fib_level: float,
):
    """
    Find a pullback entry within the forward 15m window, then require a simple
    continuation confirmation after the pullback is touched.

    Long:
      - first find the highest high after signal
      - target pullback = high - fib * (high - signal_price)
      - require a later 15m close above previous 15m high

    Short:
      - first find the lowest low after signal
      - target pullback = low + fib * (signal_price - low)
      - require a later 15m close below previous 15m low
    """
    if m15_window.empty:
        return None

    if signal == 1:
        running_extreme = signal_price
        pullback_touched_idx = None

        for j in range(len(m15_window)):
            row = m15_window.iloc[j]
            running_extreme = max(running_extreme, row["high"])
            move = running_extreme - signal_price
            if move <= 0:
                continue

            target_price = running_extreme - fib_level * move

            if row["low"] <= target_price:
                pullback_touched_idx = j
                break

        if pullback_touched_idx is None:
            return None

        for j in range(pullback_touched_idx + 1, len(m15_window)):
            if m15_window.iloc[j]["close"] > m15_window.iloc[j - 1]["high"]:
                return j

        return None

    else:
        running_extreme = signal_price
        pullback_touched_idx = None

        for j in range(len(m15_window)):
            row = m15_window.iloc[j]
            running_extreme = min(running_extreme, row["low"])
            move = signal_price - running_extreme
            if move <= 0:
                continue

            target_price = running_extreme + fib_level * move

            if row["high"] >= target_price:
                pullback_touched_idx = j
                break

        if pullback_touched_idx is None:
            return None

        for j in range(pullback_touched_idx + 1, len(m15_window)):
            if m15_window.iloc[j]["close"] < m15_window.iloc[j - 1]["low"]:
                return j

        return None


def simulate_trade_from_entry(
    m15: pd.DataFrame,
    entry_time: pd.Timestamp,
    signal: int,
    stop: float,
    hold_hours: int,
    spread_cost: float,
):
    hold_bars = hold_hours * 4

    entry_idx_list = m15.index[m15["datetime"] == entry_time].tolist()
    if not entry_idx_list:
        return None

    entry_idx = entry_idx_list[0]
    entry_row = m15.iloc[entry_idx]
    entry_price = entry_row["close"]

    exit_idx = min(entry_idx + hold_bars, len(m15) - 1)
    path = m15.iloc[entry_idx: exit_idx + 1].copy()
    future_close = path.iloc[-1]["close"]

    if signal == 1:
        adverse = (path["low"].min() - entry_price) / entry_price
        stop_hit = adverse < -stop
        ret = future_close / entry_price - 1
        if stop_hit:
            ret = -stop
    else:
        adverse = (path["high"].max() - entry_price) / entry_price
        stop_hit = adverse > stop
        ret = -(future_close / entry_price - 1)
        if stop_hit:
            ret = -stop

    ret -= spread_cost

    return {
        "entry_time": entry_time,
        "entry_price": entry_price,
        "return": ret,
    }


def run_backtest_for_entry_type(
    entry_type,
    threshold: float = 2.0,
    stop: float = 0.0025,
    spread_cost: float = 0.0001,
    hold_hours: int = 48,
    wait_hours: int = 6,
):
    # Best session from prior test
    allowed_hours = set(range(7, 17))  # 07:00–16:59 EET

    hourly_signals = build_hourly_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
    wait_bars = wait_hours * 4

    for _, sig in hourly_signals.iterrows():
        signal_time = sig["signal_time"]

        # one active trade at a time
        if last_exit_time is not None and signal_time < last_exit_time:
            continue

        m15_window = m15[m15["datetime"] >= signal_time].copy().reset_index(drop=True)
        if m15_window.empty:
            continue

        m15_window = m15_window.iloc[:wait_bars].copy()
        if m15_window.empty:
            continue

        if entry_type == "immediate":
            entry_offset = get_entry_idx_immediate(m15_window)
        else:
            entry_offset = get_entry_idx_pullback(
                m15_window=m15_window,
                signal=sig["signal"],
                signal_price=sig["signal_price"],
                fib_level=float(entry_type),
            )

        if entry_offset is None:
            continue

        entry_time = m15_window.iloc[entry_offset]["datetime"]

        trade = simulate_trade_from_entry(
            m15=m15,
            entry_time=entry_time,
            signal=sig["signal"],
            stop=stop,
            hold_hours=hold_hours,
            spread_cost=spread_cost,
        )
        if trade is None:
            continue

        trade["signal_time"] = signal_time
        trade["signal"] = sig["signal"]
        trade["entry_type"] = entry_type
        trades.append(trade)

        last_exit_time = entry_time + pd.Timedelta(hours=hold_hours)

    trades = pd.DataFrame(trades)

    if trades.empty:
        return {
            "entry_type": entry_type,
            "trades": 0,
            "avg_return": None,
            "win_rate": 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),
        "avg_return": trades["return"].mean(),
        "win_rate": (trades["return"] > 0).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=== PULLBACK LADDER TEST V2 ===\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()