"""
Walk-Forward Validation Rescore Tool
=====================================
Loads the top N candidates from the sweep DB for a given asset/timeframe,
evaluates each across temporal WFO OOS folds, and reports which params are
most robust across different market regimes.

This is a temporal robustness check *within* the IS window — NOT a true holdout
test (that's the OOS dashboard at 2024-10-01+).  It answers: "Of all the IS-
optimal configs, which ones also work in 2021 (bull), 2022 (bear), 2023
(recovery), and 2024-H1 (mixed)?"

Usage:
  python3 tools/wfo_rescore.py --asset COINBASE_BTCUSD --timeframe 4H
  python3 tools/wfo_rescore.py --asset COINBASE_BTCUSD --timeframe 4H --top-n 500
  python3 tools/wfo_rescore.py --asset COINBASE_BTCUSD --timeframe 4H --save-winner
"""

import os
import sys
import argparse
import sqlite3
import io
import pandas as pd
import numpy as np
from concurrent.futures import ProcessPoolExecutor, as_completed

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import strategies.strategy_activation_scores as strategy_module
from config import (
    TRAIN_START, TRAIN_END, SCORE_START,
    RESULTS_DIR, WINNERS_DIR, WEIGHT_COLS,
    WFO_FOLDS, WFO_MIN_OOS_TRADES, WFO_MIN_VALID_FOLDS,
)

SWEEP_DB_FILE = "results/sweep_database.db"
DB_SCHEMA_VER = 3

DATA_FILE_TEMPLATE = "data/{asset}-{timeframe}.csv"


def _wfo_score_one(args):
    """
    Worker: score one candidate across all WFO folds.
    Must be at module level for ProcessPoolExecutor pickle.
    """
    params_dict, df_full_bytes, folds, min_oos_trades, min_valid_folds = args
    import io as _io
    import pandas as _pd
    import numpy as _np
    import strategies.strategy_activation_scores as _strat

    try:
        df_full = _pd.read_pickle(_io.BytesIO(df_full_bytes))
        fold_calmars = []
        fold_trades = []

        for (is_end, oos_start, oos_end) in folds:
            df_window = df_full[df_full['time'] <= _pd.to_datetime(oos_end)].copy()
            df_signals = _strat.generate_signals(df_window, **params_dict)
            df_oos = df_signals[df_signals['time'] <= _pd.to_datetime(oos_end)].copy()
            metrics = _strat.calculate_metrics(df_oos, score_start=oos_start)
            oos_trades = int(metrics.get('Total Trades', 0))
            fold_trades.append(oos_trades)
            if oos_trades >= min_oos_trades:
                fold_calmars.append(metrics.get('Calmar Ratio', -99.0))
            else:
                fold_calmars.append(None)  # not enough trades — fold excluded

        valid_calmars = [c for c in fold_calmars if c is not None]
        if len(valid_calmars) < min_valid_folds:
            wfo_score = 0.0
            wfo_min = -99.0
            wfo_neg = len(valid_calmars)
        else:
            wfo_score = float(_np.mean(valid_calmars))
            wfo_min = float(min(valid_calmars))
            wfo_neg = sum(1 for c in valid_calmars if c < 0)

        return wfo_score, fold_calmars, fold_trades, wfo_min, wfo_neg

    except Exception as e:
        return 0.0, [], [], -99.0, 0


def load_top_candidates(asset, timeframe, top_n, sort_col="composite_score"):
    """Load top N candidates from sweep DB by IS composite score."""
    if not os.path.exists(SWEEP_DB_FILE):
        print(f"Sweep DB not found: {SWEEP_DB_FILE}")
        return pd.DataFrame()
    try:
        conn = sqlite3.connect(SWEEP_DB_FILE)
        df = pd.read_sql(
            f"SELECT * FROM sweep_results "
            f"WHERE asset=? AND timeframe=? AND schema_ver=? AND {sort_col} > 0 "
            f"ORDER BY {sort_col} DESC LIMIT ?",
            conn, params=(asset, timeframe, DB_SCHEMA_VER, top_n)
        )
        conn.close()
        return df
    except Exception as e:
        print(f"DB query failed: {e}")
        return pd.DataFrame()


