#!/usr/bin/env python3
"""
Scan for **Johansen cointegration triplets** (Chan Ex 2.7–2.8) across ETF themes.

Unlike the S&P 500 stock scan, this focuses on **economically linked ETFs**
(country, commodity/producer, sector) where cointegration tends to persist.

For each triplet passing train-window filters, reports **out-of-sample** stats
on ``[--start, --end]`` using causal Johansen linear MR.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_johansen_triplet_scan.py \\
      --train-end 2016-01-04 --start 2016-01-04 --end 2026-04-02 \\
      --list-top 40
"""

from __future__ import annotations

import argparse
import itertools
import json
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any

import numpy as np
import pandas as pd

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

import universe_scanner as us  # type: ignore[import-not-found]

from RenTech.strategy_stack.chan_johansen_triplet import (
    johansen_linear_triplet,
    johansen_trace_pvalue,
    unit_portfolio_half_life,
)

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT = LOGS / "johansen_triplet_scan"

# Thematic ETF baskets (Chan-style: common macro exposure within group)
ETF_GROUPS: dict[str, list[str]] = {
    "country_dm_americas": ["EWC", "EWA", "EWZ", "SPY", "QQQ"],
    "country_dm_europe_asia": ["EWJ", "EWG", "EWU", "EFA", "EEM", "FXI", "INDA"],
    "country_commodity_export": ["EWA", "EWC", "EWZ", "FXI", "RIO", "VALE", "BHP"],
    "precious_metals": ["GLD", "SLV", "GDX", "GDXJ", "SIL", "IAU", "AGQ"],
    "energy_complex": ["USO", "XLE", "XOP", "OIH", "XOM", "CVX", "COP"],
    "commodity_broad": ["DBC", "GSG", "PDBC", "GLD", "USO", "UNG"],
    "macro_aw": ["SPY", "TLT", "IEF", "GLD", "DBC", "UUP"],
    "sector_spdr": ["XLK", "XLF", "XLE", "XLV", "XLI", "XLP", "XLY", "XLU", "XLB", "XLRE", "XLC"],
    "real_estate": ["VNQ", "IYR", "SCHH", "RWR", "O", "AMT", "PLD"],
    "bonds_credit": ["TLT", "IEF", "SHY", "LQD", "HYG", "TIP", "AGG"],
    "chan_classics": ["EWA", "EWC", "IGE", "GLD", "GDX", "USO", "RTH", "XLP"],
}

# Flat deduped universe
ALL_ETFS = sorted({t for grp in ETF_GROUPS.values() for t in grp})


@dataclass(frozen=True)
class TripletHit:
    tickers: tuple[str, str, str]
    group: str
    trace_stat: float
    crit_95: float
    half_life: float
    score: float
    weights: tuple[float, float, float]
    oos_return_pct: float
    oos_sharpe: float
    oos_max_dd_pct: float


def _close_panel(data: dict[str, pd.DataFrame]) -> pd.DataFrame:
    cols: dict[str, pd.Series] = {}
    for t, df in data.items():
        if df is None or df.empty:
            continue
        c = "close" if "close" in df.columns else "Close"
        if c not in df.columns:
            continue
        s = df[c].astype(np.float64)
        s.index = pd.to_datetime(s.index).tz_localize(None)
        cols[t] = s
    return pd.DataFrame(cols).sort_index()


def _oos_metrics(r: pd.Series, *, capital: float = 100_000.0) -> dict[str, float]:
    r = r.fillna(0.0).astype(np.float64)
    if len(r) < 20:
        return {"return_pct": float("nan"), "sharpe": float("nan"), "max_dd_pct": float("nan")}
    eq = capital * (1.0 + r).cumprod()
    tot = float(eq.iloc[-1] / capital - 1.0) * 100.0
    sd = float(r.std(ddof=1))
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")
    mdd = float((eq / eq.cummax() - 1.0).min()) * 100.0
    return {"return_pct": tot, "sharpe": sharpe, "max_dd_pct": mdd}


