"""
multi_pair_signals_v1.py
=========================
Builds trade signals and full trade logs for GBPUSD, USDJPY, and AUDUSD
using the same validated macro lead-lag framework as EURUSD.

Signal Logic (identical to EURUSD)
------------------------------------
  Primary driver : US_2Y - Foreign_2Y spread change
  Secondary      : US_10Y - Foreign_10Y spread change
  Rolling beta   : 120-day window, 20-day EWM smooth, 1-day lookahead shift
  Lag signal     : predicted_return_1d - actual_return_24h
  Z-score        : 60-bar rolling normalisation
  Threshold      : 2.75 (validated)

Rate data (daily frequency)
----------------------------
  UK  : BOE GLC Nominal spot curve  → uk2y_daily.csv  / uk10y_daily.csv
  JP  : MoF Japan JGB yields        → jp2y_daily.csv  / jp10y_daily.csv
  AU  : RBA Table F2                → au2y_daily.csv  / au10y_daily.csv

  All daily — processed by process_daily_rates.py

Sign convention
---------------
  Spread = US_rate - Foreign_rate (consistent for all pairs)
  Rolling beta learns the correct direction from price data:
    GBPUSD, AUDUSD: beta_2y < 0 (spread widens → price falls)
    USDJPY:         beta_2y > 0 (spread widens → price rises)

Validated parameters (frozen from EURUSD research)
----------------------------------------------------
  Hold hours  : 52
  Fib entry   : 0.786
  Stop        : 0.25%
  Session     : hours 7-16 (London + NY)
  Threshold   : 2.75

Output
------
  data/processed/trades/trades_GBPUSD.csv
  data/processed/trades/trades_USDJPY.csv
  data/processed/trades/trades_AUDUSD.csv

Place this file in:
  C:\\Users\\paul_\\OneDrive\\fx_macro_intraday\\src\\research\\multi_pair_signals_v1.py

Run from project root:
  python src/research/multi_pair_signals_v1.py
"""

import pandas as pd
import numpy as np
import statsmodels.api as sm
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.combined_candidate_matrix_v1 import (
    get_first_m15_idx_at_or_after,
    find_entry_pullback,
)

RATES_DIR  = BASE_PATH / "data" / "raw" / "rates"
PRICES_DIR = BASE_PATH / "data" / "raw" / "prices"
TRADES_DIR = BASE_PATH / "data" / "processed" / "trades"
TRADES_DIR.mkdir(parents=True, exist_ok=True)

# ── Validated frozen parameters ───────────────────────────────────────────────
HOLD_HOURS    = 52
FIB           = 0.786
STOP          = 0.0025
SPREAD_COST   = 0.0001
ALLOWED_HOURS = set(range(7, 17))
THRESHOLD     = 2.75
BETA_WINDOW   = 120
SMOOTH_SPAN   = 20
ZSCORE_WINDOW = 60
YEARS         = 22.0

# ── Pair configurations — DAILY rate files ────────────────────────────────────
# monthly: False = daily data, use 1-day diff directly (same as EURUSD)
PAIR_CONFIGS = {
    "GBPUSD": {
        "us_2y"  : ("us2y.csv",       "us2y"),
        "us_10y" : ("us10y.csv",      "us10y"),
        "for_2y" : ("uk2y_daily.csv",  "uk2y"),
        "for_10y": ("uk10y_daily.csv", "uk10y"),
        "monthly": False,
    },
    "USDJPY": {
        "us_2y"  : ("us2y.csv",       "us2y"),
        "us_10y" : ("us10y.csv",      "us10y"),
        "for_2y" : ("jp2y_daily.csv",  "jp2y"),
        "for_10y": ("jp10y_daily.csv", "jp10y"),
        "monthly": False,
    },
    "AUDUSD": {
        "us_2y"  : ("us2y.csv",       "us2y"),
        "us_10y" : ("us10y.csv",      "us10y"),
        "for_2y" : ("au2y_daily.csv",  "au2y"),
        "for_10y": ("au10y_daily.csv", "au10y"),
        "monthly": False,
    },
}


