#!/usr/bin/env python3
"""
Sweep liquidity floor + max weight per name on confirm_entry_4b (de-tail).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_detail_sweep.py \\
        --start 2020-01-02 --end 2026-06-26
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import replace
from pathlib import Path

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))

from RenTech.strategy_stack.alpaca_minute_loader import (
    DEFAULT_ALPACA_RTH_DIR,
    compound_intraday_to_daily,
    list_parquet_symbols,
    load_equity_panels,
)
from RenTech.strategy_stack.ma_slope_intraday_enhanced import (
    EnhancedIntradayConfig,
    EnhancedIntradayEngine,
    baseline_enhanced_config,
    liquidity_filter_symbols,
    metrics_daily,
)
from RenTech.strategy_stack.run_johansen_triplet_sp500 import load_sp500_sectors

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


def _confirm_base(**kw) -> EnhancedIntradayConfig:
    defaults: dict = {"hold_mode": "confirm_entry", "confirm_lag_bars": 4, "max_weight_per_name": 1.0}
    defaults.update(kw)
    return replace(baseline_enhanced_config(), **defaults)


def _filter_universe(
    intra: dict,
    daily: dict,
    *,
    min_price: float,
    min_adv: float,
) -> dict:
    if min_price <= 0 and min_adv <= 0:
        return intra
    ok = liquidity_filter_symbols(daily, min_price=min_price, min_adv=min_adv)
    return {k: v for k, v in intra.items() if k in ok or k == "SPY"}


def _tail_stats(port_r: pd.Series, exec_w: pd.DataFrame, ret_df: pd.DataFrame) -> dict:
    from RenTech.strategy_stack.run_ma_slope_intraday_trade_stats import _session_key

    daily = compound_intraday_to_daily(port_r).dropna()
    ses = _session_key(pd.DatetimeIndex(exec_w.index))
    rows = []
    for sym in exec_w.columns:
        w = exec_w[sym].astype(float)
        r = ret_df[sym].astype(float)
        for day, idx in w.groupby(ses).groups.items():
            wi = w.loc[idx]
            if float(wi.max()) <= 1e-9:
                continue
            ri = r.loc[idx]
            active = wi > 1e-9
            if not active.any():
                continue
            tr = float(np.prod(1.0 + ri.values[active.values]) - 1.0)
            rows.append({"name_ret": tr, "weight": float(wi.max())})
    if not rows:
        return {}
    nt = pd.DataFrame(rows)
    pos = daily[daily > 0].sort_values(ascending=False)
    log_total = float(np.log1p(daily).sum())
    k10 = min(10, len(pos))
    return {
        "median_name_ret_pct": float(nt["name_ret"].median() * 100),
        "pct_names_gt_20pct": float((nt["name_ret"] > 0.20).mean() * 100),
        "max_weight_used": float(exec_w.max().max()),
        "avg_gross_exposure": float(exec_w.sum(axis=1).groupby(ses).max().mean()),
        "top10_days_log_share_pct": float(np.log1p(pos.head(k10)).sum() / log_total * 100) if log_total > 0 else 0.0,
        "n_days_gt_10pct": int((daily > 0.10).sum()),
    }


def _run_variant(
    intra: dict,
    daily: dict,
    sector_map: dict,
    *,
    name: str,
    cfg: EnhancedIntradayConfig,
    top_n: int,
    min_price_univ: float,
    min_adv_univ: float,
    ret_start: pd.Timestamp,
    end: pd.Timestamp,
) -> dict:
    from RenTech.strategy_stack.run_ma_slope_intraday_trade_stats import run_with_weights

    sub = _filter_universe(intra, daily, min_price=min_price_univ, min_adv=min_adv_univ)
    eng = EnhancedIntradayEngine(config=cfg)
    panels = eng.build_panels(sub, sector_map)
    port_r, exec_w, ret_df = run_with_weights(eng, sub, top_n, panels)
    port_r = port_r.loc[(port_r.index >= ret_start) & (port_r.index <= end)]
    exec_w = exec_w.loc[(exec_w.index >= ret_start) & (exec_w.index <= end)]
    ret_df = ret_df.loc[(ret_df.index >= ret_start) & (ret_df.index <= end)]

    m = metrics_daily(port_r)
    t = _tail_stats(port_r, exec_w, ret_df)
    return {
        "variant": name,
        "top_n": top_n,
        "n_symbols": len(sub),
        "max_weight_per_name": cfg.max_weight_per_name,
        "min_price_entry": cfg.min_price,
        "min_price_univ": min_price_univ,
        "min_adv_univ": min_adv_univ,
        **m,
        **t,
    }


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2020-01-02")
    ap.add_argument("--end", default="2026-06-26")
    ap.add_argument("--max-tickers", type=int, default=500)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_detail_sweep")
    args = ap.parse_args()

    syms = list_parquet_symbols(DEFAULT_ALPACA_RTH_DIR)[: int(args.max_tickers)]
    if "SPY" not in syms:
        syms.insert(0, "SPY")
    sector_map = dict(
        zip(
            load_sp500_sectors()["ticker"].astype(str).str.upper(),
            load_sp500_sectors()["sector"].astype(str),
        )
    )

    print("Loading panels …", flush=True)
    intra, daily = load_equity_panels(
        syms,
        start="2019-11-01",
        end=str(args.end),
        warmup_sessions=15,
    )
    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)

    specs: list[tuple[str, EnhancedIntradayConfig, int, float, float]] = [
        ("ref_top10", _confirm_base(), 10, 0.0, 0.0),
        ("ref_top5", _confirm_base(), 5, 0.0, 0.0),
        ("cap20_top10", _confirm_base(max_weight_per_name=0.20), 10, 0.0, 0.0),
        ("cap20_top5", _confirm_base(max_weight_per_name=0.20), 5, 0.0, 0.0),
        ("liq5_adv5m_top10", _confirm_base(min_price=5.0), 10, 5.0, 5_000_000.0),
        ("liq10_adv20m_top10", _confirm_base(min_price=10.0), 10, 10.0, 20_000_000.0),
        ("cap20_liq5_adv5m_top10", _confirm_base(max_weight_per_name=0.20, min_price=5.0), 10, 5.0, 5_000_000.0),
        ("cap20_liq10_adv20m_top10", _confirm_base(max_weight_per_name=0.20, min_price=10.0), 10, 10.0, 20_000_000.0),
        ("cap20_liq10_adv20m_top5", _confirm_base(max_weight_per_name=0.20, min_price=10.0), 5, 10.0, 20_000_000.0),
        ("cap25_liq10_adv20m_top10", _confirm_base(max_weight_per_name=0.25, min_price=10.0), 10, 10.0, 20_000_000.0),
    ]

    rows = []
    ref_ret = None
    for name, cfg, top_n, mp_u, adv_u in specs:
        print(f"  {name} …", flush=True)
        row = _run_variant(
            intra,
            daily,
            sector_map,
            name=name,
            cfg=cfg,
            top_n=top_n,
            min_price_univ=mp_u,
            min_adv_univ=adv_u,
            ret_start=ret_start,
            end=end,
        )
        if name == "ref_top10":
            ref_ret = float(row.get("total_return_pct", 0))
        row["return_vs_ref_top10_pct"] = float(row.get("total_return_pct", 0)) - (ref_ret or 0)
        rows.append(row)
        print(
            f"    ret {row.get('total_return_pct', float('nan')):+7.1f}%  "
            f"Sharpe {row.get('sharpe', float('nan')):5.2f}  DD {row.get('max_dd_pct', float('nan')):6.1f}%  "
            f"avg gross {row.get('avg_gross_exposure', float('nan')):.2f}  "
            f"n>20% names {row.get('pct_names_gt_20pct', float('nan')):.2f}%",
            flush=True,
        )

    df = pd.DataFrame(rows).sort_values("total_return_pct", ascending=False)
    out = args.out_prefix.expanduser().resolve()
    csv_path = Path(f"{out}.csv")
    df.to_csv(csv_path, index=False)
    meta = {"reference_return_pct": ref_ret, "ranked": df.to_dict(orient="records")}
    Path(f"{out}_meta.json").write_text(json.dumps(meta, indent=2) + "\n")

    print(f"\nReference ref_top10: {ref_ret:+.1f}%", flush=True)
    print("\nRanked:", flush=True)
    for _, r in df.iterrows():
        print(
            f"  {r['variant']:28s}  {r['total_return_pct']:+7.1f}%  "
            f"Δ {r['return_vs_ref_top10_pct']:+6.1f}pp  Sharpe {r['sharpe']:.2f}  "
            f"DD {r['max_dd_pct']:.1f}%  gross {r.get('avg_gross_exposure', 0):.2f}",
            flush=True,
        )
    print(f"\nWrote {csv_path}", flush=True)


if __name__ == "__main__":
    main()