def _eval_triplet(
    close: pd.DataFrame,
    legs: tuple[str, str, str],
    *,
    train_mask: pd.Series,
    oos_start: pd.Timestamp,
    oos_end: pd.Timestamp | None,
    half_life_min: float,
    half_life_max: float,
) -> TripletHit | None:
    px_train = close.loc[train_mask, list(legs)].dropna(how="any")
    if len(px_train) < 200:
        return None
    try:
        trace, crit95, evec = johansen_trace_pvalue(px_train)
    except Exception:
        return None
    if trace <= crit95:
        return None
    hl = unit_portfolio_half_life(px_train, evec)
    if not np.isfinite(hl) or hl < half_life_min or hl > half_life_max:
        return None

    px_full = close[list(legs)].dropna(how="any")
    try:
        r, _ = johansen_linear_triplet(px_full, fidelity="causal", min_train=252)
    except Exception:
        return None
    r.index = pd.to_datetime(r.index).tz_localize(None)
    mask = r.index >= oos_start
    if oos_end is not None:
        mask &= r.index <= oos_end
    oos = r.loc[mask].fillna(0.0)
    om = _oos_metrics(oos)
    score = float(trace / max(crit95, 1e-9) - 0.05 * hl + 0.1 * (om["sharpe"] if np.isfinite(om["sharpe"]) else 0))
    return TripletHit(
        tickers=legs,
        group="",
        trace_stat=trace,
        crit_95=crit95,
        half_life=float(hl),
        score=score,
        weights=(float(evec[0]), float(evec[1]), float(evec[2])),
        oos_return_pct=float(om["return_pct"]),
        oos_sharpe=float(om["sharpe"]),
        oos_max_dd_pct=float(om["max_dd_pct"]),
    )


def scan_group(
    close: pd.DataFrame,
    tickers: list[str],
    group: str,
    *,
    train_mask: pd.Series,
    oos_start: pd.Timestamp,
    oos_end: pd.Timestamp | None,
    half_life_min: float,
    half_life_max: float,
) -> list[TripletHit]:
    avail = [t for t in tickers if t in close.columns]
    if len(avail) < 3:
        return []
    hits: list[TripletHit] = []
    for legs in itertools.combinations(sorted(avail), 3):
        hit = _eval_triplet(
            close,
            legs,
            train_mask=train_mask,
            oos_start=oos_start,
            oos_end=oos_end,
            half_life_min=half_life_min,
            half_life_max=half_life_max,
        )
        if hit is not None:
            hits.append(
                TripletHit(
                    tickers=hit.tickers,
                    group=group,
                    trace_stat=hit.trace_stat,
                    crit_95=hit.crit_95,
                    half_life=hit.half_life,
                    score=hit.score,
                    weights=hit.weights,
                    oos_return_pct=hit.oos_return_pct,
                    oos_sharpe=hit.oos_sharpe,
                    oos_max_dd_pct=hit.oos_max_dd_pct,
                )
            )
    return hits


