"""
Per-strategy signal quality classifier.

For each of the 5 strategies, collects all CALL/PUT signals + market context,
trains a binary classifier to predict WIN (1) vs LOSS (0).

Uses scikit-learn RandomForest + SMOTE for small sample sizes.
No TensorFlow required.
"""

import os
import json
import warnings
import logging
import numpy as np
import pandas as pd
from sklearn.ensemble import RandomForestClassifier
from sklearn.preprocessing import StandardScaler
from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, roc_auc_score
from imblearn.over_sampling import SMOTE
from collections import Counter

warnings.filterwarnings('ignore')
logging.basicConfig(level=logging.INFO, format='%(asctime)s  %(levelname)-7s %(message)s')
logger = logging.getLogger('per_strategy')

from config import Config
from strategies import STRATEGIES

# ── Feature engineering per signal ──
CONTEXT_FEATURES = ['RSI14', 'ATR14', 'MACD_hist', 'body_pct', 'hour_sin', 'hour_cos', 'dow_sin', 'dow_cos']


def build_signal_dataset(df: pd.DataFrame, strategy_fn, params: dict,
                         instrument: str, strategy_name: str, expiry_bars: int = 1) -> pd.DataFrame:
    """Extract all signals + context features + outcome labels for one strategy."""
    rows = []

    for i in range(50, len(df) - expiry_bars):
        sig = strategy_fn(df, params, i)
        if sig['signal'] is None:
            continue

        row = df.iloc[i]
        future_close = df['Close'].iloc[i + expiry_bars]
        current_close = row['Close']

        is_call = sig['signal'] == 'CALL'
        win = (is_call and future_close > current_close) or (not is_call and future_close < current_close)

        features = {
            'confidence': sig['confidence'],
        }
        for f in CONTEXT_FEATURES:
            if f in df.columns:
                features[f] = row.get(f, np.nan)

        ts = df.index[i]
        if hasattr(ts, 'hour'):
            features['hour'] = ts.hour
        else:
            features['hour'] = 0

        features['label'] = 1 if win else 0
        features['instrument'] = instrument
        features['strategy'] = strategy_name
        features['expiry'] = expiry_bars
        rows.append(features)

    if not rows:
        return pd.DataFrame()

    df_out = pd.DataFrame(rows)
    df_out = df_out.dropna()
    return df_out


def train_evaluate(df_signals: pd.DataFrame, strategy_name: str):
    """Train RandomForest on strategy signals, evaluate on test set."""
    if len(df_signals) < 20:
        return {'error': f'Too few samples: {len(df_signals)}'}

    # Features
    feat_cols = ['confidence'] + [f for f in CONTEXT_FEATURES if f in df_signals.columns]
    if 'hour' in df_signals.columns:
        feat_cols.append('hour')

    X = df_signals[feat_cols].values.astype(np.float64)
    y = df_signals['label'].values.astype(np.int64)

    # Temporal split: 70/30
    n = len(X)
    split = int(n * 0.70)
    X_train, X_test = X[:split], X[split:]
    y_train, y_test = y[:split], y[split:]

    if len(X_test) < 5:
        return {'error': f'Too few test samples: {len(X_test)}'}

    # Scale
    scaler = StandardScaler()
    X_train = scaler.fit_transform(X_train)
    X_test = scaler.transform(X_test)

    # SMOTE for class balance
    counter_before = Counter(y_train)
    if min(counter_before.values()) >= 3 and max(counter_before.values()) / min(counter_before.values()) > 1.5:
        k_neighbors = min(3, min(counter_before.values()) - 1)
        if k_neighbors >= 1:
            try:
                smote = SMOTE(random_state=42, k_neighbors=k_neighbors)
                X_train, y_train = smote.fit_resample(X_train, y_train)
            except Exception:
                pass

    # Train
    clf = RandomForestClassifier(
        n_estimators=200, max_depth=6, min_samples_leaf=5,
        class_weight='balanced', random_state=42, n_jobs=-1
    )
    clf.fit(X_train, y_train)

    # Predict
    y_pred = clf.predict(X_test)
    y_prob = clf.predict_proba(X_test)[:, 1]

    # Metrics
    metrics = {
        'n_train': int(len(y_train)), 'n_test': int(len(y_test)),
        'class_balance': {str(k): int(v) for k, v in Counter(y_test).items()},
        'accuracy': float(round(accuracy_score(y_test, y_pred), 4)),
        'precision': float(round(precision_score(y_test, y_pred, zero_division=0), 4)),
        'recall': float(round(recall_score(y_test, y_pred, zero_division=0), 4)),
        'f1': float(round(f1_score(y_test, y_pred, zero_division=0), 4)),
        'auc': float(round(roc_auc_score(y_test, y_prob), 4)) if len(set(y_test)) > 1 else 0.5,
    }

    # Feature importance — convert np int64 keys to str
    importance = {str(k): float(v) for k, v in
                  sorted(zip(feat_cols, clf.feature_importances_.round(4)), key=lambda x: x[1], reverse=True)[:8]}

    # Trading simulation: what WR if we only trade when model predicts WIN?
    wins_when_pred_win = int(sum(1 for i in range(len(y_test)) if y_pred[i] == 1 and y_test[i] == 1))
    preds_win = int(sum(y_pred))
    simulated_wr = round(wins_when_pred_win / preds_win * 100, 1) if preds_win > 0 else 0
    trades_taken = preds_win
    baseline_wr = round(float(y_test.mean()) * 100, 1)

    return {
        'strategy': strategy_name,
        'total_signals': int(len(df_signals)),
        'baseline_wr': float(baseline_wr),
        'simulated_wr': float(simulated_wr),
        'trades_taken': int(trades_taken),
        'trades_skipped': int(len(y_test) - trades_taken),
        'metrics': metrics,
        'top_features': importance,
        'wr_improvement': float(round(simulated_wr - baseline_wr, 1)),
    }


