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 features.yield_spreads import build_spread_features


def build_spot_lag_features():
    prices = load_eurusd().copy()
    spreads = build_spread_features().copy()

    spread_cols = [
        "date",
        "spread_2y",
        "spread_10y",
        "spread_2y_change_1d",
        "spread_2y_change_5d",
        "spread_2y_zscore_20d",
        "spread_10y_change_1d",
        "spread_10y_change_5d",
        "spread_10y_zscore_20d",
    ]
    spreads = spreads[spread_cols].copy()

    # Keep hourly bars in EET
    prices["date"] = prices["datetime"].dt.normalize()

    # Merge daily spread state onto hourly bars
    df = prices.merge(spreads, on="date", how="left")

    # Forward fill spread state across hourly bars
    fill_cols = [c for c in spread_cols if c != "date"]
    df[fill_cols] = df[fill_cols].ffill()

    # Spot returns
    df["eurusd_return_1h"] = df["close"].pct_change()
    df["eurusd_return_4h"] = df["close"].pct_change(4)
    df["eurusd_return_24h"] = df["close"].pct_change(24)

    # Raw lag gaps
    df["lag_gap_2y_24h"] = df["spread_2y_change_1d"] - df["eurusd_return_24h"]
    df["lag_gap_10y_24h"] = df["spread_10y_change_1d"] - df["eurusd_return_24h"]

    # Composite lag score
    df["raw_lag_score"] = (
        0.7 * df["lag_gap_2y_24h"] +
        0.3 * df["lag_gap_10y_24h"]
    )

    return df


def get_model_ready_spot_lag():
    df = build_spot_lag_features().copy()

    required_cols = [
        "close",
        "spread_2y",
        "spread_10y",
        "spread_2y_change_1d",
        "spread_10y_change_1d",
        "eurusd_return_24h",
        "lag_gap_2y_24h",
        "lag_gap_10y_24h",
        "raw_lag_score",
    ]

    model_df = df.dropna(subset=required_cols).reset_index(drop=True)
    return model_df


if __name__ == "__main__":
    raw_df = build_spot_lag_features()
    model_df = get_model_ready_spot_lag()

    cols_to_show = [
        "datetime",
        "close",
        "spread_2y",
        "spread_10y",
        "spread_2y_change_1d",
        "spread_10y_change_1d",
        "eurusd_return_24h",
        "lag_gap_2y_24h",
        "lag_gap_10y_24h",
        "raw_lag_score",
    ]

    print("\nRAW DATA SAMPLE:")
    print(raw_df[cols_to_show].head(5))
    print(raw_df[cols_to_show].tail(5))

    print("\nMODEL-READY SAMPLE:")
    print(model_df[cols_to_show].head(5))
    print(model_df[cols_to_show].tail(5))

    print("\nRAW rows:", len(raw_df))
    print("MODEL rows:", len(model_df))
    print("Dropped rows:", len(raw_df) - len(model_df))

    print("\nModel date range:", model_df["datetime"].min(), "to", model_df["datetime"].max())