#!/usr/bin/env python3
"""
bt_optimizer.py — Pillar-weight optimizer for RoaringKittyTracker

Random search over pillar weights (value / health / sentiment / quality) and
portfolio size (top-N) to maximize risk-adjusted returns on the in-sample
period, then validates the best configurations on the out-of-sample holdout.

Performance
───────────
The key optimization: build_signals_panel() is called ONCE before the trial
loop, not once per trial.  Each trial then re-runs only build_scores() with
different weights, which takes ~1–2 s.  750 trials → ~12–15 minutes total.

Without this: ~6 min × 750 = ~75 hours.  With this: ~15 minutes.

Approach
────────
  1. Load scores.csv + weekly prices as a close panel (Date × Ticker)
  2. Precompute the signals panel for all dates × all tickers (~10–15 s)
  3. Dirichlet-sample N random weight combos (always sum to 100)
  4. For each combo × top-N value: run_backtest() using the cached signals panel
  5. Rank all trials by in-sample Sharpe ratio
  6. Validate top-K on out-of-sample period
  7. Write results to data/backtest/optimizer_*.csv

Usage
─────
  python bt_optimizer.py                          # 150 random + 5 top-N = 750 trials
  python bt_optimizer.py --trials 300             # more thorough
  python bt_optimizer.py --top_n_values "10,20,30" --trials 200
  python bt_optimizer.py --update_config          # write best weights to config.yaml
"""

from __future__ import annotations

import argparse
import contextlib
import io
import sys
import time
import warnings
from pathlib import Path
from typing import Any, Dict, List, Optional

import numpy as np
import pandas as pd
import yaml

warnings.filterwarnings("ignore", category=FutureWarning)

from bt_engine import (
    DEFAULT_WEIGHTS,
    DEFAULT_GUARDRAILS,
    DEFAULT_TC_BPS,
    BENCHMARK,
    load_scores,
    load_prices_panel,
    load_quarterly_cache,
    build_signals_panel,
    build_fundamentals_panel,
    run_backtest,
    compute_metrics,
    print_metrics,
    _build_rebalance_dates,
)

# backward-compat alias
load_weekly_prices = load_prices_panel


# ── Defaults ──────────────────────────────────────────────────────────────────
DEFAULT_TRIALS       = 150
DEFAULT_TOP_N_VALUES = [10, 15, 20, 25, 30]
VALIDATE_TOP_K       = 5


# ── Weight generation ─────────────────────────────────────────────────────────

def sample_weights(
    n: int,
    pillar_names: List[str],
    seed: Optional[int] = None,
) -> List[Dict[str, float]]:
    """
    Sample n weight combinations via Dirichlet (alpha=2 per pillar) so they
    always sum to 100.  Baseline DEFAULT_WEIGHTS is always first.
    """
    rng    = np.random.default_rng(seed)
    k      = len(pillar_names)
    alpha  = np.full(k, 2.0)
    combos = [dict(DEFAULT_WEIGHTS)]  # baseline first

    raw    = rng.dirichlet(alpha, size=n - 1)
    scaled = (raw * 100).round(1)
    for row in scaled:
        row[-1] += 100 - row.sum()   # fix rounding residual
        combos.append({p: float(w) for p, w in zip(pillar_names, row)})

    return combos


# ── Optimizer ─────────────────────────────────────────────────────────────────

