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_test():
    df = get_model_ready_spot_lag_v3().copy()

    horizon = 24
    stop = 0.0025
    spread_cost = 0.0001  # ~1 pip

    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"]).copy()

    thresholds = [1.0, 1.25, 1.5, 1.75, 2.0, 2.5]

    for threshold in thresholds:
        subset = df[df["lag_zscore_24h_v3"].abs() >= threshold].copy()

        trades = []
        prev_signal = 0

        for _, row in subset.iterrows():
            signal = 0
            if row["lag_zscore_24h_v3"] >= threshold:
                signal = 1
            elif row["lag_zscore_24h_v3"] <= -threshold:
                signal = -1

            if signal != 0 and signal != prev_signal:
                entry = row["close"]

                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(ret)

            prev_signal = signal

        trades = pd.Series(trades, dtype=float)

        if len(trades) == 0:
            print(f"\nThreshold {threshold}: no trades")
            continue

        print(f"\nThreshold: {threshold}")
        print("Trades:", len(trades))
        print("Average return:", trades.mean())
        print("Win rate:", (trades > 0).mean())
        print("Total return:", trades.sum())


if __name__ == "__main__":
    run_test()