def main():
    parser = argparse.ArgumentParser(description="WFO temporal robustness rescore")
    parser.add_argument("--asset",     required=True, help="e.g. COINBASE_BTCUSD")
    parser.add_argument("--timeframe", required=True, help="e.g. 4H")
    parser.add_argument("--top-n",     type=int, default=200, help="Top N IS candidates to rescore (default: 200)")
    parser.add_argument("--data",      type=str, default=None, help="Override data CSV path")
    parser.add_argument("--save-winner", action="store_true", help="Save WFO winner CSV to results/winners/")
    args = parser.parse_args()

    data_file = args.data or DATA_FILE_TEMPLATE.format(asset=args.asset, timeframe=args.timeframe)
    if not os.path.exists(data_file):
        print(f"Data file not found: {data_file}")
        sys.exit(1)

    print(f"\n=== WFO Rescore: {args.asset} {args.timeframe} ===")
    print(f"Loading top {args.top_n} IS candidates from sweep DB...")
    df_top = load_top_candidates(args.asset, args.timeframe, args.top_n)
    if df_top.empty:
        print("No candidates found in sweep DB.")
        sys.exit(1)
    print(f"Loaded {len(df_top)} candidates  (IS Composite range: "
          f"{df_top['composite_score'].min():.4f} – {df_top['composite_score'].max():.4f})")

    print(f"Loading data from {data_file}...")
    df_data = pd.read_csv(data_file)
    df_data.columns = df_data.columns.str.lower()
    if 'time' in df_data.columns:
        df_data['time'] = pd.to_datetime(df_data['time'], utc=True).dt.tz_localize(None)
        mask = (df_data['time'] >= TRAIN_START) & (df_data['time'] <= TRAIN_END)
        df_data = df_data.loc[mask].copy()

    buf = io.BytesIO()
    df_data.to_pickle(buf)
    df_data_bytes = buf.getvalue()

    # Build param dicts from DB rows (only include cols that appear in data)
    all_param_cols = list(WEIGHT_COLS) + [
        "i_long_entry_activation_threshold", "i_long_exit_activation_threshold",
        "i_long_exit_activation_confirmation_threshold", "i_use_long_exit_confirmation",
        "i_use_long_entry_confirmation", "i_trailing_stop_threshold",
        "i_m3_momentum_period", "i_regime_window", "i_regime_entry_min_score",
        "i_mvrv_suppress_bear", "i_div_window",
    ]
    candidates = []
    for _, row in df_top.iterrows():
        p = {col: row[col] for col in all_param_cols if col in row.index and pd.notna(row[col])}
        p['_is_composite'] = row.get('composite_score', 0.0)
        candidates.append(p)

    cpu_count = os.cpu_count() or 4
    max_workers = max(1, min(16, cpu_count - 2))
    print(f"\nRunning WFO across {len(WFO_FOLDS)} folds using {max_workers} workers...")

    fold_labels = ["2021", "2022", "2023", "2024H1", "PoliRun"]
    wfo_args = [
        (p, df_data_bytes, WFO_FOLDS, WFO_MIN_OOS_TRADES, WFO_MIN_VALID_FOLDS)
        for p in candidates
    ]

    results = [None] * len(candidates)
    with ProcessPoolExecutor(max_workers=max_workers) as executor:
        futures = {executor.submit(_wfo_score_one, a): i for i, a in enumerate(wfo_args)}
        done = 0
        for fut in as_completed(futures):
            idx = futures[fut]
            wfo_score, fold_calmars, fold_trades, wfo_min, wfo_neg = fut.result()
            results[idx] = {
                'is_composite': candidates[idx]['_is_composite'],
                'wfo_score': wfo_score,
                'wfo_min_fold': wfo_min,
                'wfo_neg_folds': wfo_neg,
                'fold_calmars': fold_calmars,
                'fold_trades': fold_trades,
                **{k: v for k, v in candidates[idx].items() if not k.startswith('_')},
            }
            done += 1
            if done % 50 == 0 or done == len(candidates):
                print(f"  {done}/{len(candidates)} scored...", flush=True)

    df_results = pd.DataFrame(results).sort_values(['wfo_min_fold', 'wfo_score'], ascending=False)

    n_positive = (df_results['wfo_score'] > 0).sum()
    n_all_positive = (df_results['wfo_neg_folds'] == 0).sum()
    n_negative = (df_results['wfo_score'] < 0).sum()
    n_zero = (df_results['wfo_score'] == 0).sum()
    print(f"\n{'='*75}")
    print(f"  WFO Results: {n_positive} positive mean ({n_positive/len(df_results)*100:.0f}%)  "
          f"| {n_all_positive} all-folds-positive  "
          f"| {n_negative} negative  | {n_zero} excluded (<{WFO_MIN_VALID_FOLDS} valid folds)")
    print(f"{'='*75}")

    # Print top 10 WFO candidates (sorted by min-fold then mean)
    print(f"\n{'Rank':<5} {'IS Comp':>8} {'WFO Mean':>9} {'WFO Min':>8} {'NegF':>5}  {'2021':>7} {'2022':>7} {'2023':>7} {'2024H1':>7} {'PoliRun':>8}")
    print("-" * 75)
    for rank, (_, row) in enumerate(df_results.head(10).iterrows(), 1):
        fc = row['fold_calmars']
        fold_str = "  ".join(
            f"{c:>7.3f}" if c is not None else f"{'N/A':>7}"
            for c in fc
        )
        print(f"{rank:<5} {row['is_composite']:>8.4f} {row['wfo_score']:>9.4f} {row['wfo_min_fold']:>8.4f} {int(row['wfo_neg_folds']):>5}  {fold_str}")

    # IS-optimal winner for comparison
    is_winner = df_results.sort_values('is_composite', ascending=False).iloc[0]
    wfo_winner = df_results.iloc[0]

    print(f"\n  IS winner  : IS={is_winner['is_composite']:.4f}  WFO_Mean={is_winner['wfo_score']:.4f}  WFO_Min={is_winner['wfo_min_fold']:.4f}")
    print(f"  WFO winner : IS={wfo_winner['is_composite']:.4f}  WFO_Mean={wfo_winner['wfo_score']:.4f}  WFO_Min={wfo_winner['wfo_min_fold']:.4f}  NegFolds={int(wfo_winner['wfo_neg_folds'])}")

    if args.save_winner and wfo_winner['wfo_min_fold'] >= 0:
        out_path = os.path.join(WINNERS_DIR, f"optimization_winner_activation_scores_{args.asset}_{args.timeframe}_wfo.csv")
        param_cols = [c for c in wfo_winner.index if c not in ('is_composite', 'wfo_score', 'fold_calmars', 'fold_trades')]
        pd.DataFrame([wfo_winner[param_cols]]).to_csv(out_path, index=False)
        print(f"\n  WFO winner saved → {out_path}")

    print()


if __name__ == "__main__":
    main()