def run_optimizer(
    raw_features: pd.DataFrame,
    close_panel:  pd.DataFrame,
    signals_panel: pd.DataFrame,
    train_start: str,
    train_end:   str,
    n_trials:    int,
    top_n_values: List[int],
    seed: int = 42,
    fundamentals_panel: Optional[pd.DataFrame] = None,
    tc_bps: float = DEFAULT_TC_BPS,
    strict_pit: bool = False,
) -> pd.DataFrame:
    """
    Run random-search optimisation using precomputed panels.

    Both signals_panel and fundamentals_panel (if provided) are reused across
    all trials — only build_scores() re-runs per trial (~0.4s each).

    Returns a DataFrame of all trial results sorted by in-sample Sharpe.
    """
    pillars = ["value", "health", "sentiment", "quality"]
    combos  = sample_weights(n_trials, pillars, seed=seed)
    total   = len(combos) * len(top_n_values)

    pit_label = "[PIT fundamentals]" if fundamentals_panel is not None else "[static fundamentals]"
    print(f"\nOptimizer: {len(combos)} weight combos × {len(top_n_values)} top-N "
          f"= {total} trials  {pit_label}")
    print(f"  Train period : {train_start} → {train_end}\n")

    all_rows  = []
    trial_num = 0
    t_global  = time.time()

    for top_n in top_n_values:
        for weights in combos:
            trial_num += 1

            if trial_num % 25 == 0 or trial_num == 1:
                elapsed = time.time() - t_global
                rate    = trial_num / max(elapsed, 1)
                eta_s   = (total - trial_num) / rate if rate > 0 else 0
                print(f"  Trial {trial_num:>4}/{total}  "
                      f"[{elapsed/60:.1f} min elapsed, ETA ≈ {eta_s/60:.1f} min]")

            buf = io.StringIO()
            with contextlib.redirect_stdout(buf):
                try:
                    results = run_backtest(
                        raw_features        = raw_features,
                        prices_dict         = close_panel,
                        start_date          = train_start,
                        end_date            = train_end,
                        top_n               = top_n,
                        weights             = weights,
                        guardrails          = DEFAULT_GUARDRAILS,
                        signals_panel       = signals_panel,
                        fundamentals_panel  = fundamentals_panel,
                        tc_bps              = tc_bps,
                        strict_pit          = strict_pit,
                    )
                    m = compute_metrics(results)
                except Exception as exc:
                    m = {}

            if not m or m.get("n_quarters", 0) == 0:
                continue

            all_rows.append({
                "top_n":          top_n,
                "w_value":        weights["value"],
                "w_health":       weights["health"],
                "w_sentiment":    weights["sentiment"],
                "w_quality":      weights["quality"],
                "train_sharpe":   m.get("sharpe",         np.nan),
                "train_alpha_ann": m.get("alpha_ann",     np.nan),
                "train_cagr":     m.get("port_cagr",      np.nan),
                "train_max_dd":   m.get("max_drawdown",   np.nan),
                "train_win_rate": m.get("win_rate_vs_spy", np.nan),
                "train_ir":       m.get("info_ratio",     np.nan),
                "n_quarters":     m.get("n_quarters",     0),
            })

    elapsed_total = time.time() - t_global
    print(f"\n  Search complete — {trial_num} trials in {elapsed_total/60:.1f} min "
          f"({elapsed_total/trial_num:.2f}s/trial)")

    df = pd.DataFrame(all_rows)
    if df.empty:
        return df

    df = df.sort_values("train_sharpe", ascending=False).reset_index(drop=True)
    df["rank"] = df.index + 1
    return df


# ── Out-of-sample validation ───────────────────────────────────────────────────