def _dedupe_hits(hits: list[TripletHit]) -> list[TripletHit]:
    """Keep best score per sorted ticker tuple."""
    best: dict[tuple[str, ...], TripletHit] = {}
    for h in hits:
        key = tuple(sorted(h.tickers))
        if key not in best or h.score > best[key].score:
            best[key] = h
    return list(best.values())


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--train-end", default="2016-01-04", help="Last day of Johansen training window")
    ap.add_argument("--train-days", type=int, default=504, help="~2y training lookback")
    ap.add_argument("--start", default="2016-01-04", help="OOS backtest start")
    ap.add_argument("--end", default="2026-04-02")
    ap.add_argument("--half-life-min", type=float, default=1.0)
    ap.add_argument("--half-life-max", type=float, default=60.0)
    ap.add_argument("--lookback-years", type=float, default=15.0)
    ap.add_argument("--list-top", type=int, default=50)
    ap.add_argument("--min-oos-sharpe", type=float, default=0.0, help="Filter OOS Sharpe")
    ap.add_argument("--out", type=Path, default=DEFAULT_OUT)
    ap.add_argument("--refresh-cache", action="store_true")
    ap.add_argument("--tickers", default="", help="Comma list override (single group)")
    args = ap.parse_args()

    train_end = pd.Timestamp(args.train_end).normalize()
    train_start = train_end - pd.Timedelta(days=int(args.train_days * 1.6))
    oos_start = pd.Timestamp(args.start).normalize()
    oos_end = pd.Timestamp(args.end).normalize() if str(args.end).strip() else None

    if str(args.tickers).strip():
        tickers = [t.strip().upper() for t in str(args.tickers).split(",") if t.strip()]
        groups = {"custom": tickers}
    else:
        groups = ETF_GROUPS
        tickers = sorted({t for grp in groups.values() for t in grp})

    print(f"Johansen triplet ETF scan  train={train_start.date()}→{train_end.date()}  OOS from {oos_start.date()}", flush=True)
    print(f"  universe: {len(tickers)} ETFs  groups={len(groups)}", flush=True)

    raw = us.download_and_cache_data(
        tickers,
        timeframe="1d",
        lookback_years=float(args.lookback_years),
        max_age_hours=0.0 if args.refresh_cache else 24.0,
    )
    close = _close_panel(raw)
    train_mask = (close.index >= train_start) & (close.index <= train_end)
    print(f"  panel: {close.shape[1]} tickers × {len(close)} days", flush=True)

    all_hits: list[TripletHit] = []
    for gname, gtickers in groups.items():
        found = scan_group(
            close,
            gtickers,
            gname,
            train_mask=train_mask,
            oos_start=oos_start,
            oos_end=oos_end,
            half_life_min=float(args.half_life_min),
            half_life_max=float(args.half_life_max),
        )
        print(f"  {gname:28s}  {len(found):4d} triplets", flush=True)
        all_hits.extend(found)

    hits = _dedupe_hits(all_hits)
    hits = [h for h in hits if np.isfinite(h.oos_sharpe) and h.oos_sharpe >= float(args.min_oos_sharpe)]
    hits.sort(key=lambda h: (h.oos_sharpe, h.score), reverse=True)

    print(f"\n=== Top {min(args.list_top, len(hits))} triplets (deduped, sorted by OOS Sharpe) ===", flush=True)
    rows: list[dict[str, Any]] = []
    for i, h in enumerate(hits[: int(args.list_top)], 1):
        leg_s = "-".join(h.tickers)
        w_s = ",".join(f"{w:.3f}" for w in h.weights)
        print(
            f"{i:3d}. {leg_s:22s}  grp={h.group:22s}  "
            f"trace={h.trace_stat:5.1f}/{h.crit_95:4.1f}  hl={h.half_life:4.1f}d  "
            f"OOS ret={h.oos_return_pct:6.1f}%  Sharpe={h.oos_sharpe:5.2f}  DD={h.oos_max_dd_pct:5.1f}%  "
            f"w=[{w_s}]",
            flush=True,
        )
        rows.append(
            {
                "rank": i,
                "ticker_a": h.tickers[0],
                "ticker_b": h.tickers[1],
                "ticker_c": h.tickers[2],
                "triplet": leg_s,
                "group": h.group,
                "trace_stat": h.trace_stat,
                "crit_95": h.crit_95,
                "half_life_days": h.half_life,
                "weight_a": h.weights[0],
                "weight_b": h.weights[1],
                "weight_c": h.weights[2],
                "oos_return_pct": h.oos_return_pct,
                "oos_sharpe": h.oos_sharpe,
                "oos_max_dd_pct": h.oos_max_dd_pct,
                "score": h.score,
            }
        )

    out = args.out.expanduser().resolve()
    out.parent.mkdir(parents=True, exist_ok=True)
    csv_path = Path(f"{out}_ranked.csv")
    pd.DataFrame(rows).to_csv(csv_path, index=False)
    meta_path = Path(f"{out}_meta.json")
    meta_path.write_text(
        json.dumps(
            {
                "command": " ".join(sys.argv),
                "train_window": [str(train_start.date()), str(train_end.date())],
                "oos_window": [str(oos_start.date()), str(oos_end.date()) if oos_end else None],
                "n_candidates": len(hits),
                "groups": list(groups.keys()),
            },
            indent=2,
        )
    )
    print(f"\nWrote {csv_path}  ({len(rows)} rows)", flush=True)
    print(f"Wrote {meta_path}", flush=True)


if __name__ == "__main__":
    main()