# ── Generic price loaders ─────────────────────────────────────────────────────
def _normalise_price_df(df: pd.DataFrame) -> pd.DataFrame:
    """
    Normalises column names across different broker export formats.
    Handles 'Time (EET)', trailing spaces, dot-format datetimes.
    """
    df.columns = [c.strip().lower() for c in df.columns]

    for candidate in ["time (eet)", "time", "date", "timestamp", "time(eet)"]:
        if candidate in df.columns and "datetime" not in df.columns:
            df = df.rename(columns={candidate: "datetime"})
            break

    df["datetime"] = pd.to_datetime(
        df["datetime"].astype(str).str.replace(
            r"(\d{4})\.(\d{2})\.(\d{2})", r"\1-\2-\3", regex=True
        ),
        format="mixed", dayfirst=False,
    )

    rename_map = {}
    for col in df.columns:
        if col in ["open", "high", "low", "close", "volume"]:
            continue
        if col.startswith("open"):
            rename_map[col] = "open"
        elif col.startswith("high"):
            rename_map[col] = "high"
        elif col.startswith("low"):
            rename_map[col] = "low"
        elif col.startswith("close"):
            rename_map[col] = "close"
        elif col.startswith("vol"):
            rename_map[col] = "volume"
    if rename_map:
        df = df.rename(columns=rename_map)

    return df


def load_price_1h(pair: str) -> pd.DataFrame:
    path = PRICES_DIR / pair / f"{pair}_1H_2003_2026.csv"
    df   = pd.read_csv(path)
    df   = _normalise_price_df(df)
    return df.sort_values("datetime").reset_index(drop=True)


def load_price_15m(pair: str) -> pd.DataFrame:
    path = PRICES_DIR / pair / f"{pair}_15M_2003_2026.csv"
    df   = pd.read_csv(path)
    df   = _normalise_price_df(df)
    df["datetime"] = pd.to_datetime(
        df["datetime"].astype(str).str.replace(".", "-", regex=False),
        format="mixed", dayfirst=False,
    )
    return df.sort_values("datetime").reset_index(drop=True)


# ── Rate loaders ──────────────────────────────────────────────────────────────
def load_rate_file(filename: str, col: str) -> pd.DataFrame:
    path = RATES_DIR / filename
    df   = pd.read_csv(path)
    df.columns = [c.lower() for c in df.columns]
    if "date" not in df.columns:
        df = df.rename(columns={df.columns[0]: "date"})
    if col not in df.columns:
        df = df.rename(columns={df.columns[1]: col})
    df["date"] = pd.to_datetime(df["date"], errors="coerce")
    df[col]    = pd.to_numeric(df[col], errors="coerce")
    return df[["date", col]].dropna().sort_values("date").reset_index(drop=True)


def expand_to_daily(df: pd.DataFrame, date_col: str, value_col: str,
                    start: str, end: str) -> pd.DataFrame:
    """Forward-fills any sparse rate series to a complete daily date range."""
    daily_dates = pd.DataFrame({
        date_col: pd.date_range(start=start, end=end, freq="D")
    })
    df = df.copy()
    df[date_col] = pd.to_datetime(df[date_col])
    merged = pd.merge_asof(
        daily_dates,
        df[[date_col, value_col]].sort_values(date_col),
        on=date_col,
        direction="backward",
    )
    return merged


# ── Spread builder ────────────────────────────────────────────────────────────
def build_spread_features(config: dict, date_start: str, date_end: str) -> pd.DataFrame:
    """
    Builds daily spread features for one pair.
    Spread = US_rate - foreign_rate.
    monthly=False: daily 1-day diff (same logic as EURUSD).
    monthly=True : forward-fill monthly data, use persisted diff.
    """
    us_2y_raw   = load_rate_file(*config["us_2y"])
    us_10y_raw  = load_rate_file(*config["us_10y"])
    for_2y_raw  = load_rate_file(*config["for_2y"])
    for_10y_raw = load_rate_file(*config["for_10y"])

    us_2y_col   = config["us_2y"][1]
    us_10y_col  = config["us_10y"][1]
    for_2y_col  = config["for_2y"][1]
    for_10y_col = config["for_10y"][1]

    us_2y   = expand_to_daily(us_2y_raw,   "date", us_2y_col,  date_start, date_end)
    us_10y  = expand_to_daily(us_10y_raw,  "date", us_10y_col, date_start, date_end)
    for_2y  = expand_to_daily(for_2y_raw,  "date", for_2y_col, date_start, date_end)
    for_10y = expand_to_daily(for_10y_raw, "date", for_10y_col,date_start, date_end)

    df = us_2y.merge(us_10y,  on="date", how="inner")
    df = df.merge(for_2y,     on="date", how="left")
    df = df.merge(for_10y,    on="date", how="left")
    df = df.dropna().reset_index(drop=True)

    df["spread_2y"]  = df[us_2y_col]  - df[for_2y_col]
    df["spread_10y"] = df[us_10y_col] - df[for_10y_col]

    if config["monthly"]:
        # Monthly data: forward-fill the last non-zero change
        df["spread_2y_change_1d"]  = df["spread_2y"].diff(1).replace(0, np.nan).ffill()
        df["spread_10y_change_1d"] = df["spread_10y"].diff(1).replace(0, np.nan).ffill()
    else:
        # Daily data: use 1-day change directly (identical to EURUSD)
        df["spread_2y_change_1d"]  = df["spread_2y"].diff(1)
        df["spread_10y_change_1d"] = df["spread_10y"].diff(1)

    df["spread_2y_change_5d"]  = df["spread_2y"]  - df["spread_2y"].shift(5)
    df["spread_10y_change_5d"] = df["spread_10y"] - df["spread_10y"].shift(5)

    return df[["date", "spread_2y", "spread_10y",
               "spread_2y_change_1d", "spread_10y_change_1d",
               "spread_2y_change_5d", "spread_10y_change_5d"]].copy()