def validate_top_k(
    top_results:   pd.DataFrame,
    raw_features:  pd.DataFrame,
    close_panel:   pd.DataFrame,
    signals_panel: pd.DataFrame,
    eval_start:    str,
    eval_end:      str,
    k:             int = VALIDATE_TOP_K,
    fundamentals_panel: Optional[pd.DataFrame] = None,
    tc_bps: float = DEFAULT_TC_BPS,
    strict_pit: bool = False,
) -> pd.DataFrame:
    """Run the top-K in-sample configurations on the out-of-sample test period."""
    rows = []
    print(f"\nValidating top-{k} configs out-of-sample "
          f"({eval_start} → {eval_end})...\n")

    for _, cfg in top_results.head(k).iterrows():
        weights = {
            "value":     float(cfg["w_value"]),
            "health":    float(cfg["w_health"]),
            "sentiment": float(cfg["w_sentiment"]),
            "quality":   float(cfg["w_quality"]),
        }
        top_n = int(cfg["top_n"])

        buf = io.StringIO()
        with contextlib.redirect_stdout(buf):
            try:
                results = run_backtest(
                    raw_features        = raw_features,
                    prices_dict         = close_panel,
                    start_date          = eval_start,
                    end_date            = eval_end,
                    top_n               = top_n,
                    weights             = weights,
                    guardrails          = DEFAULT_GUARDRAILS,
                    signals_panel       = signals_panel,
                    fundamentals_panel  = fundamentals_panel,
                    tc_bps              = tc_bps,
                    strict_pit          = strict_pit,
                )
                m = compute_metrics(results)
            except Exception:
                m = {}

        row = {
            "train_rank":      int(cfg["rank"]),
            "top_n":           top_n,
            "w_value":         weights["value"],
            "w_health":        weights["health"],
            "w_sentiment":     weights["sentiment"],
            "w_quality":       weights["quality"],
            "train_sharpe":    float(cfg["train_sharpe"]),
            "train_alpha_ann": float(cfg["train_alpha_ann"]),
            "train_cagr":      float(cfg["train_cagr"]),
            "test_sharpe":     m.get("sharpe",          np.nan),
            "test_alpha_ann":  m.get("alpha_ann",        np.nan),
            "test_cagr":       m.get("port_cagr",        np.nan),
            "test_max_dd":     m.get("max_drawdown",     np.nan),
            "test_win_rate":   m.get("win_rate_vs_spy",  np.nan),
            "test_quarters":   m.get("n_quarters",       0),
        }
        rows.append(row)

        print(f"  Rank #{int(cfg['rank'])}  top-{top_n}  "
              f"V={weights['value']:.0f}% H={weights['health']:.0f}% "
              f"S={weights['sentiment']:.0f}% Q={weights['quality']:.0f}%")
        print(f"    In-sample   Sharpe={float(cfg['train_sharpe']):>5.2f}  "
              f"alpha={float(cfg['train_alpha_ann']):>+6.1%}  "
              f"CAGR={float(cfg['train_cagr']):>+6.1%}")
        print(f"    Out-sample  Sharpe={row['test_sharpe']:>5.2f}  "
              f"alpha={row['test_alpha_ann']:>+6.1%}  "
              f"CAGR={row['test_cagr']:>+6.1%}\n")

    return pd.DataFrame(rows)


def update_config(best: pd.Series, config_path: str = "config.yaml") -> None:
    """Write the best weight configuration back into config.yaml."""
    path = Path(config_path)
    cfg  = yaml.safe_load(path.read_text(encoding="utf-8")) or {}

    new_weights = {
        "value":              round(float(best["w_value"]),     1),
        "financial_health":   round(float(best["w_health"]),    1),
        "sentiment_crowding": round(float(best["w_sentiment"]), 1),
        "quality_momentum":   round(float(best["w_quality"]),   1),
    }
    cfg["scoring_weights"]         = new_weights
    cfg["_optimized_top_n"]        = int(best["top_n"])
    cfg["_optimized_train_sharpe"] = round(float(best["train_sharpe"]), 3)
    cfg["_optimized_test_sharpe"]  = round(float(best.get("test_sharpe", float("nan"))), 3)

    path.write_text(yaml.dump(cfg, default_flow_style=False, sort_keys=False), encoding="utf-8")
    print(f"\n  config.yaml updated with optimized weights: {new_weights}")


# ── CLI ────────────────────────────────────────────────────────────────────────

