#!/usr/bin/env python3
"""
**CNN-LSTM SPY Direction Signal** — walk-forward fund-scale overlay.

Trains a CNN-LSTM classifier to predict next-day SPY direction (up/down)
using a strict walk-forward protocol (no lookahead):

  - Training window : 2 years (~504 trading sessions)
  - Retraining      : every quarter (~63 sessions)
  - Lookback        : 20 sessions of lagged features per sample
  - Target          : 1 if next-day SPY return ≥ 0, else 0

Features (all strictly lagged to avoid lookahead):
  ``spy_ret_1d``, ``spy_ret_5d``, ``spy_ret_20d``, ``spy_rsi_14``,
  ``spy_bb_pos``, ``spy_sma50_ratio``, ``vix_level``, ``vix_ret_5d``,
  ``tlt_ret_1d``

Predicted probability → suggested fund scale:
  - prob ≥ 0.55 → 2.0×  (bullish confidence)
  - 0.45 ≤ prob < 0.55 → 1.5×  (neutral)
  - prob < 0.45 → 1.0×  (bearish)

The dynamic scale replaces the fixed ``--fund-scale`` on the stock-only book.
The ``--apply-to-daily`` option reads an existing combine daily CSV, re-applies
the dynamic scale, and prints a before/after yearly comparison.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_cnn_lstm_spy_signal.py \\
      --start 2016-01-04 --end 2025-12-31 \\
      --out-prefix RenTech/data/logs/cnn_lstm_spy_signal

    # Apply to baseline stock-only combine output:
    .venv/bin/python RenTech/strategy_stack/run_cnn_lstm_spy_signal.py \\
      --apply-to-daily RenTech/data/logs/stock_only_baseline_plus_stock_only_plus_sp500_dip_plus_tactical_aw_plus_tsmom_plus_johansen_etf_plus_vol_edge_plus_fund_plus_nav_q_mtm_daily.csv \\
      --out-prefix RenTech/data/logs/cnn_lstm_spy_signal

Outputs (``--out-prefix``):
  ``*_signal.csv``   — date, prob_up, fund_scale
  ``*_metrics.json`` — walk-forward accuracy, scale distribution, yearly impact
"""
from __future__ import annotations

import argparse
import json
import math
import os
import sys
import warnings
from pathlib import Path

import numpy as np
import pandas as pd
import yfinance as yf

warnings.filterwarnings("ignore", category=FutureWarning)
os.environ.setdefault("TF_CPP_MIN_LOG_LEVEL", "2")

_REPO = Path(__file__).resolve().parents[2]
if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT_PREFIX = LOGS / "cnn_lstm_spy_signal"

TRAIN_WINDOW = 504    # ~2 trading years
RETRAIN_EVERY = 63   # ~1 quarter
SEQ_LEN = 20         # look-back window per sample
SCALE_HIGH = 2.0     # scale when prob_up ≥ 0.55
SCALE_MID  = 1.5     # scale when 0.45 ≤ prob_up < 0.55
SCALE_LOW  = 1.0     # scale when prob_up < 0.45


# ─────────────────────────── feature engineering ─────────────────────────────

def _rsi(series: pd.Series, period: int = 14) -> pd.Series:
    delta = series.diff()
    gain = delta.clip(lower=0).rolling(period).mean()
    loss = (-delta.clip(upper=0)).rolling(period).mean()
    rs = gain / loss.replace(0, np.nan)
    return 100 - 100 / (1 + rs)


def _bb_position(series: pd.Series, window: int = 20) -> pd.Series:
    ma = series.rolling(window).mean()
    sd = series.rolling(window).std()
    lower = ma - 2 * sd
    upper = ma + 2 * sd
    rng = upper - lower
    pos = (series - lower) / rng.replace(0, np.nan)
    return pos.clip(0, 1)


