"""
Utilities Module
===============
Helper functions for the user clustering pipeline.
"""

import logging
import json
import pandas as pd
from typing import Any, Optional, List
from pathlib import Path


def setup_logging(output_dir: Optional[str] = None, log_file: str = "user_clustering.log") -> None:
    """
    Set up logging configuration.

    Args:
        output_dir: Directory to save log file. If None, uses current directory
        log_file: Name of the log file
    """
    # Determine log file path
    if output_dir:
        Path(output_dir).mkdir(parents=True, exist_ok=True)
        log_path = Path(output_dir) / log_file
    else:
        log_path = log_file
    
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
        handlers=[
            logging.FileHandler(log_path),
            logging.StreamHandler(),  # Also log to console
        ],
    )

    # Set specific log levels for libraries
    logging.getLogger("matplotlib").setLevel(logging.WARNING)
    logging.getLogger("sklearn").setLevel(logging.WARNING)

    logger = logging.getLogger(__name__)
    logger.info("Logging initialized")


def parse_metadata_field(metadata_str: Any) -> bool:
    """
    Parse metadata JSON string to check for gpt_description_prompt.

    Args:
        metadata_str: Metadata field which can be a string, dict, or None

    Returns:
        Boolean indicating if gpt_description_prompt exists
    """
    if metadata_str is None:
        return False

    try:
        # Handle case where metadata is already a dict
        if isinstance(metadata_str, dict):
            metadata = metadata_str
        elif isinstance(metadata_str, str):
            # Try to parse as JSON
            metadata = json.loads(metadata_str)
        else:
            return False

        # Check for gpt_description_prompt
        if "gpt_description_prompt" in metadata:
            prompt = metadata["gpt_description_prompt"]
            # Check if prompt is not None and not empty
            if prompt is not None and str(prompt).strip():
                return True

        return False
    except:
        return False


def safe_divide(numerator: float, denominator: float, default: float = 0.0) -> float:
    """
    Safely divide two numbers, returning default if denominator is zero.

    Args:
        numerator: The numerator
        denominator: The denominator
        default: Default value if division by zero

    Returns:
        Result of division or default
    """
    if denominator == 0:
        return default
    return numerator / denominator


def memory_usage(df: pd.DataFrame) -> str:
    """
    Calculate and format memory usage of a DataFrame.

    Args:
        df: DataFrame to analyze

    Returns:
        Formatted string with memory usage
    """
    mem_usage = df.memory_usage(deep=True).sum()

    # Convert to appropriate unit
    if mem_usage < 1024:
        return f"{mem_usage} bytes"
    elif mem_usage < 1024**2:
        return f"{mem_usage / 1024:.2f} KB"
    elif mem_usage < 1024**3:
        return f"{mem_usage / (1024 ** 2):.2f} MB"
    else:
        return f"{mem_usage / (1024 ** 3):.2f} GB"


def validate_dataframe_columns(
    df: pd.DataFrame, required_columns: list, df_name: str
) -> None:
    """
    Validate that a DataFrame contains required columns.

    Args:
        df: DataFrame to validate
        required_columns: List of required column names
        df_name: Name of the DataFrame for error messages

    Raises:
        ValueError: If required columns are missing
    """
    missing_columns = set(required_columns) - set(df.columns)

    if missing_columns:
        raise ValueError(
            f"{df_name} is missing required columns: {sorted(missing_columns)}"
        )


def get_feature_importance(
    feature_values: pd.Series, overall_mean: float, overall_std: float
) -> float:
    """
    Calculate feature importance as standardized difference from overall mean.

    Args:
        feature_values: Feature values for a cluster
        overall_mean: Overall mean across all clusters
        overall_std: Overall standard deviation

    Returns:
        Standardized importance score
    """
    if overall_std == 0:
        return 0.0

    cluster_mean = feature_values.mean()
    z_score = (cluster_mean - overall_mean) / overall_std

    return abs(z_score)


def format_duration(seconds: float) -> str:
    """
    Format duration in seconds to human-readable string.

    Args:
        seconds: Duration in seconds

    Returns:
        Formatted duration string
    """
    if seconds < 60:
        return f"{seconds:.1f} seconds"
    elif seconds < 3600:
        minutes = seconds / 60
        return f"{minutes:.1f} minutes"
    else:
        hours = seconds / 3600
        return f"{hours:.1f} hours"


def create_feature_summary(df: pd.DataFrame) -> pd.DataFrame:
    """
    Create a summary of features including basic statistics.

    Args:
        df: DataFrame with features

    Returns:
        DataFrame with feature summaries
    """
    # Select numeric columns
    numeric_cols = df.select_dtypes(include=["number"]).columns

    # Calculate statistics
    summary = pd.DataFrame(
        {
            "feature": numeric_cols,
            "mean": df[numeric_cols].mean(),
            "std": df[numeric_cols].std(),
            "min": df[numeric_cols].min(),
            "max": df[numeric_cols].max(),
            "skewness": df[numeric_cols].skew(),
            "missing_pct": (df[numeric_cols].isna().sum() / len(df) * 100),
        }
    )

    return summary.round(3)
