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_horizon(
    horizon: int,
    threshold: float = 2.0,
    stop: float = 0.0025,
    spread_cost: float = 0.0001,
):
    df = get_model_ready_spot_lag_v3().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,
            }
        )

        # one trade at a time
        i += horizon

    trades = pd.DataFrame(trades)

    if trades.empty:
        return {
            "horizon": horizon,
            "trades": 0,
            "win_rate": None,
            "avg_return": None,
            "std_return": None,
            "sharpe_proxy": None,
            "final_equity": None,
            "max_drawdown": None,
            "max_losing_streak": None,
            "worst_day": None,
            "best_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

    max_dd = trades["drawdown"].min()
    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

    # 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

    trades["entry_date"] = pd.to_datetime(trades["entry_time"]).dt.date
    daily_returns = trades.groupby("entry_date")["return"].sum()

    return {
        "horizon": horizon,
        "trades": len(trades),
        "win_rate": win_rate,
        "avg_return": avg_return,
        "std_return": std_return,
        "sharpe_proxy": sharpe_proxy,
        "final_equity": trades["equity_curve"].iloc[-1],
        "max_drawdown": max_dd,
        "max_losing_streak": max_losing_streak,
        "worst_day": daily_returns.min(),
        "best_day": daily_returns.max(),
    }


def run_test():
    horizons = [24, 36, 48, 72]
    results = []

    for h in horizons:
        res = run_backtest_for_horizon(horizon=h)
        results.append(res)

    result_df = pd.DataFrame(results)

    print("\n=== EXIT HORIZON 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()