def build_features(spy: pd.Series, vix: pd.Series, tlt: pd.Series) -> pd.DataFrame:
    """Return aligned feature DataFrame (all values are lagged 1 day — no lookahead)."""
    ret1  = spy.pct_change(1)
    ret5  = spy.pct_change(5)
    ret20 = spy.pct_change(20)
    rsi   = _rsi(spy, 14)
    bb    = _bb_position(spy, 20)
    sma50 = spy.rolling(50).mean()
    sma50_ratio = (spy / sma50 - 1).replace([np.inf, -np.inf], np.nan)

    vix_norm   = (vix - vix.rolling(252).mean()) / vix.rolling(252).std()
    vix_ret5   = vix.pct_change(5)
    tlt_ret1   = tlt.pct_change(1)

    feat = pd.DataFrame({
        "spy_ret_1d":     ret1,
        "spy_ret_5d":     ret5,
        "spy_ret_20d":    ret20,
        "spy_rsi_14":     rsi,
        "spy_bb_pos":     bb,
        "spy_sma50_ratio": sma50_ratio,
        "vix_norm":        vix_norm,
        "vix_ret_5d":      vix_ret5,
        "tlt_ret_1d":      tlt_ret1,
    }, index=spy.index)

    # Lag all features by 1 day so today's prediction uses only yesterday's data
    feat = feat.shift(1)
    return feat.dropna()


# ─────────────────────────── model construction ───────────────────────────────

def build_model(n_features: int, seq_len: int = SEQ_LEN, fast: bool = False):
    """Build CNN-LSTM binary classifier using Keras 3."""
    import keras
    from keras import layers

    filters   = 16 if fast else 32
    lstm_units = 16 if fast else 32
    dense_units = 8 if fast else 16

    model = keras.Sequential([
        # Local pattern extraction
        layers.Conv1D(filters, kernel_size=3, activation="relu", padding="same",
                      input_shape=(seq_len, n_features)),
        layers.MaxPooling1D(pool_size=2),
        # Temporal dependencies
        layers.LSTM(lstm_units, dropout=0.2, recurrent_dropout=0.1),
        # Classification head
        layers.Dense(dense_units, activation="relu"),
        layers.Dropout(0.2),
        layers.Dense(1, activation="sigmoid"),
    ])
    model.compile(
        optimizer=keras.optimizers.Adam(learning_rate=1e-3),
        loss="binary_crossentropy",
        metrics=["accuracy"],
    )
    return model


# ─────────────────────────── sequence builder ─────────────────────────────────

def make_sequences(X: np.ndarray, y: np.ndarray, seq_len: int):
    """Slide a window of length seq_len over X to create (samples, seq_len, features)."""
    xs, ys = [], []
    for i in range(seq_len, len(X)):
        xs.append(X[i - seq_len : i])
        ys.append(y[i])
    return np.array(xs, dtype=np.float32), np.array(ys, dtype=np.float32)


# ─────────────────────────── walk-forward engine ─────────────────────────────

