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 build_benchmark_trades(
    threshold: float = 2.0,
    stop: float = 0.0025,
    horizon: int = 48,
    spread_cost: float = 0.0001,
):
    """
    Frozen benchmark:
    - signal layer = existing lag_zscore_24h_v3
    - session filter = London + NY only
    - one trade at a time
    - immediate entry on benchmark timeframe
    - fixed 48h hold
    - fixed stop
    """

    df = get_model_ready_spot_lag_v3().copy()
    df = df.sort_values("datetime").reset_index(drop=True)

    # London + NY only
    df["hour"] = pd.to_datetime(df["datetime"]).dt.hour
    df = df[df["hour"].isin(range(7, 17))].copy().reset_index(drop=True)

    # forward path on the SAME benchmark dataframe
    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

    return pd.DataFrame(trades)


def analyze_trades(trades: pd.DataFrame):
    if trades.empty:
        print("No trades found.")
        return

    trades = trades.copy()
    trades["equity_curve"] = (1 + trades["return"]).cumprod()
    trades["running_peak"] = trades["equity_curve"].cummax()
    trades["drawdown"] = trades["equity_curve"] / trades["running_peak"] - 1

    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

    trades["entry_date"] = pd.to_datetime(trades["entry_time"]).dt.date
    daily_returns = trades.groupby("entry_date")["return"].sum()

    print("\n=== BENCHMARK REBUILD V1 ===")
    print("Trades:", len(trades))
    print("Win rate:", round(win_rate, 6))
    print("Average trade return:", round(avg_return, 6))
    print("Std dev of trade returns:", round(std_return, 6))
    print("Sharpe proxy:", round(sharpe_proxy, 6))
    print("Final equity multiple:", round(final_equity, 6))
    print("Max drawdown:", round(max_dd, 6))
    print("Max losing streak:", max_losing_streak)
    print("Worst day:", round(daily_returns.min(), 6))
    print("Best day:", round(daily_returns.max(), 6))

    print("\nLAST 10 TRADES:")
    print(trades.tail(10).to_string(index=False))


if __name__ == "__main__":
    trades = build_benchmark_trades(
        threshold=2.0,
        stop=0.0025,
        horizon=48,
        spread_cost=0.0001,
    )
    analyze_trades(trades)