"""
MLP Backtest Results Table
==========================
Runs the MLP strategy backtest for every timeframe that has both a data file
and a winner CSV, then prints a comparison table matching TradingView's
Strategy Tester columns:

    TF | PnL % | Max DD % | Win Rate

Usage
-----
    python3 tools/mlp_results_table.py              # all TFs
    python3 tools/mlp_results_table.py --tf 4H 8H  # specific TFs
"""

from __future__ import annotations

import argparse
import os
import sys
from pathlib import Path

import pandas as pd

REPO = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(REPO))

import strategies.strategy_mlp_scores as strat
from config import OOS_START, SCORE_START, TRAIN_END

ASSET = "COINBASE_BTCUSD"  # overridden by --asset arg

_TF_PERIODS = {"4H": "240", "6H": "360", "8H": "480", "12H": "720", "1D": "1D"}


def _tf_map(asset: str) -> dict[str, str]:
    return {tf: f"data/mlp/{asset}, {p}.csv" for tf, p in _TF_PERIODS.items()}


TF_MAP = _tf_map(ASSET)

WINNER_DIR = Path("results/winners")
COMMISSION = 0.005  # 0.5% per side, matches TradingView default
def winner_path(tf: str, asset: str | None = None) -> Path:
    a = asset if asset is not None else ASSET
    return WINNER_DIR / f"optimization_winner_strategy_mlp_scores_{a}_{tf}.csv"


def run_backtest(df_sig: pd.DataFrame) -> dict:
    """Trade simulation matching validate_strategy.py logic."""
    equity = 100.0
    peak = 100.0
    max_dd_frac = 0.0
    in_pos = False
    entry_price = 0.0
    total = 0
    wins = 0

    for _, row in df_sig.iterrows():
        if row["execute_entry"] and not in_pos:
            in_pos = True
            entry_price = row["close"]
            total += 1
            continue

        if in_pos:
            low_ret = (row["low"] - entry_price) / entry_price
            floating = equity * (1 + low_ret)
            dd = (peak - floating) / peak
            if dd > max_dd_frac:
                max_dd_frac = dd

            if row["execute_exit"]:
                in_pos = False
                eff_entry = entry_price * (1 + COMMISSION)
                eff_exit = row["close"] * (1 - COMMISSION)
                pnl = (eff_exit - eff_entry) / eff_entry
                if pnl > 0:
                    wins += 1
                equity *= (1 + pnl)
                if equity > peak:
                    peak = equity

    total_pnl = (equity - 100.0) / 100.0 * 100.0
    win_pct = (wins / total * 100.0) if total > 0 else 0.0
    return {
        "pnl": total_pnl,
        "max_dd": -max_dd_frac * 100.0,
        "win_pct": win_pct,
        "wins": wins,
        "total": total,
    }


def evaluate_tf(tf: str, start: str | None = None, end: str | None = None,
                 asset: str | None = None) -> dict | None:
    # Resolve the asset once and derive BOTH the data file and the winner CSV
    # path from that single value. Do NOT read the module-global TF_MAP here:
    # TF_MAP and ASSET are two separate mutable globals that must be
    # reassigned together (see main()), and a caller who updates one without
    # the other would otherwise silently pair one asset's price data with a
    # different asset's optimized winner params (see tests/test_mlp_results_table.py).
    a = asset if asset is not None else ASSET
    data_file = _tf_map(a).get(tf)
    wp = winner_path(tf, a)

    if not data_file or not os.path.exists(data_file):
        print(f"  {tf}: data file missing — skip", file=sys.stderr)
        return None
    if not wp.exists():
        print(f"  {tf}: no winner CSV — skip", file=sys.stderr)
        return None

    params = pd.read_csv(wp).iloc[0].to_dict()

    df = pd.read_csv(data_file)
    df.columns = df.columns.str.lower()
    df["time"] = pd.to_datetime(df["time"], utc=True).dt.tz_localize(None)

    # Preserve full warm-up history, then gate entries/exits exactly like Pine.
    is_params = dict(params)
    is_params["_pine_time_start"] = start
    is_params["_pine_time_end"] = end
    df_sig_is = strat.generate_signals(df.copy(), **is_params)
    if start:
        df_sig_is = df_sig_is[df_sig_is["time"] >= pd.Timestamp(start)]
    if end:
        df_sig_is = df_sig_is[df_sig_is["time"] <= pd.Timestamp(end)]
    df_sig_is = df_sig_is.reset_index(drop=True)
    result = run_backtest(df_sig_is)

    # Full range: same start, no end cap — includes OOS bars
    full_params = dict(params)
    full_params["_pine_time_start"] = start
    df_sig_full = strat.generate_signals(df.copy(), **full_params)
    if start:
        df_sig_full = df_sig_full[df_sig_full["time"] >= pd.Timestamp(start)]
    df_sig_full = df_sig_full.reset_index(drop=True)
    latest_date = df_sig_full["time"].max().strftime("%Y-%m-%d") if not df_sig_full.empty else "?"
    result_full = run_backtest(df_sig_full)

    # OOS-only: same signal generation, backtest starts fresh at OOS_START
    df_sig_oos = df_sig_full[df_sig_full["time"] >= pd.Timestamp(OOS_START)].reset_index(drop=True)
    result_oos = run_backtest(df_sig_oos)

    result["tf"] = tf
    result["pnl_full"] = result_full["pnl"]
    result["pnl_oos"]  = result_oos["pnl"]
    result["total_oos"] = result_oos["total"]
    result["latest_date"] = latest_date
    return result