def run_walk_forward(
    feat: pd.DataFrame,
    target: pd.Series,
    start: str,
    end: str,
    *,
    train_window: int = TRAIN_WINDOW,
    retrain_every: int = RETRAIN_EVERY,
    seq_len: int = SEQ_LEN,
    fast: bool = False,
    verbose: bool = True,
) -> pd.DataFrame:
    """
    Strict walk-forward: at each retraining point, fit on the past
    ``train_window`` sessions only, then predict the next ``retrain_every``
    sessions.  Returns a DataFrame with ``prob_up`` and ``fund_scale`` indexed
    by date.
    """
    import keras

    t0 = pd.Timestamp(start)
    t1 = pd.Timestamp(end)

    all_dates = feat.index
    backtest_idx = all_dates[(all_dates >= t0) & (all_dates <= t1)]

    results: dict[pd.Timestamp, dict] = {}
    total_correct = 0
    total_pred    = 0
    n_features    = feat.shape[1]

    epochs = 30 if fast else 100
    patience = 5 if fast else 15

    feat_arr   = feat.values.astype(np.float32)
    target_arr = target.reindex(feat.index).values.astype(np.float32)

    retraining_dates = []
    i = 0
    while i < len(backtest_idx):
        dt = backtest_idx[i]
        pos = all_dates.get_loc(dt)
        if pos < train_window + seq_len:
            i += 1
            continue
        retraining_dates.append((i, dt, pos))
        i += retrain_every

    n_events = len(retraining_dates)
    if verbose:
        print(f"Walk-forward: {n_events} retraining events, "
              f"train_window={train_window}, retrain_every={retrain_every}")

    from sklearn.preprocessing import StandardScaler

    for ev_idx, (bt_i, dt, pos) in enumerate(retraining_dates):
        train_start = pos - train_window
        train_end   = pos

        X_train_raw = feat_arr[train_start : train_end]
        y_train     = target_arr[train_start : train_end]

        # Normalize using only training data
        scaler = StandardScaler()
        X_train_norm = scaler.fit_transform(X_train_raw)

        X_tr, y_tr = make_sequences(X_train_norm, y_train, seq_len)
        if len(X_tr) < 50:
            i += 1
            continue

        # Validation: last 15% of training set
        val_split = max(int(len(X_tr) * 0.85), 50)
        X_val, y_val = X_tr[val_split:], y_tr[val_split:]
        X_tr2, y_tr2 = X_tr[:val_split], y_tr[:val_split]

        # Build fresh model each retraining
        model = build_model(n_features, seq_len, fast=fast)

        es = keras.callbacks.EarlyStopping(
            monitor="val_loss", patience=patience, restore_best_weights=True, verbose=0
        )

        model.fit(
            X_tr2, y_tr2,
            validation_data=(X_val, y_val),
            epochs=epochs,
            batch_size=32,
            callbacks=[es],
            verbose=0,
        )

        # Predict the next retrain_every days
        pred_end = min(bt_i + retrain_every, len(backtest_idx))
        pred_dates = backtest_idx[bt_i : pred_end]

        for j, pred_dt in enumerate(pred_dates):
            pred_pos = all_dates.get_loc(pred_dt)
            if pred_pos < seq_len:
                continue
            X_raw = feat_arr[pred_pos - seq_len : pred_pos]
            X_norm = scaler.transform(X_raw)
            X_seq = X_norm[np.newaxis, :, :]          # (1, seq_len, n_features)
            prob = float(model.predict(X_seq, verbose=0)[0, 0])

            actual = int(target_arr[pred_pos]) if pred_pos < len(target_arr) else None
            if actual is not None:
                total_correct += int((prob >= 0.5) == actual)
                total_pred += 1

            scale = SCALE_HIGH if prob >= 0.55 else (SCALE_LOW if prob < 0.45 else SCALE_MID)
            results[pred_dt] = {"prob_up": prob, "fund_scale": scale, "actual_up": actual}

        if verbose:
            acc_so_far = total_correct / max(total_pred, 1) * 100
            print(f"  [{ev_idx+1}/{n_events}] {dt.date()} → acc {acc_so_far:.1f}%  "
                  f"(n_pred={total_pred})", flush=True)

        keras.backend.clear_session()

    sig_df = pd.DataFrame(results).T
    sig_df.index.name = "date"
    sig_df = sig_df.sort_index()

    if verbose:
        final_acc = total_correct / max(total_pred, 1) * 100
        scale_counts = sig_df["fund_scale"].value_counts().to_dict()
        print(f"\n=== Walk-forward complete ===")
        print(f"Accuracy (next-day direction): {final_acc:.1f}%  "
              f"(n={total_pred}, random baseline=50%)")
        print(f"Scale distribution: {scale_counts}")

    return sig_df


# ─────────────────────────── dynamic scale application ────────────────────────

