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


def run_backtest_for_session(
    session_name: str,
    allowed_hours: set,
    threshold: float = 2.0,
    stop: float = 0.0025,
    horizon: int = 48,
    spread_cost: float = 0.0001,
):
    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["future_close"] = df["close"].shift(-horizon)
    df["future_min"] = df["close"].rolling(horizon).min().shift(-horizon)
    df["future_max"] = df["close"].rolling(horizon).max().shift(-horizon)

    df = df.dropna(
        subset=["future_close", "future_min", "future_max", "lag_zscore_24h_v3"]
    ).reset_index(drop=True)

    trades = []
    i = 0

    while i < len(df):
        row = df.iloc[i]

        signal = 0
        if row["lag_zscore_24h_v3"] >= threshold:
            signal = 1
        elif row["lag_zscore_24h_v3"] <= -threshold:
            signal = -1

        if signal == 0:
            i += 1
            continue

        entry = row["close"]
        entry_time = row["datetime"]

        if signal == 1:
            stop_hit = (row["future_min"] - entry) / entry < -stop
            ret = row["future_close"] / entry - 1
            if stop_hit:
                ret = -stop
        else:
            stop_hit = (row["future_max"] - entry) / entry > stop
            ret = -(row["future_close"] / entry - 1)
            if stop_hit:
                ret = -stop

        ret -= spread_cost

        trades.append(
            {
                "entry_time": entry_time,
                "signal": signal,
                "return": ret,
            }
        )

        i += horizon  # one trade at a time

    trades = pd.DataFrame(trades)

    if trades.empty:
        return {
            "session": session_name,
            "trades": 0,
            "win_rate": None,
            "avg_return": None,
            "final_equity": None,
            "max_drawdown": None,
            "worst_day": None,
        }

    trades["equity_curve"] = (1 + trades["return"]).cumprod()
    trades["running_peak"] = trades["equity_curve"].cummax()
    trades["drawdown"] = trades["equity_curve"] / trades["running_peak"] - 1

    trades["entry_date"] = pd.to_datetime(trades["entry_time"]).dt.date
    daily_returns = trades.groupby("entry_date")["return"].sum()

    return {
        "session": session_name,
        "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(),
        "worst_day": daily_returns.min(),
    }


def run_test():
    session_defs = {
        "all_hours": None,
        "london_only": set(range(7, 12)),         # 07:00–11:59 EET
        "london_ny": set(range(7, 17)),           # 07:00–16:59 EET
        "exclude_dead_hours": set(list(range(1, 22))),  # exclude 22,23,00
    }

    results = []

    for name, hours in session_defs.items():
        results.append(run_backtest_for_session(name, hours))

    result_df = pd.DataFrame(results)

    print("\n=== SESSION FILTER TEST V3 ===\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()