# ── Daily price aggregation ───────────────────────────────────────────────────
def build_daily_prices(pair: str) -> pd.DataFrame:
    prices_1h = load_price_1h(pair)
    prices_1h["date"] = prices_1h["datetime"].dt.normalize()
    daily = (
        prices_1h
        .groupby("date", as_index=False)
        .agg(close=("close", "last"))
        .sort_values("date")
        .reset_index(drop=True)
    )
    daily["return_1d"] = daily["close"].pct_change()
    return daily


# ── Rolling beta model ────────────────────────────────────────────────────────
def build_rolling_beta(
    daily_prices : pd.DataFrame,
    spread_df    : pd.DataFrame,
    window       : int = BETA_WINDOW,
    smooth_span  : int = SMOOTH_SPAN,
) -> pd.DataFrame:
    df = daily_prices.merge(spread_df, on="date", how="inner")
    df = df.dropna(subset=["return_1d", "spread_2y_change_1d",
                            "spread_10y_change_1d"]).reset_index(drop=True)

    n        = len(df)
    b2y_raw  = np.full(n, np.nan)
    b10y_raw = np.full(n, np.nan)

    for i in range(window, n):
        sample = df.iloc[i - window:i].copy()
        sample = sample[
            (sample["spread_2y_change_1d"].abs() > 0) |
            (sample["spread_10y_change_1d"].abs() > 0)
        ]
        if len(sample) < 30:
            continue
        X = sm.add_constant(
            sample[["spread_2y_change_1d", "spread_10y_change_1d"]]
        )
        try:
            model        = sm.OLS(sample["return_1d"], X).fit()
            b2y_raw[i]   = model.params.get("spread_2y_change_1d", np.nan)
            b10y_raw[i]  = model.params.get("spread_10y_change_1d", np.nan)
        except Exception:
            pass

    df["beta_2y_raw"]  = b2y_raw
    df["beta_10y_raw"] = b10y_raw

    df["beta_2y_smooth"]  = df["beta_2y_raw"].ewm(span=smooth_span, adjust=False).mean()
    df["beta_10y_smooth"] = df["beta_10y_raw"].ewm(span=smooth_span, adjust=False).mean()

    df["beta_2y_smooth"]  = df["beta_2y_smooth"].clip(-0.10, 0.10)
    df["beta_10y_smooth"] = df["beta_10y_smooth"].clip(-0.08, 0.08)

    df["beta_2y"]  = df["beta_2y_smooth"].shift(1)
    df["beta_10y"] = df["beta_10y_smooth"].shift(1)

    df["predicted_return_1d"] = (
        df["beta_2y"]  * df["spread_2y_change_1d"] +
        df["beta_10y"] * df["spread_10y_change_1d"]
    )

    keep_cols = ["date", "beta_2y", "beta_10y", "predicted_return_1d",
                 "spread_2y_change_1d", "spread_10y_change_1d"]
    return df[keep_cols].dropna().reset_index(drop=True)


# ── Lag signal builder ────────────────────────────────────────────────────────
def build_lag_signal(
    prices_1h : pd.DataFrame,
    beta_df   : pd.DataFrame,
) -> pd.DataFrame:
    prices = prices_1h.copy()
    prices["date"] = prices["datetime"].dt.normalize()

    df = prices.merge(beta_df, on="date", how="left")

    fill_cols = ["beta_2y", "beta_10y", "predicted_return_1d",
                 "spread_2y_change_1d", "spread_10y_change_1d"]
    df[fill_cols] = df[fill_cols].ffill()

    df["return_24h"]  = df["close"].pct_change(24)
    df["lag_gap_24h"] = df["predicted_return_1d"] - df["return_24h"]

    lag_mean = df["lag_gap_24h"].rolling(ZSCORE_WINDOW, min_periods=ZSCORE_WINDOW).mean()
    lag_std  = df["lag_gap_24h"].rolling(ZSCORE_WINDOW, min_periods=ZSCORE_WINDOW).std()
    df["lag_zscore_24h"] = (df["lag_gap_24h"] - lag_mean) / lag_std

    df["signal_direction"] = 0
    df.loc[df["lag_zscore_24h"] >  1.0, "signal_direction"] =  1
    df.loc[df["lag_zscore_24h"] < -1.0, "signal_direction"] = -1

    required = ["close", "predicted_return_1d", "return_24h",
                "lag_gap_24h", "lag_zscore_24h"]
    return df.dropna(subset=required).reset_index(drop=True)


