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 ingestion.price_loader import load_eurusd
from ingestion.price_loader_15m import load_eurusd_15m
from models.rolling_beta_model import build_rolling_beta_model


def build_hourly_signal_layer(
    threshold: float = 2.0,
    allowed_hours: set | None = None,
):
    """
    Clean hourly signal layer:
    - 1H price data
    - daily predicted return from rolling beta model
    - lag signal built on 1H bars
    - session filter applied ONLY to signal timestamps
    """
    h1 = load_eurusd().copy()
    daily_model = build_rolling_beta_model(window=120, smooth_span=20).copy()

    h1 = h1.sort_values("datetime").reset_index(drop=True)
    h1["date"] = h1["datetime"].dt.normalize()

    daily_cols = ["date", "predicted_return_1d"]
    daily_model = daily_model[daily_cols].copy()

    df = h1.merge(daily_model, on="date", how="left")
    df["predicted_return_1d"] = df["predicted_return_1d"].ffill()

    # actual 24h move on 1H bars
    df["eurusd_return_24h"] = df["close"].pct_change(24)

    # lag signal
    df["lag_gap_24h_v3"] = df["predicted_return_1d"] - df["eurusd_return_24h"]

    lag_mean = df["lag_gap_24h_v3"].rolling(60, min_periods=60).mean()
    lag_std = df["lag_gap_24h_v3"].rolling(60, min_periods=60).std()

    df["lag_zscore_24h_v3"] = (df["lag_gap_24h_v3"] - lag_mean) / lag_std

    df = df.dropna(subset=["lag_zscore_24h_v3"]).reset_index(drop=True)

    df["hour"] = 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)

    return signals


def get_first_m15_idx_at_or_after(m15: pd.DataFrame, ts: pd.Timestamp):
    """
    Find the first 15m bar index with datetime >= ts.
    """
    idx = m15["datetime"].searchsorted(ts, side="left")
    if idx >= len(m15):
        return None
    return int(idx)


def simulate_trade_on_m15_path(
    m15: pd.DataFrame,
    entry_idx: int,
    signal: int,
    hold_hours: int,
    stop: float,
    spread_cost: float,
):
    """
    Truth-path simulation:
    - entry on 15m close at entry_idx
    - hold for hold_hours worth of contiguous 15m bars
    - stop checked against the full 15m high/low path
    """
    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

    entry_row = path.iloc[0]
    exit_row = path.iloc[-1]

    entry_price = entry_row["close"]
    exit_price = exit_row["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_row["datetime"],
        "exit_time": exit_row["datetime"],
        "entry_price": entry_price,
        "exit_price": exit_price,
        "signal": signal,
        "return": ret,
    }


def run_backtest(
    threshold: float = 2.0,
    stop: float = 0.0025,
    hold_hours: int = 48,
    spread_cost: float = 0.0001,
):
    # Best session from prior test
    allowed_hours = set(range(7, 17))  # 07:00–16:59 EET

    signals = build_hourly_signal_layer(
        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 position at a time
        if last_exit_time is not None and signal_time < last_exit_time:
            continue

        entry_idx = get_first_m15_idx_at_or_after(m15, signal_time)
        if entry_idx is None:
            continue

        trade = simulate_trade_on_m15_path(
            m15=m15,
            entry_idx=entry_idx,
            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:
        print("No trades found.")
        return trades

    trades["equity_curve"] = (1 + trades["return"]).cumprod()
    trades["running_peak"] = trades["equity_curve"].cummax()
    trades["drawdown"] = trades["equity_curve"] / trades["running_peak"] - 1

    return trades


def analyze_results(trades: pd.DataFrame):
    if trades.empty:
        print("No trades to analyze.")
        return

    win_rate = (trades["return"] > 0).mean()
    avg_return = trades["return"].mean()
    std_return = trades["return"].std()
    sharpe_proxy = avg_return / std_return if std_return != 0 else 0
    final_equity = trades["equity_curve"].iloc[-1]
    max_dd = trades["drawdown"].min()

    # losing streak
    max_losing_streak = 0
    current_streak = 0
    for r in trades["return"]:
        if r <= 0:
            current_streak += 1
            max_losing_streak = max(max_losing_streak, current_streak)
        else:
            current_streak = 0

    # daily aggregation
    trades["entry_date"] = pd.to_datetime(trades["entry_time"]).dt.date
    daily_returns = trades.groupby("entry_date")["return"].sum()

    print("\n=== TRUTH BASELINE BACKTEST V1 ===")
    print("Trades:", len(trades))
    print("Win rate:", round(win_rate, 4))
    print("Average trade return:", round(avg_return, 6))
    print("Std dev of trade returns:", round(std_return, 6))
    print("Sharpe proxy:", round(sharpe_proxy, 4))
    print("Final equity multiple:", round(final_equity, 4))
    print("Max drawdown:", round(max_dd, 4))
    print("Max losing streak:", max_losing_streak)
    print("Worst day:", round(daily_returns.min(), 4))
    print("Best day:", round(daily_returns.max(), 4))

    print("\nLAST 10 TRADES:")
    print(
        trades[
            ["entry_time", "exit_time", "signal", "return", "equity_curve", "drawdown"]
        ].tail(10)
    )


if __name__ == "__main__":
    trades = run_backtest(
        threshold=2.0,
        stop=0.0025,
        hold_hours=48,
        spread_cost=0.0001,
    )
    analyze_results(trades)