def apply_dynamic_scale(
    combine_daily_csv: Path,
    signal_df: pd.DataFrame,
    base_scale: float = 1.5,
    capital: float = 100_000.0,
    verbose: bool = True,
) -> pd.DataFrame:
    """
    Re-apply dynamic fund scale to an existing combine daily CSV.

    The combine CSV was produced at a fixed scale (base_scale). We rescale
    the daily returns by (dynamic_scale / base_scale), then recompute equity.

    Column heuristic: reads ``daily_return_mtm`` if present, else the first
    column containing ``return`` in its name.
    """
    daily = pd.read_csv(combine_daily_csv, parse_dates=["date"]).set_index("date")
    daily.index = pd.to_datetime(daily.index).tz_localize(None)

    # Find the daily return column
    ret_col = None
    for col in daily.columns:
        if "daily_return" in col.lower() or col.lower() == "daily_return_mtm":
            ret_col = col
            break
    if ret_col is None:
        raise ValueError(f"No daily return column found in {combine_daily_csv}")

    sig_idx = signal_df.index.tz_localize(None) if signal_df.index.tz else signal_df.index
    scale_series = signal_df["fund_scale"].copy()
    scale_series.index = sig_idx
    scale_series = scale_series.reindex(daily.index).ffill().fillna(base_scale)

    base_ret = daily[ret_col].astype(float)
    # Rescale: multiply daily return by (dynamic / base)
    dynamic_ret = base_ret * (scale_series / base_scale)

    eq_base    = capital * (1 + base_ret).cumprod()
    eq_dynamic = capital * (1 + dynamic_ret).cumprod()

    result = pd.DataFrame({
        "date":          daily.index,
        "ret_base":      base_ret.values,
        "ret_dynamic":   dynamic_ret.values,
        "eq_base":       eq_base.values,
        "eq_dynamic":    eq_dynamic.values,
        "fund_scale":    scale_series.values,
    }).set_index("date")

    if verbose:
        def _stats(eq: pd.Series, ret: pd.Series, label: str):
            n = len(ret)
            yrs = n / 252
            cagr = (float(eq.iloc[-1]) / capital) ** (1 / yrs) - 1
            dd   = float((eq / eq.cummax() - 1).min()) * 100
            sh   = float(ret.mean() / ret.std(ddof=1) * math.sqrt(252))
            print(f"  {label}: CAGR {cagr*100:.1f}%  Sharpe {sh:.2f}  "
                  f"MaxDD {dd:.1f}%  End ${float(eq.iloc[-1]):,.0f}")

        print(f"\n=== Dynamic scale applied to {combine_daily_csv.name} ===")
        _stats(eq_base, base_ret, "Base (fixed scale)")
        _stats(eq_dynamic, dynamic_ret, "Dynamic (CNN-LSTM)")

        print("\nYearly comparison (base vs dynamic):")
        hdr = f"{'year':>6}  {'base':>7}  {'dynamic':>8}  {'delta':>6}"
        print(hdr)
        eq_cur_b, eq_cur_d = capital, capital
        for yr, g in base_ret.groupby(base_ret.index.year):
            ret_b = float((1 + g).prod() - 1) * 100
            ret_d_yr = dynamic_ret.loc[dynamic_ret.index.year == yr]
            ret_d = float((1 + ret_d_yr).prod() - 1) * 100
            delta = ret_d - ret_b
            flag = " ✓" if ret_d >= 10 and ret_b < 10 else (
                   " ✗" if ret_d < 10 and ret_b >= 10 else ""
            )
            print(f"  {yr}  {ret_b:+6.1f}%  {ret_d:+7.1f}%  {delta:+5.1f}pp{flag}")

    return result


# ─────────────────────────── data download ────────────────────────────────────

def _download(start_fetch: str, end_fetch: str) -> tuple[pd.Series, pd.Series, pd.Series]:
    raw = yf.download(
        ["SPY", "^VIX", "TLT"],
        start=start_fetch, end=end_fetch,
        auto_adjust=True, progress=False,
    )
    if isinstance(raw.columns, pd.MultiIndex):
        close = raw["Close"]
    else:
        close = raw[["Close"]]
    close.index = pd.to_datetime(close.index).tz_localize(None)
    spy = close["SPY"].dropna()
    vix = close["^VIX"].dropna()
    tlt = close["TLT"].dropna()
    common = spy.index.intersection(vix.index).intersection(tlt.index)
    return spy.reindex(common), vix.reindex(common), tlt.reindex(common)


