"""
Оптимизация MoE v12: удаление худших признаков и переобучение моделей.

Принцип:
    1. Загружаем существующее сохранение MoE (moe_v12_rr1x2.joblib)
    2. У XGBoost Long/Short берём `feature_importances_` первых n_base столбцов —
       это вклад каждого базового признака в финальную модель.
    3. Сортируем базовые признаки по сумме важности Long+Short (imp_total).
    4. Оставляем top-K% (по умолчанию 50%), с защитой min_features.
    5. Делаем бэкап оригинальной модели в backup_top<pct>_<date>/.
    6. Запускаем train_moe_for(ticker, feature_cols_override=top_features).
       Внутри train_moe_for monkey-patchим _get_feature_cols / _crypto_feature_cols —
       поэтому LSTM-эксперты и XGBoost обучаются только на выбранном подмножестве.
    7. Сохраняем финальную модель поверх старой, логируем before/after.

Запуск:
    python train_all_rr1x2.py --optimize-features 50 --ticker SBER
    python train_all_rr1x2.py --optimize-features 50             # 4 тикера по умолчанию
    python train_all_rr1x2.py --optimize-features 60 --min-features 50 --ticker ALL

Альтернативно (как модуль):
    from analysis.optimize_features import run_optimization
    results = run_optimization(['SBER','GAZP'], keep_pct=0.5, min_features=40)
"""
from __future__ import annotations

import argparse
import os
import shutil
import time
from datetime import datetime
from typing import Any

import joblib
import numpy as np

# Важно: import тиков осуществляется через train_all_rr1x2 (там уже вызов config.set_device).
# На прямую не вызываем train_moe_for здесь — это делает вызывающий код,
# но мы реэкспортируем его ниже для удобства.

SAVE_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
                        'models', 'saved')


def _ticker_pathname(ticker: str) -> str:
    return ticker.lower()


def _model_path(ticker: str) -> str:
    return os.path.join(SAVE_DIR, f'{_ticker_pathname(ticker)}_moe_v12_rr1x2.joblib')


def _fetch_base_importances(saved: dict) -> dict[str, dict]:
    """Извлекает важность базовых признаков из сохранённой модели.

    Returns:
        {feature_name: {'imp_long': float, 'imp_short': float, 'imp_total': float}}
    """
    rf_long = saved.get('rf_long')
    rf_short = saved.get('rf_short')
    n_base = saved['n_base_features']
    feature_cols = saved['feature_cols']

    imp_long = rf_long.feature_importances_[:n_base]
    has_short = rf_short is not None
    imp_short = rf_short.feature_importances_[:n_base] if has_short else np.zeros(n_base)

    out = {}
    for i, name in enumerate(feature_cols[:n_base]):
        out[name] = {
            'imp_long': float(imp_long[i]),
            'imp_short': float(imp_short[i]),
            'imp_total': float(imp_long[i] + imp_short[i]),
        }
    return out


def select_top_features(saved: dict, keep_pct: float = 0.5,
                         min_features: int = 40) -> tuple[list[str], dict]:
    """Возвращает список top-K% базовых признаков по imp_total.

    Args:
        saved: joblib-загруженная MoE.
        keep_pct: Доля признаков для сохранения (0.5 = 50%).
        min_features: Минимум признаков (failsafe).

    Returns:
        (top_features, stats)
    """
    importances = _fetch_base_importances(saved)
    sorted_feats = sorted(importances.items(), key=lambda x: x[1]['imp_total'], reverse=True)
    n_total = len(sorted_feats)
    n_keep = max(int(n_total * keep_pct), min_features)
    if n_keep > n_total:
        n_keep = n_total
    top_features = [name for name, _ in sorted_feats[:n_keep]]
    kept_imp = sum(importances[f]['imp_total'] for f in top_features)
    total_imp = sum(v['imp_total'] for v in importances.values())
    coverage = kept_imp / max(total_imp, 1e-10)
    stats = {
        'n_total_base': n_total,
        'n_keep_base': n_keep,
        'keep_pct': keep_pct,
        'coverage_imp': coverage,
        'removed_count': n_total - n_keep,
        'removed_imp_share': 1.0 - coverage,
        # демонстрационные prunings
        'top_features': top_features,
        'removed_features': [name for name, _ in sorted_feats[n_keep:]],
        # первые 10 dropped и оставленные
        'top_5_kept': [name for name, _ in sorted_feats[:5]],
        'top_5_dropped': [name for name, _ in sorted_feats[n_keep:n_keep + 5]],
    }
    return top_features, stats


def _backup_model(ticker: str, backup_dir: str) -> str:
    """Копирует модель в backup_dir/{ticker}_moe_v12_rr1x2.joblib."""
    src = _model_path(ticker)
    dst_dir = os.path.join(SAVE_DIR, backup_dir)
    os.makedirs(dst_dir, exist_ok=True)
    dst = os.path.join(dst_dir, os.path.basename(src))
    shutil.copy2(src, dst)
    return dst


