"""
Regression Controller
Handles HTTP layer for ML model training, scoring, and prediction.
Follows BaseController(ApiResponseMixin) pattern from prophet_controller.py.
"""

from typing import Dict, Any, Optional

from app.controllers.base_controller import BaseController
from app.services.regression_service import RegressionService
from app.services.oms_regression_data_service import OMSRegressionDataService
from app.utils.helpers import get_logger


class RegressionController(BaseController):
    """Controller for XGBoost regression and classification operations."""

    def __init__(self):
        super().__init__()
        self.service = RegressionService()
        self.data_service = OMSRegressionDataService()
        self.logger = get_logger("regression_controller")

    # ------------------------------------------------------------------
    # Training
    # ------------------------------------------------------------------

    def train_collection_risk(self, data: Dict[str, Any]) -> tuple:
        """Train collection risk classifier for a company."""
        try:
            if hasattr(data, 'model_dump'):
                data = data.model_dump()

            company_id = data.get('company_id')
            if not company_id:
                return self.bad_request(message="company_id is required")

            model_id = f"collection_risk_co{company_id}"

            df = self.data_service.get_collection_risk_training_data(int(company_id))
            if df.empty or len(df) < 5:
                return self.unprocessable_entity(
                    message=f"Insufficient training data for company {company_id} "
                            f"(need ≥5 clients with history, got {len(df)})"
                )

            features = self.data_service.COLLECTION_FEATURES
            X = df[features]
            y = df['is_high_risk']

            result = self.service.train(
                model_id=model_id,
                X=X,
                y=y,
                model_type='classifier',
                config=data.get('config'),
            )
            return self.created(result, "Collection risk model trained successfully")

        except Exception as e:
            self.logger.error(f"train_collection_risk error: {str(e)}", exc_info=True)
            return self.server_error(message=f"Training failed: {str(e)}")

    def train_sales_target(self, data: Dict[str, Any]) -> tuple:
        """Train sales target attainment regressor for a company."""
        try:
            if hasattr(data, 'model_dump'):
                data = data.model_dump()

            company_id = data.get('company_id')
            if not company_id:
                return self.bad_request(message="company_id is required")

            model_id = f"sales_target_co{company_id}"

            df = self.data_service.get_sales_target_training_data(int(company_id))
            if df.empty or len(df) < 5:
                return self.unprocessable_entity(
                    message=f"Insufficient training data for company {company_id} "
                            f"(need ≥5 district×season rows, got {len(df)})"
                )

            features = self.data_service.SALES_TARGET_FEATURES
            X = df[[f for f in features if f in df.columns]]
            y = df['attainment_pct']

            result = self.service.train(
                model_id=model_id,
                X=X,
                y=y,
                model_type='regressor',
                config=data.get('config'),
            )
            return self.created(result, "Sales target model trained successfully")

        except Exception as e:
            self.logger.error(f"train_sales_target error: {str(e)}", exc_info=True)
            return self.server_error(message=f"Training failed: {str(e)}")

    # ------------------------------------------------------------------
    # Scoring / Prediction
    # ------------------------------------------------------------------

    def score_collection_risk(self, data: Dict[str, Any]) -> tuple:
        """Score clients for collection payment delay risk."""
        try:
            if hasattr(data, 'model_dump'):
                data = data.model_dump()

            company_id = data.get('company_id')
            if not company_id:
                return self.bad_request(message="company_id is required")

            client_ids = data.get('client_ids') or None
            model_id = f"collection_risk_co{company_id}"

            features_df = self.data_service.get_collection_risk_scoring_features(
                int(company_id), client_ids
            )
            if features_df.empty:
                return self.ok([], message="No scoreable clients found")

            probabilities = self.service.predict_proba(model_id, features_df)

            results = []
            for i, (_, row) in enumerate(features_df.iterrows()):
                prob = round(float(probabilities[i]), 4)
                results.append({
                    'client_id': int(row['client_id']),
                    'delay_probability': prob,
                    'risk_level': RegressionService.collection_risk_level(prob),
                })

            return self.ok(results, f"Scored {len(results)} clients")

        except ValueError as e:
            return self.unprocessable_entity(message=str(e))
        except Exception as e:
            self.logger.error(f"score_collection_risk error: {str(e)}", exc_info=True)
            return self.server_error(message=f"Scoring failed: {str(e)}")

    def predict_sales_target(self, data: Dict[str, Any]) -> tuple:
        """Predict attainment % for district×product_head combinations."""
        try:
            if hasattr(data, 'model_dump'):
                data = data.model_dump()

            company_id = data.get('company_id')
            account_year_id = data.get('account_year_id')
            if not company_id or not account_year_id:
                return self.bad_request(message="company_id and account_year_id are required")

            model_id = f"sales_target_co{company_id}"
            features_df = self.data_service.get_sales_target_prediction_features(
                int(company_id),
                int(account_year_id),
                product_head_ids=data.get('product_head_ids'),
                district_ids=data.get('district_ids'),
            )
            if features_df.empty:
                return self.ok([], message="No prediction targets found for this account year")

            feature_cols = self.data_service.SALES_TARGET_FEATURES
            X = features_df[[f for f in feature_cols if f in features_df.columns]]
            predictions = self.service.predict(model_id, X)

            results = []
            for i, (_, row) in enumerate(features_df.iterrows()):
                pct = round(float(predictions[i]), 2)
                results.append({
                    'product_head_id': int(row['product_head_id']),
                    'district_id': int(row['district_id']),
                    'account_year_id': int(row['account_year_id']),
                    'predicted_attainment_pct': pct,
                    'risk_level': RegressionService.sales_target_risk_level(pct),
                })

            return self.ok(results, f"Predicted {len(results)} combinations")

        except ValueError as e:
            return self.unprocessable_entity(message=str(e))
        except Exception as e:
            self.logger.error(f"predict_sales_target error: {str(e)}", exc_info=True)
            return self.server_error(message=f"Prediction failed: {str(e)}")

    # ------------------------------------------------------------------
    # Model management
    # ------------------------------------------------------------------

    def get_model_info(self, model_id: str) -> tuple:
        try:
            info = self.service.get_model_info(model_id)
            if info is None:
                return self.not_found(message=f"Model '{model_id}' not found")
            return self.ok(info)
        except Exception as e:
            return self.server_error(message=str(e))

    def list_models(self) -> tuple:
        try:
            models = self.service.list_models()
            return self.ok(models, f"{len(models)} regression models found")
        except Exception as e:
            return self.server_error(message=str(e))

    def delete_model(self, model_id: str) -> tuple:
        try:
            deleted = self.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")
        except Exception as e:
            return self.server_error(message=str(e))