# ── Trade simulator ───────────────────────────────────────────────────────────
def simulate_trade_full(
    m15         : pd.DataFrame,
    entry_time  : pd.Timestamp,
    entry_price : float,
    signal      : int,
    hold_hours  : int = HOLD_HOURS,
    stop        : float = STOP,
    spread_cost : float = SPREAD_COST,
) -> dict | None:
    entry_idx = get_first_m15_idx_at_or_after(m15, entry_time)
    if entry_idx is None:
        return None

    hold_bars = hold_hours * 4
    exit_idx  = min(entry_idx + hold_bars, len(m15) - 1)
    path      = m15.iloc[entry_idx : exit_idx + 1].copy()
    if path.empty:
        return None

    exit_price = float(path.iloc[-1]["close"])

    if signal == 1:
        mae      = (path["low"].min()  - entry_price) / entry_price
        mfe      = (path["high"].max() - entry_price) / entry_price
        stop_hit = mae < -stop
        raw_ret  = -stop if stop_hit else (exit_price / entry_price - 1)
    else:
        mae      = -((path["high"].max() - entry_price) / entry_price)
        mfe      = (entry_price - path["low"].min()) / entry_price
        stop_hit = (path["high"].max() - entry_price) / entry_price > stop
        raw_ret  = -stop if stop_hit else -(exit_price / entry_price - 1)

    return {
        "entry_time" : entry_time,
        "exit_time"  : path.iloc[-1]["datetime"],
        "signal"     : signal,
        "entry_price": entry_price,
        "exit_price" : exit_price,
        "stop_hit"   : stop_hit,
        "mae"        : mae,
        "mfe"        : mfe,
        "return"     : raw_ret - spread_cost,
    }