def main() -> None:
    ap = argparse.ArgumentParser(description="Optimize scoring weights via random search")
    ap.add_argument("--scores",         default="data/outputs/scores.csv")
    ap.add_argument("--weekly_dir",     default="data/weekly")
    ap.add_argument("--quarterly_dir",  default="data/quarterly",
                    help="Dir of quarterly JSON files; enables PIT mode if present")
    ap.add_argument("--out_dir",        default="data/backtest")
    ap.add_argument("--train_start",    default="2015-01-01")
    ap.add_argument("--train_end",      default="2021-12-31")
    ap.add_argument("--eval_start",     default="2022-01-01")
    ap.add_argument("--eval_end",       default=None)
    ap.add_argument("--trials",         type=int, default=DEFAULT_TRIALS)
    ap.add_argument("--top_n_values",   default=None,
                    help="Comma-separated list, e.g. '10,15,20,25,30'")
    ap.add_argument("--validate_k",     type=int, default=VALIDATE_TOP_K)
    ap.add_argument("--update_config",  action="store_true",
                    help="Write best weights to config.yaml")
    ap.add_argument("--no_pit",         action="store_true",
                    help="Disable PIT fundamentals even if quarterly_dir exists")
    ap.add_argument("--tc_bps",         type=float, default=DEFAULT_TC_BPS,
                    help=f"One-way transaction cost in bps (default {DEFAULT_TC_BPS:.0f}; 0 = frictionless)")
    ap.add_argument("--strict_pit",     action="store_true",
                    help="Exclude tickers without PIT coverage (no static fallback)")
    ap.add_argument("--seed",           type=int, default=42)
    args = ap.parse_args()

    out_dir  = Path(args.out_dir)
    out_dir.mkdir(parents=True, exist_ok=True)
    eval_end = args.eval_end or pd.Timestamp.today().strftime("%Y-%m-%d")
    top_n_values = (
        [int(x) for x in args.top_n_values.split(",")]
        if args.top_n_values
        else DEFAULT_TOP_N_VALUES
    )

    # ── Load scores ────────────────────────────────────────────────────────────
    scores_path = Path(args.scores)
    if not scores_path.exists():
        sys.exit(f"Scores file not found: {scores_path}\nRun rk_tracker.py first.")

    print(f"\nLoading scores  ({scores_path})...")
    raw = load_scores(scores_path)
    print(f"  {len(raw):,} tickers")

    # ── Load weekly prices ────────────────────────────────────────────────────
    weekly_dir = Path(args.weekly_dir)
    if not weekly_dir.exists():
        sys.exit(f"Weekly price dir not found: {weekly_dir}\n"
                 "Run download_weekly_history.py first.")

    print(f"Loading weekly prices ({weekly_dir}/)...")
    t0 = time.time()
    close_panel = load_prices_panel(weekly_dir, list(raw.index))

    spy_path = weekly_dir / f"{BENCHMARK}.csv"
    if BENCHMARK not in close_panel.columns and spy_path.exists():
        spy_df = pd.read_csv(spy_path, parse_dates=["Date"], index_col="Date")
        close_panel[BENCHMARK] = spy_df["Close"].dropna()
        close_panel = close_panel.sort_index()

    print(f"  {close_panel.shape[1]:,} tickers  {close_panel.shape[0]:,} weeks  "
          f"({time.time()-t0:.1f}s)")

    if BENCHMARK not in close_panel.columns:
        sys.exit(f"SPY benchmark not in panel. Check {spy_path}")

    # ── Precompute signals panel ONCE ─────────────────────────────────────────
    all_dates = _build_rebalance_dates(args.train_start, eval_end)
    print(f"\nPrecomputing signals panel for {len(all_dates)} rebalance dates × "
          f"{close_panel.shape[1]:,} tickers ...")
    t0 = time.time()
    signals = build_signals_panel(close_panel, all_dates)
    print(f"  Done in {time.time()-t0:.1f}s  —  {len(signals):,} (date, ticker) rows")

    # ── Precompute fundamentals panel ONCE (PIT mode) ─────────────────────────
    fund_panel: Optional[pd.DataFrame] = None
    quarterly_dir = Path(args.quarterly_dir)
    if not args.no_pit and quarterly_dir.exists():
        n_json = len(list(quarterly_dir.glob("*.json")))
        if n_json > 0:
            print(f"\nLoading quarterly cache ({n_json:,} tickers)...")
            t0 = time.time()
            q_cache = load_quarterly_cache(quarterly_dir, list(raw.index))
            print(f"  {len(q_cache):,} tickers loaded ({time.time()-t0:.1f}s)")

            print(f"Building PIT fundamentals panel (2–4 min)...")
            t0 = time.time()
            fund_panel = build_fundamentals_panel(
                tickers=list(raw.index),
                quarterly_cache=q_cache,
                close_panel=close_panel,
                rebalance_dates=all_dates,
                static_features=raw,
            )
            print(f"  Done in {time.time()-t0:.1f}s  —  "
                  f"{len(fund_panel):,} rows  [PIT mode ON]")
        else:
            print(f"\nNo quarterly JSON in {quarterly_dir}/ — static fundamentals.")
    else:
        if not args.no_pit:
            print(f"\nNo quarterly dir ({quarterly_dir}/) — static fundamentals.")

    print()

    # ── Run optimizer (both panels reused across all trials) ──────────────────
    all_results = run_optimizer(
        raw_features        = raw,
        close_panel         = close_panel,
        signals_panel       = signals,
        train_start         = args.train_start,
        train_end           = args.train_end,
        n_trials            = args.trials,
        top_n_values        = top_n_values,
        seed                = args.seed,
        fundamentals_panel  = fund_panel,
        tc_bps              = args.tc_bps,
        strict_pit          = args.strict_pit,
    )

    if all_results.empty:
        print("No valid results — check that scores.csv and weekly data are present.")
        return

    all_results.to_csv(out_dir / "optimizer_search.csv", index=False)

    # ── Print top-10 in-sample ─────────────────────────────────────────────────
    print(f"\nTop-10 by in-sample Sharpe:")
    print(f"{'Rank':>5}  {'N':>4}  {'Value':>6}  {'Health':>7}  "
          f"{'Senti':>6}  {'Qual':>5}  {'Sharpe':>7}  {'Alpha':>7}  {'CAGR':>7}")
    print("─" * 70)
    for _, row in all_results.head(10).iterrows():
        print(f"  {int(row['rank']):>3}  {int(row['top_n']):>4}  "
              f"{row['w_value']:>6.1f}  {row['w_health']:>7.1f}  "
              f"{row['w_sentiment']:>6.1f}  {row['w_quality']:>5.1f}  "
              f"{row['train_sharpe']:>7.2f}  "
              f"{row['train_alpha_ann']:>+7.1%}  "
              f"{row['train_cagr']:>+7.1%}")

    # ── Validate top-K out-of-sample ───────────────────────────────────────────
    validated = validate_top_k(
        top_results         = all_results,
        raw_features        = raw,
        close_panel         = close_panel,
        signals_panel       = signals,
        eval_start          = args.eval_start,
        eval_end            = eval_end,
        k                   = args.validate_k,
        fundamentals_panel  = fund_panel,
        tc_bps              = args.tc_bps,
        strict_pit          = args.strict_pit,
    )
    validated.to_csv(out_dir / "optimizer_validated.csv", index=False)

    # Best = highest out-of-sample Sharpe
    best_row = validated.sort_values("test_sharpe", ascending=False).iloc[0]

    print(f"\n{'═'*60}")
    print(f"  BEST CONFIGURATION  (by out-of-sample Sharpe)")
    print(f"{'═'*60}")
    print(f"  top_n       : {int(best_row['top_n'])}")
    print(f"  value       : {best_row['w_value']:.1f}%")
    print(f"  health      : {best_row['w_health']:.1f}%")
    print(f"  sentiment   : {best_row['w_sentiment']:.1f}%")
    print(f"  quality     : {best_row['w_quality']:.1f}%")
    print(f"  Train Sharpe: {best_row['train_sharpe']:.2f}   "
          f"Test Sharpe: {best_row['test_sharpe']:.2f}")
    print(f"  Train alpha : {best_row['train_alpha_ann']:+.1%}   "
          f"Test alpha: {best_row['test_alpha_ann']:+.1%}")
    print(f"{'═'*60}")

    if args.update_config:
        update_config(best_row)

    print(f"\nResults → {out_dir}/")
    print(f"  optimizer_search.csv    — all {len(all_results):,} trial results")
    print(f"  optimizer_validated.csv — top-{args.validate_k} with test validation")
    print(f"\nTo apply best weights:")
    print(f"  python bt_optimizer.py --update_config")
    print(f"  python bt_engine.py "
          f"--weights 'value={best_row['w_value']:.1f},health={best_row['w_health']:.1f},"
          f"sentiment={best_row['w_sentiment']:.1f},quality={best_row['w_quality']:.1f}' "
          f"--top_n {int(best_row['top_n'])}\n")


if __name__ == "__main__":
    main()