# ─────────────────────────── main ─────────────────────────────────────────────

def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2016-01-04",
                    help="First backtest session (default 2016-01-04)")
    ap.add_argument("--end", default="2025-12-31",
                    help="Last backtest session (default 2025-12-31)")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX,
                    help="Output file prefix")
    ap.add_argument("--fast", action="store_true",
                    help="Faster smoke test: smaller model, fewer epochs")
    ap.add_argument(
        "--apply-to-daily", type=Path, default=None,
        help="Existing combine daily CSV to apply dynamic scale to. "
             "If omitted, only the signal CSV is written.",
    )
    ap.add_argument("--base-scale", type=float, default=1.5,
                    help="Fixed scale used to generate the combine CSV (default 1.5)")
    ap.add_argument("--capital", type=float, default=100_000.0,
                    help="Starting capital for before/after comparison (default 100000)")
    args = ap.parse_args()

    t0 = pd.Timestamp(args.start)
    t1 = pd.Timestamp(args.end)
    fetch_start = (t0 - pd.DateOffset(years=3)).strftime("%Y-%m-%d")
    fetch_end   = t1.strftime("%Y-%m-%d")

    print(f"Downloading SPY / VIX / TLT from {fetch_start} → {fetch_end} …", flush=True)
    spy, vix, tlt = _download(fetch_start, fetch_end)

    print("Engineering features …", flush=True)
    feat   = build_features(spy, vix, tlt)
    target = (spy.pct_change(1) >= 0).astype(int).reindex(feat.index).dropna()
    feat   = feat.reindex(target.index).dropna()
    target = target.reindex(feat.index)

    print(f"Feature matrix: {feat.shape}  (features: {list(feat.columns)})", flush=True)

    sig_df = run_walk_forward(
        feat, target,
        start=args.start,
        end=args.end,
        fast=args.fast,
        verbose=True,
    )

    out_prefix = Path(args.out_prefix)
    out_prefix.parent.mkdir(parents=True, exist_ok=True)
    signal_path = Path(f"{out_prefix}_signal.csv")
    sig_out = sig_df.reset_index()
    sig_out["date"] = sig_out["date"].dt.strftime("%Y-%m-%d")
    sig_out.to_csv(signal_path, index=False)
    print(f"\nSignal CSV → {signal_path}", flush=True)

    acc = float((sig_df["prob_up"] >= 0.5).eq(sig_df["actual_up"]).mean()) * 100
    scale_dist = sig_df["fund_scale"].value_counts().to_dict()
    meta = {
        "start": str(t0.date()),
        "end": str(t1.date()),
        "n_sessions": len(sig_df),
        "accuracy_pct": round(acc, 2),
        "scale_distribution": {str(k): int(v) for k, v in scale_dist.items()},
        "scale_mean": round(float(sig_df["fund_scale"].mean()), 3),
        "signal_csv": str(signal_path),
    }

    if args.apply_to_daily and args.apply_to_daily.is_file():
        result_df = apply_dynamic_scale(
            args.apply_to_daily,
            sig_df,
            base_scale=args.base_scale,
            capital=args.capital,
            verbose=True,
        )
        result_path = Path(f"{out_prefix}_applied_daily.csv")
        result_df.reset_index().to_csv(result_path, index=False)
        print(f"Applied daily CSV → {result_path}", flush=True)
        meta["applied_to"] = str(args.apply_to_daily)
        meta["applied_daily_csv"] = str(result_path)

    metrics_path = Path(f"{out_prefix}_metrics.json")
    with open(metrics_path, "w") as fh:
        json.dump(meta, fh, indent=2)
    print(f"Metrics → {metrics_path}", flush=True)


if __name__ == "__main__":
    main()
