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 research.entry_refinement_threshold_test_v1 import run_backtest_for_threshold


def build_current_benchmark_trade_log():
    """
    Rebuild the current truthful benchmark:
    threshold = 2.75
    0.786 pullback entry
    London + NY
    48h hold
    0.25% stop
    """
    # Reuse the threshold test logic, but we need the actual trade log version.
    # So import and mirror the same settings from the current benchmark result source.
    from research.entry_refinement_threshold_test_v1 import (
        build_frozen_signals,
        find_entry_0786,
        simulate_trade,
    )
    from ingestion.price_loader_15m import load_eurusd_15m

    allowed_hours = set(range(7, 17))
    threshold = 2.75
    hold_hours = 48
    stop = 0.0025
    spread_cost = 0.0001
    wait_hours = 6

    signals = build_frozen_signals(
        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"]

        if last_exit_time is not None and signal_time < last_exit_time:
            continue

        entry = find_entry_0786(
            m15=m15,
            signal_row=sig,
            wait_hours=wait_hours,
        )
        if entry is None:
            continue

        trade = simulate_trade(
            m15=m15,
            entry_time=entry["entry_time"],
            entry_price=entry["entry_price"],
            signal=int(sig["signal"]),
            hold_hours=hold_hours,
            stop=stop,
            spread_cost=spread_cost,
        )
        if trade is None:
            continue

        trade["signal_time"] = sig["datetime"]
        trade["lag_zscore"] = sig["lag_zscore_24h_v3"]
        trade["trade_date"] = pd.to_datetime(trade["entry_time"]).date()

        trades.append(trade)
        last_exit_time = trade["exit_time"]

    return pd.DataFrame(trades)


def run_oi_filter_test():
    trades = build_current_benchmark_trade_log()

    oi_path = BASE_PATH / "data" / "raw" / "oi" / "eurusd_oi_labels.csv"
    oi = pd.read_csv(oi_path)

    oi["date"] = pd.to_datetime(oi["date"]).dt.date

    df = trades.merge(
        oi,
        left_on="trade_date",
        right_on="date",
        how="left"
    )

    print("\nTotal trades in benchmark:", len(df))
    print("Trades with OI label:", df["oi_label"].notna().sum())

    results = []

    for label_set_name, allowed_labels in {
        "all_labeled": ["supportive", "neutral", "hostile"],
        "supportive_only": ["supportive"],
        "supportive_neutral": ["supportive", "neutral"],
        "exclude_hostile": ["supportive", "neutral"],
        "hostile_only": ["hostile"],
    }.items():

        subset = df[df["oi_label"].isin(allowed_labels)].copy()

        if subset.empty:
            continue

        subset["equity_curve"] = (1 + subset["return"]).cumprod()
        subset["running_peak"] = subset["equity_curve"].cummax()
        subset["drawdown"] = subset["equity_curve"] / subset["running_peak"] - 1

        results.append({
            "filter": label_set_name,
            "trades": len(subset),
            "win_rate": (subset["return"] > 0).mean(),
            "avg_return": subset["return"].mean(),
            "final_equity": subset["equity_curve"].iloc[-1],
            "max_drawdown": subset["drawdown"].min(),
        })

    result_df = pd.DataFrame(results)

    print("\n=== OI FILTER TEST V1 ===\n")
    if result_df.empty:
        print("No matching labeled trades found.")
    else:
        print(result_df.to_string(index=False))


if __name__ == "__main__":
    run_oi_filter_test()