"""
Utility functions for feature extraction.
"""

import pandas as pd
import numpy as np
from typing import Dict, List, Optional, Any


def extract_categorical_features(
    df: pd.DataFrame,
    user_col: str,
    cat_col: str,
    prefix: str,
    top_n: int = 3,
    include_diversity: bool = True
) -> pd.DataFrame:
    """
    Extract features from categorical columns efficiently.
    
    Args:
        df: DataFrame containing the data
        user_col: Column name for user IDs
        cat_col: Column name for categorical data
        prefix: Prefix for generated feature names
        top_n: Number of top categories to include
        include_diversity: Whether to include diversity metrics
        
    Returns:
        DataFrame with extracted features
    """
    if cat_col not in df.columns:
        return pd.DataFrame()
    
    # Create pivot table
    pivot = pd.crosstab(df[user_col], df[cat_col], normalize="index")
    
    # Get top categories
    top_cats = df[cat_col].value_counts().head(top_n).index
    
    # Create feature dict
    features = {}
    
    # Add top category fractions
    for i, cat in enumerate(top_cats):
        if cat in pivot.columns:
            features[f"{prefix}_top{i+1}_{str(cat)[:20]}"] = pivot[cat]
    
    if include_diversity:
        # Calculate diversity (entropy)
        pivot_safe = pivot + 1e-10
        entropy = -(pivot_safe * np.log2(pivot_safe)).sum(axis=1)
        features[f"{prefix}_diversity"] = entropy
        
        # Count of unique values
        features[f"{prefix}_unique_count"] = (pivot > 0).sum(axis=1)
        
        # Concentration (Herfindahl index)
        features[f"{prefix}_concentration"] = (pivot**2).sum(axis=1)
        
        # Dominant category encoded
        if len(pivot.columns) > 0:
            dominant = pivot.idxmax(axis=1)
            features[f"{prefix}_dominant_encoded"] = dominant.apply(
                lambda x: hash(str(x)) % 1000
            )
    
    return pd.DataFrame(features).reset_index()


def safe_merge_features(
    features_df: pd.DataFrame,
    new_features_df: pd.DataFrame,
    on: str = "user_id",
    check_existing: bool = True,
    verbose: bool = True
) -> pd.DataFrame:
    """
    Safely merge new features into existing dataframe, avoiding duplicates.
    
    Args:
        features_df: Main features dataframe
        new_features_df: New features to merge
        on: Column to merge on
        check_existing: Whether to check for existing columns
        verbose: Whether to print messages
        
    Returns:
        Updated features dataframe
    """
    if check_existing:
        # Get columns to merge (excluding the merge key)
        new_cols = [col for col in new_features_df.columns if col != on]
        
        # Check for existing columns
        existing_cols = [col for col in new_cols if col in features_df.columns]
        
        if existing_cols and verbose:
            print(f"   ⚠️ Found existing columns: {existing_cols[:5]}{'...' if len(existing_cols) > 5 else ''}")
        
        # Keep only new columns
        cols_to_merge = [col for col in new_cols if col not in features_df.columns]
        
        if not cols_to_merge:
            if verbose:
                print("   ⚠️ All features already exist, skipping merge")
            return features_df
        
        # Merge only new columns
        features_df = features_df.merge(
            new_features_df[[on] + cols_to_merge],
            on=on,
            how="left"
        )
    else:
        # Merge all columns
        features_df = features_df.merge(
            new_features_df,
            on=on,
            how="left"
        )
    
    # Fill NaN values for numeric columns
    numeric_cols = features_df.select_dtypes(include=[np.number]).columns
    numeric_cols = [col for col in numeric_cols if col != on]
    features_df[numeric_cols] = features_df[numeric_cols].fillna(0)
    
    return features_df


def print_feature_summary(
    features_df: pd.DataFrame,
    title: str = "Feature Summary",
    show_sample: int = 10
) -> None:
    """
    Print a summary of features in the dataframe.
    
    Args:
        features_df: DataFrame to summarize
        title: Title for the summary
        show_sample: Number of columns to show as sample
    """
    print(f"\n{'='*60}")
    print(f"{title}")
    print(f"{'='*60}")
    print(f"Shape: {features_df.shape}")
    print(f"Total features: {len(features_df.columns) - 1}")  # Exclude user_id
    
    # Show sample columns
    if show_sample > 0:
        print(f"\nSample columns:")
        cols = [col for col in features_df.columns if col != "user_id"]
        for i, col in enumerate(cols[:show_sample]):
            print(f"  {i+1}. {col}")
        if len(cols) > show_sample:
            print(f"  ... and {len(cols) - show_sample} more columns")
    
    # Basic statistics
    print(f"\nBasic statistics:")
    print(f"  Users: {len(features_df):,}")
    print(f"  Memory usage: {features_df.memory_usage().sum() / 1024**2:.2f} MB") 