def _fmt_pnl(val: float | None) -> str:
    if val is None:
        return "—"
    return f"+{val:,.1f}%" if val >= 0 else f"{val:,.1f}%"


def format_markdown(rows: list[dict], asset: str, start: str, end: str) -> str:
    """Return a GitHub-renderable markdown table for the given result rows."""
    from datetime import datetime, timezone
    now    = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M UTC")
    latest = rows[0].get("latest_date", "?") if rows else "?"
    lines  = [
        f"# MLP Performance — {asset}",
        "",
        f"*Updated: {now} · IS: {start} → {end} · OOS: {OOS_START}+*",
        "",
        f"| TF | IS PnL % | Max DD % | Win Rate | Thru {latest} | OOS PnL % | OOS Trades |",
        "|---|---:|---:|---:|---:|---:|---:|",
    ]
    for r in rows:
        lines.append(
            f"| {r['tf']} "
            f"| {_fmt_pnl(r['pnl'])} "
            f"| {r['max_dd']:.1f}% "
            f"| {r['win_pct']:.0f}% ({r['wins']}/{r['total']}) "
            f"| {_fmt_pnl(r.get('pnl_full'))} "
            f"| {_fmt_pnl(r.get('pnl_oos'))} "
            f"| {r.get('total_oos', 0)} |"
        )
    lines += ["", "*Refresh: `python3 tools/mlp_results_table.py --save`*"]
    return "\n".join(lines)


def print_table(rows: list[dict]) -> None:
    if not rows:
        print("No results.")
        return

    latest     = rows[0].get("latest_date", "?") if rows else "?"
    full_hdr   = f"PnL % (thru {latest})"
    oos_hdr    = f"OOS PnL % ({OOS_START}+)"

    header = (f"{'TF':<5} {'PnL %':>12} {'Max DD %':>10} {'Win Rate':>18}"
              f"  {'':>4}  {full_hdr:>22}"
              f"  {'':>4}  {oos_hdr:>26}")
    sep    = (f"{'-'*5} {'-'*12} {'-'*10} {'-'*18}"
              f"  {'':>4}  {'-'*22}"
              f"  {'':>4}  {'-'*26}")
    print()
    print(header)
    print(sep)
    for r in rows:
        pnl_str      = _fmt_pnl(r["pnl"])
        dd_str       = f"{r['max_dd']:.1f}%"
        win_str      = f"{r['win_pct']:.0f}% ({r['wins']}/{r['total']})"
        full_pnl_str = _fmt_pnl(r.get("pnl_full"))
        oos_trades   = r.get("total_oos", 0)
        oos_str      = f"{_fmt_pnl(r.get('pnl_oos'))} ({oos_trades} trades)"
        print(f"{r['tf']:<5} {pnl_str:>12} {dd_str:>10} {win_str:>18}"
              f"  {'':>4}  {full_pnl_str:>22}"
              f"  {'':>4}  {oos_str:>26}")
    print()


DEFAULT_START = SCORE_START
DEFAULT_END   = TRAIN_END


def main() -> None:
    global ASSET, TF_MAP
    parser = argparse.ArgumentParser(description="MLP backtest results table")
    parser.add_argument("--asset", default=ASSET,
                        help=f"Asset to evaluate (default: {ASSET})")
    parser.add_argument("--tf", nargs="*", default=list(_TF_PERIODS.keys()),
                        choices=list(_TF_PERIODS.keys()),
                        help="Timeframes to include (default: all)")
    parser.add_argument("--start", default=DEFAULT_START,
                        help=f"Start date inclusive (default: {DEFAULT_START})")
    parser.add_argument("--end", default=DEFAULT_END,
                        help=f"End date inclusive (default: {DEFAULT_END})")
    parser.add_argument("--save", nargs="?",
                        const=str(REPO / "results" / "mlp_performance.md"),
                        default=str(REPO / "results" / "mlp_performance.md"),
                        metavar="PATH",
                        help="Write markdown table to PATH (default: results/mlp_performance.md; pass --no-save to skip)")
    parser.add_argument("--no-save", dest="save", action="store_const", const=None,
                        help="Skip writing the markdown artifact")
    args = parser.parse_args()

    ASSET  = args.asset
    TF_MAP = _tf_map(ASSET)

    print(f"Running MLP backtest for {ASSET} [{args.start} → {args.end}]...", file=sys.stderr)
    rows = []
    for tf in args.tf:
        result = evaluate_tf(tf, start=args.start, end=args.end)
        if result:
            rows.append(result)

    print_table(rows)

    if args.save:
        save_path = Path(args.save)
        save_path.parent.mkdir(parents=True, exist_ok=True)
        save_path.write_text(format_markdown(rows, ASSET, args.start, args.end))
        print(f"Saved: {save_path.relative_to(REPO)}", file=sys.stderr)


if __name__ == "__main__":
    main()
