"""
Prophet Controller
Handles API endpoints for Facebook Prophet forecasting operations
"""

from typing import Dict, Any, Optional
from flask import request, g
from app.controllers.base_controller import BaseController
from app.services.prophet_service import ProphetService
from app.services.oms_data_service import OMSDataService
from app.utils.helpers import get_logger


class ProphetController(BaseController):
    """Controller for Prophet forecasting operations"""

    def __init__(self):
        super().__init__()
        self.prophet_service = ProphetService()
        self.oms_data_service = OMSDataService()
        self.logger = get_logger("prophet_controller")

    def train_model(self, data: Dict[str, Any]) -> tuple:
        """
        Train a Prophet model

        Args:
            data: Request data containing time series data and configuration

        Returns:
            API response tuple
        """
        try:
            if hasattr(data, 'model_dump'):
                data = data.model_dump()

            # Validate required parameters
            validation_error = self.validate_required_params(data, ['data'])
            if validation_error:
                return validation_error

            time_series_data = data['data']
            model_id = data.get('model_id')
            config = data.get('config') or {}

            # Validate data format
            is_valid, error_msg = self.prophet_service.validate_data_format(time_series_data)
            if not is_valid:
                return self.bad_request({"message": "Invalid data format", "error": error_msg})

            # Train the model
            result = self.prophet_service.train_model(
                data=time_series_data,
                model_id=model_id,
                config=config
            )

            return self.created(result, "Model trained successfully")

        except ValueError as e:
            return self.bad_request({"message": str(e)})
        except Exception as e:
            self.logger.error(f"Error training model: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to train model", "error": str(e)})

    def generate_forecast(self, model_id: str, params: Dict[str, Any]) -> tuple:
        """
        Generate forecast for a trained model

        Args:
            model_id: ID of the trained model
            params: Forecast parameters

        Returns:
            API response tuple
        """
        try:
            if hasattr(params, 'model_dump'):
                params = params.model_dump()

            periods = params.get('periods', 30)
            freq = params.get('freq', 'D')
            include_history = params.get('include_history', True)

            # Validate periods
            if not isinstance(periods, int) or periods < 1 or periods > 365:
                return self.bad_request({"message": "Periods must be an integer between 1 and 365"})

            # Validate frequency
            valid_freqs = ['D', 'H', 'W', 'M', 'ME', 'Q', 'QE', 'Y', 'YE', 'A']
            if freq not in valid_freqs:
                return self.bad_request({
                    "message": f"Invalid frequency. Must be one of: {', '.join(valid_freqs)}"
                })

            # Generate forecast
            result = self.prophet_service.generate_forecast(
                model_id=model_id,
                periods=periods,
                freq=freq,
                include_history=include_history
            )

            return self.ok(result, "Forecast generated successfully")

        except ValueError as e:
            return self.bad_request({"message": str(e)})
        except Exception as e:
            self.logger.error(f"Error generating forecast: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to generate forecast", "error": str(e)})

    def get_forecast_plot(self, model_id: str, params: Dict[str, Any]) -> tuple:
        """
        Generate forecast plot for a trained model

        Args:
            model_id: ID of the trained model
            params: Plot parameters

        Returns:
            API response tuple
        """
        try:
            if hasattr(params, 'model_dump'):
                params = params.model_dump()

            periods = params.get('periods', 30)
            freq = params.get('freq', 'D')
            width = params.get('width', 800)
            height = params.get('height', 600)

            # Validate parameters
            if not isinstance(periods, int) or periods < 1 or periods > 365:
                return self.bad_request({"message": "Periods must be an integer between 1 and 365"})

            if not isinstance(width, int) or width < 400 or width > 2000:
                return self.bad_request({"message": "Width must be an integer between 400 and 2000"})

            if not isinstance(height, int) or height < 300 or height > 1500:
                return self.bad_request({"message": "Height must be an integer between 300 and 1500"})

            # Generate plot
            plot_data = self.prophet_service.generate_plot(
                model_id=model_id,
                periods=periods,
                freq=freq,
                width=width,
                height=height
            )

            return self.ok({"plot": plot_data}, "Plot generated successfully")

        except ValueError as e:
            return self.bad_request({"message": str(e)})
        except Exception as e:
            self.logger.error(f"Error generating plot: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to generate plot", "error": str(e)})

    def get_model_info(self, model_id: str) -> tuple:
        """
        Get information about a trained model

        Args:
            model_id: ID of the model

        Returns:
            API response tuple
        """
        try:
            model_info = self.prophet_service.get_model_info(model_id)

            if model_info is None:
                return self.not_found({"message": f"Model {model_id} not found"})

            return self.ok(model_info, "Model information retrieved successfully")

        except Exception as e:
            self.logger.error(f"Error getting model info: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to get model information", "error": str(e)})

    def list_models(self) -> tuple:
        """
        List all trained models

        Returns:
            API response tuple
        """
        try:
            models = self.prophet_service.list_models()
            return self.ok({"models": models, "count": len(models)}, "Models listed successfully")

        except Exception as e:
            self.logger.error(f"Error listing models: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to list models", "error": str(e)})

    def delete_model(self, model_id: str) -> tuple:
        """
        Delete a trained model

        Args:
            model_id: ID of the model to delete

        Returns:
            API response tuple
        """
        try:
            deleted = self.prophet_service.delete_model(model_id)

            if not deleted:
                return self.not_found({"message": f"Model {model_id} not found"})

            return self.ok({"model_id": model_id}, "Model deleted successfully")

        except Exception as e:
            self.logger.error(f"Error deleting model: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to delete model", "error": str(e)})

    def validate_data(self, data: Dict[str, Any]) -> tuple:
        """
        Validate time series data format

        Args:
            data: Request data containing time series data

        Returns:
            API response tuple
        """
        try:
            if hasattr(data, 'model_dump'):
                data = data.model_dump()

            # Validate required parameters
            validation_error = self.validate_required_params(data, ['data'])
            if validation_error:
                return validation_error

            time_series_data = data['data']

            # Validate data format
            is_valid, error_msg = self.prophet_service.validate_data_format(time_series_data)

            result = {
                "valid": is_valid,
                "data_points": len(time_series_data) if is_valid else 0
            }

            if not is_valid:
                result["error"] = error_msg
                return self.bad_request(result)

            # Compute date_range and missing_values from valid data
            try:
                import pandas as pd
                df = pd.DataFrame(time_series_data)
                df['ds'] = pd.to_datetime(df['ds'])
                df['y'] = pd.to_numeric(df['y'], errors='coerce')
                result["date_range"] = {
                    "start": df['ds'].min().strftime('%Y-%m-%d'),
                    "end": df['ds'].max().strftime('%Y-%m-%d')
                }
                result["missing_values"] = int(df['y'].isna().sum())
            except Exception:
                pass

            return self.ok(result, "Data validation successful")

        except Exception as e:
            self.logger.error(f"Error validating data: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to validate data", "error": str(e)})

    def train_from_database(self, data: Dict[str, Any]) -> tuple:
        """
        Train a Prophet model using data from OMS database

        Args:
            data: Request data containing table configuration

        Returns:
            API response tuple
        """
        try:
            # Validate required parameters
            validation_error = self.validate_required_params(data, ['table_key'])
            if validation_error:
                return validation_error

            table_key = data['table_key']
            start_date = data.get('start_date')
            end_date = data.get('end_date')
            aggregation = data.get('aggregation', 'daily')
            filters = data.get('filters', {})
            model_id = data.get('model_id')
            config = data.get('config', {})

            # Extract data from database
            time_series_data = self.oms_data_service.get_time_series_data(
                table_key=table_key,
                start_date=start_date,
                end_date=end_date,
                filters=filters,
                aggregation=aggregation
            )

            if not time_series_data:
                return self.bad_request({"message": "No data found for the specified criteria"})

            # Validate data format
            is_valid, error_msg = self.prophet_service.validate_data_format(time_series_data)
            if not is_valid:
                return self.bad_request({"message": "Invalid data format", "error": error_msg})

            # Train the model
            result = self.prophet_service.train_model(
                data=time_series_data,
                model_id=model_id,
                config=config
            )

            result['data_source'] = {
                'table_key': table_key,
                'start_date': start_date,
                'end_date': end_date,
                'aggregation': aggregation,
                'filters': filters
            }

            return self.created(result, "Model trained from database successfully")

        except ValueError as e:
            return self.bad_request({"message": str(e)})
        except Exception as e:
            self.logger.error(f"Error training model from database: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to train model from database", "error": str(e)})

    def get_data_preview(self, table_key: str, params: Dict[str, Any]) -> tuple:
        """
        Get a preview of time series data from database

        Args:
            table_key: Table configuration key
            params: Preview parameters

        Returns:
            API response tuple
        """
        try:
            limit = params.get('limit', 10)
            start_date = params.get('start_date')
            end_date = params.get('end_date')

            # Validate limit
            if not isinstance(limit, int) or limit < 1 or limit > 100:
                return self.bad_request({"message": "Limit must be an integer between 1 and 100"})

            # Get data preview
            preview_data = self.oms_data_service.get_data_preview(
                table_key=table_key,
                limit=limit,
                start_date=start_date,
                end_date=end_date
            )

            return self.ok({
                "table_key": table_key,
                "preview_data": preview_data,
                "count": len(preview_data)
            }, "Data preview retrieved successfully")

        except ValueError as e:
            return self.bad_request({"message": str(e)})
        except Exception as e:
            self.logger.error(f"Error getting data preview: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to get data preview", "error": str(e)})

    def get_available_tables(self) -> tuple:
        """
        Get list of available OMS tables for time series analysis

        Returns:
            API response tuple
        """
        try:
            tables = self.oms_data_service.get_available_tables()
            table_info = {}

            for table_key in tables:
                info = self.oms_data_service.get_table_info(table_key)
                if info:
                    table_info[table_key] = info

            return self.ok({
                "tables": tables,
                "table_info": table_info
            }, "Available tables retrieved successfully")

        except Exception as e:
            self.logger.error(f"Error getting available tables: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to get available tables", "error": str(e)})

    def validate_table_config(self, table_key: str) -> tuple:
        """
        Validate a table configuration

        Args:
            table_key: Table configuration key

        Returns:
            API response tuple
        """
        try:
            is_valid, error_msg = self.oms_data_service.validate_table_config(table_key)

            result = {
                "table_key": table_key,
                "is_valid": is_valid
            }

            if not is_valid:
                result["error"] = error_msg
                return self.bad_request(result)

            return self.ok(result, "Table configuration is valid")

        except Exception as e:
            self.logger.error(f"Error validating table config: {str(e)}", exc_info=True)
            return self.server_error({"message": "Failed to validate table configuration", "error": str(e)})