#!/usr/bin/env python3
"""
Per-calendar-year returns for each ``build_catalog_100()`` spec (sequential engine).

Use this to find sleeves that **compound in most years** vs. one-off regime hits:
  - ``pos_year_frac`` = fraction of years with strictly positive intra-year equity change
  - ``min_year_ret_pct`` / ``std_year_ret_pct`` — dispersion of yearly outcomes

One Theta prepare, then 100 × :func:`run_signal_backtest`.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
      RenTech/strategy_stack/yearly_lit_catalog_consistency.py \\
      --start 2016-01-04 --end 2026-04-02 --capital 100000 \\
      --out-csv RenTech/data/logs/lit_catalog_yearly_consistency_2016_2026.csv
"""
from __future__ import annotations

import argparse
import csv
import math
import statistics
import sys
from pathlib import Path

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

import pandas as pd

from RenTech.core.theta_chunks_loader import theta_chunks_date_bounds
from RenTech.strategy_stack import research_literature_theta_strategies as L
from RenTech.strategy_stack.literature_search_agent import _compile_signal, _compile_trade
from RenTech.strategy_stack.literature_strategy_catalog import build_catalog_100


def _yearly_returns(eq: pd.Series) -> dict[int, float]:
    """First → last observation inside each calendar year (trading index only)."""
    out: dict[int, float] = {}
    for y in sorted(eq.index.year.unique()):
        sub = eq[eq.index.year == y]
        if len(sub) < 8:
            continue
        a, b = float(sub.iloc[0]), float(sub.iloc[-1])
        if a <= 0 or not math.isfinite(a) or not math.isfinite(b):
            continue
        out[y] = (b / a - 1.0) * 100.0
    return out


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("--theta-dir", type=Path, default=_REPO / "RenTech/data/theta_chunks")
    ap.add_argument("--start", type=str, default="2016-01-04")
    ap.add_argument("--end", type=str, default="")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument(
        "--out-csv",
        type=Path,
        default=_REPO / "RenTech/data/logs/lit_catalog_yearly_consistency_default.csv",
    )
    ap.add_argument("--top", type=int, default=20)
    args = ap.parse_args()

    end = str(args.end).strip()
    if not end:
        _, d1 = theta_chunks_date_bounds(args.theta_dir.expanduser().resolve())
        end = d1.strftime("%Y-%m-%d")
        print(f"--end omitted → {end}", flush=True)

    cap = float(args.capital)
    print(f"Preparing Theta {args.start}..{end} …", flush=True)
    days, panel, get_chain, iv_atm, skew_put_minus_call_iv, n_contracts, spy_wide = (
        L.prepare_theta_research_context(
            theta_dir=args.theta_dir,
            capital=cap,
            start=str(args.start),
            end=end,
            max_days=0,
        )
    )
    print(f"  sessions={len(days)}  capital=${cap:,.0f}", flush=True)

    rows: list[dict[str, object]] = []
    for spec in build_catalog_100():
        sig = _compile_signal(spec, panel, iv_atm, skew_put_minus_call_iv, n_contracts, spy_wide)
        tfn = _compile_trade(spec)
        ex, pnls, ntr = L.run_signal_backtest(
            days, get_chain, panel, sig, int(spec.hold), tfn, spec.trade_params
        )
        eq, sh = L.equity_curve_from_realized(ex, pnls, days, cap)
        yr = _yearly_returns(eq)
        vals = list(yr.values())
        n_y = len(vals)
        n_pos = sum(1 for v in vals if v > 0)
        pos_frac = n_pos / n_y if n_y else float("nan")
        mn = statistics.mean(vals) if vals else float("nan")
        std = statistics.stdev(vals) if len(vals) > 1 else 0.0
        mn_yr = min(vals) if vals else float("nan")
        mx_yr = max(vals) if vals else float("nan")

        rows.append(
            {
                "sid": spec.sid,
                "family": spec.family,
                "description": spec.description,
                "hold": int(spec.hold),
                "trades": int(ntr),
                "sharpe": float(sh) if math.isfinite(sh) else float("nan"),
                "n_years": n_y,
                "n_pos_years": n_pos,
                "pos_year_frac": round(pos_frac, 4) if math.isfinite(pos_frac) else float("nan"),
                "mean_year_ret_pct": round(mn, 3) if math.isfinite(mn) else float("nan"),
                "std_year_ret_pct": round(std, 3),
                "min_year_ret_pct": round(mn_yr, 3) if math.isfinite(mn_yr) else float("nan"),
                "max_year_ret_pct": round(mx_yr, 3) if math.isfinite(mx_yr) else float("nan"),
            }
        )

    out = args.out_csv.expanduser().resolve()
    out.parent.mkdir(parents=True, exist_ok=True)
    keys = list(rows[0].keys()) if rows else []
    with out.open("w", newline="", encoding="utf-8") as f:
        w = csv.DictWriter(f, fieldnames=keys)
        w.writeheader()
        w.writerows(rows)
    print(f"Wrote {out}", flush=True)

    top_n = int(args.top)
    # Rank: must have enough years and trades; prefer high pos_year_frac then high min_year (floor)
    cand = [
        r
        for r in rows
        if int(r["n_years"]) >= 5
        and int(r["trades"]) >= 40
        and math.isfinite(float(r["pos_year_frac"]))
    ]
    by_consistency = sorted(
        cand,
        key=lambda r: (float(r["pos_year_frac"]), float(r["min_year_ret_pct"]), float(r["sharpe"])),
        reverse=True,
    )[:top_n]

    print(f"\nTop {top_n} by (pos_year_frac, then min_year_ret_pct, then Sharpe); "
          f"n_years>=5 trades>=40; pool n={len(cand)}", flush=True)
    for r in by_consistency:
        shv = float(r["sharpe"])
        shs = f"{shv:.3f}" if math.isfinite(shv) else "nan"
        print(
            f"  {float(r['pos_year_frac']):.2f}  minYr={r['min_year_ret_pct']:>7.2f}%  "
            f"stdYr={r['std_year_ret_pct']:>6.2f}%  Sharpe={shs:>7}  "
            f"trades={int(r['trades']):4d}  {r['sid']}  {str(r['family']):16s}  {str(r['description'])[:48]}",
            flush=True,
        )


if __name__ == "__main__":
    main()
