"""
Feature analysis: correlation analysis для выявления и удаления мультиколлинеарных признаков.
"""

import pandas as pd
import numpy as np
from utils.logger import logger


def remove_collinear_features(
    X: pd.DataFrame, y: pd.Series, threshold: float = 0.90
) -> list[str]:
    """
    Удаляет высоко коррелированные признаки (мультиколлинеарность).

    Алгоритм:
    1. Считает корреляцию каждого признака с таргетом
    2. Считает матрицу корреляций между фичами
    3. Ищет пары с корреляцией > threshold
    4. Если пара найдена, оставляет признак с большей корреляцией с таргетом

    Args:
        X: Feature matrix (n_samples, n_features)
        y: Target series (n_samples,)
        threshold: Порог корреляции для удаления (по умолчанию 0.90)

    Returns:
        list[str]: Список признаков, которые нужно оставить

    Example:
        >>> keep_features = remove_collinear_features(X, y, threshold=0.90)
        >>> X_clean = X[keep_features]
    """
    logger.info(f"Starting collinearity check (threshold={threshold})...")

    # Проверка входных данных
    if len(X) != len(y):
        raise ValueError(f"X and y must have the same length: X={len(X)}, y={len(y)}")

    # Считаем корреляцию с таргетом (по модулю)
    logger.info("Computing correlations with target...")
    correlations_with_target = X.corrwith(y).abs().fillna(0)

    # Считаем матрицу корреляций между фичами
    logger.info("Computing feature correlation matrix...")
    corr_matrix = X.corr().abs()

    # Берем верхний треугольник матрицы (без диагонали)
    upper = corr_matrix.where(np.triu(np.ones(corr_matrix.shape), k=1).astype(bool))

    to_drop = set()
    pairs_dropped = []

    # Идем по колонкам верхнего треугольника
    for col in upper.columns:
        # Находим фичи, которые сильно коррелируют с текущей
        highly_correlated = upper.index[upper[col] > threshold].tolist()

        if highly_correlated:
            # Сравниваем корреляцию с таргетом
            col_corr = correlations_with_target[col]
            corr_col_corr = correlations_with_target[highly_correlated[0]]

            if col_corr >= corr_col_corr:
                to_drop.add(highly_correlated[0])
                pairs_dropped.append((col, highly_correlated[0], col_corr, corr_col_corr))
            else:
                to_drop.add(col)
                pairs_dropped.append((highly_correlated[0], col, corr_col_corr, col_corr))

    keep = [col for col in X.columns if col not in to_drop]

    logger.info(
        f"Removed {len(to_drop)} collinear features. Kept {len(keep)} ({len(keep)/len(X.columns)*100:.1f}%)."
    )

    if pairs_dropped:
        logger.info(f"Most correlated pairs dropped:")
        for i, (removed, kept, corr_removed, corr_kept) in enumerate(
            pairs_dropped[:5], 1
        ):
            logger.info(
                f"  {i}. Dropped '{removed}' (corr={corr_removed:.3f}) in favor of '{kept}' (corr={corr_kept:.3f})"
            )

    return keep


def compute_correlation_matrix(
    X: pd.DataFrame, method: str = "pearson"
) -> pd.DataFrame:
    """
    Вычисляет матрицу корреляций между фичами.

    Args:
        X: Feature matrix (n_samples, n_features)
        method: Метод корреляции ('pearson', 'spearman', 'kendall')

    Returns:
        pd.DataFrame: Матрица корреляций (n_features, n_features)

    Example:
        >>> corr = compute_correlation_matrix(X)
        >>> high_corr = corr[corr > 0.9].stack()
    """
    logger.info(f"Computing correlation matrix (method={method})...")
    corr_matrix = X.corr(method=method)
    return corr_matrix


def get_highly_correlated_pairs(
    X: pd.DataFrame, threshold: float = 0.90
) -> pd.DataFrame:
    """
    Находит пары признаков с высокой корреляцией.

    Args:
        X: Feature matrix (n_samples, n_features)
        threshold: Порог корреляции (по умолчанию 0.90)

    Returns:
        pd.DataFrame: DataFrame с парами коррелирующих фичей

    Example:
        >>> pairs = get_highly_correlated_pairs(X, threshold=0.90)
        >>> print(pairs)
    """
    logger.info(f"Finding highly correlated pairs (threshold={threshold})...")

    corr_matrix = X.corr().abs()
    upper = corr_matrix.where(np.triu(np.ones(corr_matrix.shape), k=1).astype(bool))

    # Находим пары с корреляцией выше порога
    highly_correlated = []
    for col in upper.columns:
        correlated_cols = upper.index[upper[col] > threshold].tolist()
        for corr_col in correlated_cols:
            highly_correlated.append((col, corr_col, upper.loc[col, corr_col]))

    if not highly_correlated:
        logger.info("No highly correlated pairs found.")
        return pd.DataFrame(
            columns=["Feature 1", "Feature 2", "Correlation"]
        )

    df = pd.DataFrame(
        highly_correlated, columns=["Feature 1", "Feature 2", "Correlation"]
    )
    df = df.sort_values("Correlation", ascending=False)

    logger.info(f"Found {len(df)} highly correlated pairs.")
    return df


def plot_correlation_heatmap(X: pd.DataFrame, figsize: tuple = (20, 20)):
    """
    Строит тепловую карту корреляций между фичами.

    Args:
        X: Feature matrix (n_samples, n_features)
        figsize: Размер фигуры (по умолчанию (20, 20))

    Example:
        >>> plot_correlation_heatmap(X, figsize=(15, 15))
    """
    import matplotlib.pyplot as plt

    logger.info("Plotting correlation heatmap...")

    corr_matrix = X.corr().abs()

    # Масштабируем
    if len(corr_matrix.columns) > 20:
        # Выбираем топ-20 коррелирующих фичей
        top_corr = corr_matrix.unstack().sort_values(ascending=False)
        top_corr = top_corr[top_corr < 1.0].head(20)
        feature_names = set(top_corr.index.get_level_values(0)) | set(
            top_corr.index.get_level_values(1)
        )
        corr_matrix = corr_matrix.loc[list(feature_names), list(feature_names)]

    plt.figure(figsize=figsize)
    im = plt.imshow(
        corr_matrix, cmap="coolwarm", aspect="auto", vmin=0, vmax=1
    )
    plt.colorbar(im, label="Correlation")
    plt.xticks(
        ticks=range(len(corr_matrix.columns)),
        labels=corr_matrix.columns,
        rotation=45,
        ha="right",
    )
    plt.yticks(
        ticks=range(len(corr_matrix.index)),
        labels=corr_matrix.index,
        rotation=0,
    )
    plt.title("Feature Correlation Matrix", fontsize=14, pad=20)
    plt.tight_layout()

    return plt.gcf()