def run_all(expiry_bars: int = 1):
    """Train and evaluate per-strategy classifiers for all instruments."""
    Config.setup_dirs()
    data = {}

    for instr in Config.INSTRUMENTS:
        path = os.path.join(Config.DATA_DIR, f'{instr}_H1_indicators.csv')
        if not os.path.exists(path):
            logger.warning("No data for %s, skipping", instr)
            continue
        df = pd.read_csv(path, index_col=0, parse_dates=True)
        data[instr] = df

    # Add cyclical time features
    for instr, df in data.items():
        df['hour_sin'] = np.sin(2 * np.pi * df.index.hour / 24)
        df['hour_cos'] = np.cos(2 * np.pi * df.index.hour / 24)
        df['dow_sin'] = np.sin(2 * np.pi * df.index.dayofweek / 7)
        df['dow_cos'] = np.cos(2 * np.pi * df.index.dayofweek / 7)

    results = []

    for instr, df in data.items():
        for name, fn in STRATEGIES.items():
            params = Config.STRATEGY_PARAMS[name]
            signals_df = build_signal_dataset(df, fn, params, instr, name, expiry_bars)

            if len(signals_df) < 20:
                logger.info("%-25s %s: %3d signals — too few, skipping", name, instr, len(signals_df))
                continue

            result = train_evaluate(signals_df, f'{instr}_{name}')
            if 'error' in result:
                logger.info("%-25s %s: %s", name, instr, result['error'])
                continue

            logger.info("%-25s %s: N=%4d  base=%.1f%%  sim=%.1f%%  Δ=%+.1f%%  auc=%.3f",
                        name, instr, result['total_signals'],
                        result['baseline_wr'], result['simulated_wr'],
                        result['wr_improvement'], result['metrics']['auc'])
            results.append(result)

    # Summary table
    results.sort(key=lambda r: r['wr_improvement'], reverse=True)

    print(f"\n{'='*90}")
    print(f"  PER-STRATEGY SIGNAL FILTER — exp{expiry_bars}h")
    print(f"{'='*90}")
    print(f"  {'Strategy':30s} {'N':>5s} {'Base WR':>8s} {'Filter WR':>9s} {'Δ':>6s} {'Trades':>7s} {'AUC':>6s}")
    print(f"  {'-'*75}")

    for r in results:
        print(f"  {r['strategy']:30s} {r['total_signals']:5d} {r['baseline_wr']:7.1f}% "
              f"{r['simulated_wr']:8.1f}% {r['wr_improvement']:+5.1f}% "
              f"{r['trades_taken']:4d}/{r['total_signals']:4d} {r['metrics']['auc']:.3f}")

    # Save
    path = os.path.join(Config.REPORT_DIR, f'per_strategy_filter_exp{expiry_bars}.json')
    with open(path, 'w') as f:
        json.dump(results, f, indent=2, default=str)
    logger.info("Saved to %s", path)

    return results


def main():
    import argparse
    parser = argparse.ArgumentParser(description='Per-strategy signal quality classifier')
    parser.add_argument('--expiry', type=int, default=1, help='Expiry in hours (default: 1)')
    parser.add_argument('--all', action='store_true', help='Test all expiries 1-5')
    args = parser.parse_args()

    if args.all:
        for exp in [1, 2, 3, 4, 5]:
            logger.info("=== Expiry %dh ===", exp)
            run_all(expiry_bars=exp)
    else:
        run_all(expiry_bars=args.expiry)


if __name__ == '__main__':
    main()
