"""
Фреймворк для сравнения моделей на Walk-Forward кросс-валидации.

Позволяет запустить один раз и получить таблицу метрик для всех моделей
на одних и тех же данных и сплитах.
"""

import pandas as pd
import numpy as np
from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, confusion_matrix
from utils.cv import WalkForwardSplit
from utils.logger import logger


class BenchmarkFramework:
    """
    Фреймворк для сравнения моделей на одинаковых Walk-Forward сплитах.

    Args:
        n_splits: Количество фолдов для Walk-Forward CV
        test_size: Размер тестового окна (float 0-1 или int)
        gap: Зазор между train и test в баров

    Example:
        >>> benchmark = BenchmarkFramework(n_splits=5, test_size=0.2, gap=5)
        >>> benchmark.evaluate_model('RandomForest', lambda: RandomForestClassifier(), X, y)
        >>> benchmark.evaluate_model('XGBoost', lambda: XGBClassifier(), X, y)
        >>> benchmark.print_summary()
    """

    def __init__(self, n_splits: int = 5, test_size: float = 0.2, gap: int = 5):
        """
        Инициализация BenchmarkFramework.

        Args:
            n_splits: Количество фолдов
            test_size: Размер тестового окна (0.2 = 20%)
            gap: Зазор между train и test (баров)
        """
        self.cv = WalkForwardSplit(n_splits=n_splits, test_size=test_size, gap=gap)
        self.results = []
        self.model_params = []

    def evaluate_model(self, model_name: str, model_factory, X, y: np.ndarray, y_true: np.ndarray = None):
        """
        Запускает кросс-валидацию для одной модели.
        X может быть np.ndarray (для sklearn) или pd.DataFrame (для LSTM экспертов).

        Args:
            model_name: Имя модели для отчета
            model_factory: Функция, возвращающая свежий инстанс модели (вызывается на каждом фолде)
            X: Feature matrix (n_samples, n_features) или DataFrame (для LSTM)
            y: Target vector or matrix (n_samples, n_targets) or (n_samples,)
            y_true: Optional ground truth for additional analysis
        """
        logger.info(f"--- Benchmarking: {model_name} ---")

        # Проверяем multi-output
        is_multioutput = len(y.shape) > 1 and y.shape[1] > 1
        target_names = ['Long', 'Short'] if is_multioutput else ['Target']

        # Сохраняем параметры модели
        self.model_params.append({
            'Model': model_name,
            'Params': str(model_factory())
        })

        for fold, (train_idx, test_idx) in enumerate(self.cv.split(X)):
            # Поддержка DataFrame и Numpy
            if isinstance(X, pd.DataFrame):
                X_train, X_test = X.iloc[train_idx], X.iloc[test_idx]
            else:
                X_train, X_test = X[train_idx], X[test_idx]

            y_train, y_test = y[train_idx], y[test_idx]

            # Создаем новую модель
            model = model_factory()

            # Обучаем (поддерживаем и sklearn, и кастомные API)
            if hasattr(model, 'fit'):
                try:
                    model.fit(X_train, y_train)
                except Exception as e:
                    logger.error(f"Fold {fold+1} fit error: {e}")
                    continue
            else:
                logger.error(f"Модель {model_name} не имеет метода fit")
                continue

            # Предсказываем
            if hasattr(model, 'predict_proba'):
                try:
                    if is_multioutput and isinstance(model, (list, tuple)):
                        # MultiOutputClassifier sklearn
                        probs = model.predict_proba(X_test)
                        if isinstance(probs, list):
                            # Превращаем list в numpy array (N x 2)
                            probs_array = np.column_stack([p[:, 1] for p in probs])
                            y_pred = (probs_array >= 0.5).astype(int)
                        else:
                            y_pred = (probs[:, 1] >= 0.5).astype(int)
                    else:
                        probs = model.predict_proba(X_test)
                        if isinstance(probs, list):
                            # list → numpy array (N x 2)
                            probs = np.column_stack(probs)
                        if probs.ndim == 1:
                            # Single output
                            y_pred = (probs >= 0.5).astype(int)
                        else:
                            y_pred = (probs[:, 1] >= 0.5).astype(int)
                except Exception as e:
                    logger.warning(f"Fold {fold+1} predict_proba error: {e}, trying predict")
                    y_pred = model.predict(X_test)
            else:
                y_pred = model.predict(X_test)

            # Преобразуем y_pred в нужную форму для multi-output
            if is_multioutput and y_pred.ndim == 1:
                y_pred = np.column_stack([y_pred, y_pred])

            # Считаем метрики
            for i, name in enumerate(target_names):
                y_t = y_test[:, i] if is_multioutput else y_test
                y_p = y_pred[:, i] if is_multioutput else y_pred

                # Проверяем, что метрики можно посчитать
                unique_classes = np.unique(y_t)
                if len(unique_classes) == 1:
                    logger.debug(f"Fold {fold+1}: {name} has only one class, skipping metrics")
                    continue

                try:
                    acc = accuracy_score(y_t, y_p)
                    prec = precision_score(y_t, y_p, zero_division=0)
                    rec = recall_score(y_t, y_p, zero_division=0)
                    f1 = f1_score(y_t, y_p, zero_division=0)

                    # Confusion matrix
                    tn, fp, fn, tp = confusion_matrix(y_t, y_p, labels=[0, 1]).ravel()

                    self.results.append({
                        'Model': model_name,
                        'Fold': fold + 1,
                        'Target': name,
                        'Accuracy': acc,
                        'Precision': prec,
                        'Recall': rec,
                        'F1': f1,
                        'TN': tn,
                        'TP': tp,
                        'FP': fp,
                        'FN': fn
                    })
                except Exception as e:
                    logger.warning(f"Fold {fold+1}: metrics error: {e}")
                    continue

        logger.info(f"--- Completed: {model_name} ---\n")

    def get_results(self) -> pd.DataFrame:
        """Возвращает DataFrame с результатами."""
        return pd.DataFrame(self.results)

    def get_summary(self) -> pd.DataFrame:
        """Возвращает DataFrame с суммарными метриками по фолдам."""
        if not self.results:
            return pd.DataFrame()

        df = self.get_results()
        summary = df.groupby(['Model', 'Target']).agg({
            'Accuracy': ['mean', 'std', 'min', 'max'],
            'Precision': ['mean', 'std'],
            'Recall': ['mean', 'std'],
            'F1': ['mean', 'std']
        }).round(4)

        # Переименуем уровни индекса для красивого вывода
        summary.columns = ['_'.join(col) for col in summary.columns.values]
        summary = summary.reset_index()

        return summary

    def print_summary(self):
        """Выводит красивую таблицу метрик в консоль."""
        df = self.get_summary()
        if df.empty:
            logger.warning("No results to summarize.")
            return

        logger.info("\n" + "=" * 70)
        logger.info("BENCHMARK SUMMARY (Mean across folds)")
        logger.info("=" * 70)
        print(df.to_string())

        logger.info("\n" + "=" * 70)
        logger.info("Total evaluations: %d folds × %d models × %d targets",
                   self.cv.n_splits, len(self.model_params), df.index.nunique())
        logger.info("=" * 70)

    def get_confusion_matrices(self) -> dict:
        """
        Возвращает confusion matrix для каждой модели и фолда.

        Returns:
            dict: {'Model': {'Fold N': confusion_matrix}}
        """
        if self.results.empty:
            return {}

        cm_dict = {}
        for _, row in self.results.iterrows():
            model = row['Model']
            fold = row['Fold']
            key = (model, fold)

            if key not in cm_dict:
                cm_dict[key] = []

            cm_dict[key].append(row[['TN', 'TP', 'FP', 'FN']].values[0])

        return cm_dict

    def compare_models(self) -> pd.DataFrame:
        """
        Сравнивает модели и возвращает таблицу с лучшими показателями.

        Returns:
            pd.DataFrame: Сравнение моделей
        """
        summary = self.get_results()

        # Фильтруем только Long targets
        summary_long = summary[summary['Target'] == 'Long']

        # Берем максимальные значения для каждой модели
        best = summary_long.groupby('Model').agg({
            'Accuracy': 'max',
            'Precision': 'max',
            'Recall': 'max',
            'F1': 'max'
        }).reset_index()
        best.columns = ['Model', 'Best Accuracy', 'Best Precision', 'Best Recall', 'Best F1']

        return best