def _restore_from_backup(ticker: str, backup_dir: str) -> bool:
    """Восстанавливает модель из бэкапа (если оптимизация не удалась)."""
    src = os.path.join(SAVE_DIR, backup_dir, os.path.basename(_model_path(ticker)))
    dst = _model_path(ticker)
    if not os.path.isfile(src):
        return False
    shutil.copy2(src, dst)
    return True


def _get_metrics(saved: dict) -> dict:
    """Извлекает metrics из сохранённой модели для сравнения."""
    from config import get_target_config
    th = saved.get('optimal_thresholds') or {}
    return {
        'val_acc': float(saved.get('val_acc', 0.0)),
        'long_th': float(th.get('long', 0.0)) if th else 0.0,
        'short_th': float(th.get('short', 0.0)) if th else 0.0,
        'n_base_features': int(saved.get('n_base_features', 0)),
        'total_feature_size': len(saved.get('feature_cols', [])),
        'thresholds': th,
    }


def _estimate_net_pnl_on_val(saved: dict, ticker: str) -> dict:
    """Грубая оценка Net PnL Long по OOF threshold на val части (по архивным данным модели).

    Возвращает {'net_long': float, 'net_short': float, 'long_thr': float, 'short_thr': float}.
    Если не получилось — все 0.0.
    """
    # Используем первые найденные поля (модель MoE не хранит OOF probs, только thresholds/val_acc)
    return {'net_long': 0.0, 'net_short': 0.0}


def _run_optimization_single(ticker: str,
                              keep_pct: float = 0.5,
                              min_features: int = 40,
                              backup_dir: str | None = None,
                              verbose: bool = True) -> dict:
    """Запускает оптимизацию для одного тикера."""
    from train_all_rr1x2 import train_moe_for

    path = _model_path(ticker)
    if not os.path.isfile(path):
        return {'status': 'error', 'ticker': ticker, 'error': f'model not found: {path}'}

    if verbose:
        print(f'\n[{ticker}] Загрузка текущей модели для анализа importance...')
    saved_old = joblib.load(path)
    metrics_before = _get_metrics(saved_old)
    top_features, selection_stats = select_top_features(saved_old, keep_pct=keep_pct,
                                                        min_features=min_features)

    if verbose:
        print(f'  ➤ Base features: {selection_stats["n_total_base"]} → '
              f'{selection_stats["n_keep_base"]} (удалено {selection_stats["removed_count"]})')
        print(f'  ➤ Сохранённая важность (coverage): {selection_stats["coverage_imp"]:.1%}')
        print(f'  ➤ Top-5 оставленных: {selection_stats["top_5_kept"]}')
        print(f'  ➤ Top-5 удалённых:  {selection_stats["top_5_dropped"]}')

    # Бэкап
    if backup_dir is None:
        date = datetime.now().strftime('%Y-%m-%d_%H-%M')
        backup_dir = f'backup_top{int(keep_pct * 100)}_{date}'
    backup_path = _backup_model(ticker, backup_dir)
    if verbose:
        print(f'  ➤ Бэкап: {backup_path}')

    # Переобучение с override
    if verbose:
        print(f'  ➤ Переобучение MoE с feature_cols_override ({len(top_features)} признаков)...')
    try:
        result = train_moe_for(ticker, feature_cols_override=top_features)
    except Exception as e:
        # Восстанавливаем бэкап,если обучение упало
        if verbose:
            print(f'  ✗ train_moe_for упал: {e}. Восстанавливаю бэкап...')
        _restore_from_backup(ticker, backup_dir)
        import traceback; traceback.print_exc()
        return {'status': 'error', 'ticker': ticker, 'error': str(e)}

    if not result or result.get('status') != 'ok':
        if verbose:
            print(f'  ✗ train_moe_for вернул не ok. Восстанавливаю бэкап...')
        _restore_from_backup(ticker, backup_dir)
        return {'status': 'error', 'ticker': ticker, 'error': 'train_moe_for returned non-ok',
                'train_result': result}

    # Считаем metrics после
    saved_new = joblib.load(_model_path(ticker))
    metrics_after = _get_metrics(saved_new)

    # Решение: оставляем новую версию только если НЕ хуже старой по валидационному accuracy,
    # либо если разница < 3% (допускаем снижение за компактность модели).
    va_delta = metrics_after['val_acc'] - metrics_before['val_acc']
    accept_threshold = -0.03  # допускаем падение до 3%
    accept = va_delta >= accept_threshold

    if not accept:
        if verbose:
            print(f'  ⚠ val_acc снизился на {-va_delta:.1%} (>{-accept_threshold:.0%}). '
                  f'Восстанавливаю бэкап...')
        _restore_from_backup(ticker, backup_dir)
        return {
            'status': 'rejected',
            'ticker': ticker,
            'reason': f'val_acc decreased by {-va_delta:.1%}',
            'backup_dir': backup_dir,
            **_interop_metrics(ticker, metrics_before, metrics_after, selection_stats, result),
        }

    return {
        'status': 'ok',
        'ticker': ticker,
        'backup_dir': backup_dir,
        **_interop_metrics(ticker, metrics_before, metrics_after, selection_stats, result),
    }


