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 ingestion.price_loader import load_eurusd
from models.rolling_beta_model import build_rolling_beta_model


def build_spot_lag_v3():
    prices = load_eurusd().copy()
    daily_model = build_rolling_beta_model(window=120, smooth_span=20).copy()

    # Keep EET session structure
    prices["date"] = prices["datetime"].dt.normalize()

    # Only keep daily fields we need
    daily_cols = [
        "date",
        "spread_2y_change_1d",
        "spread_10y_change_1d",
        "beta_2y",
        "beta_10y",
        "predicted_return_1d",
    ]
    daily_model = daily_model[daily_cols].copy()

    # Merge daily predicted move onto hourly bars
    df = prices.merge(daily_model, on="date", how="left")

    # Forward fill daily state across the intraday bars
    fill_cols = [c for c in daily_cols if c != "date"]
    df[fill_cols] = df[fill_cols].ffill()

    # Actual spot move over the last 24h
    df["eurusd_return_24h"] = df["close"].pct_change(24)

    # Core lag signal
    df["lag_gap_24h_v3"] = df["predicted_return_1d"] - df["eurusd_return_24h"]

    # Normalize lag signal
    lag_mean = df["lag_gap_24h_v3"].rolling(60, min_periods=60).mean()
    lag_std = df["lag_gap_24h_v3"].rolling(60, min_periods=60).std()

    df["lag_zscore_24h_v3"] = (df["lag_gap_24h_v3"] - lag_mean) / lag_std

    # Simple directional state
    df["signal_direction_v3"] = 0
    df.loc[df["lag_zscore_24h_v3"] > 1.0, "signal_direction_v3"] = 1
    df.loc[df["lag_zscore_24h_v3"] < -1.0, "signal_direction_v3"] = -1

    return df


def get_model_ready_spot_lag_v3():
    df = build_spot_lag_v3().copy()

    required_cols = [
        "close",
        "predicted_return_1d",
        "eurusd_return_24h",
        "lag_gap_24h_v3",
        "lag_zscore_24h_v3",
    ]

    model_df = df.dropna(subset=required_cols).reset_index(drop=True)
    return model_df


if __name__ == "__main__":
    df = get_model_ready_spot_lag_v3()

    cols = [
        "datetime",
        "close",
        "predicted_return_1d",
        "eurusd_return_24h",
        "lag_gap_24h_v3",
        "lag_zscore_24h_v3",
        "signal_direction_v3",
    ]

    print(df[cols].head(10))
    print(df[cols].tail(10))

    print("\nRows:", len(df))
    print("Date range:", df["datetime"].min(), "to", df["datetime"].max())

    print("\nSignal counts:")
    print(df["signal_direction_v3"].value_counts(dropna=False).sort_index())