"""
Data Loading Module
==================
Handles loading CSV and pickle files and data validation.
"""

import pandas as pd
import logging
from pathlib import Path
from typing import Dict

logger = logging.getLogger(__name__)


def load_all_dataframes(input_dir: str, files_to_load: Dict[str, str]) -> Dict[str, pd.DataFrame]:
    """
    Load all required dataframes from the input directory.
    Supports both CSV and pickle (.pkl) files.
    
    Args:
        input_dir: Directory containing input files
        files_to_load: Dictionary mapping df names to file names
        
    Returns:
        Dictionary with dataframe names as keys and loaded dataframes as values
    """
    input_path = Path(input_dir)
    dfs = {}
    
    for df_name, file_name in files_to_load.items():
        file_path = input_path / file_name
        
        # Check if file exists with either extension if no extension provided
        if not file_path.exists() and not file_path.suffix:
            # Try with .csv extension
            csv_path = file_path.with_suffix('.csv')
            pkl_path = file_path.with_suffix('.pkl')
            
            if csv_path.exists():
                file_path = csv_path
            elif pkl_path.exists():
                file_path = pkl_path
            else:
                logger.warning(f"File not found: {file_path} (tried .csv and .pkl)")
                continue
        elif not file_path.exists():
            logger.warning(f"File not found: {file_path}")
            continue
            
        try:
            # Determine file type by extension
            file_extension = file_path.suffix.lower()
            
            if file_extension == '.csv':
                logger.info(f"Loading CSV file: {file_path.name}")
                df = pd.read_csv(file_path)
            elif file_extension in ['.pkl', '.pickle']:
                logger.info(f"Loading pickle file: {file_path.name}")
                df = pd.read_pickle(file_path)
            else:
                # Try to infer from content
                logger.info(f"Unknown extension {file_extension}, attempting to load: {file_path.name}")
                try:
                    # Try CSV first
                    df = pd.read_csv(file_path)
                    logger.info(f"Successfully loaded as CSV")
                except:
                    # Try pickle
                    df = pd.read_pickle(file_path)
                    logger.info(f"Successfully loaded as pickle")
            
            dfs[df_name] = df
            logger.info(f"Loaded {len(df)} rows from {file_path.name}")
            
        except Exception as e:
            logger.error(f"Error loading {file_path.name}: {str(e)}")
            raise
    
    return dfs


def validate_data(dfs: Dict[str, pd.DataFrame]) -> None:
    """
    Validate the loaded dataframes for required columns and data quality.
    
    Args:
        dfs: Dictionary of loaded dataframes
        
    Raises:
        ValueError: If validation fails
    """
    logger.info("Validating data")
    
    # Define required columns for each dataframe
    required_columns = {
        'boosts_action_df': ['clip_id', 'created_at'],
        'reaction_df': ['clip_id', 'user_id'],
        'total_clip_df': ['id', 'user_id', 'created_at'],
        'playlist_clip_df': ['clip_id', 'playlist_id']
    }
    
    # Check required dataframes exist
    required_dfs = set(required_columns.keys())
    loaded_dfs = set(dfs.keys())
    missing_dfs = required_dfs - loaded_dfs
    
    if missing_dfs:
        raise ValueError(f"Missing required dataframes: {missing_dfs}")
    
    # Validate columns
    for df_name, columns in required_columns.items():
        if df_name in dfs:
            df = dfs[df_name]
            missing_cols = set(columns) - set(df.columns)
            
            if missing_cols:
                raise ValueError(f"{df_name} missing required columns: {missing_cols}")
    
    # Validate total_clip_df has users
    if 'total_clip_df' in dfs:
        total_clip_df = dfs['total_clip_df']
        
        # Check for null user_ids
        null_users = total_clip_df['user_id'].isnull().sum()
        if null_users > 0:
            logger.warning(f"Found {null_users} clips with null user_id")
            
        # Check we have users
        unique_users = total_clip_df['user_id'].nunique()
        if unique_users == 0:
            raise ValueError("No users found in total_clip_df")
            
        logger.info(f"Found {unique_users} unique users in total_clip_df")
    
    # Check clip_id consistency
    if 'total_clip_df' in dfs and 'boosts_action_df' in dfs:
        total_clip_ids = set(dfs['total_clip_df']['id'].dropna())
        boost_clip_ids = set(dfs['boosts_action_df']['clip_id'].dropna())
        
        # Some boosts clips might not be in total_clip_df (deleted clips)
        orphan_boosts = boost_clip_ids - total_clip_ids
        if orphan_boosts:
            logger.info(f"Found {len(orphan_boosts)} boost clips not in total_clip_df")
    
    logger.info("Data validation complete") 