def _interop_metrics(ticker: str,
                     before: dict, after: dict,
                     selection_stats: dict,
                     train_result: dict) -> dict:
    """Сводка метрик before/after для отчёта."""
    return {
        'n_features_before': before['total_feature_size'],
        'n_features_after': after['total_feature_size'],
        'n_base_before': before['n_base_features'],
        'n_base_after': after['n_base_features'],
        'val_acc_before': before['val_acc'],
        'val_acc_after': after['val_acc'],
        'long_th_before': before['long_th'],
        'long_th_after': after['long_th'],
        'short_th_before': before['short_th'],
        'short_th_after': after['short_th'],
        'coverage_imp': selection_stats['coverage_imp'],
        'removed_count': selection_stats['removed_count'],
        'top_5_kept': selection_stats['top_5_kept'],
        'top_5_dropped': selection_stats['top_5_dropped'],
    }


def run_optimization(tickers: list[str],
                     keep_pct: float = 0.5,
                     min_features: int = 40,
                     verbose: bool = True) -> list[dict]:
    """Параллельно-последовательная оптимизация для списка тикеров.

    Args:
        tickers: список тикеров (SBER, GAZP, ...)
        keep_pct: доля признаков для сохранения (0.5 = 50%)
        min_features: минимально допустимое количество признаков
        verbose: печатать прогресс

    Returns:
        Список dict с результатами per-ticker
    """
    # Фиксируем seed для воспроизводимости
    from models.experts import seed_everything
    seed_everything()

    date = datetime.now().strftime('%Y-%m-%d_%H-%M')
    backup_dir = f'backup_top{int(keep_pct * 100)}_{date}'
    output = []
    for ticker in tickers:
        t0 = time.time()
        try:
            r = _run_optimization_single(ticker, keep_pct=keep_pct,
                                        min_features=min_features,
                                        backup_dir=backup_dir,
                                        verbose=verbose)
            r['elapsed_min'] = round((time.time() - t0) / 60, 1)
            output.append(r)
        except Exception as e:
            import traceback
            traceback.print_exc()
            output.append({'status': 'error', 'ticker': ticker, 'error': str(e)})
    return output


# ---------------------------------------------------------------------------
# CLI (standalone)
# ---------------------------------------------------------------------------
def main():
    import sys
    parser = argparse.ArgumentParser(description='Optimize MoE models: keep top-N% base features')
    parser.add_argument('--tickers', type=str, default='SBER,GAZP,BITCOIN,EURUSD',
                        help='Tickers; comma-separated (default: 4 tickers analyzed)')
    parser.add_argument('--keep-pct', type=float, default=0.5,
                        help='Доля признаков для сохранения (0.5 = 50%)')
    parser.add_argument('--min-features', type=int, default=40)
    args = parser.parse_args()

    tickers = [t.strip().upper() for t in args.tickers.split(',') if t.strip()]
    results = run_optimization(tickers, keep_pct=args.keep_pct,
                                min_features=args.min_features,
                                verbose=True)

    print('\n=== Результаты оптимизации ===')
    print(f'{"Тикер":>8} {"was":>20} {"now":>20} {"val_acc":>20} {"Long th":>16} {"Verdict":>10}')
    print('-' * 100)
    for r in results:
        if r.get('status') == 'ok':
            print(f"{r['ticker']:>8} | {r['n_features_before']:3d} ({r['n_base_before']:2d} base) "
                  f"| {r['n_features_after']:3d} ({r['n_base_after']:2d} base) "
                  f"| {r['val_acc_before']:.1%} → {r['val_acc_after']:.1%} "
                  f"| L≥{r['long_th_before']:.2f}→{r['long_th_after']:.2f} "
                  f"| ACCEPTED")
        elif r.get('status') == 'rejected':
            print(f"{r['ticker']:>8} | {r['n_features_before']:3d} ({r['n_base_before']:2d} base) "
                  f"| reverted | {r['val_acc_before']:.1%} vs {r['val_acc_after']:.1%} "
                  f"| REJECTED: {r.get('reason','')}")
        else:
            print(f"{r.get('ticker','?'):>8} | ERROR: {r.get('error','?')}")


if __name__ == '__main__':
    main()