"""Feature selection for ML models.

Supports two modes:
1. FAST (default): Uses tree-based `feature_importances_` — instant, no extra deps.
2. DEEP (opt-in):  Uses SHAP TreeExplainer — thorough but slow (2-5 min per asset).

In production signal scanning (where speed matters), always use FAST mode.
DEEP mode is useful for one-off analysis or when you want to understand *why*
a feature matters (interaction effects, non-linear contributions).
"""
from __future__ import annotations
import logging
import numpy as np
import pandas as pd
from typing import Optional

logger = logging.getLogger(__name__)


def _cfg(key: str, fallback):
    try:
        from flask import current_app
        return current_app.config.get(key, fallback)
    except RuntimeError:
        return fallback


class FeatureSelector:
    """Feature importance-based selection with fast & deep modes."""

    def __init__(self, top_k: int | None = None):
        self.top_k = top_k
        self._selected_features: dict[str, list[str]] = {}
        self._importance_scores: dict[str, dict[str, float]] = {}

    # ------------------------------------------------------------------
    # Public API
    # ------------------------------------------------------------------

    def select_features(
        self, df_or_model, target_or_X=None, top_k: int | None = None,
        cache_key: str = 'default', method: str = 'fast',
    ) -> list[str]:
        """Select top_k most important features.

        Supports two calling conventions:
        1. select_features(dataframe_with_target, target_col_name)
           — Trains a quick GradientBoosting model internally
        2. select_features(trained_model, X_dataframe)
           — Uses the provided trained model directly

        Args:
            df_or_model: Either a DataFrame (with target column) or a trained model
            target_or_X: Either target column name (str) or feature DataFrame
            top_k: Number of features to keep (default from config or __init__)
            cache_key: Key for caching results
            method: 'fast' (feature_importances_) or 'deep' (SHAP)

        Returns:
            List of selected feature column names
        """
        if top_k is None:
            top_k = self.top_k or _cfg('ML_FEATURE_TOP_K', 30)

        if cache_key in self._selected_features:
            return self._selected_features[cache_key]

        # Determine calling convention
        if isinstance(df_or_model, pd.DataFrame) and isinstance(target_or_X, str):
            # Convention 1: DataFrame + target column name
            return self._select_from_dataframe(
                df_or_model, target_or_X, top_k, cache_key, method
            )
        else:
            # Convention 2: model + X DataFrame
            if method == 'deep':
                return self._select_with_shap(
                    df_or_model, target_or_X, top_k, cache_key
                )
            else:
                return self._select_with_importances(
                    df_or_model, target_or_X, top_k, cache_key
                )

    # ------------------------------------------------------------------
    # Internal
    # ------------------------------------------------------------------

    def _select_from_dataframe(
        self, df: pd.DataFrame, target_col: str,
        top_k: int, cache_key: str, method: str,
    ) -> list[str]:
        """Train a quick model and use it for feature selection."""
        X = df.drop(columns=[target_col]).select_dtypes(include=[np.number])
        y = df[target_col].values

        mask = ~(np.isnan(X.values).any(axis=1) | np.isnan(y))
        X, y = X[mask], y[mask]

        if len(X) < 50:
            return X.columns.tolist()

        try:
            from sklearn.ensemble import GradientBoostingRegressor
            model = GradientBoostingRegressor(
                n_estimators=100, max_depth=6, learning_rate=0.05,
                subsample=0.8, random_state=42,
            )
            model.fit(X.values, y)

            if method == 'deep':
                return self._select_with_shap(model, X, top_k, cache_key)
            else:
                return self._select_with_importances(model, X, top_k, cache_key)

        except Exception as e:
            logger.warning(f'Feature selection failed: {e}')
            return X.columns.tolist()

    def _select_with_importances(
        self, model, X: pd.DataFrame, top_k: int, cache_key: str,
    ) -> list[str]:
        """FAST mode: Use tree-based feature_importances_ (instant).

        Works with any sklearn-compatible tree model: GradientBoosting,
        XGBoost, LightGBM, CatBoost, RandomForest, etc.
        """
        try:
            importances = model.feature_importances_
            feature_importance = dict(zip(X.columns.tolist(), importances.tolist()))

            sorted_features = sorted(
                feature_importance.items(), key=lambda x: x[1], reverse=True
            )
            selected = [f[0] for f in sorted_features[:top_k]]

            self._selected_features[cache_key] = selected
            self._importance_scores[cache_key] = feature_importance

            logger.info(
                f'Fast feature selection: {len(X.columns)} → {len(selected)} features '
                f'(top: {selected[:5]})'
            )
            return selected

        except AttributeError:
            logger.warning(
                'Model has no feature_importances_, falling back to all features'
            )
            return X.columns.tolist()

    def _select_with_shap(
        self, model, X: pd.DataFrame, top_k: int, cache_key: str,
    ) -> list[str]:
        """DEEP mode: Use SHAP values for thorough feature analysis.

        Slower (2-5 min) but captures interaction effects & non-linear
        contributions that feature_importances_ can miss.
        """
        try:
            import shap

            # Sample for speed (SHAP on full dataset is very slow)
            sample_size = min(200, len(X))
            X_sample = X.iloc[-sample_size:]

            # Use appropriate explainer
            try:
                explainer = shap.TreeExplainer(model)
                shap_values = explainer.shap_values(X_sample)
            except Exception:
                explainer = shap.Explainer(model, X_sample)
                shap_values = explainer(X_sample).values

            # Mean absolute SHAP values per feature
            importance = np.abs(shap_values).mean(axis=0)
            feature_importance = dict(zip(X.columns.tolist(), importance.tolist()))

            # Sort by importance
            sorted_features = sorted(
                feature_importance.items(), key=lambda x: x[1], reverse=True
            )

            # Keep top_k
            selected = [f[0] for f in sorted_features[:top_k]]

            self._selected_features[cache_key] = selected
            self._importance_scores[cache_key] = feature_importance

            logger.info(
                f'SHAP feature selection: {len(X.columns)} → {len(selected)} features'
            )
            return selected

        except ImportError:
            logger.warning('SHAP not installed, falling back to fast mode')
            return self._select_with_importances(model, X, top_k, cache_key)
        except Exception as e:
            logger.warning(f'SHAP failed: {e}, falling back to fast mode')
            return self._select_with_importances(model, X, top_k, cache_key)

    def get_importance_scores(self, cache_key: str = 'default') -> dict[str, float]:
        """Get cached feature importance scores."""
        return self._importance_scores.get(cache_key, {})

    def filter_dataframe(
        self, df: pd.DataFrame, cache_key: str = 'default',
        keep_cols: list[str] | None = None,
    ) -> pd.DataFrame:
        """Filter DataFrame to keep only selected features + specified columns."""
        selected = self._selected_features.get(cache_key)
        if not selected:
            return df

        cols_to_keep = list(selected)
        if keep_cols:
            cols_to_keep.extend([c for c in keep_cols if c in df.columns])

        available = [c for c in cols_to_keep if c in df.columns]
        return df[available]
