"""
Prophet Data Utilities
Utilities for preprocessing time series data from various sources
"""

import pandas as pd
from typing import List, Dict, Any, Optional, Union
from datetime import datetime
import csv
import io
from app.utils.helpers import get_logger


class ProphetDataUtils:
    """Utilities for preparing time series data for Prophet"""

    def __init__(self):
        self.logger = get_logger("prophet_data_utils")

    def prepare_from_json(self, data: List[Dict], date_column: str = 'ds', value_column: str = 'y') -> pd.DataFrame:
        """
        Prepare DataFrame from JSON data

        Args:
            data: List of dictionaries with date and value columns
            date_column: Name of the date column
            value_column: Name of the value column

        Returns:
            Prepared pandas DataFrame
        """
        try:
            df = pd.DataFrame(data)

            # Ensure date column exists and is datetime
            if date_column not in df.columns:
                raise ValueError(f"Date column '{date_column}' not found in data")

            df[date_column] = pd.to_datetime(df[date_column])

            # Ensure value column exists and is numeric
            if value_column not in df.columns:
                raise ValueError(f"Value column '{value_column}' not found in data")

            df[value_column] = pd.to_numeric(df[value_column], errors='coerce')

            # Remove rows with NaN values
            df = df.dropna()

            # Sort by date
            df = df.sort_values(date_column)

            self.logger.info(f"Prepared DataFrame from JSON: {len(df)} rows")
            return df

        except Exception as e:
            self.logger.error(f"Error preparing DataFrame from JSON: {str(e)}")
            raise

    def prepare_from_csv_string(self, csv_content: str, date_column: str = 'ds', value_column: str = 'y') -> pd.DataFrame:
        """
        Prepare DataFrame from CSV string content

        Args:
            csv_content: CSV content as string
            date_column: Name of the date column
            value_column: Name of the value column

        Returns:
            Prepared pandas DataFrame
        """
        try:
            # Parse CSV from string
            csv_reader = csv.DictReader(io.StringIO(csv_content))
            data = list(csv_reader)

            if not data:
                raise ValueError("CSV content is empty")

            # Convert to DataFrame
            df = pd.DataFrame(data)

            # Ensure required columns exist
            if date_column not in df.columns:
                raise ValueError(f"Date column '{date_column}' not found in CSV")

            if value_column not in df.columns:
                raise ValueError(f"Value column '{value_column}' not found in CSV")

            # Convert date column
            df[date_column] = pd.to_datetime(df[date_column])

            # Convert value column to numeric
            df[value_column] = pd.to_numeric(df[value_column], errors='coerce')

            # Remove rows with NaN values
            df = df.dropna()

            # Sort by date
            df = df.sort_values(date_column)

            self.logger.info(f"Prepared DataFrame from CSV: {len(df)} rows")
            return df

        except Exception as e:
            self.logger.error(f"Error preparing DataFrame from CSV: {str(e)}")
            raise

    def prepare_from_database_query(self, query_results: List[Dict], date_column: str = 'ds', value_column: str = 'y') -> pd.DataFrame:
        """
        Prepare DataFrame from database query results

        Args:
            query_results: List of dictionaries from database query
            date_column: Name of the date column
            value_column: Name of the value column

        Returns:
            Prepared pandas DataFrame
        """
        try:
            if not query_results:
                raise ValueError("Database query returned no results")

            df = pd.DataFrame(query_results)

            # Ensure required columns exist
            if date_column not in df.columns:
                raise ValueError(f"Date column '{date_column}' not found in query results")

            if value_column not in df.columns:
                raise ValueError(f"Value column '{value_column}' not found in query results")

            # Convert date column
            df[date_column] = pd.to_datetime(df[date_column])

            # Convert value column to numeric
            df[value_column] = pd.to_numeric(df[value_column], errors='coerce')

            # Remove rows with NaN values
            df = df.dropna()

            # Sort by date
            df = df.sort_values(date_column)

            self.logger.info(f"Prepared DataFrame from database query: {len(df)} rows")
            return df

        except Exception as e:
            self.logger.error(f"Error preparing DataFrame from database query: {str(e)}")
            raise

    def aggregate_daily(self, df: pd.DataFrame, date_column: str = 'ds', value_column: str = 'y', agg_func: str = 'sum') -> pd.DataFrame:
        """
        Aggregate data to daily level

        Args:
            df: Input DataFrame
            date_column: Name of the date column
            value_column: Name of the value column
            agg_func: Aggregation function ('sum', 'mean', 'count', etc.)

        Returns:
            Daily aggregated DataFrame
        """
        try:
            # Ensure date column is datetime
            df[date_column] = pd.to_datetime(df[date_column])

            # Set date as index for resampling
            df = df.set_index(date_column)

            # Resample to daily and aggregate
            if agg_func == 'sum':
                daily_df = df[value_column].resample('D').sum()
            elif agg_func == 'mean':
                daily_df = df[value_column].resample('D').mean()
            elif agg_func == 'count':
                daily_df = df[value_column].resample('D').count()
            elif agg_func == 'max':
                daily_df = df[value_column].resample('D').max()
            elif agg_func == 'min':
                daily_df = df[value_column].resample('D').min()
            else:
                raise ValueError(f"Unsupported aggregation function: {agg_func}")

            # Reset index and rename columns
            daily_df = daily_df.reset_index()
            daily_df.columns = [date_column, value_column]

            # Remove NaN values
            daily_df = daily_df.dropna()

            self.logger.info(f"Aggregated to daily data: {len(daily_df)} rows")
            return daily_df

        except Exception as e:
            self.logger.error(f"Error aggregating to daily: {str(e)}")
            raise

    def aggregate_weekly(self, df: pd.DataFrame, date_column: str = 'ds', value_column: str = 'y', agg_func: str = 'sum') -> pd.DataFrame:
        """
        Aggregate data to weekly level

        Args:
            df: Input DataFrame
            date_column: Name of the date column
            value_column: Name of the value column
            agg_func: Aggregation function

        Returns:
            Weekly aggregated DataFrame
        """
        try:
            df[date_column] = pd.to_datetime(df[date_column])
            df = df.set_index(date_column)

            if agg_func == 'sum':
                weekly_df = df[value_column].resample('W').sum()
            elif agg_func == 'mean':
                weekly_df = df[value_column].resample('W').mean()
            elif agg_func == 'count':
                weekly_df = df[value_column].resample('W').count()
            else:
                raise ValueError(f"Unsupported aggregation function: {agg_func}")

            weekly_df = weekly_df.reset_index()
            weekly_df.columns = [date_column, value_column]
            weekly_df = weekly_df.dropna()

            self.logger.info(f"Aggregated to weekly data: {len(weekly_df)} rows")
            return weekly_df

        except Exception as e:
            self.logger.error(f"Error aggregating to weekly: {str(e)}")
            raise

    def aggregate_monthly(self, df: pd.DataFrame, date_column: str = 'ds', value_column: str = 'y', agg_func: str = 'sum') -> pd.DataFrame:
        """
        Aggregate data to monthly level

        Args:
            df: Input DataFrame
            date_column: Name of the date column
            value_column: Name of the value column
            agg_func: Aggregation function

        Returns:
            Monthly aggregated DataFrame
        """
        try:
            df[date_column] = pd.to_datetime(df[date_column])
            df = df.set_index(date_column)

            if agg_func == 'sum':
                monthly_df = df[value_column].resample('M').sum()
            elif agg_func == 'mean':
                monthly_df = df[value_column].resample('M').mean()
            elif agg_func == 'count':
                monthly_df = df[value_column].resample('M').count()
            else:
                raise ValueError(f"Unsupported aggregation function: {agg_func}")

            monthly_df = monthly_df.reset_index()
            monthly_df.columns = [date_column, value_column]
            monthly_df = monthly_df.dropna()

            self.logger.info(f"Aggregated to monthly data: {len(monthly_df)} rows")
            return monthly_df

        except Exception as e:
            self.logger.error(f"Error aggregating to monthly: {str(e)}")
            raise

    def fill_missing_dates(self, df: pd.DataFrame, date_column: str = 'ds', value_column: str = 'y', fill_value: Union[float, str] = 0) -> pd.DataFrame:
        """
        Fill missing dates in time series

        Args:
            df: Input DataFrame
            date_column: Name of the date column
            value_column: Name of the value column
            fill_value: Value to fill missing dates (0, 'mean', 'median', 'ffill', 'bfill')

        Returns:
            DataFrame with missing dates filled
        """
        try:
            df[date_column] = pd.to_datetime(df[date_column])
            df = df.set_index(date_column)

            # Create complete date range
            date_range = pd.date_range(start=df.index.min(), end=df.index.max(), freq='D')
            df_complete = df.reindex(date_range)

            # Fill missing values
            if fill_value == 'mean':
                df_complete[value_column] = df_complete[value_column].fillna(df[value_column].mean())
            elif fill_value == 'median':
                df_complete[value_column] = df_complete[value_column].fillna(df[value_column].median())
            elif fill_value == 'ffill':
                df_complete[value_column] = df_complete[value_column].fillna(method='ffill')
            elif fill_value == 'bfill':
                df_complete[value_column] = df_complete[value_column].fillna(method='bfill')
            else:
                df_complete[value_column] = df_complete[value_column].fillna(float(fill_value))

            df_complete = df_complete.reset_index()
            df_complete.columns = [date_column, value_column]

            self.logger.info(f"Filled missing dates: {len(df_complete)} rows")
            return df_complete

        except Exception as e:
            self.logger.error(f"Error filling missing dates: {str(e)}")
            raise

    def validate_time_series(self, df: pd.DataFrame, date_column: str = 'ds', value_column: str = 'y') -> Dict[str, Any]:
        """
        Validate time series data quality

        Args:
            df: DataFrame to validate
            date_column: Name of the date column
            value_column: Name of the value column

        Returns:
            Validation results
        """
        validation = {
            'is_valid': True,
            'issues': [],
            'stats': {}
        }

        try:
            # Check for required columns
            if date_column not in df.columns:
                validation['is_valid'] = False
                validation['issues'].append(f"Missing date column: {date_column}")
                return validation

            if value_column not in df.columns:
                validation['is_valid'] = False
                validation['issues'].append(f"Missing value column: {value_column}")
                return validation

            # Check data types
            if not pd.api.types.is_datetime64_any_dtype(df[date_column]):
                validation['issues'].append(f"Date column {date_column} is not datetime type")

            if not pd.api.types.is_numeric_dtype(df[value_column]):
                validation['issues'].append(f"Value column {value_column} is not numeric type")

            # Check for missing values
            missing_dates = df[date_column].isnull().sum()
            missing_values = df[value_column].isnull().sum()

            if missing_dates > 0:
                validation['issues'].append(f"Found {missing_dates} missing dates")

            if missing_values > 0:
                validation['issues'].append(f"Found {missing_values} missing values")

            # Basic statistics
            validation['stats'] = {
                'total_rows': len(df),
                'date_range': {
                    'start': df[date_column].min().isoformat() if len(df) > 0 else None,
                    'end': df[date_column].max().isoformat() if len(df) > 0 else None
                },
                'value_stats': {
                    'min': float(df[value_column].min()) if len(df) > 0 else None,
                    'max': float(df[value_column].max()) if len(df) > 0 else None,
                    'mean': float(df[value_column].mean()) if len(df) > 0 else None,
                    'std': float(df[value_column].std()) if len(df) > 0 else None
                }
            }

            # Check for duplicates
            duplicates = df.duplicated(subset=[date_column]).sum()
            if duplicates > 0:
                validation['issues'].append(f"Found {duplicates} duplicate dates")

            # Check for monotonic dates
            if not df[date_column].is_monotonic_increasing:
                validation['issues'].append("Dates are not in chronological order")

        except Exception as e:
            validation['is_valid'] = False
            validation['issues'].append(f"Validation error: {str(e)}")

        return validation