# ── Main pair runner ──────────────────────────────────────────────────────────
def run_pair(pair: str, config: dict) -> pd.DataFrame | None:
    print(f"\n{'─'*65}")
    print(f"PAIR: {pair}")
    print(f"{'─'*65}")

    for key in ["for_2y", "for_10y"]:
        rate_path = RATES_DIR / config[key][0]
        if not rate_path.exists():
            print(f"  [ERROR] Missing rate file: {rate_path}")
            print(f"  Run process_daily_rates.py first.")
            return None

    print("  Loading prices...")
    try:
        prices_1h = load_price_1h(pair)
        m15       = load_price_15m(pair)
    except Exception as e:
        print(f"  [ERROR] Failed to load price data: {e}")
        return None

    m15 = m15.sort_values("datetime").reset_index(drop=True)
    print(f"  1H bars : {len(prices_1h):,}  "
          f"({prices_1h['datetime'].min().date()} to {prices_1h['datetime'].max().date()})")
    print(f"  15M bars: {len(m15):,}")

    print("  Building spread features...")
    date_start = str(prices_1h["datetime"].min().date())
    date_end   = str(prices_1h["datetime"].max().date())
    try:
        spread_df = build_spread_features(config, date_start, date_end)
    except Exception as e:
        print(f"  [ERROR] Spread build failed: {e}")
        return None
    print(f"  Spread rows: {len(spread_df):,}")

    print("  Building rolling beta model...")
    daily_px = build_daily_prices(pair)
    try:
        beta_df = build_rolling_beta(daily_px, spread_df)
    except Exception as e:
        print(f"  [ERROR] Rolling beta failed: {e}")
        return None
    print(f"  Beta rows: {len(beta_df):,}  "
          f"(beta_2y median: {beta_df['beta_2y'].median():.4f})")

    print("  Building lag signal...")
    try:
        signal_df = build_lag_signal(prices_1h, beta_df)
    except Exception as e:
        print(f"  [ERROR] Lag signal failed: {e}")
        return None
    print(f"  Signal rows: {len(signal_df):,}")

    signal_df["hour"] = signal_df["datetime"].dt.hour
    signal_df = signal_df[signal_df["hour"].isin(ALLOWED_HOURS)].copy()

    signal_df["signal"] = 0
    signal_df.loc[signal_df["lag_zscore_24h"] >=  THRESHOLD, "signal"] =  1
    signal_df.loc[signal_df["lag_zscore_24h"] <= -THRESHOLD, "signal"] = -1
    signals = signal_df[signal_df["signal"] != 0].copy().reset_index(drop=True)
    print(f"  Signals above threshold {THRESHOLD}: {len(signals)}")

    print("  Running trade simulation...")
    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_pullback(
            m15=m15, signal_row=sig, fib=FIB, wait_hours=6,
        )
        if entry is None:
            continue

        trade = simulate_trade_full(
            m15        =m15,
            entry_time =entry["entry_time"],
            entry_price=entry["entry_price"],
            signal     =int(sig["signal"]),
        )
        if trade is None:
            continue

        trade["zscore_abs"] = abs(float(sig["lag_zscore_24h"]))
        trades.append(trade)
        last_exit_time = trade["exit_time"]

    if not trades:
        print(f"  [WARNING] No trades generated for {pair}")
        return None

    t             = pd.DataFrame(trades)
    t["equity"]   = (1 + t["return"]).cumprod()
    t["peak"]     = t["equity"].cummax()
    t["drawdown"] = t["equity"] / t["peak"] - 1
    t["win"]      = (t["return"] > 0).astype(int)
    t["trade_num"]= range(1, len(t) + 1)
    t["candidate"]= pair

    streak = max_streak = 0
    for r in t["return"]:
        if r <= 0:
            streak    += 1
            max_streak = max(max_streak, streak)
        else:
            streak = 0

    n_per_yr = len(t) / YEARS
    print(f"\n  Results for {pair}:")
    print(f"    Trades          : {len(t)}  ({n_per_yr:.1f}/year)")
    print(f"    Win rate        : {t['win'].mean():.4%}")
    print(f"    Avg return      : {t['return'].mean():.6f}")
    print(f"    Final equity    : {t['equity'].iloc[-1]:.6f}")
    print(f"    Max drawdown    : {t['drawdown'].min():.4%}")
    print(f"    Max lose streak : {max_streak}")
    print(f"    Stop hit rate   : {t['stop_hit'].mean():.4%}")
    print(f"    Avg MAE         : {t['mae'].mean()*100:.3f}%")
    print(f"    Avg MFE         : {t['mfe'].mean()*100:.3f}%")

    out_path = TRADES_DIR / f"trades_{pair}.csv"
    t.to_csv(out_path, index=False)
    print(f"\n  Saved: {out_path}")

    return t


# ── Main ──────────────────────────────────────────────────────────────────────
def main():
    print("=" * 65)
    print("MULTI PAIR SIGNALS V1")
    print(f"  Threshold: {THRESHOLD}  |  Hold: {HOLD_HOURS}h  |  "
          f"Fib: {FIB}  |  Stop: {STOP:.2%}")
    print(f"  Rate data: DAILY frequency (BOE / MoF Japan / RBA)")
    print("=" * 65)

    results = {}
    for pair, config in PAIR_CONFIGS.items():
        result = run_pair(pair, config)
        if result is not None:
            results[pair] = result

    print(f"\n{'='*65}")
    print("CROSS-PAIR SUMMARY")
    print(f"{'='*65}")

    print(f"\n  {'Pair':<12}{'Trades':>8}{'Per/yr':>8}{'WR':>8}"
          f"{'AvgRet':>10}{'FinalEq':>10}{'MaxDD':>9}")
    print(f"  {'─'*65}")

    total_trades = 0
    for pair, t in results.items():
        n = len(t)
        total_trades += n
        print(f"  {pair:<12}{n:>8}{n/YEARS:>8.1f}"
              f"{t['win'].mean():>7.2%}"
              f"{t['return'].mean():>10.6f}"
              f"{t['equity'].iloc[-1]:>10.4f}"
              f"{t['drawdown'].min():>8.2%}")

    if results:
        print(f"\n  {'EURUSD (ref)':<12}{'1408':>8}{'64.0':>8}"
              f"{'30.89%':>8}{'0.000740':>10}{'2.7625':>10}{'-7.05%':>9}")
        print(f"\n  Total trades across all new pairs: {total_trades}")
        print(f"  Total trades per year (new pairs) : {total_trades/YEARS:.1f}")
        print(f"  Combined with EURUSD              : {(total_trades+1408)/YEARS:.1f}/year")

    print(f"\n  Next step: run portfolio_ftmo_v1.py")


if __name__ == "__main__":
    main()
